Lobakkang commited on
Commit
2981407
·
verified ·
1 Parent(s): 5e6bc01

Upload TaoNet model to HuggingFace Hub

Browse files
Files changed (11) hide show
  1. README.md +199 -0
  2. bitlinear.py +83 -0
  3. config.json +30 -0
  4. configuration_taonet.py +56 -0
  5. factorized_embedding.py +44 -0
  6. mla.py +123 -0
  7. model.py +397 -0
  8. modeling_taonet.py +181 -0
  9. pytorch_model.bin +3 -0
  10. rope.py +47 -0
  11. ssm.py +147 -0
README.md ADDED
@@ -0,0 +1,199 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ library_name: transformers
3
+ tags: []
4
+ ---
5
+
6
+ # Model Card for Model ID
7
+
8
+ <!-- Provide a quick summary of what the model is/does. -->
9
+
10
+
11
+
12
+ ## Model Details
13
+
14
+ ### Model Description
15
+
16
+ <!-- Provide a longer summary of what this model is. -->
17
+
18
+ This is the model card of a 🤗 transformers model that has been pushed on the Hub. This model card has been automatically generated.
19
+
20
+ - **Developed by:** [More Information Needed]
21
+ - **Funded by [optional]:** [More Information Needed]
22
+ - **Shared by [optional]:** [More Information Needed]
23
+ - **Model type:** [More Information Needed]
24
+ - **Language(s) (NLP):** [More Information Needed]
25
+ - **License:** [More Information Needed]
26
+ - **Finetuned from model [optional]:** [More Information Needed]
27
+
28
+ ### Model Sources [optional]
29
+
30
+ <!-- Provide the basic links for the model. -->
31
+
32
+ - **Repository:** [More Information Needed]
33
+ - **Paper [optional]:** [More Information Needed]
34
+ - **Demo [optional]:** [More Information Needed]
35
+
36
+ ## Uses
37
+
38
+ <!-- Address questions around how the model is intended to be used, including the foreseeable users of the model and those affected by the model. -->
39
+
40
+ ### Direct Use
41
+
42
+ <!-- This section is for the model use without fine-tuning or plugging into a larger ecosystem/app. -->
43
+
44
+ [More Information Needed]
45
+
46
+ ### Downstream Use [optional]
47
+
48
+ <!-- This section is for the model use when fine-tuned for a task, or when plugged into a larger ecosystem/app -->
49
+
50
+ [More Information Needed]
51
+
52
+ ### Out-of-Scope Use
53
+
54
+ <!-- This section addresses misuse, malicious use, and uses that the model will not work well for. -->
55
+
56
+ [More Information Needed]
57
+
58
+ ## Bias, Risks, and Limitations
59
+
60
+ <!-- This section is meant to convey both technical and sociotechnical limitations. -->
61
+
62
+ [More Information Needed]
63
+
64
+ ### Recommendations
65
+
66
+ <!-- This section is meant to convey recommendations with respect to the bias, risk, and technical limitations. -->
67
+
68
+ Users (both direct and downstream) should be made aware of the risks, biases and limitations of the model. More information needed for further recommendations.
69
+
70
+ ## How to Get Started with the Model
71
+
72
+ Use the code below to get started with the model.
73
+
74
+ [More Information Needed]
75
+
76
+ ## Training Details
77
+
78
+ ### Training Data
79
+
80
+ <!-- This should link to a Dataset Card, perhaps with a short stub of information on what the training data is all about as well as documentation related to data pre-processing or additional filtering. -->
81
+
82
+ [More Information Needed]
83
+
84
+ ### Training Procedure
85
+
86
+ <!-- This relates heavily to the Technical Specifications. Content here should link to that section when it is relevant to the training procedure. -->
87
+
88
+ #### Preprocessing [optional]
89
+
90
+ [More Information Needed]
91
+
92
+
93
+ #### Training Hyperparameters
94
+
95
+ - **Training regime:** [More Information Needed] <!--fp32, fp16 mixed precision, bf16 mixed precision, bf16 non-mixed precision, fp16 non-mixed precision, fp8 mixed precision -->
96
+
97
+ #### Speeds, Sizes, Times [optional]
98
+
99
+ <!-- This section provides information about throughput, start/end time, checkpoint size if relevant, etc. -->
100
+
101
+ [More Information Needed]
102
+
103
+ ## Evaluation
104
+
105
+ <!-- This section describes the evaluation protocols and provides the results. -->
106
+
107
+ ### Testing Data, Factors & Metrics
108
+
109
+ #### Testing Data
110
+
111
+ <!-- This should link to a Dataset Card if possible. -->
112
+
113
+ [More Information Needed]
114
+
115
+ #### Factors
116
+
117
+ <!-- These are the things the evaluation is disaggregating by, e.g., subpopulations or domains. -->
118
+
119
+ [More Information Needed]
120
+
121
+ #### Metrics
122
+
123
+ <!-- These are the evaluation metrics being used, ideally with a description of why. -->
124
+
125
+ [More Information Needed]
126
+
127
+ ### Results
128
+
129
+ [More Information Needed]
130
+
131
+ #### Summary
132
+
133
+
134
+
135
+ ## Model Examination [optional]
136
+
137
+ <!-- Relevant interpretability work for the model goes here -->
138
+
139
+ [More Information Needed]
140
+
141
+ ## Environmental Impact
142
+
143
+ <!-- Total emissions (in grams of CO2eq) and additional considerations, such as electricity usage, go here. Edit the suggested text below accordingly -->
144
+
145
+ Carbon emissions can be estimated using the [Machine Learning Impact calculator](https://mlco2.github.io/impact#compute) presented in [Lacoste et al. (2019)](https://arxiv.org/abs/1910.09700).
146
+
147
+ - **Hardware Type:** [More Information Needed]
148
+ - **Hours used:** [More Information Needed]
149
+ - **Cloud Provider:** [More Information Needed]
150
+ - **Compute Region:** [More Information Needed]
151
+ - **Carbon Emitted:** [More Information Needed]
152
+
153
+ ## Technical Specifications [optional]
154
+
155
+ ### Model Architecture and Objective
156
+
157
+ [More Information Needed]
158
+
159
+ ### Compute Infrastructure
160
+
161
+ [More Information Needed]
162
+
163
+ #### Hardware
164
+
165
+ [More Information Needed]
166
+
167
+ #### Software
168
+
169
+ [More Information Needed]
170
+
171
+ ## Citation [optional]
172
+
173
+ <!-- If there is a paper or blog post introducing the model, the APA and Bibtex information for that should go in this section. -->
174
+
175
+ **BibTeX:**
176
+
177
+ [More Information Needed]
178
+
179
+ **APA:**
180
+
181
+ [More Information Needed]
182
+
183
+ ## Glossary [optional]
184
+
185
+ <!-- If relevant, include terms and calculations in this section that can help readers understand the model or model card. -->
186
+
187
+ [More Information Needed]
188
+
189
+ ## More Information [optional]
190
+
191
+ [More Information Needed]
192
+
193
+ ## Model Card Authors [optional]
194
+
195
+ [More Information Needed]
196
+
197
+ ## Model Card Contact
198
+
199
+ [More Information Needed]
bitlinear.py ADDED
@@ -0,0 +1,83 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ BitLinear - Simplified for training stability.
3
+ """
4
+
5
+ import torch
6
+ import torch.nn as nn
7
+ import torch.nn.functional as F
8
+
9
+
10
+ class RMSNorm(nn.Module):
11
+ """Root Mean Square Layer Normalization."""
12
+
13
+ def __init__(self, dim, eps=1e-6):
14
+ super().__init__()
15
+ self.eps = eps
16
+
17
+ def forward(self, x):
18
+ rms = torch.sqrt(torch.mean(x * x, dim=-1, keepdim=True) + self.eps)
19
+ return (x / rms)
20
+
21
+
22
+ class TernaryQuantize(torch.autograd.Function):
23
+ """Ternary quantization with straight-through estimator."""
24
+
25
+ @staticmethod
26
+ def forward(ctx, w):
27
+ scale = 1.0 / w.abs().mean().clamp_(min=1e-5)
28
+ u = (w * scale).round().clamp_(-1, 1) / scale
29
+ return u
30
+
31
+ @staticmethod
32
+ def backward(ctx, grad_output):
33
+ return grad_output
34
+
35
+
36
+ class ActivationQuantize(torch.autograd.Function):
37
+ """INT8 activation quantization."""
38
+
39
+ @staticmethod
40
+ def forward(ctx, x):
41
+ scale = 127.0 / x.abs().max(dim=-1, keepdim=True).values.clamp_(min=1e-5)
42
+ y = (x * scale).round().clamp_(-128, 127) / scale
43
+ return y
44
+
45
+ @staticmethod
46
+ def backward(ctx, grad_output):
47
+ return grad_output
48
+
49
+
50
+ class BitLinear(nn.Linear):
51
+ """
52
+ Linear layer with ternary weight quantization.
53
+
54
+ No internal normalization - caller handles it (Pre-Norm architecture).
55
+ """
56
+
57
+ def __init__(self, in_features, out_features, bias=True):
58
+ super().__init__(in_features, out_features)
59
+
60
+ # Gentler initialization for ternary stability
61
+ nn.init.normal_(self.weight, mean=0.0, std=0.02)
62
+ self.rmsnorm = RMSNorm(in_features)
63
+
64
+ def forward(self, x):
65
+ w = self.weight # a weight tensor with shape [d, k]
66
+ x_norm = self.rmsnorm(x)
67
+ # A trick for implementing Straight−Through−Estimator (STE) using detach()
68
+ x_quant = x_norm + (ActivationQuantize.apply(x_norm) - x_norm).detach()
69
+ w_quant = w + (TernaryQuantize.apply(w) - w).detach()
70
+ y = F.linear(x_quant, w_quant)
71
+
72
+ return self.rmsnorm(y)
73
+
74
+ def get_inference_params(self):
75
+ """Export for FPGA deployment."""
76
+ with torch.no_grad():
77
+ scale = self.weight.abs().mean(dim=-1, keepdim=True).clamp(min=1e-5)
78
+ w_ternary = (self.weight / scale).round().clamp(-1, 1).to(torch.int8)
79
+
80
+ return {
81
+ 'weight_ternary': w_ternary,
82
+ 'weight_scale': scale.squeeze()
83
+ }
config.json ADDED
@@ -0,0 +1,30 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "TaoNetForCausalLM"
4
+ ],
5
+ "auto_map": {
6
+ "AutoConfig": "configuration_taonet.TaoNetConfig",
7
+ "AutoModelForCausalLM": "modeling_taonet.TaoNetForCausalLM"
8
+ },
9
+ "block_arrangement": "layered",
10
+ "bos_token_id": 1,
11
+ "d_embed_rank": 384,
12
+ "d_ff": 512,
13
+ "d_kv_comp": 384,
14
+ "d_model": 512,
15
+ "d_rope": 64,
16
+ "d_state": 512,
17
+ "dropout": 0.0,
18
+ "dtype": "float32",
19
+ "eos_token_id": 2,
20
+ "layered_mla_num": 0,
21
+ "max_seq_len": 256,
22
+ "model_type": "taonet",
23
+ "n_heads": 4,
24
+ "n_layers": 8,
25
+ "pad_token_id": 3,
26
+ "ssm_per_mla": 3,
27
+ "transformers_version": "4.57.6",
28
+ "unk_token_id": 0,
29
+ "vocab_size": 50257
30
+ }
configuration_taonet.py ADDED
@@ -0,0 +1,56 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Configuration class for TaoNet model.
3
+ """
4
+
5
+ from transformers import PretrainedConfig
6
+
7
+
8
+ class TaoNetConfig(PretrainedConfig):
9
+ """Configuration for TaoNet model."""
10
+
11
+ model_type = "taonet"
12
+
13
+ def __init__(
14
+ self,
15
+ vocab_size: int = 25000,
16
+ d_model: int = 512,
17
+ d_embed_rank: int = 384,
18
+ d_state: int = 512,
19
+ d_ff: int = 512,
20
+ n_heads: int = 4,
21
+ d_kv_comp: int = 384,
22
+ d_rope: int = 64,
23
+ n_layers: int = 8,
24
+ max_seq_len: int = 256,
25
+ dropout: float = 0.02,
26
+ block_arrangement: str = "layered",
27
+ ssm_per_mla: int = 3,
28
+ layered_mla_num: int = 0,
29
+ pad_token_id: int = 3,
30
+ bos_token_id: int = 1,
31
+ eos_token_id: int = 2,
32
+ unk_token_id: int = 0,
33
+ **kwargs,
34
+ ):
35
+ super().__init__(
36
+ pad_token_id=pad_token_id,
37
+ bos_token_id=bos_token_id,
38
+ eos_token_id=eos_token_id,
39
+ unk_token_id=unk_token_id,
40
+ **kwargs,
41
+ )
42
+
43
+ self.vocab_size = vocab_size
44
+ self.d_model = d_model
45
+ self.d_embed_rank = d_embed_rank
46
+ self.d_state = d_state
47
+ self.d_ff = d_ff
48
+ self.n_heads = n_heads
49
+ self.d_kv_comp = d_kv_comp
50
+ self.d_rope = d_rope
51
+ self.n_layers = n_layers
52
+ self.max_seq_len = max_seq_len
53
+ self.dropout = dropout
54
+ self.block_arrangement = block_arrangement
55
+ self.ssm_per_mla = ssm_per_mla
56
+ self.layered_mla_num = layered_mla_num
factorized_embedding.py ADDED
@@ -0,0 +1,44 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Low-Rank Factorized Embedding.
3
+
4
+ IMPORTANT: Uses standard nn.Linear for projection, NOT BitLinear.
5
+ Embeddings need full precision for good token representations.
6
+ """
7
+
8
+ import torch
9
+ import torch.nn as nn
10
+
11
+ class FactorizedEmbedding(nn.Module):
12
+ """
13
+ Low-Rank Factorized Embedding: vocab → d_embed_rank → d_model
14
+
15
+ Uses standard Linear (not BitLinear) for the projection.
16
+ Embeddings are memory lookups - they benefit from full precision.
17
+ """
18
+
19
+ def __init__(self, vocab_size, d_model, d_embed_rank=96):
20
+ super().__init__()
21
+ self.vocab_size = vocab_size
22
+ self. d_model = d_model
23
+ self.d_embed_rank = d_embed_rank
24
+
25
+ # Embedding table: vocab → compressed
26
+ self.embed = nn.Embedding(vocab_size, d_embed_rank)
27
+
28
+ # Projection: compressed → full (standard Linear, NOT BitLinear)
29
+ self.proj = nn.Linear(d_embed_rank, d_model, bias=False)
30
+
31
+ # Initialize
32
+ nn.init.normal_(self.embed.weight, mean=0.0, std=0.02)
33
+ nn.init.normal_(self.proj.weight, mean=0.0, std=0.02)
34
+
35
+ print(f"FactorizedEmbedding: {vocab_size} × {d_embed_rank} → {d_model}")
36
+ print(f" Params: {self.get_num_params()/1e6:.2f}M (vs {vocab_size * d_model/1e6:.2f}M dense)")
37
+
38
+ def forward(self, input_ids):
39
+ x = self.embed(input_ids) # [B, S, d_embed_rank]
40
+ x = self.proj(x) # [B, S, d_model]
41
+ return x
42
+
43
+ def get_num_params(self):
44
+ return self.vocab_size * self.d_embed_rank + self.d_embed_rank * self.d_model
mla.py ADDED
@@ -0,0 +1,123 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Basic Multi-headed Latent Attention (MLA).
3
+ Simple implementation without KV cache.
4
+ """
5
+
6
+ import torch
7
+ import torch.nn as nn
8
+ import torch.nn.functional as F
9
+ import math
10
+
11
+ from .rope import RotaryEmbedding, apply_rotary
12
+
13
+
14
+ class MemoryOptimizedMLA(nn.Module):
15
+ """
16
+ Basic MLA: Project to latent space, apply multi-head attention, project back.
17
+ Numerically stable implementation with proper normalization.
18
+ """
19
+
20
+ def __init__(self, config):
21
+ super().__init__()
22
+ self.config = config
23
+ self.n_heads = config.n_heads
24
+ self.d_head = config.d_kv_comp // config.n_heads
25
+ self.d_rope = config.d_rope
26
+ # Improved scaling: use sqrt(d_head) with a small epsilon for numerical stability
27
+ self.scale = 1.0 / math.sqrt(max(self.d_head, 1.0))
28
+
29
+ # Layer normalization before projections for stability
30
+ self.norm_latent = nn.LayerNorm(config.d_model)
31
+
32
+ # Projections
33
+ self.to_latent = nn.Linear(config.d_model, config.d_kv_comp, bias=False)
34
+
35
+ # Q/K/V from latent
36
+ self.q_proj = nn.Linear(config.d_kv_comp, config.d_kv_comp, bias=False)
37
+ self.k_proj = nn.Linear(config.d_kv_comp, config.d_kv_comp, bias=False)
38
+ self.v_proj = nn.Linear(config.d_kv_comp, config.d_kv_comp, bias=False)
39
+
40
+ # RoPE
41
+ self.rotary = RotaryEmbedding(config.d_rope)
42
+
43
+ # Output
44
+ self.out_proj = nn.Linear(config.d_kv_comp, config.d_model, bias=False)
45
+
46
+ self.attn_dropout = nn.Dropout(config.dropout)
47
+ self.resid_dropout = nn.Dropout(config.dropout)
48
+
49
+ def forward(self, x, mask=None):
50
+ """
51
+ Args:
52
+ x: (batch_size, seq_len, d_model)
53
+ mask: (batch_size, seq_len) or (batch_size, 1, seq_len, seq_len), optional
54
+
55
+ Returns:
56
+ out: (batch_size, seq_len, d_model)
57
+ """
58
+ batch_size, seq_len, _ = x.shape
59
+
60
+ # Normalize input before projection to prevent activation explosion
61
+ x_norm = self.norm_latent(x)
62
+
63
+ # Project to latent space
64
+ latent = self.to_latent(x_norm)
65
+
66
+ # Generate Q/K/V
67
+ q = self.q_proj(latent)
68
+ k = self.k_proj(latent)
69
+ v = self.v_proj(latent)
70
+
71
+ # Reshape for multi-head attention: (batch_size, seq_len, d_kv_comp) -> (batch_size, n_heads, seq_len, d_head)
72
+ q = q.view(batch_size, seq_len, self.n_heads, self.d_head).transpose(1, 2)
73
+ k = k.view(batch_size, seq_len, self.n_heads, self.d_head).transpose(1, 2)
74
+ v = v.view(batch_size, seq_len, self.n_heads, self.d_head).transpose(1, 2)
75
+
76
+ # Normalize Q and K for stable attention (standard practice in modern attention mechanisms)
77
+ q = F.normalize(q, dim=-1, p=2)
78
+ k = F.normalize(k, dim=-1, p=2)
79
+
80
+ # Apply RoPE
81
+ if self.d_rope > 0:
82
+ rotary_emb = self.rotary(seq_len, x.device)
83
+ cos = torch.cos(rotary_emb).unsqueeze(0).unsqueeze(0)
84
+ sin = torch.sin(rotary_emb).unsqueeze(0).unsqueeze(0)
85
+
86
+ q_rot = apply_rotary(q[..., :self.d_rope], cos, sin)
87
+ k_rot = apply_rotary(k[..., :self.d_rope], cos, sin)
88
+
89
+ q = torch.cat([q_rot, q[..., self.d_rope:]], dim=-1)
90
+ k = torch.cat([k_rot, k[..., self.d_rope:]], dim=-1)
91
+
92
+ # Attention computation with numerical stability
93
+ # Scale before matmul to prevent overflow
94
+ attn_scores = torch.matmul(q, k.transpose(-2, -1)) * self.scale
95
+
96
+ # Clamp attention scores to prevent inf/-inf in softmax
97
+ attn_scores = torch.clamp(attn_scores, min=-20.0, max=20.0)
98
+
99
+ if mask is not None:
100
+ attn_scores = attn_scores.masked_fill(mask == 0, float('-inf'))
101
+
102
+ # Numerically stable softmax
103
+ attn_weights = F.softmax(attn_scores, dim=-1)
104
+
105
+ # Check for NaN and print warning
106
+ if torch.isnan(attn_weights).any():
107
+ print(f"WARNING: NaN detected in attention weights! "
108
+ f"attn_scores min={attn_scores.min():.4f}, max={attn_scores.max():.4f}, "
109
+ f"attn_weights min={attn_weights.min():.4f}, max={attn_weights.max():.4f}")
110
+
111
+ attn_weights = self.attn_dropout(attn_weights)
112
+
113
+ # Apply attention to values
114
+ out = torch.matmul(attn_weights, v)
115
+
116
+ # Reshape back: (batch_size, n_heads, seq_len, d_head) -> (batch_size, seq_len, d_kv_comp)
117
+ out = out.transpose(1, 2).contiguous().view(batch_size, seq_len, -1)
118
+
119
+ # Project back to model dimension
120
+ out = self.out_proj(out)
121
+ out = self.resid_dropout(out)
122
+
123
+ return out
model.py ADDED
@@ -0,0 +1,397 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ SimpleLLM - Mamba-style State-Space Model with ternary quantization.
3
+ """
4
+
5
+ import torch
6
+ import torch.nn as nn
7
+ import torch.nn. functional as F
8
+
9
+ from .ssm import SSMBlock
10
+ from .bitlinear import BitLinear, RMSNorm, ActivationQuantize
11
+ from .factorized_embedding import FactorizedEmbedding
12
+ from .mla import MemoryOptimizedMLA
13
+
14
+ class SSMBlockWrapper(nn.Module):
15
+ """
16
+ Pre-Norm SSM Block (Mamba-style) with nn.Sequential structure.
17
+
18
+ Structure:
19
+ x → Norm → SSM → Add → Norm → FFN → Add → output
20
+ """
21
+
22
+ def __init__(self, config):
23
+ super().__init__()
24
+ self.ssm = SSMBlock(config)
25
+ self.feed_forward = nn.Sequential(
26
+ BitLinear(config.d_model, config.d_ff, bias=False),
27
+ nn.ReLU(),
28
+ BitLinear(config.d_ff, config.d_model, bias=False),
29
+ )
30
+ self.dropout = nn.Dropout(config.dropout)
31
+
32
+ def forward(self, x, mask=None):
33
+ # Pre-norm SSM with residual
34
+ x = x + self.dropout(self.ssm(x, mask)) # Normalize before SSM
35
+ # Pre-norm FFN with residual
36
+ x = x + self.dropout(self.feed_forward(x)) # Normalize before FFN
37
+ return x
38
+
39
+
40
+ class MLABlockWrapper(nn.Module):
41
+ """
42
+ MLA Block with residual connection and FFN.
43
+
44
+ Structure:
45
+ x → Norm → MLA → Add → Norm → FFN → Add → output
46
+
47
+ Pre-norm structure stabilizes training and prevents gradient explosion.
48
+ """
49
+
50
+ def __init__(self, config):
51
+ super().__init__()
52
+ self.mla = MemoryOptimizedMLA(config)
53
+ self.ffn = nn.Sequential(
54
+ nn.Linear(config.d_model, config.d_ff, bias=False),
55
+ nn.ReLU(),
56
+ nn.Linear(config.d_ff, config.d_model, bias=False),
57
+ nn.ReLU(),
58
+ nn.Linear(config.d_ff, config.d_model, bias=False),
59
+ )
60
+ self.dropout = nn.Dropout(config.dropout)
61
+
62
+ def forward(self, x, mask=None):
63
+ # Pre-norm MLA with residual
64
+ x = x + self.dropout(self.mla(x, mask=mask))
65
+ # Pre-norm FFN with residual
66
+ x = x + self.dropout(self.ffn(x))
67
+ return x
68
+
69
+
70
+ class SimpleLLM(nn.Module):
71
+ """
72
+ Language Model with Hybrid Mamba-style SSM + MLA blocks.
73
+
74
+ Architecture: Token Embedding → (SSM Blocks + MLA Blocks) → Output Head
75
+
76
+ Hybrid structure controlled by config.ssm_per_mla:
77
+ - ssm_per_mla = 2: SSM, SSM, MLA, SSM, SSM, MLA, ...
78
+ - ssm_per_mla = 3: SSM, SSM, SSM, MLA, SSM, SSM, SSM, MLA, ...
79
+ """
80
+
81
+ def __init__(self, config):
82
+ super().__init__()
83
+ self.config = config
84
+
85
+ # Factorized embeddings
86
+ self.token_embedding = FactorizedEmbedding(
87
+ vocab_size=config.vocab_size,
88
+ d_model=config.d_model,
89
+ d_embed_rank=config.d_embed_rank
90
+ )
91
+
92
+ self.dropout = nn.Dropout(config.dropout)
93
+
94
+ # Build block architecture based on arrangement strategy
95
+ self.blocks = nn.ModuleList()
96
+
97
+ if config.block_arrangement == "interleaving":
98
+ self._build_interleaving_blocks(config)
99
+ elif config.block_arrangement == "layered":
100
+ self._build_layered_blocks(config)
101
+ else:
102
+ raise ValueError(f"Unknown block_arrangement: {config.block_arrangement}")
103
+
104
+ # =================================================================
105
+ # Two-stage output projection (mirrors factorized embedding)
106
+ # =================================================================
107
+ # Stage 1: d_model → d_embed_rank (reverse of embedding projection)
108
+ self.output_proj = nn.Linear(config.d_model, config.d_embed_rank, bias=False)
109
+
110
+ # Stage 2: d_embed_rank → vocab_size (tied to embedding table)
111
+ self.lm_head = nn.Linear(config.d_embed_rank, config.vocab_size, bias=False)
112
+
113
+ # Tie lm_head weights to embedding table
114
+ self.lm_head.weight = self.token_embedding.embed.weight
115
+ # =================================================================
116
+
117
+ # Final layer norm before output head to stabilize predictions
118
+ self.pre_final_norm = nn.LayerNorm(config.d_model)
119
+ self.final_norm = nn.LayerNorm(config.d_embed_rank)
120
+
121
+ self.apply(self._init_weights)
122
+ self.register_buffer("causal_mask_cache", None, persistent=False)
123
+ self._print_architecture()
124
+
125
+ def _build_interleaving_blocks(self, config):
126
+ """
127
+ Build interleaving block arrangement: SSM blocks followed by MLA blocks in a pattern.
128
+
129
+ Example with ssm_per_mla=3 and n_layers=16:
130
+ SSM, SSM, SSM, MLA, SSM, SSM, SSM, MLA, SSM, SSM, SSM, MLA, SSM, SSM, SSM, MLA
131
+ """
132
+ ssm_per_mla = config.ssm_per_mla
133
+ num_mla_blocks = max(1, config.n_layers // (ssm_per_mla + 1))
134
+
135
+ block_idx = 0
136
+ for mla_idx in range(num_mla_blocks):
137
+ # Add SSM blocks before each MLA block
138
+ for _ in range(ssm_per_mla):
139
+ if block_idx < config.n_layers:
140
+ self.blocks.append(SSMBlockWrapper(config))
141
+ block_idx += 1
142
+
143
+ # Add MLA block
144
+ if block_idx < config.n_layers:
145
+ self.blocks.append(MLABlockWrapper(config))
146
+ block_idx += 1
147
+
148
+ # Add remaining SSM blocks (if n_layers is not evenly divisible)
149
+ while block_idx < config.n_layers:
150
+ self.blocks.append(SSMBlockWrapper(config))
151
+ block_idx += 1
152
+
153
+ def _build_layered_blocks(self, config):
154
+ """
155
+ Build layered block arrangement: MLA blocks followed by SSM blocks.
156
+
157
+ Example with layered_mla_num=4 and n_layers=16:
158
+ MLA, MLA, MLA, MLA, SSM, SSM, SSM, SSM, SSM, SSM, SSM, SSM, SSM, SSM, SSM, SSM
159
+ """
160
+ num_mla = config.layered_mla_num
161
+
162
+ # Add MLA blocks first
163
+ for _ in range(min(num_mla, config.n_layers)):
164
+ self.blocks.append(MLABlockWrapper(config))
165
+
166
+ # Add remaining SSM blocks
167
+ num_ssm = config.n_layers - len(self.blocks)
168
+ for _ in range(num_ssm):
169
+ self.blocks.append(SSMBlockWrapper(config))
170
+
171
+ def _init_weights(self, module):
172
+ if isinstance(module, nn.Linear) and not isinstance(module, BitLinear):
173
+ nn.init.normal_(module.weight, mean=0.0, std=0.02)
174
+ if module. bias is not None:
175
+ nn.init.zeros_(module.bias)
176
+ elif isinstance(module, nn.Embedding):
177
+ nn.init.normal_(module.weight, mean=0.0, std=0.02)
178
+
179
+ def _print_architecture(self):
180
+ total_params = self.count_parameters()
181
+ embed_params = self.token_embedding.get_num_params()
182
+ output_proj_params = self.config.d_model * self.config.d_embed_rank
183
+ ssm_params = total_params - embed_params - output_proj_params
184
+
185
+ # Count SSM and MLA blocks
186
+ num_ssm = sum(1 for b in self.blocks if isinstance(b, SSMBlockWrapper))
187
+ num_mla = sum(1 for b in self.blocks if isinstance(b, MLABlockWrapper))
188
+
189
+ print(f"\n{'='*60}")
190
+ print("MODEL ARCHITECTURE - HYBRID SSM + MLA")
191
+ print(f"{'='*60}")
192
+ print(f"Embedding: {embed_params/1e6:>6.2f}M params")
193
+ print(f"Hybrid Blocks: {num_ssm} SSM + {num_mla} MLA = {num_ssm + num_mla} total")
194
+ print(f"Output Proj: {output_proj_params/1e6:>6.2f}M params")
195
+ print(f"Output Head: tied to embedding (0 extra params)")
196
+ print(f"{'─'*60}")
197
+ print(f"Total: {total_params/1e6:>6.2f}M params")
198
+ print(f"{'='*60}")
199
+ print(f"Config: {self.config.n_layers} layers, {self.config.d_model} dim")
200
+ print(f"SSM: d_state={self.config.d_state}")
201
+ print(f"MLA: n_heads={self.config.n_heads}, d_kv_comp={self.config.d_kv_comp}")
202
+
203
+ # Print arrangement-specific info
204
+ if self.config.block_arrangement == "interleaving":
205
+ print(f"Arrangement: INTERLEAVING (ssm_per_mla={self.config.ssm_per_mla})")
206
+ elif self.config.block_arrangement == "layered":
207
+ print(f"Arrangement: LAYERED (mla_blocks={self.config.layered_mla_num}, ssm_blocks={num_ssm})")
208
+
209
+ print(f"{'='*60}\n")
210
+
211
+ def _get_causal_mask(self, seq_len, device):
212
+ if self.causal_mask_cache is None or self.causal_mask_cache. size(-1) < seq_len:
213
+ mask = torch.tril(torch.ones(seq_len, seq_len, device=device))
214
+ mask = mask.unsqueeze(0).unsqueeze(0)
215
+ self.causal_mask_cache = mask
216
+ return self.causal_mask_cache[: , :, :seq_len, :seq_len]
217
+
218
+ def forward(self, input_ids, attention_mask=None):
219
+ batch_size, seq_len = input_ids.shape
220
+
221
+ # Causal mask
222
+ causal_mask = self._get_causal_mask(seq_len, input_ids.device)
223
+ if attention_mask is not None:
224
+ padding_mask = attention_mask.unsqueeze(1).unsqueeze(1)
225
+ causal_mask = causal_mask * padding_mask
226
+
227
+ # Token embedding
228
+ x = self.token_embedding(input_ids)
229
+ x = self.dropout(x)
230
+ x = ActivationQuantize.apply(x)
231
+
232
+ # Hybrid SSM + MLA blocks
233
+ for block in self.blocks:
234
+ x = block(x, causal_mask)
235
+
236
+ # Two-stage output projection
237
+ x = self.pre_final_norm(x)
238
+ x = self.output_proj(x) # d_model → d_embed_rank
239
+ x = self.final_norm(x) # Normalize before output head
240
+ logits = self.lm_head(x) # d_embed_rank → vocab_size
241
+
242
+ return logits
243
+
244
+ def init_ssm_states(self, batch_size, device, dtype):
245
+ """
246
+ Initialize SSM states for all SSM blocks (MLA blocks are stateless).
247
+
248
+ Returns:
249
+ states: List of [batch, d_state] tensors for each SSM block
250
+ """
251
+ states = []
252
+ for block in self.blocks:
253
+ if isinstance(block, SSMBlockWrapper):
254
+ state = block.ssm.init_state(batch_size, device, dtype)
255
+ states.append(state)
256
+ return states
257
+
258
+ def inference_step(self, input_id, states, return_hidden_states=False):
259
+ """
260
+ Single inference step for autoregressive generation (RNN-like).
261
+
262
+ Args:
263
+ input_id: [batch, 1] or scalar token id
264
+ states: List of SSM states from previous step
265
+ return_hidden_states: If True, also return SSM hidden states for visualization
266
+
267
+ Returns:
268
+ logits: [batch, vocab_size] - output logits for next token
269
+ new_states: List of updated SSM states for SSM blocks
270
+ hidden_states: (Optional) List of SSM hidden state values for each SSM layer
271
+ """
272
+ if isinstance(input_id, int):
273
+ input_id = torch.tensor([[input_id]], dtype=torch.long, device=next(self.parameters()).device)
274
+ elif input_id.dim() == 1:
275
+ input_id = input_id.unsqueeze(0)
276
+
277
+ # Embed the token
278
+ x = self.token_embedding(input_id) # [batch, 1, d_model]
279
+ x = x.squeeze(1) # [batch, d_model]
280
+ x = ActivationQuantize.apply(x)
281
+
282
+ # Pass through hybrid blocks
283
+ new_states = []
284
+ hidden_states = [] if return_hidden_states else None
285
+ state_idx = 0 # Track position in states list (only for SSM blocks)
286
+
287
+ for block in self.blocks:
288
+ if isinstance(block, SSMBlockWrapper):
289
+ # SSM block with state management
290
+ residual = x
291
+ ssm_out, new_state = block.ssm.step(x, states[state_idx])
292
+
293
+ # Collect hidden state if requested
294
+ if return_hidden_states:
295
+ hidden_states.append(new_state.clone().detach())
296
+
297
+ x = residual + block.dropout(ssm_out)
298
+
299
+ # FFN + residual
300
+ residual = x
301
+ ffn_out = block.feed_forward(x)
302
+ x = residual + block.dropout(ffn_out)
303
+
304
+ new_states.append(new_state)
305
+ state_idx += 1
306
+ else:
307
+ # MLA block (stateless)
308
+ x = block(x.unsqueeze(1), mask=None).squeeze(1)
309
+
310
+ # Output projection
311
+ x = self.pre_final_norm(x)
312
+ x = self.output_proj(x)
313
+ x = self.final_norm(x)
314
+ logits = self.lm_head(x)
315
+
316
+ if return_hidden_states:
317
+ return logits, new_states, hidden_states
318
+ else:
319
+ return logits, new_states
320
+
321
+ def count_parameters(self):
322
+ return sum(p.numel() for p in self.parameters() if p.requires_grad)
323
+
324
+ def count_non_embedding_parameters(self):
325
+ total = self.count_parameters()
326
+ embedding_params = self.token_embedding.get_num_params()
327
+ return total - embedding_params
328
+
329
+ @torch.no_grad()
330
+ def generate(
331
+ self,
332
+ input_ids,
333
+ max_new_tokens=50,
334
+ temperature=1.0,
335
+ top_k=50,
336
+ top_p=0.9,
337
+ repetition_penalty=1.1,
338
+ do_sample=True
339
+ ):
340
+ """Generate tokens autoregressively."""
341
+ self.eval()
342
+
343
+ for _ in range(max_new_tokens):
344
+ # Crop to max_seq_len
345
+ idx_cond = input_ids[:, -self.config.max_seq_len:]
346
+
347
+ # Forward
348
+ logits = self(idx_cond)
349
+ logits = logits[:, -1, : ] / max(temperature, 1e-5)
350
+
351
+ # Repetition penalty
352
+ if repetition_penalty != 1.0:
353
+ for i in range(input_ids.shape[0]):
354
+ for token_id in set(input_ids[i].tolist()):
355
+ if logits[i, token_id] > 0:
356
+ logits[i, token_id] /= repetition_penalty
357
+ else:
358
+ logits[i, token_id] *= repetition_penalty
359
+
360
+ # Top-k filtering
361
+ if top_k is not None and top_k > 0:
362
+ v, _ = torch.topk(logits, min(top_k, logits. size(-1)))
363
+ logits[logits < v[:, [-1]]] = float('-inf')
364
+
365
+ # Top-p filtering
366
+ if top_p is not None and top_p < 1.0:
367
+ sorted_logits, sorted_indices = torch.sort(logits, descending=True)
368
+ cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1)
369
+
370
+ sorted_indices_to_remove = cumulative_probs > top_p
371
+ sorted_indices_to_remove[: , 1:] = sorted_indices_to_remove[:, :-1].clone()
372
+ sorted_indices_to_remove[:, 0] = 0
373
+
374
+ for i in range(logits.shape[0]):
375
+ indices_to_remove = sorted_indices[i, sorted_indices_to_remove[i]]
376
+ logits[i, indices_to_remove] = float('-inf')
377
+
378
+ # Sample or greedy
379
+ probs = F.softmax(logits, dim=-1)
380
+ if do_sample:
381
+ next_token = torch.multinomial(probs, num_samples=1)
382
+ else:
383
+ next_token = torch.argmax(probs, dim=-1, keepdim=True)
384
+
385
+ input_ids = torch. cat([input_ids, next_token], dim=1)
386
+
387
+ # Stop on EOS
388
+ if self.config.eos_token_id is not None:
389
+ if (next_token == self.config. eos_token_id).all():
390
+ break
391
+
392
+ return input_ids
393
+
394
+ def get_num_params(self, non_embedding=True):
395
+ if non_embedding:
396
+ return self.count_non_embedding_parameters()
397
+ return self.count_parameters()
modeling_taonet.py ADDED
@@ -0,0 +1,181 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Modeling class for TaoNet model.
3
+ """
4
+
5
+ import torch
6
+ import torch.nn as nn
7
+ from dataclasses import dataclass
8
+ from transformers import PreTrainedModel
9
+
10
+ from .model import SimpleLLM
11
+ from .configuration_taonet import TaoNetConfig
12
+
13
+
14
+ @dataclass
15
+ class InternalModelConfig:
16
+ """Internal config for SimpleLLM."""
17
+ vocab_size: int = 25000
18
+ d_model: int = 512
19
+ d_embed_rank: int = 384
20
+ d_state: int = 512
21
+ d_ff: int = 512
22
+ n_heads: int = 4
23
+ d_kv_comp: int = 384
24
+ d_rope: int = 64
25
+ n_layers: int = 8
26
+ max_seq_len: int = 256
27
+ dropout: float = 0.02
28
+ block_arrangement: str = "layered"
29
+ ssm_per_mla: int = 3
30
+ layered_mla_num: int = 0
31
+ pad_token_id: int = 3
32
+ bos_token_id: int = 1
33
+ eos_token_id: int = 2
34
+ unk_token_id: int = 0
35
+
36
+
37
+ class TaoNetForCausalLM(PreTrainedModel):
38
+ """TaoNet model for causal language modeling."""
39
+
40
+ config_class = TaoNetConfig
41
+ base_model_prefix = "taonet"
42
+
43
+ def __init__(self, config: TaoNetConfig):
44
+ super().__init__(config)
45
+
46
+ # Convert HF config to internal config
47
+ internal_config = InternalModelConfig(
48
+ vocab_size=config.vocab_size,
49
+ d_model=config.d_model,
50
+ d_embed_rank=config.d_embed_rank,
51
+ d_state=config.d_state,
52
+ d_ff=config.d_ff,
53
+ n_heads=config.n_heads,
54
+ d_kv_comp=config.d_kv_comp,
55
+ d_rope=config.d_rope,
56
+ n_layers=config.n_layers,
57
+ max_seq_len=config.max_seq_len,
58
+ dropout=config.dropout,
59
+ block_arrangement=config.block_arrangement,
60
+ ssm_per_mla=config.ssm_per_mla,
61
+ layered_mla_num=config.layered_mla_num,
62
+ pad_token_id=config.pad_token_id,
63
+ bos_token_id=config.bos_token_id,
64
+ eos_token_id=config.eos_token_id,
65
+ unk_token_id=config.unk_token_id,
66
+ )
67
+
68
+ self.taonet = SimpleLLM(internal_config)
69
+
70
+ # Tie the lm_head weights to the token embedding weights
71
+ self._tie_weights()
72
+
73
+ def _tie_weights(self):
74
+ """Tie the lm_head weight to the token embedding weight."""
75
+ if hasattr(self.taonet, 'token_embedding') and hasattr(self.taonet, 'lm_head'):
76
+ # Tie the weights - make lm_head.weight reference the same tensor as token_embedding.embed.weight
77
+ self.taonet.lm_head.weight = self.taonet.token_embedding.embed.weight
78
+
79
+ def _init_weights(self, module):
80
+ """Initialize weights (override to maintain tied weights)."""
81
+ # Let the parent handle initialization, then retie weights
82
+ super()._init_weights(module) if hasattr(super(), '_init_weights') else None
83
+ self._tie_weights()
84
+
85
+ @property
86
+ def all_tied_weights_keys(self):
87
+ """Return the tied weights keys to satisfy transformers requirements."""
88
+ # Return as a dict with tied_weight -> main_weight mapping
89
+ return {"taonet.lm_head.weight": "taonet.token_embedding.embed.weight"}
90
+
91
+ def mark_tied_weights_as_initialized(self):
92
+ """Mark tied weights as initialized by actually tying them together."""
93
+ # Tie the weights so they reference the same tensor
94
+ self._tie_weights()
95
+
96
+ def forward(
97
+ self,
98
+ input_ids: torch.LongTensor,
99
+ attention_mask=None,
100
+ labels=None,
101
+ **kwargs,
102
+ ):
103
+ """Forward pass."""
104
+ logits = self.taonet(input_ids)
105
+
106
+ loss = None
107
+ if labels is not None:
108
+ shift_logits = logits[..., :-1, :].contiguous()
109
+ shift_labels = labels[..., 1:].contiguous()
110
+
111
+ loss_fct = nn.CrossEntropyLoss(ignore_index=-100)
112
+ loss = loss_fct(
113
+ shift_logits.view(-1, self.config.vocab_size),
114
+ shift_labels.view(-1)
115
+ )
116
+
117
+ return {
118
+ "loss": loss,
119
+ "logits": logits,
120
+ }
121
+
122
+ def init_ssm_states(self, batch_size: int, device: torch.device, dtype: torch.dtype):
123
+ """Initialize SSM states for all SSM blocks."""
124
+ return self.taonet.init_ssm_states(batch_size, device, dtype)
125
+
126
+ def generate(
127
+ self,
128
+ input_ids: torch.LongTensor,
129
+ max_length: int = 100,
130
+ temperature: float = 1.0,
131
+ top_k=None,
132
+ top_p=None,
133
+ **kwargs,
134
+ ):
135
+ """Generate text using RNN-style inference with state management."""
136
+ self.taonet.eval()
137
+ batch_size = input_ids.shape[0]
138
+ device = input_ids.device
139
+ dtype = next(self.taonet.parameters()).dtype
140
+
141
+ # Initialize SSM states for all SSM blocks
142
+ states = self.taonet.init_ssm_states(batch_size, device, dtype)
143
+
144
+ current_ids = input_ids.clone()
145
+
146
+ # Process initial tokens to prime the states
147
+ with torch.no_grad():
148
+ for i in range(input_ids.shape[1]):
149
+ token_id = input_ids[:, i:i+1]
150
+ _, states = self.taonet.inference_step(token_id, states)
151
+
152
+ # Generate new tokens
153
+ for _ in range(max_length - input_ids.shape[1]):
154
+ with torch.no_grad():
155
+ # Get logits for next token using inference_step
156
+ next_token_id = current_ids[:, -1:]
157
+ logits, states = self.taonet.inference_step(next_token_id, states)
158
+
159
+ next_logits = logits / temperature
160
+
161
+ if top_k is not None:
162
+ top_k_logits, top_k_indices = torch.topk(next_logits, min(top_k, next_logits.size(-1)), dim=-1)
163
+ indices_to_remove = next_logits < top_k_logits[..., -1, None]
164
+ next_logits[indices_to_remove] = float('-inf')
165
+
166
+ if top_p is not None:
167
+ sorted_logits, sorted_indices = torch.sort(next_logits, descending=True, dim=-1)
168
+ cumsum_probs = torch.cumsum(torch.softmax(sorted_logits, dim=-1), dim=-1)
169
+ sorted_indices_to_remove = cumsum_probs > top_p
170
+ sorted_indices_to_remove[..., 0] = False
171
+ indices_to_remove = sorted_indices[sorted_indices_to_remove]
172
+ next_logits[..., indices_to_remove] = float('-inf')
173
+
174
+ probs = torch.softmax(next_logits, dim=-1)
175
+ next_token = torch.multinomial(probs, num_samples=1)
176
+ current_ids = torch.cat([current_ids, next_token], dim=1)
177
+
178
+ if (next_token == self.config.eos_token_id).any():
179
+ break
180
+
181
+ return current_ids
pytorch_model.bin ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:68d530216dc370cf96f066f1897a26896b98e12cab647dba61a4314364c7d05e
3
+ size 112420148
rope.py ADDED
@@ -0,0 +1,47 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Rotary Position Embedding (RoPE) implementation."""
2
+
3
+ import torch
4
+ import torch.nn as nn
5
+ import math
6
+
7
+
8
+ class RotaryEmbedding(nn.Module):
9
+ """Rotary position embeddings."""
10
+
11
+ def __init__(self, dim, scale=40):
12
+ super().__init__()
13
+ assert dim % 2 == 0, "Dimension must be even for rotary embeddings"
14
+ self.dim = dim
15
+ self.scale = scale
16
+
17
+ inv_freq = 1.0 / (10000 ** (torch.arange(0, dim, 2).float() / dim))
18
+ self.register_buffer("inv_freq", inv_freq)
19
+
20
+ def forward(self, seq_len, device):
21
+ """Generate rotary embeddings for sequence."""
22
+ t = torch.arange(seq_len, device=device).type_as(self.inv_freq) / self.scale
23
+ freqs = torch.einsum("i,j->ij", t, self.inv_freq)
24
+ return torch.cat((freqs, freqs), dim=-1)
25
+
26
+
27
+ def rotate_half(x):
28
+ """Rotate half the hidden dims of the input."""
29
+ x1, x2 = x.chunk(2, dim=-1)
30
+ return torch.cat((-x2, x1), dim=-1)
31
+
32
+
33
+ def apply_rotary(x, cos, sin):
34
+ """Apply rotary embeddings to input tensor."""
35
+ # Handle case where cos/sin may be shorter than x
36
+ cos = cos[..., :x.shape[-1]]
37
+ sin = sin[..., :x.shape[-1]]
38
+
39
+ # Split x based on cos dimensions
40
+ x_rot = x[..., :cos.shape[-1]]
41
+ x_base = x[..., cos.shape[-1]:]
42
+
43
+ # Apply rotation
44
+ x_rot = (x_rot * cos) + (rotate_half(x_rot) * sin)
45
+
46
+ # Concatenate rotated and base parts
47
+ return torch.cat([x_rot, x_base], dim=-1) if x_base.shape[-1] > 0 else x_rot
ssm.py ADDED
@@ -0,0 +1,147 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Ternary Quantized Diagonal State-Space Model (Parallel)
3
+ """
4
+
5
+ import torch
6
+ import torch.nn as nn
7
+ import torch.nn.functional as F
8
+
9
+ from .bitlinear import BitLinear, RMSNorm
10
+
11
+ class Q88Quantize(torch.autograd.Function):
12
+ """Q8.8 fixed-point quantization with straight-through estimator."""
13
+
14
+ @staticmethod
15
+ def forward(ctx, x):
16
+ """
17
+ Quantize to Q8.8 format (8 integer bits, 8 fractional bits)
18
+ Range: [-128, 127.99609375]
19
+ """
20
+ scale = 2**8 # 256
21
+ # Quantize: scale -> round -> clamp to int16 range -> dequantize
22
+ x_scaled = x * scale
23
+ x_int = torch.clamp(torch.round(x_scaled), -32768, 32767)
24
+ x_quant = x_int / scale
25
+ return x_quant
26
+
27
+ @staticmethod
28
+ def backward(ctx, grad_output):
29
+ # Straight-through estimator: pass gradients unchanged
30
+ return grad_output
31
+
32
+ class SSMBlock(nn.Module):
33
+ """
34
+ Diagonal (convolutional) SSM Block with ternary BitLinear projections.
35
+
36
+ Architecture:
37
+ Input → B projection → diagonal SSM convolution → C projection → Output
38
+
39
+ State dynamics (training, parallel):
40
+ s_t = sum_{k=0}^t (B x_k)
41
+
42
+ State dynamics (inference, step-wise):
43
+ s_t = s_{t-1} + B x_t
44
+ y_t = C s_t
45
+ """
46
+
47
+ def __init__(self, config):
48
+ super().__init__()
49
+ self.config = config
50
+ self.d_model = config.d_model
51
+ self.d_state = config.d_state
52
+
53
+ # =====================================================================
54
+ # Stationary ternary projections
55
+ # =====================================================================
56
+ self.b_proj = BitLinear(self.d_model, self.d_state, bias=False)
57
+ self.c_proj = BitLinear(self.d_state, self.d_model, bias=False)
58
+ # A matrix: identity matrix scaled by a single scalar decay factor
59
+ self.register_buffer("a_log", torch.log(torch.tensor(0.9)))
60
+
61
+ self.dropout = nn.Dropout(config.dropout)
62
+
63
+ # ---------------------------------------------------------------------
64
+ # Training / parallel forward
65
+ # ---------------------------------------------------------------------
66
+ def forward(self, x, mask=None):
67
+ """
68
+ Args:
69
+ x: [batch, seq_len, d_model]
70
+ mask: unused (SSM is causal by construction)
71
+
72
+ Returns:
73
+ y: [batch, seq_len, d_model]
74
+ """
75
+ B, L, _ = x.shape
76
+
77
+ # Input projection
78
+ u = self.b_proj(x)
79
+
80
+ # Compute decay with Q8.8 quantization
81
+ decay = torch.exp(self.a_log) # scalar
82
+ decay_quant = decay + (Q88Quantize.apply(decay) - decay).detach()
83
+
84
+ L = u.size(1)
85
+ device = u.device
86
+ dtype = u.dtype
87
+
88
+ # Decay powers with quantization: [L, 1] (broadcasts across d_state)
89
+ t = torch.arange(L, device=device, dtype=dtype).unsqueeze(1) # [L, 1]
90
+ decay_pows = decay_quant ** t # [L, 1]
91
+ #decay_pows = decay_pows + (Q88Quantize.apply(decay_pows) - decay_pows).detach()
92
+
93
+ inv_decay_pows = decay_pows.reciprocal()
94
+ #inv_decay_pows = inv_decay_pows + (Q88Quantize.apply(inv_decay_pows) - inv_decay_pows).detach()
95
+
96
+ # Reweight, cumsum, reweight back
97
+ s = torch.cumsum(u * inv_decay_pows.unsqueeze(0), dim=1) # [B, L, d_state]
98
+ s = s * decay_pows.unsqueeze(0)
99
+
100
+ # Output projection
101
+ y = self.c_proj(s)
102
+ y = self.dropout(y)
103
+
104
+ return y
105
+
106
+ # ---------------------------------------------------------------------
107
+ # Autoregressive single-step inference
108
+ # ---------------------------------------------------------------------
109
+ def step(self, x, state):
110
+ """
111
+ Single timestep SSM update (for autoregressive decoding).
112
+
113
+ Args:
114
+ x: [batch, d_model]
115
+ state: [batch, d_state]
116
+
117
+ Returns:
118
+ output: [batch, d_model]
119
+ new_state: [batch, d_state]
120
+ """
121
+ decay = torch.exp(self.a_log) # scalar
122
+ new_state = decay * state + self.b_proj(x) # [batch, d_state]
123
+ output = self.c_proj(new_state)
124
+ return output, new_state
125
+
126
+ # ---------------------------------------------------------------------
127
+ # State utilities
128
+ # ---------------------------------------------------------------------
129
+ def init_state(self, batch_size, device, dtype):
130
+ """Initialize hidden state."""
131
+ return torch.zeros(batch_size, self.d_state, device=device, dtype=dtype)
132
+
133
+ # ---------------------------------------------------------------------
134
+ # Export parameters for inference / FPGA
135
+ # ---------------------------------------------------------------------
136
+ def get_inference_params(self):
137
+ """
138
+ Export parameters for deployment.
139
+
140
+ Returns:
141
+ dict with quantized projections and diagonal A
142
+ """
143
+ with torch.no_grad():
144
+ return {
145
+ "b_proj": self.b_proj.get_inference_params(),
146
+ "c_proj": self.c_proj.get_inference_params(),
147
+ }