File size: 6,441 Bytes
fd448dd
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
"""GammaNet language model using GammaSpaceModel blocks as the sequence mixer."""

from __future__ import annotations

import site
import sys
from typing import Any, Optional

import torch
import torch.nn as nn
import torch.nn.functional as F

from taoTrain.config import ModelConfig
from taoTrain.core import BaseModel

from .embeddings import FactorizedEmbedding
from .registry import register_architecture


def _import_gamma_space_block():
    """Import GammaSpaceBlock, retrying with user site-packages if needed."""
    try:
        from gamma_space_model import GammaSpaceBlock
        return GammaSpaceBlock
    except ModuleNotFoundError as exc:
        user_site = site.getusersitepackages()
        if user_site and user_site not in sys.path:
            sys.path.append(user_site)
            try:
                from gamma_space_model import GammaSpaceBlock
                return GammaSpaceBlock
            except ModuleNotFoundError:
                pass
        raise ModuleNotFoundError(
            "gamma_net requires the GammaSpaceModel package. Install it with "
            "`pip install \"git+https://github.com/Taotern/GammaSpaceModel.git\"`."
        ) from exc


class GammaMixerBlock(nn.Module):
    """Thin adapter from TaoTrain config/model flow to GammaSpaceBlock."""

    def __init__(self, config: ModelConfig):
        super().__init__()
        gamma_block_cls = _import_gamma_space_block()
        self.block = gamma_block_cls(
            d_model=config.hidden_dim,
            hidden_dim=config.gamma_hidden_dim,
            dt_min=config.gamma_dt_min,
            dt_max=config.gamma_dt_max,
            dt_init=config.gamma_dt_init,
            discretization=config.gamma_discretization,
            prenorm=config.gamma_prenorm,
            residual_scale=config.gamma_residual_scale,
            dropout=config.dropout,
            activation=config.gamma_activation,
            gate=config.gamma_gate,
            use_D=config.gamma_use_D,
            kernel_mode=config.gamma_kernel_mode,
            kernel_threshold=config.gamma_kernel_threshold,
            use_output_linear=config.gamma_use_output_linear,
            gate_bias=config.gamma_gate_bias,
            input_gate=config.gamma_input_gate,
            input_gate_bias=config.gamma_input_gate_bias,
            layer_scale_init=config.gamma_layer_scale_init,
        )

    def forward(

        self,

        x: torch.Tensor,

        attention_mask: Optional[torch.Tensor] = None,

    ) -> torch.Tensor:
        # GammaSpaceBlock expects a padding mask of shape [batch, seq].
        mask = attention_mask
        if mask is not None and mask.dim() > 2:
            mask = None
        y, _ = self.block(x, mask=mask, return_state=False)
        return y


@register_architecture("gamma_net")
class GammaNetLM(BaseModel):
    """Causal LM with TaoNet shell and Gamma Space Model residual blocks."""

    def __init__(self, config: ModelConfig):
        super().__init__(config)

        self.vocab_size = config.vocab_size
        self.d_model = config.hidden_dim
        self.n_layers = config.num_layers
        self.dropout = config.dropout
        self.d_ff = config.hidden_dim_ff if config.hidden_dim_ff is not None else (self.d_model * 4)
        self.use_factorized_embedding = getattr(config, "use_factorized_embedding", False)
        self.d_embed_rank = getattr(config, "d_embed_rank", 96)

        if self.use_factorized_embedding:
            self.token_embedding = FactorizedEmbedding(
                self.vocab_size,
                self.d_model,
                self.d_embed_rank,
            )
        else:
            self.token_embedding = nn.Embedding(self.vocab_size, self.d_model)

        self.embedding_dropout = nn.Dropout(self.dropout)
        self.blocks = nn.ModuleList([GammaMixerBlock(config) for _ in range(self.n_layers)])
        self.final_norm = nn.LayerNorm(self.d_model)
        self.output_head = nn.Linear(self.d_model, self.vocab_size, bias=False)

        self._init_shell_weights()
        self._print_architecture()

    def _init_shell_weights(self) -> None:
        """Initialize only TaoTrain-owned LM shell weights."""
        if isinstance(self.token_embedding, nn.Embedding):
            nn.init.normal_(self.token_embedding.weight, mean=0.0, std=self.config.init_std)
        nn.init.normal_(self.output_head.weight, mean=0.0, std=self.config.init_std)

    def _print_architecture(self) -> None:
        total_params = sum(p.numel() for p in self.parameters())
        trainable_params = sum(p.numel() for p in self.parameters() if p.requires_grad)

        print(f"\n{'=' * 70}")
        print("MODEL ARCHITECTURE - GAMMA_NET (Gamma Space Model)")
        print(f"{'=' * 70}")
        print(f"Embedding type:      {'Factorized' if self.use_factorized_embedding else 'Standard'}")
        print(f"Gamma hidden_dim:    {self.config.gamma_hidden_dim}")
        print(f"Gamma discretization:{self.config.gamma_discretization}")
        print(f"Gamma kernel_mode:   {self.config.gamma_kernel_mode}")
        print(f"Layers:              {self.n_layers}")
        print(f"d_model:             {self.d_model}")
        print(f"Dropout:             {self.dropout}")
        print(f"Max sequence length: {self.config.max_seq_length}")
        print(f"Total Parameters:    {total_params/1e6:>8.2f}M (trainable: {trainable_params/1e6:.2f}M)")
        print(f"{'=' * 70}\n")

    def forward(

        self,

        input_ids: torch.Tensor,

        attention_mask: Optional[torch.Tensor] = None,

        labels: Optional[torch.Tensor] = None,

    ) -> dict[str, Any]:
        x = self.token_embedding(input_ids)
        x = self.embedding_dropout(x)

        for block in self.blocks:
            x = block(x, attention_mask=attention_mask)

        x = self.final_norm(x)
        logits = self.output_head(x)

        loss = None
        if labels is not None:
            logits_flat = logits.view(-1, logits.size(-1))
            labels_flat = labels.view(-1)
            loss = F.cross_entropy(
                logits_flat,
                labels_flat,
                reduction="mean",
                ignore_index=-100,
            )

        return {
            "logits": logits,
            "loss": loss,
        }