cloverxion commited on
Commit
e1e0b31
·
0 Parent(s):

Duplicate from cloverx-id/XoneLM-1.0-Paper

Browse files
Files changed (15) hide show
  1. .gitattributes +37 -0
  2. .gitignore +12 -0
  3. LICENSE +17 -0
  4. LuminaV.pdf +3 -0
  5. README.md +201 -0
  6. XoneLM.pdf +3 -0
  7. config.json +22 -0
  8. generate.py +140 -0
  9. luminav.py +530 -0
  10. modeling_xonelm.py +2056 -0
  11. requirements.txt +6 -0
  12. sft_example.py +150 -0
  13. tokenize_example.py +51 -0
  14. tokenizer.py +182 -0
  15. train_example.py +91 -0
.gitattributes ADDED
@@ -0,0 +1,37 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ *.7z filter=lfs diff=lfs merge=lfs -text
2
+ *.arrow filter=lfs diff=lfs merge=lfs -text
3
+ *.bin filter=lfs diff=lfs merge=lfs -text
4
+ *.bz2 filter=lfs diff=lfs merge=lfs -text
5
+ *.ckpt filter=lfs diff=lfs merge=lfs -text
6
+ *.ftz filter=lfs diff=lfs merge=lfs -text
7
+ *.gz filter=lfs diff=lfs merge=lfs -text
8
+ *.h5 filter=lfs diff=lfs merge=lfs -text
9
+ *.joblib filter=lfs diff=lfs merge=lfs -text
10
+ *.lfs.* filter=lfs diff=lfs merge=lfs -text
11
+ *.mlmodel filter=lfs diff=lfs merge=lfs -text
12
+ *.model filter=lfs diff=lfs merge=lfs -text
13
+ *.msgpack filter=lfs diff=lfs merge=lfs -text
14
+ *.npy filter=lfs diff=lfs merge=lfs -text
15
+ *.npz filter=lfs diff=lfs merge=lfs -text
16
+ *.onnx filter=lfs diff=lfs merge=lfs -text
17
+ *.ot filter=lfs diff=lfs merge=lfs -text
18
+ *.parquet filter=lfs diff=lfs merge=lfs -text
19
+ *.pb filter=lfs diff=lfs merge=lfs -text
20
+ *.pickle filter=lfs diff=lfs merge=lfs -text
21
+ *.pkl filter=lfs diff=lfs merge=lfs -text
22
+ *.pt filter=lfs diff=lfs merge=lfs -text
23
+ *.pth filter=lfs diff=lfs merge=lfs -text
24
+ *.rar filter=lfs diff=lfs merge=lfs -text
25
+ *.safetensors filter=lfs diff=lfs merge=lfs -text
26
+ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
27
+ *.tar.* filter=lfs diff=lfs merge=lfs -text
28
+ *.tar filter=lfs diff=lfs merge=lfs -text
29
+ *.tflite filter=lfs diff=lfs merge=lfs -text
30
+ *.tgz filter=lfs diff=lfs merge=lfs -text
31
+ *.wasm filter=lfs diff=lfs merge=lfs -text
32
+ *.xz filter=lfs diff=lfs merge=lfs -text
33
+ *.zip filter=lfs diff=lfs merge=lfs -text
34
+ *.zst filter=lfs diff=lfs merge=lfs -text
35
+ *tfevents* filter=lfs diff=lfs merge=lfs -text
36
+ LuminaV.pdf filter=lfs diff=lfs merge=lfs -text
37
+ XoneLM.pdf filter=lfs diff=lfs merge=lfs -text
.gitignore ADDED
@@ -0,0 +1,12 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ __pycache__/
2
+ *.pyc
3
+ *.pyo
4
+ *.pyd
5
+ .Python
6
+ env/
7
+ venv/
8
+ .venv/
9
+ *.pt
10
+ *.bin
11
+ *.safetensors
12
+ .DS_Store
LICENSE ADDED
@@ -0,0 +1,17 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ Apache License
2
+ Version 2.0, January 2004
3
+ http://www.apache.org/licenses/
4
+
5
+ Copyright 2026 Lumina Moon and Contributors
6
+
7
+ Licensed under the Apache License, Version 2.0 (the "License");
8
+ you may not use this file except in compliance with the License.
9
+ You may obtain a copy of the License at
10
+
11
+ http://www.apache.org/licenses/LICENSE-2.0
12
+
13
+ Unless required by applicable law or agreed to in writing, software
14
+ distributed under the License is distributed on an "AS IS" BASIS,
15
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
16
+ See the License for the specific language governing permissions and
17
+ limitations under the License.
LuminaV.pdf ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:4d15a04a19df9b39b89f3e7c59bfead35ffc1a9fc1767118f2d132d52899bca5
3
+ size 329711
README.md ADDED
@@ -0,0 +1,201 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: apache-2.0
3
+ language:
4
+ - en
5
+ tags:
6
+ - pytorch
7
+ - causal-lm
8
+ - working-memory
9
+ - non-euclidean
10
+ - custom-optimizer
11
+ - 3am-engineering
12
+ pipeline_tag: text-generation
13
+ ---
14
+
15
+
16
+ <div align="center">
17
+ <h2>XoneLM & LuminaV Research Papers</h2>
18
+ <p><em>(What Happens When You Code at 3 AM: The Architecture & The Optimizer)</em></p>
19
+ <h2>Check the 'Files and versions' tab to find the source code.</h2>
20
+
21
+
22
+ <table>
23
+ <tr>
24
+ <td align="center" width="50%">
25
+ <h3>1. Architecture Paper</h3>
26
+ <a href="https://huggingface.co/cloverx-id/XoneLM-1.0-Papper/blob/main/XoneLM.pdf">
27
+ <img src="https://cdn-uploads.huggingface.co/production/uploads/6835b520186e98712b0386a6/Y6fNRAdPDN2NGZOemdTKP.png" alt="XoneLM Paper Preview" width="380" style="border-radius: 6px; box-shadow: 0 4px 12px rgba(0,0,0,0.15);" />
28
+ </a>
29
+ <br /><br />
30
+ <strong>XoneLM Architecture (54M)</strong><br />
31
+ <a href="https://huggingface.co/cloverx-id/XoneLM-1.0-Papper/blob/main/XoneLM.pdf">Read XoneLM.pdf</a>
32
+ </td>
33
+ <td align="center" width="50%">
34
+ <h3>2. Optimizer Paper</h3>
35
+ <a href="https://huggingface.co/cloverx-id/XoneLM-1.0-Papper/blob/main/LuminaV.pdf">
36
+ <img src="https://cdn-uploads.huggingface.co/production/uploads/6835b520186e98712b0386a6/GfVA7g5xUHwopzOSRR5QD.png" alt="LuminaV Paper Preview" width="380" style="border-radius: 6px; box-shadow: 0 4px 12px rgba(0,0,0,0.15);" />
37
+ </a>
38
+ <br /><br />
39
+ <strong>LuminaV Optimizer Theory</strong><br />
40
+ <a href="https://huggingface.co/cloverx-id/XoneLM-1.0-Papper/blob/main/LuminaV.pdf">Read LuminaV.pdf</a>
41
+ </td>
42
+ </tr>
43
+ </table>
44
+ <p><em>Click on each image to read or download its official PDF file.</em></p>
45
+ </div>
46
+
47
+ ## Research Paper Overview
48
+
49
+ ### Abstract
50
+ XoneLM is an experimental 54.07M-parameter language model constructed by integrating Stiefel manifold QR decomposition, degree-180 Chebyshev polynomial positional encodings, topological soliton wave tracking, rational power-law attention decay, Poincare hyperbolic routing, and low-rank latent Key-Value compression (MLA).
51
+
52
+ The entire system was trained from scratch on 32.01M tokens using the custom LuminaV tanh-bounded optimizer on a single consumer GPU in under two hours. The training run achieved monotonic convergence (Final Loss: 1.7355, Perplexity: 5.67) with zero loss spikes, zero gradient explosions, and zero arithmetic underflow errors.
53
+
54
+ ### Empirical Telemetry & Training Specs
55
+ | Metric | Value |
56
+ | :--- | :--- |
57
+ | Total Active Parameters | 54,073,344 (54.07M) |
58
+ | Architecture Backbone | 12 Layers, 8 Attention Heads, Dim 512 |
59
+ | Key-Value Latent Dimension | 64 (Low-Rank Joint MLA) |
60
+ | Working Memory Slots | 512 (256 Static + 256 Dynamic) |
61
+ | Optimizer | LuminaV-2B (Learning Rate: 8e-4) |
62
+ | Precision | Mixed Precision FP16 |
63
+ | Training Tokens | 32,009,639 Tokens (TinyStories) |
64
+ | Wall-Clock Training Time | 117.88 minutes (1.96 hours) |
65
+ | Peak Training Throughput | 9,868 tokens/second |
66
+ | Peak VRAM Usage | 7.29 GB / 14.56 GB (Tesla T4) |
67
+ | Hub Diversity Z-Loss | 0.0014 |
68
+ | Final Loss / Perplexity | 1.7355 / 5.67 |
69
+
70
+ ### Qualitative Analysis: The Box vs. Ball Case Study
71
+ During unconditioned zero-shot evaluation on the base pre-trained model:
72
+ * Prompt: "Once upon a time, Lily found a Box."
73
+ * Output: "It was a big, round ball. She was so excited to play with it..."
74
+
75
+ The model produced syntactically perfect English and dialogue quotation marks without infinite looping. However, the pre-training distribution prior of the TinyStories corpus (which heavily features children playing with balls in parks) overrode the prompt keyword. This highlights that unconditioned base models act as probabilistic continuation engines, and Supervised Fine-Tuning (SFT) is necessary for strict instruction adherence.
76
+
77
+ ---
78
+
79
+ ## Architectural Components
80
+
81
+ 1. Stiefel QR Working Memory Hub: An orthogonal working memory bank initialized via QR factorization on Stiefel manifolds to enforce metric stability from step zero.
82
+ 2. PolyHoPE Positional Encodings: Degree-180 Chebyshev polynomials of the first kind evaluated on normalized token intervals to guarantee continuous variance without periodic decay.
83
+ 3. Topological Soliton State Tracking: Discrete nonlinear sech-squared wave updates derived from collisionless plasma dynamics to maintain latent memory state stability.
84
+ 4. LinHoPE Attention Decay: Heavy-tail Cauchy power-law decay combined with geometric recency bias to mitigate early token context amnesia.
85
+ 5. Poincare Hyperbolic Routing: Episodic memory cluster assignment on Riemannian conformal unit disks.
86
+ 6. Latent KV Compression: Low-rank Key-Value joint compression (dkv = 64) minimizing memory bandwidth during autoregressive decoding.
87
+ 7. LuminaV Optimizer: Master-weight-free parameter optimization featuring a hyperbolic tangent bounding envelope, central innovation variance tracking, and cautious directional masking.
88
+
89
+ ---
90
+
91
+ ## Quickstart: Training and Inference
92
+
93
+ ### 1. Installation
94
+ ```bash
95
+ git clone https://huggingface.co/cloverx-id/XoneLM-1.0-Papper
96
+ cd XoneLM-1.0-Papper
97
+ pip install -r requirements.txt
98
+ ```
99
+
100
+ ### 2. Running a Training Step with LuminaV
101
+ ```python
102
+ import torch
103
+ from tokenizer import build_xonelm_tokenizer
104
+ from modeling_xonelm import XoneLM
105
+ from luminav import LuminaV
106
+
107
+ device = "cuda" if torch.cuda.is_available() else "cpu"
108
+
109
+ tokenizer = build_xonelm_tokenizer()
110
+ vocab_size = len(tokenizer)
111
+
112
+ model = XoneLM(
113
+ vocab_size=vocab_size,
114
+ dim=512,
115
+ num_layers=12,
116
+ num_heads=8,
117
+ kv_latent_dim=64,
118
+ hub_size=512,
119
+ num_terminals=32,
120
+ slots_per_terminal=16
121
+ ).to(device)
122
+
123
+ optimizer = LuminaV(
124
+ model.parameters(),
125
+ lr=8e-4,
126
+ betas=(0.9, 0.999),
127
+ tau=0.8,
128
+ buffer=2,
129
+ cautious=True,
130
+ execution="auto"
131
+ )
132
+
133
+ dummy_tokens = torch.randint(0, vocab_size, (2, 512), device=device)
134
+ dummy_labels = torch.randint(0, vocab_size, (2, 512), device=device)
135
+
136
+ optimizer.zero_grad()
137
+ output = model(dummy_tokens, labels=dummy_labels)
138
+ loss = output.loss
139
+ loss.backward()
140
+ optimizer.step()
141
+
142
+ print(f"Training step successful. Loss: {loss.item():.4f}")
143
+ ```
144
+
145
+ ### 3. Running Autoregressive Generation (DRY + Min-P)
146
+ ```python
147
+ from generate import generate_response
148
+
149
+ prompt = "Once upon a time, Lily found a Box."
150
+ result = generate_response(
151
+ model=model,
152
+ tokenizer=tokenizer,
153
+ prompt_or_messages=prompt,
154
+ max_new_tokens=64,
155
+ temperature=0.45,
156
+ min_p=0.08
157
+ )
158
+ print("Output:", result)
159
+ ```
160
+
161
+ ---
162
+
163
+ ## Citation
164
+
165
+ If you use this model, optimizer, or refer to our research, please cite our work as follows:
166
+
167
+ ### Primary Citation (Model & Paper)
168
+ ```bibtex
169
+ @misc{luminamoon2026xonelm,
170
+ author = {{Silver Moon (cloverxion)}},
171
+ organization = {Lumina Moon},
172
+ title = {{XoneLM: An Over-Engineered 54M Language Model with Non-Euclidean Memory, Polynomial Encodings, and Bounded Optimizers}},
173
+ year = {2026},
174
+ publisher = {Hugging Face},
175
+ doi = {10.57967/hf/10270},
176
+ howpublished = {\url{https://huggingface.co/cloverx-id/XoneLM-1.0-Paper}},
177
+ url = {https://huggingface.co/cloverx-id/XoneLM-1.0-Paper},
178
+ note = {Hugging Face Model Hub}
179
+ }
180
+ ```
181
+
182
+ ### LuminaV Optimizer
183
+ ```bibtex
184
+ @misc{luminamoon2026luminav,
185
+ author = {{Silver Moon (cloverxion)}},
186
+ organization = {Lumina Moon},
187
+ title = {{LuminaV: We Were Too Broke for AdamW So We Trapped Gradients in a Hyperbolic Straitjacket and Hired a Traffic Cop to Slap Them}},
188
+ year = {2026},
189
+ publisher = {Hugging Face},
190
+ doi = {10.57967/hf/10270},
191
+ howpublished = {\url{https://huggingface.co/cloverx-id/XoneLM-1.0-Paper}},
192
+ url = {https://huggingface.co/cloverx-id/XoneLM-1.0-Paper},
193
+ note = {Hugging Face Repository}
194
+ }
195
+ ```
196
+
197
+ ## License
198
+
199
+ All code and architecture assets (including PDFs) are released under the **Apache-2.0 License.**
200
+
201
+ ---
XoneLM.pdf ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:93b2efc2aceef4ffe7c729a6a449478c0be6043bff261c45dc7d59e77c3961c6
3
+ size 385462
config.json ADDED
@@ -0,0 +1,22 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "XoneLM"
4
+ ],
5
+ "model_type": "xonelm",
6
+ "vocab_size": 32000,
7
+ "dim": 512,
8
+ "num_layers": 12,
9
+ "num_heads": 8,
10
+ "d_head": 64,
11
+ "kv_latent_dim": 64,
12
+ "hub_size": 512,
13
+ "num_specialized_hubs": 12,
14
+ "num_terminals": 32,
15
+ "slots_per_terminal": 16,
16
+ "max_episodic": 64,
17
+ "chunk_size": 1024,
18
+ "alpha_anchor": 0.1,
19
+ "polyhope_degree": 180,
20
+ "torch_dtype": "float16",
21
+ "transformers_version": "5.15.1"
22
+ }
generate.py ADDED
@@ -0,0 +1,140 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Any, Dict, List, Optional, Union
2
+ import torch
3
+ import torch.nn.functional as F
4
+ from modeling_xonelm import XoneLM, HardwareContext
5
+ from tokenizer import SpecialTokenConfig, MultiTurnConversationFormatter
6
+
7
+ def apply_dry_repetition_penalty(
8
+ logits: torch.Tensor,
9
+ generated_tokens: List[int],
10
+ dry_multiplier: float = 0.8,
11
+ dry_base: float = 1.75,
12
+ dry_allowed_length: int = 2,
13
+ ) -> torch.Tensor:
14
+ gen_len = len(generated_tokens)
15
+ if gen_len < dry_allowed_length:
16
+ return logits
17
+
18
+ for match_len in range(dry_allowed_length, min(gen_len, 40)):
19
+ target_ngram = generated_tokens[-match_len:]
20
+ for i in range(gen_len - match_len):
21
+ if generated_tokens[i : i + match_len] == target_ngram:
22
+ next_token = generated_tokens[i + match_len]
23
+ penalty = dry_multiplier * (dry_base ** (match_len - dry_allowed_length))
24
+ logits[0, next_token] -= penalty
25
+
26
+ return logits
27
+
28
+ @torch.no_grad()
29
+ def generate_response(
30
+ model: XoneLM,
31
+ tokenizer: Any,
32
+ prompt_or_messages: Union[str, List[Dict[str, str]]],
33
+ max_new_tokens: int = 128,
34
+ temperature: float = 0.45,
35
+ top_k: int = 40,
36
+ top_p: float = 0.90,
37
+ min_p: float = 0.08,
38
+ dry_multiplier: float = 0.8,
39
+ dry_base: float = 1.75,
40
+ dry_allowed_length: int = 2,
41
+ token_config: Optional[SpecialTokenConfig] = None,
42
+ ) -> str:
43
+ raw_model = getattr(model, "_orig_mod", model)
44
+ raw_model.eval()
45
+ device = next(raw_model.parameters()).device
46
+ cfg = token_config or SpecialTokenConfig()
47
+
48
+ if isinstance(prompt_or_messages, str):
49
+ enc = tokenizer.encode(prompt_or_messages)
50
+ input_ids = enc.ids if hasattr(enc, "ids") else enc["input_ids"]
51
+ else:
52
+ formatter = MultiTurnConversationFormatter(tokenizer, cfg)
53
+ formatted = formatter.format_conversation(prompt_or_messages)
54
+ input_ids = formatted["input_ids"]
55
+ if len(input_ids) > 0 and input_ids[-1] == cfg.eod_token_id:
56
+ input_ids.pop()
57
+ asst_header = "<|im_start|>assistant\n"
58
+ enc_asst = tokenizer.encode(asst_header)
59
+ asst_ids = enc_asst.ids if hasattr(enc_asst, "ids") else enc_asst["input_ids"]
60
+ input_ids.extend(asst_ids)
61
+
62
+ generated = torch.tensor([input_ids], dtype=torch.long, device=device)
63
+ prompt_len = generated.shape[1]
64
+
65
+ hub = raw_model.extract_hub(generated)
66
+ out = raw_model.forward(generated, override_hub=hub)
67
+ past_kv = out.past_key_values
68
+
69
+ stop_tokens = {cfg.im_end_id, cfg.eod_token_id, cfg.eos_token_id}
70
+ stop_tokens.discard(None)
71
+
72
+ for step_i in range(max_new_tokens):
73
+ if step_i == 0:
74
+ logits = out.logits[:, -1, :].clone()
75
+ else:
76
+ step_out = raw_model.forward(cur_token, past_key_values=past_kv)
77
+ past_kv = step_out.past_key_values
78
+ logits = step_out.logits[:, -1, :].clone()
79
+
80
+ token_history = generated[0, prompt_len:].tolist()
81
+ if token_history:
82
+ logits = apply_dry_repetition_penalty(
83
+ logits=logits,
84
+ generated_tokens=token_history,
85
+ dry_multiplier=dry_multiplier,
86
+ dry_base=dry_base,
87
+ dry_allowed_length=dry_allowed_length,
88
+ )
89
+
90
+ if temperature > 0:
91
+ logits = logits / max(temperature, 1e-5)
92
+
93
+ if top_k > 0:
94
+ v_top, _ = torch.topk(logits, min(top_k, logits.size(-1)))
95
+ logits[logits < v_top[:, [-1]]] = -float("Inf")
96
+
97
+ if min_p > 0.0:
98
+ probs = F.softmax(logits, dim=-1)
99
+ top_prob, _ = torch.max(probs, dim=-1, keepdim=True)
100
+ scaled_min_p = min_p * top_prob
101
+ logits[probs < scaled_min_p] = -float("Inf")
102
+
103
+ if top_p < 1.0:
104
+ sorted_logits, sorted_indices = torch.sort(logits, descending=True)
105
+ cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1)
106
+ sorted_indices_to_remove = cumulative_probs > top_p
107
+ sorted_indices_to_remove[:, 1:] = sorted_indices_to_remove[:, :-1].clone()
108
+ sorted_indices_to_remove[:, 0] = 0
109
+ indices_to_remove = sorted_indices_to_remove.scatter(1, sorted_indices, sorted_indices_to_remove)
110
+ logits[indices_to_remove] = -float("Inf")
111
+
112
+ probs_final = F.softmax(logits, dim=-1)
113
+ cur_token = torch.multinomial(probs_final, num_samples=1)
114
+ else:
115
+ cur_token = torch.argmax(logits, dim=-1, keepdim=True)
116
+
117
+ if cur_token.item() in stop_tokens:
118
+ break
119
+
120
+ generated = torch.cat([generated, cur_token], dim=1)
121
+
122
+ new_token_ids = generated[0, prompt_len:].tolist()
123
+ if len(new_token_ids) > 0 and new_token_ids[-1] in stop_tokens:
124
+ new_token_ids.pop()
125
+
126
+ if hasattr(tokenizer, "decode"):
127
+ return tokenizer.decode(new_token_ids).strip()
128
+ return tokenizer.decode(new_token_ids, skip_special_tokens=True).strip()
129
+
130
+ if __name__ == "__main__":
131
+ from tokenizer import build_xonelm_tokenizer
132
+
133
+ dev = HardwareContext.get_optimal_device()
134
+ tok = build_xonelm_tokenizer()
135
+ lm = XoneLM(vocab_size=len(tok), dim=512, num_layers=12, num_heads=8, kv_latent_dim=64).to(dev)
136
+
137
+ test_prompt = "Once upon a time, in a small garden, Lily found a Box."
138
+ res = generate_response(lm, tok, test_prompt, max_new_tokens=32)
139
+ print("Prompt :", test_prompt)
140
+ print("Output :", res)
luminav.py ADDED
@@ -0,0 +1,530 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2026 Lumina Moon and Contributors.
2
+ # Licensed under the Apache License, Version 2.0 (see LICENSE for details).
3
+
4
+ """
5
+ LuminaV: We Were Too Broke for AdamW So We Trapped Gradients in a
6
+ Hyperbolic Straitjacket and Hired a Traffic Cop to Slap Them
7
+
8
+ Paper PDF : https://huggingface.co/cloverx-id/XoneLM-1.0-Paper/blob/main/LuminaV.pdf
9
+ DOI : 10.57967/hf/10270
10
+ """
11
+
12
+ from __future__ import annotations
13
+
14
+ import logging
15
+ import math
16
+ from typing import Callable, Dict, List, Optional, Tuple
17
+
18
+ import torch
19
+ from torch import Tensor
20
+ from torch.optim import Optimizer
21
+
22
+ __all__ = ["LuminaV"]
23
+ __version__ = "1.0.0"
24
+ __author__ = "Silver Moon (cloverxion), Lumina Moon"
25
+ __license__ = "Apache-2.0"
26
+
27
+ logger = logging.getLogger("LuminaV")
28
+
29
+
30
+ HAS_TRITON = False
31
+ try:
32
+ import triton
33
+ import triton.language as tl
34
+ HAS_TRITON = True
35
+ except ImportError:
36
+ HAS_TRITON = False
37
+
38
+ if HAS_TRITON:
39
+ @triton.jit
40
+ def _triton_tanh_fast(x):
41
+ return 2.0 * tl.sigmoid(2.0 * x) - 1.0
42
+
43
+ @triton.jit
44
+ def _lumina_v2_pass1_kernel(
45
+ grad_ptr, exp_avg_ptr, exp_avg_sq_ptr, mask_sum_ptr,
46
+ n_elements, beta1, beta2, c1, c2, BLOCK_SIZE: tl.constexpr
47
+ ):
48
+ pid = tl.program_id(axis=0)
49
+ offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
50
+ mask = offsets < n_elements
51
+
52
+ g = tl.load(grad_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
53
+ m = tl.load(exp_avg_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
54
+ v = tl.load(exp_avg_sq_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
55
+
56
+ m_new = beta1 * m + (1.0 - beta1) * g
57
+ nes_m = beta1 * m_new + (1.0 - beta1) * g
58
+ diff = g - m_new
59
+ v_new = beta2 * v + (1.0 - beta2) * (diff * diff)
60
+
61
+ sigma = tl.sqrt(tl.maximum(v_new, 0.0)) * c1 + c2
62
+ u = _triton_tanh_fast(nes_m / sigma)
63
+ m_mask = tl.where((u * g) > 0.0, 1.0, 0.0)
64
+
65
+ tl.store(exp_avg_ptr + offsets, m_new, mask=mask)
66
+ tl.store(exp_avg_sq_ptr + offsets, v_new, mask=mask)
67
+
68
+ block_sum = tl.sum(tl.where(mask, m_mask, 0.0), axis=0)
69
+ tl.atomic_add(mask_sum_ptr, block_sum)
70
+
71
+ @triton.jit
72
+ def _lumina_v2_pass2_kernel(
73
+ p_ptr, grad_ptr, exp_avg_ptr, exp_avg_sq_ptr, mask_sum_ptr,
74
+ n_elements, beta1, c1, c2, lr, weight_decay, clamp_min, BLOCK_SIZE: tl.constexpr
75
+ ):
76
+ pid = tl.program_id(axis=0)
77
+ offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
78
+ mask = offsets < n_elements
79
+
80
+ m_sum = tl.load(mask_sum_ptr)
81
+ m_bar = tl.minimum(tl.maximum(m_sum / n_elements, clamp_min), 1.0)
82
+
83
+ p = tl.load(p_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
84
+ g = tl.load(grad_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
85
+ m = tl.load(exp_avg_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
86
+ v = tl.load(exp_avg_sq_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
87
+
88
+ nes_m = beta1 * m + (1.0 - beta1) * g
89
+ sigma = tl.sqrt(tl.maximum(v, 0.0)) * c1 + c2
90
+ u = _triton_tanh_fast(nes_m / sigma)
91
+ m_mask = tl.where((u * g) > 0.0, 1.0, 0.0)
92
+
93
+ if weight_decay != 0.0:
94
+ p = p * (1.0 - lr * weight_decay)
95
+
96
+ delta_theta = (u * m_mask) / m_bar
97
+ p_updated = p - lr * delta_theta
98
+ tl.store(p_ptr + offsets, p_updated, mask=mask)
99
+
100
+ @triton.jit
101
+ def _lumina_v2_single_pass_kernel(
102
+ p_ptr, grad_ptr, exp_avg_ptr, exp_avg_sq_ptr,
103
+ n_elements, beta1, beta2, c1, c2, lr, weight_decay, BLOCK_SIZE: tl.constexpr
104
+ ):
105
+ pid = tl.program_id(axis=0)
106
+ offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
107
+ mask = offsets < n_elements
108
+
109
+ p = tl.load(p_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
110
+ g = tl.load(grad_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
111
+ m = tl.load(exp_avg_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
112
+ v = tl.load(exp_avg_sq_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
113
+
114
+ m_new = beta1 * m + (1.0 - beta1) * g
115
+ nes_m = beta1 * m_new + (1.0 - beta1) * g
116
+ diff = g - m_new
117
+ v_new = beta2 * v + (1.0 - beta2) * (diff * diff)
118
+
119
+ sigma = tl.sqrt(tl.maximum(v_new, 0.0)) * c1 + c2
120
+ u = _triton_tanh_fast(nes_m / sigma)
121
+
122
+ tl.store(exp_avg_ptr + offsets, m_new, mask=mask)
123
+ tl.store(exp_avg_sq_ptr + offsets, v_new, mask=mask)
124
+
125
+ if weight_decay != 0.0:
126
+ p = p * (1.0 - lr * weight_decay)
127
+
128
+ p_updated = p - lr * u
129
+ tl.store(p_ptr + offsets, p_updated, mask=mask)
130
+
131
+ @triton.jit
132
+ def _lumina_v1_pass1_kernel(
133
+ grad_ptr, exp_avg_ptr, rms_sum_ptr,
134
+ n_elements, beta1, BLOCK_SIZE: tl.constexpr
135
+ ):
136
+ pid = tl.program_id(axis=0)
137
+ offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
138
+ mask = offsets < n_elements
139
+
140
+ g = tl.load(grad_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
141
+ m = tl.load(exp_avg_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
142
+
143
+ m_new = beta1 * m + (1.0 - beta1) * g
144
+ nes_m = beta1 * m_new + (1.0 - beta1) * g
145
+ tl.store(exp_avg_ptr + offsets, m_new, mask=mask)
146
+
147
+ sq_val = nes_m * nes_m
148
+ block_sq_sum = tl.sum(tl.where(mask, sq_val, 0.0), axis=0)
149
+ tl.atomic_add(rms_sum_ptr, block_sq_sum)
150
+
151
+ @triton.jit
152
+ def _lumina_v1_pass2_kernel(
153
+ grad_ptr, exp_avg_ptr, rms_sum_ptr, mask_sum_ptr,
154
+ n_elements, beta1, tau, eps, c2, alpha_ss, BLOCK_SIZE: tl.constexpr
155
+ ):
156
+ pid = tl.program_id(axis=0)
157
+ offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
158
+ mask = offsets < n_elements
159
+
160
+ sq_sum = tl.load(rms_sum_ptr)
161
+ rms = tl.sqrt(sq_sum / n_elements + eps)
162
+ sigma = tau * rms + c2
163
+
164
+ g = tl.load(grad_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
165
+ m = tl.load(exp_avg_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
166
+
167
+ nes_m = beta1 * m + (1.0 - beta1) * g
168
+ z = nes_m / sigma
169
+ z_soft = z / (1.0 + alpha_ss * tl.abs(z))
170
+ u = _triton_tanh_fast(z_soft)
171
+ m_mask = tl.where((u * g) > 0.0, 1.0, 0.0)
172
+
173
+ block_sum = tl.sum(tl.where(mask, m_mask, 0.0), axis=0)
174
+ tl.atomic_add(mask_sum_ptr, block_sum)
175
+
176
+ @triton.jit
177
+ def _lumina_v1_pass3_update_kernel(
178
+ p_ptr, grad_ptr, exp_avg_ptr, rms_sum_ptr, mask_sum_ptr,
179
+ n_elements, beta1, tau, eps, c2, alpha_ss, lr, weight_decay, clamp_min,
180
+ cautious: tl.constexpr, BLOCK_SIZE: tl.constexpr
181
+ ):
182
+ pid = tl.program_id(axis=0)
183
+ offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
184
+ mask = offsets < n_elements
185
+
186
+ sq_sum = tl.load(rms_sum_ptr)
187
+ rms = tl.sqrt(sq_sum / n_elements + eps)
188
+ sigma = tau * rms + c2
189
+
190
+ m_sum = tl.load(mask_sum_ptr)
191
+ m_bar = tl.minimum(tl.maximum(m_sum / n_elements, clamp_min), 1.0)
192
+
193
+ p = tl.load(p_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
194
+ g = tl.load(grad_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
195
+ m = tl.load(exp_avg_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
196
+
197
+ nes_m = beta1 * m + (1.0 - beta1) * g
198
+ z = nes_m / sigma
199
+ z_soft = z / (1.0 + alpha_ss * tl.abs(z))
200
+ u = _triton_tanh_fast(z_soft)
201
+ m_mask = tl.where((u * g) > 0.0, 1.0, 0.0)
202
+
203
+ if weight_decay != 0.0:
204
+ p = p * (1.0 - lr * weight_decay)
205
+
206
+ delta_theta = (u * m_mask) / m_bar if cautious else u
207
+ p_updated = p - lr * delta_theta
208
+ tl.store(p_ptr + offsets, p_updated, mask=mask)
209
+
210
+
211
+ class LuminaV(Optimizer):
212
+ def __init__(
213
+ self,
214
+ params,
215
+ lr: float = 8e-4,
216
+ betas: Tuple[float, float] = (0.9, 0.999),
217
+ eps: float = 1e-8,
218
+ weight_decay: float = 8e-2,
219
+ tau: float = 0.8,
220
+ alpha_ss: float = 0.5,
221
+ cautious: bool = True,
222
+ cautious_clamp_min: float = 0.2,
223
+ buffer: int = 2,
224
+ execution: str = "auto",
225
+ ):
226
+ if lr < 0.0:
227
+ raise ValueError(f"Invalid learning rate: {lr}")
228
+ if not 0.0 <= betas[0] < 1.0:
229
+ raise ValueError(f"Invalid beta1 parameter: {betas[0]}")
230
+ if not 0.0 <= betas[1] < 1.0:
231
+ raise ValueError(f"Invalid beta2 parameter: {betas[1]}")
232
+ if eps <= 0.0:
233
+ raise ValueError(f"Invalid epsilon value: {eps}")
234
+ if weight_decay < 0.0:
235
+ raise ValueError(f"Invalid weight_decay value: {weight_decay}")
236
+ if tau <= 0.0:
237
+ raise ValueError(f"Invalid tau parameter: {tau}")
238
+ if not 0.0 < cautious_clamp_min <= 1.0:
239
+ raise ValueError(f"Invalid cautious_clamp_min: {cautious_clamp_min}")
240
+ if buffer not in (1, 2):
241
+ raise ValueError(f"Buffer count must be 1 (Single) or 2 (Dual), got: {buffer}")
242
+
243
+ defaults = dict(
244
+ lr=lr, betas=betas, eps=eps, weight_decay=weight_decay,
245
+ tau=tau, alpha_ss=alpha_ss, cautious=cautious,
246
+ cautious_clamp_min=cautious_clamp_min, buffer=buffer, execution=execution
247
+ )
248
+ super().__init__(params, defaults)
249
+ self._scratch_tensors: Dict[torch.device, torch.Tensor] = {}
250
+
251
+ def _get_scratch_buffer(self, device: torch.device, count: int) -> torch.Tensor:
252
+ if device not in self._scratch_tensors or self._scratch_tensors[device].numel() < count:
253
+ self._scratch_tensors[device] = torch.zeros((count,), device=device, dtype=torch.float32)
254
+ buf = self._scratch_tensors[device][:count]
255
+ buf.zero_()
256
+ return buf
257
+
258
+ @torch.no_grad()
259
+ def step(self, closure: Optional[Callable[[], float]] = None) -> Optional[float]:
260
+ loss = None
261
+ if closure is not None:
262
+ with torch.enable_grad():
263
+ loss = closure()
264
+
265
+ for group in self.param_groups:
266
+ params_with_grad = []
267
+ grads = []
268
+ exp_avgs = []
269
+ exp_avg_sqs = []
270
+ steps = []
271
+ buf_count = group["buffer"]
272
+
273
+ for p in group["params"]:
274
+ if p.grad is None:
275
+ continue
276
+ if p.grad.is_sparse:
277
+ raise RuntimeError("LuminaV does not support sparse gradients.")
278
+
279
+ params_with_grad.append(p)
280
+ grads.append(p.grad)
281
+ state = self.state[p]
282
+
283
+ if len(state) == 0:
284
+ state["step"] = 0
285
+ state["exp_avg"] = torch.zeros_like(p, memory_format=torch.preserve_format)
286
+ if buf_count == 2:
287
+ state["exp_avg_sq"] = torch.zeros_like(p, memory_format=torch.preserve_format)
288
+
289
+ exp_avgs.append(state["exp_avg"])
290
+ if buf_count == 2:
291
+ if "exp_avg_sq" not in state:
292
+ state["exp_avg_sq"] = torch.zeros_like(p, memory_format=torch.preserve_format)
293
+ exp_avg_sqs.append(state["exp_avg_sq"])
294
+
295
+ state["step"] += 1
296
+ steps.append(state["step"])
297
+
298
+ if not params_with_grad:
299
+ continue
300
+
301
+ lr = group["lr"]
302
+ beta1, beta2 = group["betas"]
303
+ eps = group["eps"]
304
+ weight_decay = group["weight_decay"]
305
+ tau = group["tau"]
306
+ alpha_ss = group["alpha_ss"]
307
+ cautious = group["cautious"]
308
+ clamp_min = group["cautious_clamp_min"]
309
+ exec_mode = group["execution"]
310
+
311
+ all_cuda = all(p.is_cuda for p in params_with_grad)
312
+
313
+ if exec_mode == "auto":
314
+ if HAS_TRITON and all_cuda:
315
+ exec_mode = "triton"
316
+ elif hasattr(torch, "_foreach_mul_"):
317
+ exec_mode = "foreach"
318
+ else:
319
+ exec_mode = "single"
320
+
321
+ if exec_mode == "triton":
322
+ if not (HAS_TRITON and all_cuda):
323
+ exec_mode = "foreach"
324
+
325
+ if exec_mode == "triton" and HAS_TRITON and all_cuda:
326
+ try:
327
+ if buf_count == 2:
328
+ self._triton_step_v2(params_with_grad, grads, exp_avgs, exp_avg_sqs, steps, lr, beta1, beta2, eps, weight_decay, tau, cautious, clamp_min)
329
+ else:
330
+ self._triton_step_v1(params_with_grad, grads, exp_avgs, steps, lr, beta1, eps, weight_decay, tau, alpha_ss, cautious, clamp_min)
331
+ except Exception as e:
332
+ logger.warning(f"Triton kernel fallback to C++ foreach engine: {e}")
333
+ self._grouped_foreach_step(params_with_grad, grads, exp_avgs, exp_avg_sqs, steps, lr, beta1, beta2, eps, weight_decay, tau, alpha_ss, cautious, clamp_min, buf_count)
334
+ elif exec_mode == "foreach" or (exec_mode == "triton" and not all_cuda):
335
+ self._grouped_foreach_step(params_with_grad, grads, exp_avgs, exp_avg_sqs, steps, lr, beta1, beta2, eps, weight_decay, tau, alpha_ss, cautious, clamp_min, buf_count)
336
+ else:
337
+ if buf_count == 2:
338
+ self._single_step_v2(params_with_grad, grads, exp_avgs, exp_avg_sqs, steps, lr, beta1, beta2, eps, weight_decay, tau, cautious, clamp_min)
339
+ else:
340
+ self._single_step_v1(params_with_grad, grads, exp_avgs, steps, lr, beta1, eps, weight_decay, tau, alpha_ss, cautious, clamp_min)
341
+
342
+ return loss
343
+
344
+ def _grouped_foreach_step(self, params, grads, exp_avgs, exp_avg_sqs, steps, lr, beta1, beta2, eps, weight_decay, tau, alpha_ss, cautious, clamp_min, buf_count):
345
+ groups: Dict[Tuple[torch.device, torch.dtype], List[int]] = {}
346
+ for idx, p in enumerate(params):
347
+ key = (p.device, p.dtype)
348
+ if key not in groups:
349
+ groups[key] = []
350
+ groups[key].append(idx)
351
+
352
+ for (dev, dt), indices in groups.items():
353
+ sub_params = [params[i] for i in indices]
354
+ sub_grads = [grads[i] for i in indices]
355
+ sub_exp_avgs = [exp_avgs[i] for i in indices]
356
+ sub_steps = [steps[i] for i in indices]
357
+
358
+ if buf_count == 2:
359
+ sub_exp_avg_sqs = [exp_avg_sqs[i] for i in indices]
360
+ self._foreach_step_v2(sub_params, sub_grads, sub_exp_avgs, sub_exp_avg_sqs, sub_steps, lr, beta1, beta2, eps, weight_decay, tau, cautious, clamp_min)
361
+ else:
362
+ self._foreach_step_v1(sub_params, sub_grads, sub_exp_avgs, sub_steps, lr, beta1, eps, weight_decay, tau, alpha_ss, cautious, clamp_min)
363
+
364
+ def _triton_step_v2(self, params, grads, exp_avgs, exp_avg_sqs, steps, lr, beta1, beta2, eps, weight_decay, tau, cautious, clamp_min):
365
+ BLOCK_SIZE = 1024
366
+ device = params[0].device
367
+ scratch = self._get_scratch_buffer(device, len(params))
368
+
369
+ for i in range(len(params)):
370
+ p, grad, exp_avg, exp_avg_sq, step = params[i], grads[i], exp_avgs[i], exp_avg_sqs[i], steps[i]
371
+ is_orig_contig = p.is_contiguous()
372
+ p_contig = p if is_orig_contig else p.contiguous()
373
+ grad_contig = grad if grad.is_contiguous() else grad.contiguous()
374
+
375
+ bc1 = 1.0 - (beta1**step)
376
+ bc2 = 1.0 - (beta2**step)
377
+ c1 = (bc1 * tau) / math.sqrt(bc2)
378
+ c2 = eps * bc1 * tau
379
+
380
+ n_elements = p.numel()
381
+ grid = (triton.cdiv(n_elements, BLOCK_SIZE),)
382
+
383
+ if not cautious:
384
+ _lumina_v2_single_pass_kernel[grid](p_contig, grad_contig, exp_avg, exp_avg_sq, n_elements, beta1, beta2, c1, c2, lr, weight_decay, BLOCK_SIZE=BLOCK_SIZE)
385
+ else:
386
+ mask_sum_ptr = scratch[i : i + 1]
387
+ mask_sum_ptr.zero_()
388
+ _lumina_v2_pass1_kernel[grid](grad_contig, exp_avg, exp_avg_sq, mask_sum_ptr, n_elements, beta1, beta2, c1, c2, BLOCK_SIZE=BLOCK_SIZE)
389
+ _lumina_v2_pass2_kernel[grid](p_contig, grad_contig, exp_avg, exp_avg_sq, mask_sum_ptr, n_elements, beta1, c1, c2, lr, weight_decay, clamp_min, BLOCK_SIZE=BLOCK_SIZE)
390
+
391
+ if not is_orig_contig:
392
+ p.copy_(p_contig)
393
+
394
+ def _triton_step_v1(self, params, grads, exp_avgs, steps, lr, beta1, eps, weight_decay, tau, alpha_ss, cautious, clamp_min):
395
+ BLOCK_SIZE = 1024
396
+ device = params[0].device
397
+ scratch_rms = self._get_scratch_buffer(device, len(params) * 2)
398
+
399
+ for i in range(len(params)):
400
+ p, grad, exp_avg, step = params[i], grads[i], exp_avgs[i], steps[i]
401
+ is_orig_contig = p.is_contiguous()
402
+ p_contig = p if is_orig_contig else p.contiguous()
403
+ grad_contig = grad if grad.is_contiguous() else grad.contiguous()
404
+
405
+ bc1 = 1.0 - (beta1**step)
406
+ c2 = eps * bc1 * tau
407
+
408
+ n_elements = p.numel()
409
+ grid = (triton.cdiv(n_elements, BLOCK_SIZE),)
410
+ rms_sum_ptr = scratch_rms[2 * i : 2 * i + 1]
411
+ mask_sum_ptr = scratch_rms[2 * i + 1 : 2 * i + 2]
412
+ rms_sum_ptr.zero_()
413
+ mask_sum_ptr.zero_()
414
+
415
+ _lumina_v1_pass1_kernel[grid](grad_contig, exp_avg, rms_sum_ptr, n_elements, beta1, BLOCK_SIZE=BLOCK_SIZE)
416
+ _lumina_v1_pass2_kernel[grid](grad_contig, exp_avg, rms_sum_ptr, mask_sum_ptr, n_elements, beta1, tau, eps, c2, alpha_ss, BLOCK_SIZE=BLOCK_SIZE)
417
+ _lumina_v1_pass3_update_kernel[grid](p_contig, grad_contig, exp_avg, rms_sum_ptr, mask_sum_ptr, n_elements, beta1, tau, eps, c2, alpha_ss, lr, weight_decay, clamp_min, cautious=cautious, BLOCK_SIZE=BLOCK_SIZE)
418
+
419
+ if not is_orig_contig:
420
+ p.copy_(p_contig)
421
+
422
+ def _single_step_v2(self, params, grads, exp_avgs, exp_avg_sqs, steps, lr, beta1, beta2, eps, weight_decay, tau, cautious, clamp_min):
423
+ for i in range(len(params)):
424
+ p, grad, exp_avg, exp_avg_sq, step = params[i], grads[i], exp_avgs[i], exp_avg_sqs[i], steps[i]
425
+ bc1 = 1.0 - (beta1**step)
426
+ bc2 = 1.0 - (beta2**step)
427
+ c1 = (bc1 * tau) / math.sqrt(bc2)
428
+ c2 = eps * bc1 * tau
429
+
430
+ if weight_decay != 0.0:
431
+ p.mul_(1.0 - lr * weight_decay)
432
+
433
+ exp_avg.mul_(beta1).add_(grad, alpha=1.0 - beta1)
434
+ nes_m = torch.mul(exp_avg, beta1).add_(grad, alpha=1.0 - beta1)
435
+ grad_diff = grad - exp_avg
436
+ exp_avg_sq.mul_(beta2).addcmul_(grad_diff, grad_diff, value=1.0 - beta2)
437
+
438
+ sigma = exp_avg_sq.float().sqrt().mul_(c1).add_(c2)
439
+ update = torch.tanh(nes_m.float() / sigma).to(dtype=p.dtype)
440
+
441
+ if cautious:
442
+ mask = (update * grad > 0).to(dtype=grad.dtype)
443
+ mask_scale = mask.float().mean().clamp_(min=clamp_min, max=1.0).to(dtype=grad.dtype)
444
+ update = update.mul_(mask).div_(mask_scale)
445
+
446
+ p.add_(update, alpha=-lr)
447
+
448
+ def _single_step_v1(self, params, grads, exp_avgs, steps, lr, beta1, eps, weight_decay, tau, alpha_ss, cautious, clamp_min):
449
+ for i in range(len(params)):
450
+ p, grad, exp_avg, step = params[i], grads[i], exp_avgs[i], steps[i]
451
+ if weight_decay != 0.0:
452
+ p.mul_(1.0 - lr * weight_decay)
453
+
454
+ exp_avg.mul_(beta1).add_(grad, alpha=1.0 - beta1)
455
+ nes_m = torch.mul(exp_avg, beta1).add_(grad, alpha=1.0 - beta1)
456
+ bc1 = 1.0 - (beta1**step)
457
+ c2 = eps * bc1 * tau
458
+
459
+ rms = torch.sqrt(nes_m.float().square().mean() + eps)
460
+ sigma = rms * tau + c2
461
+ z = nes_m.float() / sigma
462
+ z_soft = z / (1.0 + alpha_ss * torch.abs(z))
463
+ update = torch.tanh(z_soft).to(dtype=p.dtype)
464
+
465
+ if cautious:
466
+ mask = (update * grad > 0).to(dtype=grad.dtype)
467
+ mask_scale = mask.float().mean().clamp_(min=clamp_min, max=1.0).to(dtype=grad.dtype)
468
+ update = update.mul_(mask).div_(mask_scale)
469
+
470
+ p.add_(update, alpha=-lr)
471
+
472
+ def _foreach_step_v2(self, params, grads, exp_avgs, exp_avg_sqs, steps, lr, beta1, beta2, eps, weight_decay, tau, cautious, clamp_min):
473
+ if weight_decay != 0.0:
474
+ torch._foreach_mul_(params, 1.0 - lr * weight_decay)
475
+
476
+ torch._foreach_mul_(exp_avgs, beta1)
477
+ torch._foreach_add_(exp_avgs, grads, alpha=1.0 - beta1)
478
+
479
+ nes_m_list = torch._foreach_mul(exp_avgs, beta1)
480
+ torch._foreach_add_(nes_m_list, grads, alpha=1.0 - beta1)
481
+
482
+ grad_diff_list = torch._foreach_sub(grads, exp_avgs)
483
+ torch._foreach_mul_(exp_avg_sqs, beta2)
484
+ torch._foreach_addcmul_(exp_avg_sqs, grad_diff_list, grad_diff_list, value=1.0 - beta2)
485
+
486
+ bias_correction1 = [1.0 - (beta1**st) for st in steps]
487
+ bias_correction2 = [1.0 - (beta2**st) for st in steps]
488
+ c1_list = [(bc1 * tau) / math.sqrt(bc2) for bc1, bc2 in zip(bias_correction1, bias_correction2)]
489
+ c2_list = [eps * bc1 * tau for bc1 in bias_correction1]
490
+
491
+ updates = []
492
+ for i in range(len(params)):
493
+ sigma = exp_avg_sqs[i].float().sqrt().mul_(c1_list[i]).add_(c2_list[i])
494
+ u = torch.tanh(nes_m_list[i].float() / sigma).to(dtype=params[i].dtype)
495
+ if cautious:
496
+ mask = (u * grads[i] > 0).to(dtype=grads[i].dtype)
497
+ scale = mask.float().mean().clamp_(min=clamp_min, max=1.0).to(dtype=grads[i].dtype)
498
+ u = u.mul(mask).div(scale)
499
+ updates.append(u)
500
+
501
+ torch._foreach_add_(params, updates, alpha=-lr)
502
+
503
+ def _foreach_step_v1(self, params, grads, exp_avgs, steps, lr, beta1, eps, weight_decay, tau, alpha_ss, cautious, clamp_min):
504
+ if weight_decay != 0.0:
505
+ torch._foreach_mul_(params, 1.0 - lr * weight_decay)
506
+
507
+ torch._foreach_mul_(exp_avgs, beta1)
508
+ torch._foreach_add_(exp_avgs, grads, alpha=1.0 - beta1)
509
+
510
+ nes_m_list = torch._foreach_mul(exp_avgs, beta1)
511
+ torch._foreach_add_(nes_m_list, grads, alpha=1.0 - beta1)
512
+
513
+ bias_correction1 = [1.0 - (beta1**st) for st in steps]
514
+ c2_list = [eps * bc1 * tau for bc1 in bias_correction1]
515
+
516
+ updates = []
517
+ for i in range(len(params)):
518
+ m = nes_m_list[i]
519
+ rms = torch.sqrt(m.float().square().mean() + eps)
520
+ sigma = rms * tau + c2_list[i]
521
+ z = m.float() / sigma
522
+ z_soft = z / (1.0 + alpha_ss * torch.abs(z))
523
+ u = torch.tanh(z_soft).to(dtype=params[i].dtype)
524
+ if cautious:
525
+ mask = (u * grads[i] > 0).to(dtype=grads[i].dtype)
526
+ scale = mask.float().mean().clamp_(min=clamp_min, max=1.0).to(dtype=grads[i].dtype)
527
+ u = u.mul(mask).div(scale)
528
+ updates.append(u)
529
+
530
+ torch._foreach_add_(params, updates, alpha=-lr)
modeling_xonelm.py ADDED
@@ -0,0 +1,2056 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import importlib
2
+ import importlib.util
3
+ import logging
4
+ import math
5
+ import os
6
+ import sys
7
+ from dataclasses import dataclass
8
+ from typing import Any, Callable, Dict, List, Optional, Tuple, Union
9
+
10
+ import torch
11
+ from torch import Tensor
12
+ import torch.nn as nn
13
+ import torch.nn.functional as F
14
+ import torch.utils.checkpoint as cp
15
+
16
+ logging.basicConfig(
17
+ level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s"
18
+ )
19
+ logger = logging.getLogger("XoneLM")
20
+
21
+
22
+ class HardwareContext:
23
+
24
+ @staticmethod
25
+ def get_optimal_device() -> torch.device:
26
+ if hasattr(torch, "accelerator") and torch.accelerator.is_available():
27
+ try:
28
+ acc_device = torch.accelerator.current_accelerator()
29
+ if acc_device is not None:
30
+ idx = (
31
+ torch.accelerator.current_device_index()
32
+ if hasattr(torch.accelerator, "current_device_index")
33
+ else 0
34
+ )
35
+ return torch.device(f"{acc_device.type}:{idx}")
36
+ except Exception:
37
+ pass
38
+
39
+ if "torch_xla" in sys.modules:
40
+ try:
41
+ import torch_xla.core.xla_model as xm
42
+ return xm.xla_device()
43
+ except Exception:
44
+ pass
45
+
46
+ # 3. NVIDIA CUDA GPU
47
+ if torch.cuda.is_available():
48
+ idx = (
49
+ torch.cuda.current_device()
50
+ if hasattr(torch.cuda, "current_device")
51
+ else 0
52
+ )
53
+ return torch.device(f"cuda:{idx}")
54
+
55
+ # 4. Intel XPU / Apple MPS / CPU
56
+ if hasattr(torch, "xpu") and torch.xpu.is_available():
57
+ idx = (
58
+ torch.xpu.current_device()
59
+ if hasattr(torch.xpu, "current_device")
60
+ else 0
61
+ )
62
+ return torch.device(f"xpu:{idx}")
63
+
64
+ if hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
65
+ return torch.device("mps")
66
+
67
+ return torch.device("cpu")
68
+
69
+ @staticmethod
70
+ def get_optimal_autocast_dtype(device: torch.device) -> torch.dtype:
71
+ dev_type = device.type
72
+ if dev_type == "cuda":
73
+ if torch.cuda.is_available():
74
+ major, _ = torch.cuda.get_device_capability(device)
75
+ if major < 8:
76
+ return torch.float16
77
+ # Ampere (sm_80), Ada (sm_89), Hopper (sm_90), Blackwell (sm_100/sm_120)
78
+ if torch.cuda.is_bf16_supported():
79
+ return torch.bfloat16
80
+ return torch.float16
81
+ elif dev_type == "xla":
82
+ return torch.bfloat16
83
+ elif dev_type == "xpu":
84
+ if (
85
+ hasattr(torch.xpu, "is_bf16_supported")
86
+ and torch.xpu.is_bf16_supported()
87
+ ):
88
+ return torch.bfloat16
89
+ return torch.float16
90
+ elif dev_type == "mps":
91
+ return torch.float16
92
+ elif dev_type == "cpu":
93
+ return torch.bfloat16
94
+ return torch.float32
95
+
96
+ @staticmethod
97
+ def get_autocast_context(device: torch.device):
98
+ dev_type = device.type
99
+ target_dtype = HardwareContext.get_optimal_autocast_dtype(device)
100
+
101
+ if dev_type in ("cuda", "cpu", "xpu"):
102
+ return torch.amp.autocast(
103
+ device_type=dev_type,
104
+ dtype=target_dtype,
105
+ enabled=(target_dtype != torch.float32),
106
+ )
107
+ elif dev_type == "xla":
108
+ try:
109
+ return torch.amp.autocast(
110
+ device_type="xla", dtype=torch.bfloat16, enabled=True
111
+ )
112
+ except Exception:
113
+ return torch.nullcontext()
114
+ elif dev_type == "mps":
115
+ try:
116
+ return torch.amp.autocast(
117
+ device_type="mps", dtype=torch.float16, enabled=True
118
+ )
119
+ except Exception:
120
+ return torch.nullcontext()
121
+ return torch.nullcontext()
122
+
123
+
124
+ HAS_SDPA = hasattr(F, "scaled_dot_product_attention")
125
+
126
+ HAS_FLASH_ATTN = False
127
+ _flash_attn_func = None
128
+ for mod_name, attr_name in [
129
+ ("flash_attn", "flash_attn_func"),
130
+ ("flash_attn_3.flash_attn_interface", "flash_attn_func"),
131
+ ("flash_attn_4.flash_attn_interface", "flash_attn_func"),
132
+ ]:
133
+ try:
134
+ root_pkg = mod_name.split(".")[0]
135
+ if importlib.util.find_spec(root_pkg) is not None:
136
+ mod = importlib.import_module(mod_name)
137
+ _flash_attn_func = getattr(mod, attr_name, None)
138
+ if _flash_attn_func is not None:
139
+ HAS_FLASH_ATTN = True
140
+ break
141
+ except Exception:
142
+ continue
143
+
144
+ HAS_COMPILED_FLEX_ATTENTION = False
145
+ _compiled_flex_attention_fn = None
146
+ try:
147
+ from torch.nn.attention.flex_attention import (
148
+ create_block_mask,
149
+ flex_attention as _raw_flex_attn,
150
+ )
151
+
152
+ if torch.cuda.is_available():
153
+ major, _ = torch.cuda.get_device_capability()
154
+ if major >= 8:
155
+ _compiled_flex_attention_fn = torch.compile(_raw_flex_attn, dynamic=True)
156
+ HAS_COMPILED_FLEX_ATTENTION = True
157
+ except Exception:
158
+ HAS_COMPILED_FLEX_ATTENTION = False
159
+
160
+ HAS_FUSED_LINEAR_CE = hasattr(nn, "LinearCrossEntropyLoss") or hasattr(
161
+ F, "linear_cross_entropy"
162
+ )
163
+
164
+
165
+ def create_universal_document_boundary_mask(
166
+ x_tokens: torch.Tensor,
167
+ hub_size: int,
168
+ past_k_len: int,
169
+ eod_token_id: int,
170
+ is_dense_with_hub: bool = True,
171
+ ) -> torch.Tensor:
172
+
173
+ batch_size, text_len = x_tokens.shape
174
+ cur_seq_len = (hub_size + text_len) if is_dense_with_hub else text_len
175
+ total_k_len = (hub_size + text_len) if is_dense_with_hub else (past_k_len + text_len)
176
+ device = x_tokens.device
177
+
178
+ mask = torch.zeros((batch_size, 1, cur_seq_len, total_k_len), device=device, dtype=torch.bool)
179
+
180
+ is_eod = (x_tokens == eod_token_id).long()
181
+ doc_ids = torch.cumsum(is_eod, dim=-1)
182
+ doc_ids_shifted = torch.cat(
183
+ [torch.zeros((batch_size, 1), device=device, dtype=torch.long), doc_ids[:, :-1]], dim=-1
184
+ )
185
+ doc_mismatch = (doc_ids_shifted.unsqueeze(-1) != doc_ids_shifted.unsqueeze(-2))
186
+
187
+ rows = torch.arange(text_len, device=device).unsqueeze(1)
188
+ cols = torch.arange(text_len, device=device).unsqueeze(0)
189
+ future_mask = (cols > rows).unsqueeze(0).unsqueeze(0)
190
+
191
+ if is_dense_with_hub:
192
+ mask[:, :, :hub_size, hub_size:] = True
193
+ combined_text_mask = future_mask | doc_mismatch.unsqueeze(1)
194
+ mask[:, :, hub_size:, hub_size:] = combined_text_mask
195
+ else:
196
+ past_text_len = past_k_len - hub_size
197
+ combined_text_mask = future_mask | doc_mismatch.unsqueeze(1)
198
+ mask[:, :, :, past_k_len:] = combined_text_mask
199
+
200
+ if past_text_len > 0:
201
+ past_doc_mask = (doc_ids_shifted > 0).unsqueeze(-1).expand(-1, -1, past_text_len).unsqueeze(1)
202
+ mask[:, :, :, hub_size:past_k_len] = past_doc_mask
203
+
204
+ return mask
205
+
206
+
207
+ def resolve_head_architecture(
208
+ dim: int,
209
+ num_heads: Optional[Union[int, str]] = "auto",
210
+ d_head: Optional[Union[int, str]] = "auto",
211
+ ) -> Tuple[int, int]:
212
+ if (
213
+ isinstance(num_heads, int)
214
+ and num_heads > 0
215
+ and isinstance(d_head, int)
216
+ and d_head > 0
217
+ ):
218
+ return num_heads, d_head
219
+
220
+ if isinstance(num_heads, int) and num_heads > 0:
221
+ resolved_d_head = max(16, dim // num_heads)
222
+ return num_heads, resolved_d_head
223
+
224
+ if isinstance(d_head, int) and d_head > 0:
225
+ resolved_heads = max(1, dim // d_head)
226
+ return resolved_heads, d_head
227
+
228
+ target_d_head = 2 ** round(math.log2(max(32.0, math.sqrt(2.0 * dim))))
229
+ candidate_divisors = [d for d in range(16, dim + 1, 8) if dim % d == 0]
230
+ if candidate_divisors:
231
+ resolved_d_head = min(
232
+ candidate_divisors, key=lambda x: abs(x - target_d_head)
233
+ )
234
+ else:
235
+ resolved_d_head = 64 if dim % 64 == 0 else (32 if dim % 32 == 0 else 16)
236
+
237
+ resolved_heads = max(1, dim // resolved_d_head)
238
+ return resolved_heads, resolved_d_head
239
+
240
+
241
+ TIER_CONFIGS = {
242
+ "65M": dict(
243
+ dim=512,
244
+ num_layers=12,
245
+ num_heads=8,
246
+ d_head=64,
247
+ hub_size=256,
248
+ num_specialized_hubs=8,
249
+ num_terminals=16,
250
+ slots_per_terminal=4,
251
+ max_episodic=32,
252
+ kv_latent_dim=64,
253
+ lora_rank=32,
254
+ ),
255
+ "100M": dict(
256
+ dim=640,
257
+ num_layers=14,
258
+ num_heads=10,
259
+ d_head=64,
260
+ hub_size=288,
261
+ num_specialized_hubs=10,
262
+ num_terminals=18,
263
+ slots_per_terminal=4,
264
+ max_episodic=40,
265
+ kv_latent_dim=80,
266
+ lora_rank=40,
267
+ ),
268
+ "200M": dict(
269
+ dim=896,
270
+ num_layers=18,
271
+ num_heads=14,
272
+ d_head=64,
273
+ hub_size=384,
274
+ num_specialized_hubs=12,
275
+ num_terminals=21,
276
+ slots_per_terminal=4,
277
+ max_episodic=48,
278
+ kv_latent_dim=112,
279
+ lora_rank=56,
280
+ ),
281
+ "300M": dict(
282
+ dim=1280,
283
+ num_layers=20,
284
+ num_heads=16,
285
+ d_head=80,
286
+ hub_size=448,
287
+ num_specialized_hubs=14,
288
+ num_terminals=24,
289
+ slots_per_terminal=6,
290
+ max_episodic=56,
291
+ kv_latent_dim=160,
292
+ lora_rank=80,
293
+ ),
294
+ "500M": dict(
295
+ dim=1280,
296
+ num_layers=20,
297
+ num_heads=16,
298
+ d_head=80,
299
+ hub_size=512,
300
+ num_specialized_hubs=18,
301
+ num_terminals=28,
302
+ slots_per_terminal=6,
303
+ max_episodic=72,
304
+ kv_latent_dim=160,
305
+ lora_rank=80,
306
+ ),
307
+ "750M": dict(
308
+ dim=1536,
309
+ num_layers=22,
310
+ num_heads=16,
311
+ d_head=96,
312
+ hub_size=608,
313
+ num_specialized_hubs=22,
314
+ num_terminals=32,
315
+ slots_per_terminal=6,
316
+ max_episodic=80,
317
+ kv_latent_dim=192,
318
+ lora_rank=96,
319
+ ),
320
+ "1.0B": dict(
321
+ dim=1792,
322
+ num_layers=24,
323
+ num_heads=16,
324
+ d_head=112,
325
+ hub_size=768,
326
+ num_specialized_hubs=24,
327
+ num_terminals=32,
328
+ slots_per_terminal=8,
329
+ max_episodic=88,
330
+ kv_latent_dim=224,
331
+ lora_rank=112,
332
+ ),
333
+ "3.0B": dict(
334
+ dim=2560,
335
+ num_layers=32,
336
+ num_heads=20,
337
+ d_head=128,
338
+ hub_size=992,
339
+ num_specialized_hubs=36,
340
+ num_terminals=42,
341
+ slots_per_terminal=8,
342
+ max_episodic=136,
343
+ kv_latent_dim=320,
344
+ lora_rank=160,
345
+ ),
346
+ "7.0B": dict(
347
+ dim=4096,
348
+ num_layers=36,
349
+ num_heads=32,
350
+ d_head=128,
351
+ hub_size=1312,
352
+ num_specialized_hubs=52,
353
+ num_terminals=51,
354
+ slots_per_terminal=8,
355
+ max_episodic=192,
356
+ kv_latent_dim=512,
357
+ lora_rank=256,
358
+ ),
359
+ }
360
+
361
+
362
+ @dataclass
363
+ class XoneLMOutput:
364
+ loss: Optional[torch.Tensor] = None
365
+ logits: Optional[torch.Tensor] = None
366
+ aux_loss: Optional[torch.Tensor] = None
367
+ z_loss: Optional[torch.Tensor] = None
368
+ past_key_values: Optional[List[torch.Tensor]] = None
369
+ soliton_state: Optional[List[torch.Tensor]] = None
370
+
371
+
372
+ class RMSNorm(nn.Module):
373
+
374
+ def __init__(self, dim: int, eps: float = 1e-6):
375
+ super().__init__()
376
+ self.eps = eps
377
+ self.weight = nn.Parameter(torch.ones(dim))
378
+
379
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
380
+ input_dtype = x.dtype
381
+ x_f32 = x.to(torch.float32)
382
+ variance = x_f32.pow(2).mean(dim=-1, keepdim=True)
383
+ normed = x_f32 * torch.rsqrt(variance + self.eps)
384
+ return (normed * self.weight.to(torch.float32)).to(dtype=input_dtype)
385
+
386
+
387
+ class SwiGLU(nn.Module):
388
+
389
+ def __init__(self, dim: int, multiple_of: int = 32):
390
+ super().__init__()
391
+ hidden_dim = multiple_of * (
392
+ (int(2 * (dim * 4) / 3) + multiple_of - 1) // multiple_of
393
+ )
394
+ self.w12 = nn.Linear(dim, 2 * hidden_dim, bias=False)
395
+ self.w3 = nn.Linear(hidden_dim, dim, bias=False)
396
+
397
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
398
+ gate, value = self.w12(x).chunk(2, dim=-1)
399
+ return self.w3(F.silu(gate) * value)
400
+
401
+
402
+ class PolyHoPE(nn.Module):
403
+
404
+ def __init__(self, dim: int, degree: int = 180, max_seq_len: int = 8192):
405
+ super().__init__()
406
+ self.dim = dim
407
+ self.degree = degree
408
+ self.max_seq_len = max_seq_len
409
+ self.w_poly = nn.Parameter(torch.randn(degree + 1, dim) * 0.02)
410
+
411
+ idx = torch.arange(max_seq_len, dtype=torch.float32)
412
+ t = (2.0 * idx / (max_seq_len - 1.0)) - 1.0
413
+ t = t.clamp(-1.0, 1.0)
414
+
415
+ t_poly = [torch.ones_like(t), t]
416
+ for n in range(2, degree + 1):
417
+ t_poly.append(2.0 * t * t_poly[n - 1] - t_poly[n - 2])
418
+ t_stack = torch.stack(t_poly, dim=1)
419
+ self.register_buffer("t_stack", t_stack, persistent=False)
420
+
421
+ def forward(
422
+ self,
423
+ seq_len: int,
424
+ device: torch.device,
425
+ dtype: torch.dtype,
426
+ offset: int = 0,
427
+ ) -> torch.Tensor:
428
+ grid = self.t_stack[offset : offset + seq_len].to(
429
+ device=device, dtype=torch.float32
430
+ )
431
+ pe = torch.matmul(grid, self.w_poly.to(dtype=torch.float32)).to(dtype=dtype)
432
+ return pe.unsqueeze(0)
433
+
434
+
435
+ class LinHoPE(nn.Module):
436
+
437
+ def __init__(
438
+ self,
439
+ num_heads: int = 8,
440
+ hub_size: int = 256,
441
+ num_hub_heads: Optional[int] = None,
442
+ min_slope: float = 0.01,
443
+ max_slope: float = 0.45,
444
+ mode: str = "geometric",
445
+ learnable: bool = False,
446
+ ):
447
+ super().__init__()
448
+ self.num_heads = num_heads
449
+ self.hub_size = hub_size
450
+ self.mode = mode.lower()
451
+ self.learnable = learnable
452
+
453
+ if isinstance(num_hub_heads, int) and num_hub_heads > 0:
454
+ self.num_hub_heads = max(1, min(num_hub_heads, num_heads - 1))
455
+ else:
456
+ self.num_hub_heads = max(1, num_heads // 8)
457
+
458
+ self.num_text_heads = self.num_heads - self.num_hub_heads
459
+
460
+ hub_slopes = torch.full(
461
+ (1, self.num_hub_heads, 1, 1), 0.015, dtype=torch.float32
462
+ )
463
+ if self.num_text_heads > 1:
464
+ text_slopes = torch.linspace(
465
+ min_slope, max_slope, self.num_text_heads
466
+ ).view(1, self.num_text_heads, 1, 1)
467
+ else:
468
+ text_slopes = torch.full((1, 1, 1, 1), 0.20, dtype=torch.float32)
469
+
470
+ init_slopes = torch.cat([hub_slopes, text_slopes], dim=1)
471
+
472
+ if learnable:
473
+ self.raw_slopes = nn.Parameter(
474
+ torch.log(torch.exp(init_slopes) - 1.0 + 1e-6)
475
+ )
476
+ else:
477
+ self.register_buffer("raw_slopes", init_slopes, persistent=False)
478
+
479
+ log_hub = math.log2(max(2.0, float(hub_size)))
480
+ hub_m = torch.zeros(1, self.num_hub_heads, 1, 1)
481
+ if self.num_text_heads > 1:
482
+ text_m = torch.linspace(
483
+ 2.0 * log_hub, 8.0 * log_hub, self.num_text_heads
484
+ ).view(1, self.num_text_heads, 1, 1)
485
+ else:
486
+ text_m = torch.full((1, 1, 1, 1), 4.5 * log_hub)
487
+
488
+ hub_damping = torch.cat([hub_m, text_m], dim=1)
489
+ self.register_buffer("hub_damping", hub_damping, persistent=False)
490
+
491
+ @property
492
+ def slopes(self) -> torch.Tensor:
493
+ if self.learnable:
494
+ return F.softplus(self.raw_slopes) + 1e-4
495
+ return self.raw_slopes
496
+
497
+ def forward(
498
+ self,
499
+ seq_len: int,
500
+ num_all: int,
501
+ device: torch.device,
502
+ dtype: torch.dtype,
503
+ is_dense_with_hub: bool = True,
504
+ past_c_kv: Optional[torch.Tensor] = None,
505
+ ) -> torch.Tensor:
506
+ text_k_len = num_all - self.hub_size
507
+ hub_pos = torch.arange(
508
+ -self.hub_size, 0, device=device, dtype=torch.float32
509
+ )
510
+ text_pos_k = torch.arange(0, text_k_len, device=device, dtype=torch.float32)
511
+ pos_k = torch.cat([hub_pos, text_pos_k], dim=0).unsqueeze(0)
512
+
513
+ if is_dense_with_hub and (past_c_kv is None):
514
+ text_q_len = seq_len - self.hub_size
515
+ text_pos_q = torch.arange(
516
+ 0, text_q_len, device=device, dtype=torch.float32
517
+ )
518
+ pos_q = torch.cat([hub_pos, text_pos_q], dim=0).unsqueeze(1)
519
+ else:
520
+ pos_q = torch.arange(
521
+ text_k_len - seq_len, text_k_len, device=device, dtype=torch.float32
522
+ ).unsqueeze(1)
523
+
524
+ raw_dist = (pos_q - pos_k).clamp(min=0.0)
525
+
526
+ dist_matrix = (
527
+ raw_dist.unsqueeze(0)
528
+ .unsqueeze(0)
529
+ .repeat(1, self.num_heads, 1, 1)
530
+ .to(device=device, dtype=torch.float32)
531
+ )
532
+ dist_matrix[:, :, :, : self.hub_size] = self.hub_damping.to(
533
+ device=device, dtype=torch.float32
534
+ )
535
+
536
+ active_slopes = self.slopes.to(device=device, dtype=torch.float32)
537
+ if self.mode == "rational":
538
+ denom = 1.0 + active_slopes * dist_matrix
539
+ log_bias = -torch.log(denom.clamp(min=1e-8))
540
+ else:
541
+ log_bias = -active_slopes * dist_matrix
542
+
543
+ return log_bias.clamp(min=-10.0, max=0.0).to(dtype=dtype)
544
+
545
+
546
+ class Isomorphic3DHoPE(nn.Module):
547
+
548
+ def __init__(self, dim: int, theta: float = 10000.0):
549
+ super().__init__()
550
+ self.dim = dim
551
+ self.dim_z = 2 * (dim // 8)
552
+ rem = dim - self.dim_z
553
+ self.dim_y = 2 * (rem // 4)
554
+ self.dim_x = dim - self.dim_z - self.dim_y
555
+
556
+ self.register_buffer(
557
+ "inv_freq_z",
558
+ 1.0
559
+ / (
560
+ theta
561
+ ** (
562
+ torch.arange(0, self.dim_z, 2, dtype=torch.float32) / self.dim_z
563
+ )
564
+ ),
565
+ persistent=False,
566
+ )
567
+ self.register_buffer(
568
+ "inv_freq_y",
569
+ 1.0
570
+ / (
571
+ theta
572
+ ** (
573
+ torch.arange(0, self.dim_y, 2, dtype=torch.float32) / self.dim_y
574
+ )
575
+ ),
576
+ persistent=False,
577
+ )
578
+ self.register_buffer(
579
+ "inv_freq_x",
580
+ 1.0
581
+ / (
582
+ theta
583
+ ** (
584
+ torch.arange(0, self.dim_x, 2, dtype=torch.float32) / self.dim_x
585
+ )
586
+ ),
587
+ persistent=False,
588
+ )
589
+
590
+ def forward(
591
+ self, p_z: torch.Tensor, p_y: torch.Tensor, p_x: torch.Tensor
592
+ ) -> torch.Tensor:
593
+ target_device = self.inv_freq_z.device
594
+ p_z = p_z.to(device=target_device).float()
595
+ p_y = p_y.to(device=target_device).float()
596
+ p_x = p_x.to(device=target_device).float()
597
+
598
+ omega_z = p_z.unsqueeze(-1) * self.inv_freq_z
599
+ omega_y = p_y.unsqueeze(-1) * self.inv_freq_y
600
+ omega_x = p_x.unsqueeze(-1) * self.inv_freq_x
601
+
602
+ omega_p = torch.cat([omega_z, omega_y, omega_x], dim=-1)
603
+ return torch.cat([torch.cos(omega_p), torch.sin(omega_p)], dim=-1)
604
+
605
+
606
+ class TopologicalSolitonWaveletState(nn.Module):
607
+
608
+ def __init__(self, kv_latent_dim: int, kappa: float = 0.1):
609
+ super().__init__()
610
+ self.kv_latent_dim = kv_latent_dim
611
+ self.kappa = kappa
612
+ self.ws = nn.Linear(kv_latent_dim, kv_latent_dim, bias=False)
613
+
614
+ def forward(
615
+ self, c_kv: torch.Tensor, s_prev: Optional[torch.Tensor] = None
616
+ ) -> torch.Tensor:
617
+ batch_size = c_kv.shape[0]
618
+ if s_prev is None:
619
+ s_prev = torch.zeros(
620
+ batch_size,
621
+ 1,
622
+ self.kv_latent_dim,
623
+ device=c_kv.device,
624
+ dtype=c_kv.dtype,
625
+ )
626
+
627
+ c_kv_summary = (
628
+ c_kv.mean(dim=1, keepdim=True) if c_kv.ndim == 3 else c_kv.unsqueeze(1)
629
+ )
630
+ tanh_s = torch.tanh(s_prev)
631
+ sech_sq = (1.0 - tanh_s.pow(2)).clamp(min=1e-6)
632
+ delta_s = self.kappa * sech_sq * torch.tanh(self.ws(c_kv_summary))
633
+ return s_prev + delta_s
634
+
635
+
636
+ def compute_fisher_spectral_anisotropy(
637
+ tensor: torch.Tensor, eps: float = 1e-8
638
+ ) -> torch.Tensor:
639
+ power = tensor.pow(2)
640
+ total_power = power.sum(dim=-1, keepdim=True) + eps
641
+ p_c = power / total_power
642
+ shannon_entropy = -torch.sum(p_c * torch.log(p_c + eps), dim=-1, keepdim=True)
643
+ max_entropy = math.log(max(tensor.shape[-1], 2))
644
+ anisotropy = 1.0 - (shannon_entropy / max_entropy)
645
+ return anisotropy.clamp(0.0, 1.0)
646
+
647
+
648
+ class PoincareHyperbolicTerminalRouter(nn.Module):
649
+
650
+ def __init__(
651
+ self,
652
+ dim: int,
653
+ max_terminals: int = 16,
654
+ c: float = 1.0,
655
+ eps: float = 1e-5,
656
+ ):
657
+ super().__init__()
658
+ self.dim = dim
659
+ self.max_terminals = max_terminals
660
+ self.c = c
661
+ self.eps = eps
662
+ self.query_proj = nn.Linear(dim, dim, bias=False)
663
+ self.norm = RMSNorm(dim)
664
+
665
+ def forward(
666
+ self, chunk_summaries: torch.Tensor, hub_query: torch.Tensor
667
+ ) -> torch.Tensor:
668
+ batch_size, num_chunks, hidden_dim = chunk_summaries.shape
669
+ if num_chunks <= self.max_terminals:
670
+ return self.norm(chunk_summaries)
671
+
672
+ q_mean = hub_query.mean(dim=1, keepdim=True)
673
+ q_proj = self.query_proj(q_mean)
674
+
675
+ q_norm = q_proj.norm(p=2, dim=-1, keepdim=True) + 1e-8
676
+ q_p = q_proj * (
677
+ torch.tanh(math.sqrt(self.c) * q_norm) / (math.sqrt(self.c) * q_norm)
678
+ )
679
+
680
+ c_norm = chunk_summaries.norm(p=2, dim=-1, keepdim=True) + 1e-8
681
+ c_p = chunk_summaries * (
682
+ torch.tanh(math.sqrt(self.c) * c_norm) / (math.sqrt(self.c) * c_norm)
683
+ )
684
+
685
+ u_sq = (q_p**2).sum(dim=-1, keepdim=True)
686
+ v_sq = (c_p**2).sum(dim=-1, keepdim=True)
687
+ diff_sq = ((q_p - c_p) ** 2).sum(dim=-1)
688
+
689
+ denom = torch.clamp((1.0 - u_sq) * (1.0 - v_sq), min=self.eps).squeeze(-1)
690
+ delta = 2.0 * diff_sq / denom
691
+ dist_poincare = torch.acosh(1.0 + delta)
692
+
693
+ probs = F.softmax(-dist_poincare, dim=-1)
694
+ _, top_k_indices = torch.topk(probs, self.max_terminals, dim=-1)
695
+
696
+ anchors = torch.gather(
697
+ chunk_summaries,
698
+ 1,
699
+ top_k_indices.unsqueeze(-1).expand(-1, -1, hidden_dim),
700
+ )
701
+ anchors_norm = F.normalize(anchors, p=2, dim=-1)
702
+ c_n = F.normalize(chunk_summaries, p=2, dim=-1)
703
+ assign_sim = torch.matmul(anchors_norm, c_n.transpose(-1, -2))
704
+ soft_assignment = F.softmax(assign_sim * 10.0, dim=-1)
705
+
706
+ compacted = torch.matmul(soft_assignment, chunk_summaries)
707
+ return self.norm(compacted)
708
+
709
+
710
+ class HierarchicalEpisodicMemoryBank(nn.Module):
711
+
712
+ def __init__(
713
+ self,
714
+ dim: int,
715
+ max_l1_terminals: int = 16,
716
+ max_l2_episodic: int = 32,
717
+ tau_lock: float = 0.65,
718
+ ):
719
+ super().__init__()
720
+ self.dim = dim
721
+ self.max_l1_terminals = max_l1_terminals
722
+ self.max_l2_episodic = max_l2_episodic
723
+ self.tau_lock = tau_lock
724
+ self.norm = RMSNorm(dim)
725
+
726
+ def forward(
727
+ self, chunk_summaries: torch.Tensor, hub_query: torch.Tensor
728
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
729
+ batch_size, num_chunks, hidden_dim = chunk_summaries.shape
730
+ q_norm = F.normalize(hub_query.mean(dim=1, keepdim=True), p=2, dim=-1)
731
+ c_norm = F.normalize(chunk_summaries, p=2, dim=-1)
732
+
733
+ anisotropy = compute_fisher_spectral_anisotropy(chunk_summaries).squeeze(-1)
734
+ c_l2 = chunk_summaries.norm(p=2, dim=-1)
735
+ q_align = torch.abs(
736
+ torch.matmul(c_norm, q_norm.transpose(-1, -2)).squeeze(-1)
737
+ )
738
+
739
+ s_fisher = anisotropy * c_l2 * q_align
740
+ l1_terminals = chunk_summaries[:, : min(num_chunks, self.max_l1_terminals), :]
741
+
742
+ if num_chunks > self.max_l2_episodic:
743
+ _, l2_top_idx = torch.topk(s_fisher, self.max_l2_episodic, dim=-1)
744
+ l2_episodic = torch.gather(
745
+ chunk_summaries, 1, l2_top_idx.unsqueeze(-1).expand(-1, -1, hidden_dim)
746
+ )
747
+ else:
748
+ l2_mask = (s_fisher > self.tau_lock).unsqueeze(-1)
749
+ l2_episodic = chunk_summaries * l2_mask
750
+
751
+ return self.norm(l1_terminals), self.norm(l2_episodic)
752
+
753
+
754
+ class EpistemicTruthVerifierGate(nn.Module):
755
+
756
+ def __init__(self, dim: int, tau_contra: float = 0.30, beta: float = 1.0):
757
+ super().__init__()
758
+ self.dim = dim
759
+ self.tau_contra = tau_contra
760
+ self.beta = beta
761
+ self.wk = nn.Linear(dim, dim, bias=False)
762
+ self.wv = nn.Linear(dim, dim, bias=False)
763
+ self.norm = RMSNorm(dim)
764
+
765
+ def forward(
766
+ self, claim_states: torch.Tensor, parametric_facts: torch.Tensor
767
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
768
+ k_facts = self.wk(parametric_facts)
769
+ v_facts = self.wv(parametric_facts)
770
+
771
+ w_truth = F.softmax(
772
+ self.beta
773
+ * torch.matmul(claim_states, k_facts.transpose(-1, -2))
774
+ / math.sqrt(self.dim),
775
+ dim=-1,
776
+ )
777
+ v_expected = torch.matmul(w_truth, v_facts)
778
+
779
+ c_norm = F.normalize(claim_states, p=2, dim=-1)
780
+ v_norm = F.normalize(v_expected, p=2, dim=-1)
781
+ s_epistemic = torch.sum(c_norm * v_norm, dim=-1, keepdim=True)
782
+
783
+ gate = torch.sigmoid((s_epistemic - self.tau_contra) * 5.0)
784
+ verified_states = self.norm(claim_states * gate)
785
+ fallacy_quarantine = self.norm(claim_states * (1.0 - gate))
786
+ return verified_states, fallacy_quarantine
787
+
788
+
789
+ class DirectiveCognitiveCompass(nn.Module):
790
+
791
+ def __init__(self, dim: int):
792
+ super().__init__()
793
+ self.dim = dim
794
+ self.wdir = nn.Linear(dim, dim, bias=False)
795
+ self.norm = RMSNorm(dim)
796
+
797
+ def forward(
798
+ self, chunk_summaries: torch.Tensor, directive_tokens: torch.Tensor
799
+ ) -> torch.Tensor:
800
+ v_compass = self.norm(self.wdir(directive_tokens.mean(dim=1, keepdim=True)))
801
+ v_comp_norm = F.normalize(v_compass, p=2, dim=-1)
802
+ c_norm = F.normalize(chunk_summaries, p=2, dim=-1)
803
+
804
+ align = torch.abs(torch.matmul(c_norm, v_comp_norm.transpose(-1, -2)))
805
+ anisotropy = compute_fisher_spectral_anisotropy(chunk_summaries)
806
+ c_l2 = chunk_summaries.norm(p=2, dim=-1, keepdim=True)
807
+
808
+ s_dcc = align * anisotropy * c_l2
809
+ weighted_chunks = chunk_summaries * (1.0 + torch.sigmoid(s_dcc))
810
+ return self.norm(weighted_chunks)
811
+
812
+
813
+ class FactPullingAttention(nn.Module):
814
+
815
+ def __init__(
816
+ self,
817
+ dim: int,
818
+ num_heads: int = 8,
819
+ num_terminals: int = 16,
820
+ slots_per_terminal: int = 4,
821
+ ):
822
+ super().__init__()
823
+ self.dim = dim
824
+ self.num_heads = num_heads
825
+ self.num_terminals = num_terminals
826
+ self.slots_per_terminal = slots_per_terminal
827
+ self.total_slots = num_terminals * slots_per_terminal
828
+
829
+ self.aspect_drawers = nn.Parameter(
830
+ torch.randn(1, num_terminals, slots_per_terminal, dim)
831
+ * (1.0 / math.sqrt(dim))
832
+ )
833
+
834
+ self.q_proj = nn.Linear(dim, dim, bias=False)
835
+ self.k_ctx_proj = nn.Linear(dim, dim, bias=False)
836
+ self.v_ctx_proj = nn.Linear(dim, dim, bias=False)
837
+ self.k_param_proj = nn.Linear(dim, dim, bias=False)
838
+ self.v_param_proj = nn.Linear(dim, dim, bias=False)
839
+ self.out_proj = nn.Linear(dim, dim, bias=False)
840
+ self.norm = RMSNorm(dim)
841
+ self.gamma_p = nn.Parameter(torch.ones(1))
842
+ self.beta_p = nn.Parameter(torch.zeros(1))
843
+
844
+ def forward(
845
+ self,
846
+ terminal_base: torch.Tensor,
847
+ context_states: torch.Tensor,
848
+ parametric_facts: torch.Tensor,
849
+ ) -> torch.Tensor:
850
+ batch_size = context_states.shape[0]
851
+
852
+ q_drawers = (terminal_base.unsqueeze(2) + self.aspect_drawers).view(
853
+ batch_size, self.total_slots, self.dim
854
+ )
855
+
856
+ q = self.q_proj(q_drawers)
857
+ k_ctx = self.k_ctx_proj(context_states)
858
+ v_ctx = self.v_ctx_proj(context_states)
859
+ k_param = self.k_param_proj(parametric_facts)
860
+ v_param = self.v_param_proj(parametric_facts)
861
+
862
+ scores_ctx = torch.matmul(q, k_ctx.transpose(-1, -2)) / math.sqrt(self.dim)
863
+ scores_param = (
864
+ torch.matmul(q, k_param.transpose(-1, -2)) / math.sqrt(self.dim)
865
+ ) * self.gamma_p + self.beta_p
866
+
867
+ probs = F.softmax(
868
+ torch.cat([scores_ctx, scores_param], dim=-1), dim=-1
869
+ ).to(dtype=q.dtype)
870
+ v_comb = torch.cat([v_ctx, v_param], dim=-2)
871
+ attn_out = torch.matmul(probs, v_comb)
872
+ return self.norm(self.out_proj(attn_out))
873
+
874
+
875
+ class LatentMentalRollout(nn.Module):
876
+
877
+ def __init__(self, dim: int, num_rollout_steps: int = 2):
878
+ super().__init__()
879
+ self.dim = dim
880
+ self.num_rollout_steps = num_rollout_steps
881
+ self.w_transition = nn.Parameter(
882
+ torch.randn(dim, dim) * (0.02 / math.sqrt(dim))
883
+ )
884
+ self.norm = RMSNorm(dim)
885
+
886
+ def forward(self, h_dyn: torch.Tensor) -> torch.Tensor:
887
+ for _ in range(self.num_rollout_steps):
888
+ delta = torch.tanh(torch.matmul(h_dyn, self.w_transition))
889
+ h_dyn = self.norm(h_dyn + delta)
890
+ return h_dyn
891
+
892
+
893
+ class PhysicsCausalDiffusionAttention(nn.Module):
894
+
895
+ def __init__(
896
+ self,
897
+ dim: int,
898
+ num_heads: int = 8,
899
+ d_head: int = 64,
900
+ kv_latent_dim: int = 64,
901
+ hub_size: int = 256,
902
+ num_hub_heads: Optional[int] = None,
903
+ use_sdpa: bool = True,
904
+ ):
905
+ super().__init__()
906
+ self.dim = dim
907
+ self.num_heads = num_heads
908
+ self.d_head = d_head
909
+ self.inner_attn_dim = num_heads * d_head
910
+ self.kv_latent_dim = kv_latent_dim
911
+ self.hub_size = hub_size
912
+ self.use_sdpa = use_sdpa and HAS_SDPA
913
+
914
+ self.kv_down_proj = nn.Linear(dim, kv_latent_dim, bias=False)
915
+ self.kv_ln = RMSNorm(kv_latent_dim)
916
+ self.kv_up_proj = nn.Linear(
917
+ kv_latent_dim, 2 * self.inner_attn_dim, bias=False
918
+ )
919
+
920
+ self.q_proj = nn.Linear(dim, self.inner_attn_dim, bias=False)
921
+ self.out_proj = nn.Linear(self.inner_attn_dim, dim, bias=False)
922
+
923
+ self.q_norm = RMSNorm(self.d_head)
924
+ self.k_norm = RMSNorm(self.d_head)
925
+
926
+ self.lin_hope = LinHoPE(
927
+ num_heads=num_heads,
928
+ hub_size=hub_size,
929
+ num_hub_heads=num_hub_heads,
930
+ min_slope=0.01,
931
+ max_slope=0.45,
932
+ mode="geometric",
933
+ )
934
+ self.scale_factor = 1.0 / math.sqrt(self.d_head)
935
+
936
+ def forward(
937
+ self,
938
+ x: torch.Tensor,
939
+ attn_mask: Optional[torch.Tensor] = None,
940
+ is_causal: bool = True,
941
+ past_c_kv: Optional[torch.Tensor] = None,
942
+ soliton_state: Optional[torch.Tensor] = None,
943
+ flex_block_mask=None,
944
+ is_dense_with_hub: bool = True,
945
+ ) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]:
946
+ batch_size, seq_len, _ = x.shape
947
+ c_kv_current = self.kv_ln(self.kv_down_proj(x))
948
+
949
+ if past_c_kv is not None:
950
+ c_kv_all = torch.cat([past_c_kv, c_kv_current], dim=1)
951
+ else:
952
+ c_kv_all = c_kv_current
953
+
954
+ num_all = c_kv_all.shape[1]
955
+ k_up, v_up = torch.split(
956
+ self.kv_up_proj(c_kv_all),
957
+ [self.inner_attn_dim, self.inner_attn_dim],
958
+ dim=-1,
959
+ )
960
+
961
+ k = k_up.reshape(
962
+ batch_size, num_all, self.num_heads, self.d_head
963
+ ).permute(0, 2, 1, 3)
964
+ v = v_up.reshape(
965
+ batch_size, num_all, self.num_heads, self.d_head
966
+ ).permute(0, 2, 1, 3)
967
+ q = self.q_proj(x).reshape(
968
+ batch_size, seq_len, self.num_heads, self.d_head
969
+ ).permute(0, 2, 1, 3)
970
+
971
+ q = self.q_norm(q)
972
+ k = self.k_norm(k)
973
+
974
+ log_decay_bias = self.lin_hope(
975
+ seq_len=seq_len,
976
+ num_all=num_all,
977
+ device=x.device,
978
+ dtype=q.dtype,
979
+ is_dense_with_hub=is_dense_with_hub,
980
+ past_c_kv=past_c_kv,
981
+ )
982
+
983
+ full_mask = log_decay_bias.clone()
984
+
985
+ if attn_mask is not None:
986
+ if attn_mask.dtype == torch.bool:
987
+ full_mask = full_mask.masked_fill(attn_mask, -10000.0)
988
+ else:
989
+ full_mask = full_mask + attn_mask
990
+ elif is_causal:
991
+ text_k_len = num_all - self.hub_size
992
+ causal_m = torch.zeros(
993
+ (1, 1, seq_len, num_all), device=x.device, dtype=q.dtype
994
+ )
995
+ if is_dense_with_hub and (past_c_kv is None):
996
+ text_q_len = seq_len - self.hub_size
997
+ causal_m[:, :, : self.hub_size, self.hub_size :] = -10000.0
998
+ rows = torch.arange(text_q_len, device=x.device).unsqueeze(1)
999
+ cols = torch.arange(text_k_len, device=x.device).unsqueeze(0)
1000
+ causal_m[:, :, self.hub_size :, self.hub_size :].masked_fill_(
1001
+ cols > rows, -10000.0
1002
+ )
1003
+ else:
1004
+ pos_q = torch.arange(
1005
+ text_k_len - seq_len,
1006
+ text_k_len,
1007
+ device=x.device,
1008
+ dtype=torch.float32,
1009
+ ).unsqueeze(1)
1010
+ text_pos_k = torch.arange(
1011
+ 0, text_k_len, device=x.device, dtype=torch.float32
1012
+ ).unsqueeze(0)
1013
+ future_text = (pos_q - text_pos_k) < 0
1014
+ causal_m[:, :, :, self.hub_size :].masked_fill_(
1015
+ future_text.unsqueeze(0).unsqueeze(0), -10000.0
1016
+ )
1017
+ full_mask = full_mask + causal_m
1018
+
1019
+ if self.use_sdpa:
1020
+ attn_out = F.scaled_dot_product_attention(
1021
+ q, k, v, attn_mask=full_mask, scale=self.scale_factor
1022
+ )
1023
+ else:
1024
+ scores = (
1025
+ torch.matmul(q, k.transpose(-1, -2)) * self.scale_factor + full_mask
1026
+ )
1027
+ probs = F.softmax(scores, dim=-1, dtype=torch.float32).to(dtype=q.dtype)
1028
+ attn_out = torch.matmul(probs, v)
1029
+
1030
+ attn_out = attn_out.permute(0, 2, 1, 3).reshape(
1031
+ batch_size, seq_len, self.inner_attn_dim
1032
+ )
1033
+ final_out = self.out_proj(attn_out)
1034
+ return final_out, c_kv_all, None
1035
+
1036
+
1037
+ class XoneLMBlock(nn.Module):
1038
+
1039
+ def __init__(
1040
+ self,
1041
+ dim: int,
1042
+ num_heads: int = 8,
1043
+ d_head: int = 64,
1044
+ kv_latent_dim: int = 64,
1045
+ hub_size: int = 256,
1046
+ num_hub_heads: Optional[int] = None,
1047
+ use_sdpa: bool = True,
1048
+ ):
1049
+ super().__init__()
1050
+ self.ln1 = RMSNorm(dim)
1051
+ self.attn = PhysicsCausalDiffusionAttention(
1052
+ dim=dim,
1053
+ num_heads=num_heads,
1054
+ d_head=d_head,
1055
+ kv_latent_dim=kv_latent_dim,
1056
+ hub_size=hub_size,
1057
+ num_hub_heads=num_hub_heads,
1058
+ use_sdpa=use_sdpa,
1059
+ )
1060
+ self.ln2 = RMSNorm(dim)
1061
+ self.hub_size = hub_size
1062
+ self.ffn = SwiGLU(dim, multiple_of=32)
1063
+ self.poly_alpha = nn.Parameter(torch.tensor(0.05))
1064
+
1065
+ def forward(
1066
+ self,
1067
+ x: torch.Tensor,
1068
+ poly_pe: Optional[torch.Tensor] = None,
1069
+ attn_mask: Optional[torch.Tensor] = None,
1070
+ is_causal: bool = True,
1071
+ past_c_kv: Optional[torch.Tensor] = None,
1072
+ soliton_state: Optional[torch.Tensor] = None,
1073
+ flex_block_mask=None,
1074
+ is_dense_with_hub: bool = True,
1075
+ ) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]:
1076
+ if poly_pe is not None:
1077
+ x = x + self.poly_alpha * poly_pe
1078
+
1079
+ attn_out, new_past_c_kv, new_soliton = self.attn(
1080
+ self.ln1(x),
1081
+ attn_mask=attn_mask,
1082
+ is_causal=is_causal,
1083
+ past_c_kv=past_c_kv,
1084
+ soliton_state=soliton_state,
1085
+ is_dense_with_hub=is_dense_with_hub,
1086
+ )
1087
+
1088
+ x = x + attn_out
1089
+ x = x + self.ffn(self.ln2(x))
1090
+ return x, new_past_c_kv, new_soliton
1091
+
1092
+
1093
+ class XoneLM(nn.Module):
1094
+
1095
+ def __init__(
1096
+ self,
1097
+ tier: Optional[str] = None,
1098
+ vocab_size: int = 32000,
1099
+ dim: Optional[int] = None,
1100
+ num_layers: Optional[int] = None,
1101
+ num_heads: Optional[Union[int, str]] = "auto",
1102
+ d_head: Optional[Union[int, str]] = "auto",
1103
+ num_hub_heads: Optional[int] = None,
1104
+ hub_size: Optional[int] = None,
1105
+ num_specialized_hubs: Optional[int] = None,
1106
+ num_terminals: Optional[int] = None,
1107
+ slots_per_terminal: int = 4,
1108
+ max_terminals: Optional[int] = None,
1109
+ tokens_per_terminal: Optional[int] = None,
1110
+ max_episodic: Optional[int] = None,
1111
+ kv_latent_dim: Optional[int] = None,
1112
+ kv_lora_dim: Optional[int] = None,
1113
+ lora_rank: Optional[int] = None,
1114
+ chunk_size: int = 1024,
1115
+ alpha_anchor: float = 0.1,
1116
+ checkpoint_every_n: int = 0,
1117
+ separator_token_id: Optional[int] = None,
1118
+ config: Optional[Any] = None,
1119
+ hub_mode: str = "moh",
1120
+ use_sdpa: bool = True,
1121
+ **kwargs,
1122
+ ):
1123
+ super().__init__()
1124
+
1125
+ base_params = {}
1126
+ if tier is not None and tier in TIER_CONFIGS:
1127
+ base_params = TIER_CONFIGS[tier].copy()
1128
+
1129
+ if kv_latent_dim is None and kv_lora_dim is not None:
1130
+ kv_latent_dim = kv_lora_dim
1131
+
1132
+ self.vocab_size = (
1133
+ vocab_size
1134
+ if config is None
1135
+ else getattr(config, "vocab_size", vocab_size)
1136
+ )
1137
+
1138
+ self.dim = dim if dim is not None else base_params.get("dim", 512)
1139
+ self.num_layers = (
1140
+ num_layers
1141
+ if num_layers is not None
1142
+ else base_params.get("num_layers", 12)
1143
+ )
1144
+
1145
+ cfg_heads = base_params.get("num_heads", "auto")
1146
+ cfg_d_head = base_params.get("d_head", "auto")
1147
+ req_heads = num_heads if num_heads != "auto" else cfg_heads
1148
+ req_d_head = d_head if d_head != "auto" else cfg_d_head
1149
+
1150
+ self.num_heads, self.d_head = resolve_head_architecture(
1151
+ dim=self.dim, num_heads=req_heads, d_head=req_d_head
1152
+ )
1153
+
1154
+ self.hub_size = (
1155
+ hub_size if hub_size is not None else base_params.get("hub_size", 256)
1156
+ )
1157
+ self.num_specialized_hubs = (
1158
+ num_specialized_hubs
1159
+ if num_specialized_hubs is not None
1160
+ else base_params.get("num_specialized_hubs", 8)
1161
+ )
1162
+
1163
+ self.num_terminals = (
1164
+ num_terminals
1165
+ or max_terminals
1166
+ or base_params.get("num_terminals", 16)
1167
+ )
1168
+ self.slots_per_terminal = (
1169
+ slots_per_terminal or base_params.get("slots_per_terminal", 4)
1170
+ )
1171
+ self.total_terminal_slots = (
1172
+ self.num_terminals * self.slots_per_terminal
1173
+ )
1174
+
1175
+ self.max_episodic = (
1176
+ max_episodic
1177
+ if max_episodic is not None
1178
+ else base_params.get("max_episodic", 32)
1179
+ )
1180
+ self.kv_latent_dim = (
1181
+ kv_latent_dim
1182
+ if kv_latent_dim is not None
1183
+ else base_params.get("kv_latent_dim", max(64, self.dim // 8))
1184
+ )
1185
+ self.lora_rank = (
1186
+ lora_rank
1187
+ if lora_rank is not None
1188
+ else base_params.get("lora_rank", max(32, self.kv_latent_dim // 2))
1189
+ )
1190
+
1191
+ self.chunk_size = chunk_size
1192
+ self.alpha_anchor = alpha_anchor
1193
+ self.checkpoint_every_n = max(0, checkpoint_every_n)
1194
+ self.separator_token_id = separator_token_id
1195
+ self.hub_mode = hub_mode
1196
+ self.num_hub_heads = num_hub_heads
1197
+
1198
+ self.static_hub_slots = self.hub_size // 2
1199
+ self.dynamic_hub_slots = self.hub_size - self.static_hub_slots
1200
+
1201
+ mask_static = torch.zeros(1, self.static_hub_slots, 1, dtype=torch.float32)
1202
+ mask_dynamic = torch.ones(1, self.dynamic_hub_slots, 1, dtype=torch.float32)
1203
+ self.register_buffer(
1204
+ "dynamic_slot_mask",
1205
+ torch.cat([mask_static, mask_dynamic], dim=1),
1206
+ persistent=False,
1207
+ )
1208
+
1209
+ raw_basis = torch.randn(self.dim, self.hub_size)
1210
+ q_basis, _ = torch.linalg.qr(raw_basis)
1211
+ self.shared_hub_base = nn.Parameter(q_basis.T.unsqueeze(0).contiguous())
1212
+
1213
+ self.hub_lora_a = nn.Parameter(
1214
+ torch.randn(self.num_specialized_hubs, self.hub_size, self.lora_rank)
1215
+ * 0.02
1216
+ )
1217
+ self.hub_lora_b = nn.Parameter(
1218
+ torch.randn(self.num_specialized_hubs, self.lora_rank, self.dim) * 0.02
1219
+ )
1220
+ self.hub_router_gate = nn.Linear(
1221
+ self.dim, self.num_specialized_hubs, bias=False
1222
+ )
1223
+
1224
+ self.w_anchor = nn.Linear(self.dim, self.dim, bias=False)
1225
+ self.token_embeddings = nn.Embedding(self.vocab_size, self.dim)
1226
+ nn.init.normal_(self.token_embeddings.weight, mean=0.0, std=0.02)
1227
+
1228
+ self.poly_hope = PolyHoPE(dim=self.dim, degree=180, max_seq_len=8192)
1229
+ self.hope_3d = Isomorphic3DHoPE(dim=self.dim)
1230
+ self.terminal_router = PoincareHyperbolicTerminalRouter(
1231
+ dim=self.dim, max_terminals=self.num_terminals
1232
+ )
1233
+ self.episodic_memory = HierarchicalEpisodicMemoryBank(
1234
+ dim=self.dim,
1235
+ max_l1_terminals=self.num_terminals,
1236
+ max_l2_episodic=self.max_episodic,
1237
+ )
1238
+ self.etvg = EpistemicTruthVerifierGate(dim=self.dim)
1239
+ self.dcc = DirectiveCognitiveCompass(dim=self.dim)
1240
+ self.latent_rollout = LatentMentalRollout(dim=self.dim)
1241
+ self.slot_expander = nn.Linear(self.dim, self.dim, bias=False)
1242
+
1243
+ self.layers = nn.ModuleList([
1244
+ XoneLMBlock(
1245
+ dim=self.dim,
1246
+ num_heads=self.num_heads,
1247
+ d_head=self.d_head,
1248
+ kv_latent_dim=self.kv_latent_dim,
1249
+ hub_size=self.hub_size,
1250
+ num_hub_heads=num_hub_heads,
1251
+ use_sdpa=use_sdpa,
1252
+ )
1253
+ for _ in range(self.num_layers)
1254
+ ])
1255
+
1256
+ self.norm = RMSNorm(self.dim)
1257
+ self.head = nn.Linear(self.dim, self.vocab_size, bias=False)
1258
+ self.head.weight = self.token_embeddings.weight
1259
+
1260
+ self.fused_ce_loss = None
1261
+ if hasattr(nn, "LinearCrossEntropyLoss"):
1262
+ try:
1263
+ self.fused_ce_loss = nn.LinearCrossEntropyLoss(ignore_index=-100)
1264
+ except Exception:
1265
+ self.fused_ce_loss = None
1266
+
1267
+ self.parametric_facts = nn.Parameter(torch.randn(1, 32, self.dim) * 0.02)
1268
+
1269
+ self.fact_puller = FactPullingAttention(
1270
+ dim=self.dim,
1271
+ num_heads=self.num_heads,
1272
+ num_terminals=self.num_terminals,
1273
+ slots_per_terminal=self.slots_per_terminal,
1274
+ )
1275
+ self._block_mask_cache: Dict[Tuple[int, int, str], Any] = {}
1276
+
1277
+ @property
1278
+ def total_active_hub_slots(self) -> int:
1279
+ return self.hub_size
1280
+
1281
+ def _init_hub_base(
1282
+ self, batch_size: int, query_rep: Optional[torch.Tensor] = None
1283
+ ) -> torch.Tensor:
1284
+ base = self.shared_hub_base.expand(batch_size, -1, -1)
1285
+ if query_rep is None:
1286
+ return base
1287
+ q_vec = query_rep.mean(dim=1)
1288
+ routing_weights = F.softmax(
1289
+ self.hub_router_gate(q_vec) / math.sqrt(self.dim), dim=-1
1290
+ )
1291
+ expert_deltas = torch.bmm(self.hub_lora_a, self.hub_lora_b)
1292
+ combined_delta = torch.einsum("be,esd->bsd", routing_weights, expert_deltas)
1293
+ return base + combined_delta
1294
+
1295
+ def extract_hub(
1296
+ self,
1297
+ x: torch.Tensor,
1298
+ page_idx: Optional[torch.Tensor] = None,
1299
+ para_idx: Optional[torch.Tensor] = None,
1300
+ sent_idx: Optional[torch.Tensor] = None,
1301
+ directive_tokens: Optional[torch.Tensor] = None,
1302
+ is_sft: bool = False,
1303
+ ) -> torch.Tensor:
1304
+ batch_size, total_len = x.shape
1305
+ chunk_size = min(total_len, self.chunk_size)
1306
+ x_prompt_only = x[:, :chunk_size]
1307
+ prompt_len = x_prompt_only.shape[1]
1308
+
1309
+ if page_idx is None:
1310
+ page_idx = torch.zeros(batch_size, 1, device=x.device, dtype=torch.long)
1311
+ if para_idx is None:
1312
+ para_idx = torch.zeros(batch_size, 1, device=x.device, dtype=torch.long)
1313
+ if sent_idx is None:
1314
+ sent_idx = torch.zeros(batch_size, 1, device=x.device, dtype=torch.long)
1315
+
1316
+ hope_anchor = self.hope_3d(page_idx, para_idx, sent_idx)
1317
+ hope_anchor_scaled = F.normalize(hope_anchor, p=2, dim=-1) * 0.05
1318
+
1319
+ pe_text = self.poly_hope(
1320
+ prompt_len, x.device, self.token_embeddings.weight.dtype, offset=0
1321
+ )
1322
+ h_text_all = self.token_embeddings(x_prompt_only) + pe_text
1323
+
1324
+ hub_init = self._init_hub_base(batch_size, query_rep=h_text_all)
1325
+ c_term_base = hub_init[:, : self.num_terminals, :]
1326
+ expanded_pf = self.parametric_facts.expand(batch_size, -1, -1)
1327
+
1328
+ c_hub_out = self.fact_puller(c_term_base, h_text_all, expanded_pf)
1329
+ if directive_tokens is not None:
1330
+ c_hub_out = self.dcc(c_hub_out, directive_tokens)
1331
+ verified_terminals, _ = self.etvg(c_hub_out, expanded_pf)
1332
+
1333
+ q_norm = F.normalize(hub_init, p=2, dim=-1)
1334
+ k_norm = F.normalize(verified_terminals, p=2, dim=-1)
1335
+ attn_weights = F.softmax(
1336
+ torch.matmul(q_norm, k_norm.transpose(-1, -2)) * 8.0, dim=-1
1337
+ ) # [B, hub_size, total_terminal_slots]
1338
+ context_features = torch.matmul(
1339
+ attn_weights, verified_terminals
1340
+ ) # [B, hub_size, dim]
1341
+
1342
+ normed_base = F.normalize(hub_init, p=2, dim=-1)
1343
+ normed_context = F.normalize(context_features, p=2, dim=-1)
1344
+ hub = (
1345
+ 0.85 * normed_base + 0.15 * normed_context + hope_anchor_scaled
1346
+ ) * math.sqrt(self.dim)
1347
+
1348
+ if self.dynamic_hub_slots > 0:
1349
+ h_evolved = self.latent_rollout(hub)
1350
+ delta_evolved = h_evolved - hub
1351
+ hub = hub + delta_evolved * self.dynamic_slot_mask.to(dtype=hub.dtype)
1352
+
1353
+ return hub.to(dtype=self.token_embeddings.weight.dtype)
1354
+
1355
+ def compute_hub_diversity_loss(self, hub: torch.Tensor) -> torch.Tensor:
1356
+ active_mask = (hub.norm(dim=-1) > 1e-4).float()
1357
+ pair_mask = torch.matmul(active_mask.unsqueeze(-1), active_mask.unsqueeze(-2))
1358
+ h_norm = F.normalize(hub, p=2, dim=-1, eps=1e-8)
1359
+ sim_matrix = torch.matmul(h_norm, h_norm.transpose(-1, -2))
1360
+ identity = torch.eye(
1361
+ self.hub_size, device=hub.device, dtype=hub.dtype
1362
+ ).unsqueeze(0)
1363
+ diff_sq = (sim_matrix - identity).pow(2) * pair_mask
1364
+ return diff_sq.sum() / pair_mask.sum().clamp(min=1.0)
1365
+
1366
+ def _compute_loss_efficient(
1367
+ self, hidden_states: torch.Tensor, labels: torch.Tensor
1368
+ ) -> torch.Tensor:
1369
+ valid_mask = labels != -100
1370
+ if not valid_mask.any():
1371
+ return (hidden_states * 0.0).sum()
1372
+
1373
+ if hasattr(F, "linear_cross_entropy"):
1374
+ try:
1375
+ return F.linear_cross_entropy(
1376
+ hidden_states,
1377
+ self.head.weight,
1378
+ labels,
1379
+ ignore_index=-100,
1380
+ reduction="mean",
1381
+ )
1382
+ except Exception:
1383
+ pass
1384
+
1385
+ if self.fused_ce_loss is not None:
1386
+ try:
1387
+ return self.fused_ce_loss(hidden_states, self.head.weight, labels)
1388
+ except Exception:
1389
+ pass
1390
+
1391
+ logits = self.head(hidden_states)
1392
+ return F.cross_entropy(
1393
+ logits.reshape(-1, self.vocab_size).float(),
1394
+ labels.reshape(-1),
1395
+ ignore_index=-100,
1396
+ )
1397
+
1398
+ def _forward_dense(
1399
+ self,
1400
+ x: torch.Tensor,
1401
+ past_c_kv_list: Optional[List[torch.Tensor]] = None,
1402
+ soliton_state_list: Optional[List[torch.Tensor]] = None,
1403
+ override_hub: Optional[torch.Tensor] = None,
1404
+ return_logits: bool = True,
1405
+ past_key_values: Optional[Any] = None,
1406
+ attn_mask: Optional[torch.Tensor] = None,
1407
+ **kwargs,
1408
+ ) -> Tuple:
1409
+ if past_c_kv_list is None and past_key_values is not None:
1410
+ past_c_kv_list = past_key_values
1411
+
1412
+ batch_size, seq_len = x.shape
1413
+ target_dtype = self.token_embeddings.weight.dtype
1414
+
1415
+ if past_c_kv_list is None:
1416
+ pe_text = self.poly_hope(seq_len, x.device, target_dtype, offset=0)
1417
+ h_text = self.token_embeddings(x) + pe_text
1418
+ hub = (
1419
+ override_hub.to(dtype=target_dtype)
1420
+ if override_hub is not None
1421
+ else self._init_hub_base(batch_size, query_rep=h_text).to(
1422
+ dtype=target_dtype
1423
+ )
1424
+ )
1425
+ h_initial = torch.cat([hub, h_text], dim=1)
1426
+ h = h_initial
1427
+
1428
+ hub_zeros = torch.zeros(
1429
+ 1, self.hub_size, self.dim, device=x.device, dtype=target_dtype
1430
+ )
1431
+ poly_pe_full = torch.cat([hub_zeros, pe_text], dim=1)
1432
+
1433
+ new_past_c_kv_list = []
1434
+
1435
+ for i, layer in enumerate(self.layers):
1436
+ if (
1437
+ self.training
1438
+ and (self.checkpoint_every_n > 0)
1439
+ and (i % self.checkpoint_every_n == 0)
1440
+ ):
1441
+
1442
+ def make_checkpoint_fn(l_mod):
1443
+
1444
+ def forward_fn(hidden_states, pe, m):
1445
+ return l_mod(
1446
+ hidden_states,
1447
+ poly_pe=pe,
1448
+ attn_mask=m,
1449
+ is_causal=True,
1450
+ is_dense_with_hub=True,
1451
+ )
1452
+
1453
+ return forward_fn
1454
+
1455
+ h, layer_c_kv, _ = cp.checkpoint(
1456
+ make_checkpoint_fn(layer),
1457
+ h,
1458
+ poly_pe_full,
1459
+ attn_mask,
1460
+ use_reentrant=False,
1461
+ )
1462
+ else:
1463
+ h, layer_c_kv, _ = layer(
1464
+ h,
1465
+ poly_pe=poly_pe_full,
1466
+ attn_mask=attn_mask,
1467
+ is_causal=True,
1468
+ is_dense_with_hub=True,
1469
+ )
1470
+
1471
+ new_past_c_kv_list.append(layer_c_kv)
1472
+
1473
+ hub_l0 = h_initial[:, : self.hub_size, :]
1474
+ hub_ln = h[:, : self.hub_size, :]
1475
+ hub_anchored = hub_ln + self.alpha_anchor * self.w_anchor(
1476
+ self.norm(hub_l0)
1477
+ )
1478
+ text_evolved = self.norm(h[:, self.hub_size :, :])
1479
+
1480
+ out_logits = self.head(text_evolved) if return_logits else text_evolved
1481
+ z_loss = self.compute_hub_diversity_loss(hub_anchored)
1482
+ return (
1483
+ out_logits,
1484
+ torch.tensor(0.0, device=x.device, dtype=target_dtype),
1485
+ z_loss,
1486
+ new_past_c_kv_list,
1487
+ None,
1488
+ )
1489
+ else:
1490
+ n_past = past_c_kv_list[0].shape[1]
1491
+ pos_offset = max(0, n_past - self.hub_size)
1492
+ pe_step = self.poly_hope(
1493
+ seq_len, x.device, target_dtype, offset=pos_offset
1494
+ )
1495
+ h = self.token_embeddings(x) + pe_step
1496
+
1497
+ new_past_c_kv_list = []
1498
+
1499
+ for i, layer in enumerate(self.layers):
1500
+ past_layer_c_kv = past_c_kv_list[i]
1501
+ if (
1502
+ self.training
1503
+ and (self.checkpoint_every_n > 0)
1504
+ and (i % self.checkpoint_every_n == 0)
1505
+ ):
1506
+
1507
+ def make_checkpoint_fn_kv(l_mod, p_ckv):
1508
+
1509
+ def forward_fn(hidden_states, pe, m):
1510
+ return l_mod(
1511
+ hidden_states,
1512
+ poly_pe=pe,
1513
+ attn_mask=m,
1514
+ is_causal=True,
1515
+ past_c_kv=p_ckv,
1516
+ is_dense_with_hub=False,
1517
+ )
1518
+
1519
+ return forward_fn
1520
+
1521
+ h, layer_c_kv, _ = cp.checkpoint(
1522
+ make_checkpoint_fn_kv(layer, past_layer_c_kv),
1523
+ h,
1524
+ pe_step,
1525
+ attn_mask,
1526
+ use_reentrant=False,
1527
+ )
1528
+ else:
1529
+ h, layer_c_kv, _ = layer(
1530
+ h,
1531
+ poly_pe=pe_step,
1532
+ attn_mask=attn_mask,
1533
+ is_causal=True,
1534
+ past_c_kv=past_layer_c_kv,
1535
+ is_dense_with_hub=False,
1536
+ )
1537
+
1538
+ new_past_c_kv_list.append(layer_c_kv)
1539
+
1540
+ h_normed = self.norm(h)
1541
+ out_logits = self.head(h_normed) if return_logits else h_normed
1542
+ return (
1543
+ out_logits,
1544
+ torch.tensor(0.0, device=x.device, dtype=target_dtype),
1545
+ torch.tensor(0.0, device=x.device, dtype=target_dtype),
1546
+ new_past_c_kv_list,
1547
+ None,
1548
+ )
1549
+
1550
+ def forward(
1551
+ self,
1552
+ x: torch.Tensor,
1553
+ labels: Optional[torch.Tensor] = None,
1554
+ scaler: Optional[torch.amp.GradScaler] = None,
1555
+ grad_accum_steps: int = 1,
1556
+ past_key_values: Optional[List] = None,
1557
+ soliton_states: Optional[List] = None,
1558
+ directive_tokens: Optional[torch.Tensor] = None,
1559
+ override_hub: Optional[torch.Tensor] = None,
1560
+ is_sft: bool = False,
1561
+ execute_chunk_backward: bool = False,
1562
+ chunk_callback: Optional[Callable[[int, int], None]] = None,
1563
+ attn_mask: Optional[torch.Tensor] = None,
1564
+ **kwargs,
1565
+ ) -> XoneLMOutput:
1566
+ if past_key_values is not None or x.shape[1] == 1:
1567
+ if chunk_callback is not None:
1568
+ chunk_callback(1, 1)
1569
+ logits, aux_l, z_l, new_kv, new_sol = self._forward_dense(
1570
+ x,
1571
+ past_c_kv_list=past_key_values,
1572
+ soliton_state_list=soliton_states,
1573
+ return_logits=True,
1574
+ attn_mask=attn_mask,
1575
+ )
1576
+ return XoneLMOutput(
1577
+ logits=logits,
1578
+ aux_loss=aux_l,
1579
+ z_loss=z_l,
1580
+ past_key_values=new_kv,
1581
+ soliton_state=new_sol,
1582
+ )
1583
+
1584
+ batch_size, seq_len = x.shape
1585
+ chunk_size = self.chunk_size
1586
+
1587
+ hub = (
1588
+ override_hub
1589
+ if override_hub is not None
1590
+ else self.extract_hub(
1591
+ x, directive_tokens=directive_tokens, is_sft=is_sft
1592
+ )
1593
+ )
1594
+ target_dtype = self.token_embeddings.weight.dtype
1595
+
1596
+ if seq_len > chunk_size:
1597
+ num_chunks = math.ceil(seq_len / chunk_size)
1598
+ pad_len = (num_chunks * chunk_size) - seq_len
1599
+
1600
+ x_padded = F.pad(x, (0, pad_len), value=0) if pad_len > 0 else x
1601
+ labels_padded = (
1602
+ F.pad(labels, (0, pad_len), value=-100)
1603
+ if (labels is not None and pad_len > 0)
1604
+ else labels
1605
+ )
1606
+
1607
+ if self.training and labels is not None and execute_chunk_backward:
1608
+ total_lm_loss_val, total_z_loss_val, past_kv = 0.0, 0.0, None
1609
+ device = x.device
1610
+
1611
+ for c_idx in range(num_chunks):
1612
+ if chunk_callback is not None:
1613
+ chunk_callback(c_idx + 1, num_chunks)
1614
+
1615
+ chunk_tokens = x_padded[
1616
+ :, c_idx * chunk_size : (c_idx + 1) * chunk_size
1617
+ ]
1618
+ chunk_labels = labels_padded[
1619
+ :, c_idx * chunk_size : (c_idx + 1) * chunk_size
1620
+ ]
1621
+ is_last_chunk = c_idx == num_chunks - 1
1622
+
1623
+ with HardwareContext.get_autocast_context(device):
1624
+ if c_idx == 0:
1625
+ chunk_hidden, l_aux, l_z, past_kv, _ = self._forward_dense(
1626
+ chunk_tokens,
1627
+ past_c_kv_list=None,
1628
+ override_hub=hub,
1629
+ return_logits=False,
1630
+ attn_mask=attn_mask,
1631
+ )
1632
+ else:
1633
+ detached_past_kv = [kv.detach().clone() for kv in past_kv]
1634
+ chunk_hidden, l_aux, l_z, past_kv, _ = self._forward_dense(
1635
+ chunk_tokens,
1636
+ past_c_kv_list=detached_past_kv,
1637
+ override_hub=None,
1638
+ return_logits=False,
1639
+ attn_mask=attn_mask,
1640
+ )
1641
+
1642
+ chunk_lm = self._compute_loss_efficient(chunk_hidden, chunk_labels)
1643
+ chunk_total = (chunk_lm + 0.01 * l_z) / (
1644
+ num_chunks * grad_accum_steps
1645
+ )
1646
+
1647
+ retain_flag = not is_last_chunk
1648
+ if scaler is not None:
1649
+ scaler.scale(chunk_total).backward(retain_graph=retain_flag)
1650
+ else:
1651
+ chunk_total.backward(retain_graph=retain_flag)
1652
+
1653
+ total_lm_loss_val += chunk_lm.item()
1654
+ total_z_loss_val += l_z.item()
1655
+
1656
+ return XoneLMOutput(
1657
+ loss=torch.tensor(total_lm_loss_val / num_chunks, device=x.device),
1658
+ z_loss=torch.tensor(total_z_loss_val / num_chunks, device=x.device),
1659
+ past_key_values=None,
1660
+ soliton_state=None,
1661
+ )
1662
+ else:
1663
+ logits_chunks, past_kv = [], None
1664
+ total_moe_loss = torch.tensor(0.0, device=x.device, dtype=target_dtype)
1665
+ total_z_loss = torch.tensor(0.0, device=x.device, dtype=target_dtype)
1666
+ chunk_losses = []
1667
+
1668
+ for c_idx in range(num_chunks):
1669
+ if chunk_callback is not None:
1670
+ chunk_callback(c_idx + 1, num_chunks)
1671
+
1672
+ chunk_tokens = x_padded[
1673
+ :, c_idx * chunk_size : (c_idx + 1) * chunk_size
1674
+ ]
1675
+ is_last_chunk = c_idx == num_chunks - 1
1676
+
1677
+ if c_idx == 0:
1678
+ chunk_hidden_or_logits, l_aux, l_z, past_kv, _ = (
1679
+ self._forward_dense(
1680
+ chunk_tokens,
1681
+ past_c_kv_list=None,
1682
+ override_hub=hub,
1683
+ return_logits=(labels is None),
1684
+ attn_mask=attn_mask,
1685
+ )
1686
+ )
1687
+ else:
1688
+ detached_past_kv = [kv.detach().clone() for kv in past_kv]
1689
+ chunk_hidden_or_logits, l_aux, l_z, past_kv, _ = (
1690
+ self._forward_dense(
1691
+ chunk_tokens,
1692
+ past_c_kv_list=detached_past_kv,
1693
+ override_hub=None,
1694
+ return_logits=(labels is None),
1695
+ attn_mask=attn_mask,
1696
+ )
1697
+ )
1698
+
1699
+ if labels is not None:
1700
+ chunk_labels = labels_padded[
1701
+ :, c_idx * chunk_size : (c_idx + 1) * chunk_size
1702
+ ]
1703
+ chunk_lm = self._compute_loss_efficient(
1704
+ chunk_hidden_or_logits, chunk_labels
1705
+ )
1706
+ chunk_losses.append(chunk_lm)
1707
+
1708
+ if is_last_chunk or labels is None:
1709
+ logits_chunks.append(
1710
+ chunk_hidden_or_logits
1711
+ if labels is None
1712
+ else self.head(chunk_hidden_or_logits)
1713
+ )
1714
+
1715
+ total_moe_loss = total_moe_loss + l_aux
1716
+ total_z_loss = total_z_loss + l_z
1717
+
1718
+ final_loss = None
1719
+ if labels is not None:
1720
+ final_loss = torch.stack(chunk_losses).mean() + 0.01 * (
1721
+ total_z_loss / num_chunks
1722
+ )
1723
+
1724
+ return XoneLMOutput(
1725
+ loss=final_loss,
1726
+ logits=logits_chunks[-1] if logits_chunks else None,
1727
+ aux_loss=total_moe_loss / num_chunks,
1728
+ z_loss=(total_z_loss / num_chunks)
1729
+ + self.compute_hub_diversity_loss(hub),
1730
+ past_key_values=past_kv,
1731
+ soliton_state=None,
1732
+ )
1733
+ else:
1734
+ if chunk_callback is not None:
1735
+ chunk_callback(1, 1)
1736
+
1737
+ if is_sft and labels is not None:
1738
+ hidden_text, aux_l, z_l, new_kv, _ = self._forward_dense(
1739
+ x, override_hub=hub, return_logits=False, attn_mask=attn_mask
1740
+ )
1741
+ shift_hidden = hidden_text[..., :-1, :].contiguous()
1742
+ shift_labels = labels[..., 1:].contiguous()
1743
+ final_loss = (
1744
+ self._compute_loss_efficient(shift_hidden, shift_labels) + 0.01 * z_l
1745
+ )
1746
+ logits = None
1747
+ elif self.training and labels is not None:
1748
+ hidden_text, aux_l, z_l, new_kv, _ = self._forward_dense(
1749
+ x, override_hub=hub, return_logits=False, attn_mask=attn_mask
1750
+ )
1751
+ final_loss = (
1752
+ self._compute_loss_efficient(hidden_text, labels) + 0.01 * z_l
1753
+ )
1754
+ logits = None
1755
+ else:
1756
+ hidden_or_logits, aux_l, z_l, new_kv, _ = self._forward_dense(
1757
+ x,
1758
+ override_hub=hub,
1759
+ return_logits=(labels is None),
1760
+ attn_mask=attn_mask,
1761
+ )
1762
+ final_loss = None
1763
+ if labels is not None:
1764
+ if is_sft:
1765
+ shift_h = hidden_or_logits[..., :-1, :].contiguous()
1766
+ shift_l = labels[..., 1:].contiguous()
1767
+ final_loss = (
1768
+ self._compute_loss_efficient(shift_h, shift_l) + 0.01 * z_l
1769
+ )
1770
+ else:
1771
+ final_loss = (
1772
+ self._compute_loss_efficient(hidden_or_logits, labels)
1773
+ + 0.01 * z_l
1774
+ )
1775
+ logits = self.head(hidden_or_logits)
1776
+ else:
1777
+ logits = hidden_or_logits
1778
+
1779
+ return XoneLMOutput(
1780
+ loss=final_loss,
1781
+ logits=logits,
1782
+ aux_loss=aux_l,
1783
+ z_loss=z_l,
1784
+ past_key_values=new_kv,
1785
+ soliton_state=None,
1786
+ )
1787
+
1788
+ @torch.no_grad()
1789
+ def generate(
1790
+ self,
1791
+ prompt_tokens: torch.Tensor,
1792
+ max_new_tokens: int = 64,
1793
+ temperature: float = 0.7,
1794
+ top_k: int = 40,
1795
+ repetition_penalty: float = 1.15,
1796
+ eos_token_id: Optional[int] = None,
1797
+ ) -> torch.Tensor:
1798
+ self.eval()
1799
+ batch_size = prompt_tokens.shape[0]
1800
+
1801
+ hub = self.extract_hub(prompt_tokens)
1802
+ out = self.forward(prompt_tokens, override_hub=hub)
1803
+ past_kv = out.past_key_values
1804
+ generated = prompt_tokens.clone()
1805
+
1806
+ logits = out.logits[:, -1, :].clone() / max(temperature, 1e-5)
1807
+
1808
+ if repetition_penalty != 1.0:
1809
+ for i in range(batch_size):
1810
+ for prev_token in set(generated[i].tolist()):
1811
+ if logits[i, prev_token] < 0:
1812
+ logits[i, prev_token] *= repetition_penalty
1813
+ else:
1814
+ logits[i, prev_token] /= repetition_penalty
1815
+
1816
+ if top_k > 0:
1817
+ v_top, _ = torch.topk(logits, min(top_k, logits.size(-1)))
1818
+ logits[logits < v_top[:, [-1]]] = -float("Inf")
1819
+
1820
+ probs = F.softmax(logits, dim=-1)
1821
+ cur_token = torch.multinomial(probs, num_samples=1)
1822
+ generated = torch.cat([generated, cur_token], dim=1)
1823
+
1824
+ for _ in range(max_new_tokens - 1):
1825
+ if eos_token_id is not None and (cur_token == eos_token_id).all():
1826
+ break
1827
+
1828
+ step_out = self.forward(cur_token, past_key_values=past_kv)
1829
+ past_kv = step_out.past_key_values
1830
+
1831
+ logits = step_out.logits[:, -1, :].clone() / max(temperature, 1e-5)
1832
+
1833
+ if repetition_penalty != 1.0:
1834
+ for i in range(batch_size):
1835
+ for prev_token in set(generated[i].tolist()):
1836
+ if logits[i, prev_token] < 0:
1837
+ logits[i, prev_token] *= repetition_penalty
1838
+ else:
1839
+ logits[i, prev_token] /= repetition_penalty
1840
+
1841
+ if top_k > 0:
1842
+ v_top, _ = torch.topk(logits, min(top_k, logits.size(-1)))
1843
+ logits[logits < v_top[:, [-1]]] = -float("Inf")
1844
+
1845
+ probs = F.softmax(logits, dim=-1)
1846
+ cur_token = torch.multinomial(probs, num_samples=1)
1847
+ generated = torch.cat([generated, cur_token], dim=1)
1848
+
1849
+ return generated
1850
+
1851
+
1852
+ @dataclass
1853
+ class SpecialTokenConfig:
1854
+ pad_token_id: int = 0
1855
+ bos_token_id: int = 1
1856
+ eos_token_id: int = 2
1857
+ unk_token_id: int = 3
1858
+ eod_token_id: int = 4
1859
+ im_start_id: Optional[int] = None
1860
+ im_end_id: Optional[int] = None
1861
+ separator_token_id: Optional[int] = None
1862
+
1863
+
1864
+ class MultiTurnConversationFormatter:
1865
+
1866
+ def __init__(
1867
+ self,
1868
+ tokenizer: Any,
1869
+ token_config: Optional[SpecialTokenConfig] = None,
1870
+ ):
1871
+ self.tokenizer = tokenizer
1872
+ self.config = token_config or SpecialTokenConfig()
1873
+
1874
+ def _get_id(token_str: str) -> Optional[int]:
1875
+ if hasattr(tokenizer, "token_to_id"):
1876
+ return tokenizer.token_to_id(token_str)
1877
+ elif hasattr(tokenizer, "convert_tokens_to_ids"):
1878
+ res = tokenizer.convert_tokens_to_ids(token_str)
1879
+ return res if isinstance(res, int) and res >= 0 else None
1880
+ return None
1881
+
1882
+ if self.config.im_start_id is None:
1883
+ self.config.im_start_id = _get_id("<|im_start|>")
1884
+ if self.config.im_end_id is None:
1885
+ self.config.im_end_id = _get_id("<|im_end|>")
1886
+ if self.config.eod_token_id is None:
1887
+ self.config.eod_token_id = _get_id("[EOD]")
1888
+
1889
+ def format_conversation(
1890
+ self, messages: List[Dict[str, str]], max_len: Optional[int] = None
1891
+ ) -> Dict[str, List[int]]:
1892
+ input_ids = []
1893
+ labels = []
1894
+
1895
+ def _encode_text(t: str) -> List[int]:
1896
+ if hasattr(self.tokenizer, "encode"):
1897
+ res = self.tokenizer.encode(t)
1898
+ return res.ids if hasattr(res, "ids") else res
1899
+ elif callable(self.tokenizer):
1900
+ return self.tokenizer(t)["input_ids"]
1901
+ return []
1902
+
1903
+ for msg in messages:
1904
+ role = msg["role"]
1905
+ content = msg["content"].strip()
1906
+
1907
+ header_text = f"<|im_start|>{role}\n"
1908
+ body_text = f"{content}<|im_end|>\n"
1909
+
1910
+ header_ids = _encode_text(header_text)
1911
+ body_ids = _encode_text(body_text)
1912
+
1913
+ turn_input_ids = header_ids + body_ids
1914
+ input_ids.extend(turn_input_ids)
1915
+
1916
+ if role == "assistant":
1917
+ turn_labels = [-100] * len(header_ids) + body_ids
1918
+ labels.extend(turn_labels)
1919
+ else:
1920
+ labels.extend([-100] * len(turn_input_ids))
1921
+
1922
+ if self.config.eod_token_id is not None:
1923
+ input_ids.append(self.config.eod_token_id)
1924
+ labels.append(self.config.eod_token_id)
1925
+
1926
+ if max_len is not None:
1927
+ input_ids = input_ids[:max_len]
1928
+ labels = labels[:max_len]
1929
+
1930
+ return {"input_ids": input_ids, "labels": labels}
1931
+
1932
+
1933
+ class LumiSFTCollator:
1934
+
1935
+ def __init__(self, seq_len: int = 2048, pad_token_id: int = 0):
1936
+ self.seq_len = seq_len
1937
+ self.pad_token_id = pad_token_id
1938
+
1939
+ def __call__(
1940
+ self, samples: List[Dict[str, List[int]]]
1941
+ ) -> Dict[str, torch.Tensor]:
1942
+ packed_inputs = []
1943
+ packed_labels = []
1944
+
1945
+ cur_input_buf = []
1946
+ cur_label_buf = []
1947
+
1948
+ for item in samples:
1949
+ inp = item["input_ids"]
1950
+ lbl = item["labels"]
1951
+ doc_len = len(inp)
1952
+
1953
+ if doc_len > self.seq_len:
1954
+ for s in range(0, doc_len, self.seq_len):
1955
+ chunk_inp = inp[s : s + self.seq_len]
1956
+ chunk_lbl = lbl[s : s + self.seq_len]
1957
+ pad_sz = self.seq_len - len(chunk_inp)
1958
+ packed_inputs.append(
1959
+ torch.tensor(
1960
+ chunk_inp + [self.pad_token_id] * pad_sz, dtype=torch.long
1961
+ )
1962
+ )
1963
+ packed_labels.append(
1964
+ torch.tensor(chunk_lbl + [-100] * pad_sz, dtype=torch.long)
1965
+ )
1966
+ else:
1967
+ if len(cur_input_buf) + doc_len <= self.seq_len:
1968
+ cur_input_buf.extend(inp)
1969
+ cur_label_buf.extend(lbl)
1970
+ else:
1971
+ pad_sz = self.seq_len - len(cur_input_buf)
1972
+ packed_inputs.append(
1973
+ torch.tensor(
1974
+ cur_input_buf + [self.pad_token_id] * pad_sz, dtype=torch.long
1975
+ )
1976
+ )
1977
+ packed_labels.append(
1978
+ torch.tensor(cur_label_buf + [-100] * pad_sz, dtype=torch.long)
1979
+ )
1980
+ cur_input_buf = list(inp)
1981
+ cur_label_buf = list(lbl)
1982
+
1983
+ if cur_input_buf:
1984
+ pad_sz = self.seq_len - len(cur_input_buf)
1985
+ packed_inputs.append(
1986
+ torch.tensor(
1987
+ cur_input_buf + [self.pad_token_id] * pad_sz, dtype=torch.long
1988
+ )
1989
+ )
1990
+ packed_labels.append(
1991
+ torch.tensor(cur_label_buf + [-100] * pad_sz, dtype=torch.long)
1992
+ )
1993
+
1994
+ return {
1995
+ "input_ids": torch.stack(packed_inputs),
1996
+ "labels": torch.stack(packed_labels),
1997
+ }
1998
+
1999
+
2000
+ class LumiLoaderCollator:
2001
+
2002
+ def __init__(
2003
+ self,
2004
+ seq_len: int = 8192,
2005
+ eos_token_id: int = 2,
2006
+ pad_token_id: int = 0,
2007
+ ):
2008
+ self.seq_len = seq_len
2009
+ self.eos_token_id = eos_token_id
2010
+ self.pad_token_id = pad_token_id
2011
+
2012
+ def __call__(self, documents: List[List[int]]) -> Dict[str, torch.Tensor]:
2013
+ packed_batches = []
2014
+ packed_labels = []
2015
+ current_buf = []
2016
+ current_lbl = []
2017
+
2018
+ docs_sorted = sorted(documents, key=len, reverse=True)
2019
+
2020
+ for doc in docs_sorted:
2021
+ doc_with_eos = doc + [self.eos_token_id]
2022
+ doc_len = len(doc_with_eos)
2023
+
2024
+ if doc_len > self.seq_len:
2025
+ for start_idx in range(0, doc_len, self.seq_len):
2026
+ chunk = doc_with_eos[start_idx : start_idx + self.seq_len]
2027
+ pad_size = self.seq_len - len(chunk)
2028
+ chunk_input = chunk + [self.pad_token_id] * pad_size
2029
+ chunk_label = chunk + [-100] * pad_size
2030
+ packed_batches.append(torch.tensor(chunk_input, dtype=torch.long))
2031
+ packed_labels.append(torch.tensor(chunk_label, dtype=torch.long))
2032
+ else:
2033
+ if len(current_buf) + doc_len <= self.seq_len:
2034
+ current_buf.extend(doc_with_eos)
2035
+ current_lbl.extend(doc_with_eos)
2036
+ else:
2037
+ pad_size = self.seq_len - len(current_buf)
2038
+ buf_input = current_buf + [self.pad_token_id] * pad_size
2039
+ buf_label = current_lbl + [-100] * pad_size
2040
+ packed_batches.append(torch.tensor(buf_input, dtype=torch.long))
2041
+ packed_labels.append(torch.tensor(buf_label, dtype=torch.long))
2042
+
2043
+ current_buf = list(doc_with_eos)
2044
+ current_lbl = list(doc_with_eos)
2045
+
2046
+ if current_buf:
2047
+ pad_size = self.seq_len - len(current_buf)
2048
+ buf_input = current_buf + [self.pad_token_id] * pad_size
2049
+ buf_label = current_lbl + [-100] * pad_size
2050
+ packed_batches.append(torch.tensor(buf_input, dtype=torch.long))
2051
+ packed_labels.append(torch.tensor(buf_label, dtype=torch.long))
2052
+
2053
+ return {
2054
+ "input_ids": torch.stack(packed_batches),
2055
+ "labels": torch.stack(packed_labels),
2056
+ }
requirements.txt ADDED
@@ -0,0 +1,6 @@
 
 
 
 
 
 
 
1
+ torch>=2.11.0
2
+ tokenizers>=0.22.2
3
+ transformers>=5.15.1
4
+ numpy>=2.1.3
5
+ huggingface_hub>=1.28.0
6
+ triton>=3.6.0; platform_system == "Linux"
sft_example.py ADDED
@@ -0,0 +1,150 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import time
2
+ from typing import Dict, List
3
+ import torch
4
+ from torch.utils.data import DataLoader, Dataset
5
+ from modeling_xonelm import XoneLM, HardwareContext
6
+ from luminav import LuminaV
7
+ from tokenizer import (
8
+ build_xonelm_tokenizer,
9
+ MultiTurnConversationFormatter,
10
+ SpecialTokenConfig,
11
+ )
12
+
13
+ class SafeSFTCollator:
14
+ def __init__(self, max_seq_len: int = 512, pad_token_id: int = 0):
15
+ self.max_seq_len = max_seq_len
16
+ self.pad_token_id = pad_token_id
17
+
18
+ def __call__(self, samples: List[Dict[str, List[int]]]) -> Dict[str, torch.Tensor]:
19
+ batch_inputs = []
20
+ batch_labels = []
21
+
22
+ for item in samples:
23
+ inp = item["input_ids"][: self.max_seq_len]
24
+ lbl = item["labels"][: self.max_seq_len]
25
+ pad_len = self.max_seq_len - len(inp)
26
+
27
+ batch_inputs.append(
28
+ torch.tensor(inp + [self.pad_token_id] * pad_len, dtype=torch.long)
29
+ )
30
+ batch_labels.append(
31
+ torch.tensor(lbl + [-100] * pad_len, dtype=torch.long)
32
+ )
33
+
34
+ return {
35
+ "input_ids": torch.stack(batch_inputs),
36
+ "labels": torch.stack(batch_labels),
37
+ }
38
+
39
+ class ConversationDataset(Dataset):
40
+ def __init__(self, data: List[Dict[str, List[int]]]):
41
+ self.data = data
42
+
43
+ def __len__(self) -> int:
44
+ return len(self.data)
45
+
46
+ def __getitem__(self, idx: int) -> Dict[str, List[int]]:
47
+ return self.data[idx]
48
+
49
+ def run_sft_demo():
50
+ device = HardwareContext.get_optimal_device()
51
+ autocast_dtype = HardwareContext.get_optimal_autocast_dtype(device)
52
+
53
+ print("Compute Device :", device)
54
+ print("Autocast Dtype :", autocast_dtype)
55
+
56
+ tokenizer = build_xonelm_tokenizer()
57
+ vocab_size = len(tokenizer)
58
+
59
+ token_cfg = SpecialTokenConfig(
60
+ pad_token_id=0,
61
+ bos_token_id=1,
62
+ eos_token_id=2,
63
+ unk_token_id=3,
64
+ eod_token_id=4,
65
+ )
66
+ formatter = MultiTurnConversationFormatter(tokenizer, token_cfg)
67
+
68
+ sample_dialogues = [
69
+ [
70
+ {"role": "system", "content": "You are a precise reasoning assistant."},
71
+ {"role": "user", "content": "Lily found a wooden box. What did she open?"},
72
+ {"role": "assistant", "content": "She opened the wooden box to see what was inside."},
73
+ ],
74
+ [
75
+ {"role": "system", "content": "You are a polite companion."},
76
+ {"role": "user", "content": "Hello! How can we optimize memory bandwidth?"},
77
+ {"role": "assistant", "content": "We can compress Key-Value caches using low-rank latent projections."},
78
+ ],
79
+ [
80
+ {"role": "system", "content": "You are a creative writer."},
81
+ {"role": "user", "content": "Tell me a story about a kitten in the garden."},
82
+ {"role": "assistant", "content": "Once upon a time, a tiny kitten chased a butterfly across the grass."},
83
+ ],
84
+ ]
85
+
86
+ formatted_samples = [formatter.format_conversation(dialogue) for dialogue in sample_dialogues]
87
+
88
+ dataset = ConversationDataset(formatted_samples)
89
+ collator = SafeSFTCollator(max_seq_len=256, pad_token_id=token_cfg.pad_token_id)
90
+ loader = DataLoader(dataset, batch_size=2, shuffle=True, collate_fn=collator)
91
+
92
+ model = XoneLM(
93
+ vocab_size=vocab_size,
94
+ dim=512,
95
+ num_layers=12,
96
+ num_heads=8,
97
+ kv_latent_dim=64,
98
+ hub_size=512,
99
+ num_specialized_hubs=12,
100
+ num_terminals=32,
101
+ slots_per_terminal=16,
102
+ ).to(device)
103
+
104
+ optimizer = LuminaV(
105
+ model.parameters(),
106
+ lr=2e-4,
107
+ betas=(0.9, 0.999),
108
+ eps=1e-8,
109
+ weight_decay=1e-3,
110
+ tau=0.8,
111
+ buffer=2,
112
+ cautious=True,
113
+ execution="auto",
114
+ )
115
+
116
+ use_scaler = (device.type == "cuda" and autocast_dtype == torch.float16)
117
+ scaler = torch.amp.GradScaler("cuda", enabled=True) if use_scaler else None
118
+
119
+ model.train()
120
+ optimizer.zero_grad()
121
+ start_time = time.time()
122
+
123
+ for epoch in range(2):
124
+ for step, batch in enumerate(loader):
125
+ x = batch["input_ids"].to(device, non_blocking=True)
126
+ y = batch["labels"].to(device, non_blocking=True)
127
+
128
+ with HardwareContext.get_autocast_context(device):
129
+ output = model(x, labels=y, is_sft=True)
130
+ loss = output.loss
131
+
132
+ if scaler is not None:
133
+ scaler.scale(loss).backward()
134
+ scaler.unscale_(optimizer)
135
+ torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
136
+ scaler.step(optimizer)
137
+ scaler.update()
138
+ else:
139
+ loss.backward()
140
+ torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
141
+ optimizer.step()
142
+
143
+ optimizer.zero_grad()
144
+ print(f"Epoch [{epoch+1}/2] | Step [{step+1}/{len(loader)}] | SFT Loss: {loss.item():.4f}")
145
+
146
+ elapsed = time.time() - start_time
147
+ print(f"[+] SFT Training Demo completed successfully in {elapsed:.2f}s!")
148
+
149
+ if __name__ == "__main__":
150
+ run_sft_demo()
tokenize_example.py ADDED
@@ -0,0 +1,51 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from tokenizer import (
2
+ build_xonelm_tokenizer,
3
+ MultiTurnConversationFormatter,
4
+ SpecialTokenConfig,
5
+ )
6
+
7
+ def run_tokenizer_demo():
8
+ sample_corpus = [
9
+ "Once upon a time, Lily found a golden key in the garden.",
10
+ "Timmy and his dog Max played with a red ball.",
11
+ "def solve_quadratic(a, b, c): return (-b + (b**2 - 4*a*c)**0.5) / (2*a)",
12
+ "\\int_{0}^{\\infty} e^{-x^2} dx = \\frac{\\sqrt{\\pi}}{2}",
13
+ "The system latency is <= 10ms with async/await workers.",
14
+ ]
15
+
16
+ tokenizer = build_xonelm_tokenizer(corpus=sample_corpus, vocab_size=1000)
17
+ print("Tokenizer Vocab Size:", len(tokenizer))
18
+
19
+ text_to_encode = "Lily solved \\alpha + \\beta == 42 async await."
20
+ encoded = tokenizer.encode(text_to_encode)
21
+ token_ids = encoded.ids if hasattr(encoded, "ids") else encoded["input_ids"]
22
+ decoded = tokenizer.decode(token_ids)
23
+
24
+ print("Single Text Tokenization")
25
+ print("Input Text :", text_to_encode)
26
+ print("Token IDs :", token_ids)
27
+ print("Decoded :", decoded)
28
+
29
+ conversation = [
30
+ {"role": "system", "content": "You are a helpful and wise AI assistant."},
31
+ {"role": "user", "content": "Can you explain how it's work?"},
32
+ {"role": "assistant", "content": "No! I can't. hehe"},
33
+ ]
34
+
35
+ cfg = SpecialTokenConfig(
36
+ pad_token_id=0,
37
+ bos_token_id=1,
38
+ eos_token_id=2,
39
+ unk_token_id=3,
40
+ eod_token_id=4,
41
+ )
42
+ formatter = MultiTurnConversationFormatter(tokenizer, cfg)
43
+ formatted = formatter.format_conversation(conversation)
44
+
45
+ print("Multi-Turn ChatML Formatting")
46
+ print("Input IDs Length :", len(formatted["input_ids"]))
47
+ print("Labels Length :", len(formatted["labels"]))
48
+ print("Formatted Text :\n" + tokenizer.decode(formatted["input_ids"]))
49
+
50
+ if __name__ == "__main__":
51
+ run_tokenizer_demo()
tokenizer.py ADDED
@@ -0,0 +1,182 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ from dataclasses import dataclass
3
+ from typing import Any, Dict, List, Optional, Union
4
+ from tokenizers import Regex, Tokenizer
5
+ from tokenizers.decoders import ByteLevel as ByteLevelDecoder
6
+ from tokenizers.models import BPE
7
+ from tokenizers.normalizers import NFKC, Sequence as NormalizerSequence
8
+ from tokenizers.pre_tokenizers import (
9
+ ByteLevel,
10
+ Digits,
11
+ Sequence as PreTokenizerSequence,
12
+ Split,
13
+ )
14
+ from tokenizers.trainers import BpeTrainer
15
+ from transformers import PreTrainedTokenizerFast
16
+
17
+ SPECIAL_TOKENS = ["<s>", "<pad>", "</s>", "<unk>", "[EOD]", "<|eod|>"]
18
+
19
+ EMOJIS = [
20
+ "\U0001F602", "\U0001F62D", "\u2728", "\U0001F680", "\U0001F44D",
21
+ "\U0001F64F", "\U0001F525", "\U0001F60A", "\u2764\ufe0f", "\U0001F914",
22
+ "\U0001F923", "\U0001F60D", "\U0001F480", "\U0001F4AF", "\u26a0\ufe0f",
23
+ "\u2705", "\u274c", "\U0001F4CA", "\U0001F4BB", "\U0001F4F1",
24
+ "\U0001F623", "\U0001F970", "\U0001F605", "\U0001F606", "\U0001F979",
25
+ "\U0001F61A", "\U0001F917", "\U0001F61D", "\U0001F440",
26
+ ]
27
+
28
+ EMOTICONS = [
29
+ ":-)", ":)", ":D", ":(", ";)", "XD", "OwO", "UwU", "T_T", "QAQ", "¯\\_(ツ)_/¯",
30
+ ]
31
+
32
+ MATH_LATEX = [
33
+ "\\alpha", "\\beta", "\\gamma", "\\theta", "\\pi", "\\sigma", "\\omega",
34
+ "\\sum", "\\int", "\\approx", "\\neq", "\\le", "\\ge", "\\infty",
35
+ "\\partial", "\\nabla", "\\forall", "\\exists", "\\in", "\\notin",
36
+ "\\rightarrow", "\\Rightarrow", "\\Leftrightarrow",
37
+ ]
38
+
39
+ CODE_OPERATORS = [
40
+ "==", "!=", "<=", ">=", "+=", "-=", "*=", "/=",
41
+ "=>", "->", "&&", "||", "async", "await", "lambda",
42
+ ]
43
+
44
+ CUSTOM_TOKENS = [
45
+ "<|im_start|>",
46
+ "<|im_end|>",
47
+ "<|system|>",
48
+ "<|user|>",
49
+ "<|assistant|>",
50
+ "<think>",
51
+ "</think>",
52
+ ] + (EMOJIS + EMOTICONS + MATH_LATEX + CODE_OPERATORS)
53
+
54
+ ALL_TOKENS = SPECIAL_TOKENS + CUSTOM_TOKENS
55
+
56
+ LLM_SPLIT_REGEX = (
57
+ r"""(?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\r\n\p{L}\p{N}]?+\p{L}+|\p{N}|"""
58
+ r""" ?[^\s\p{L}\p{N}]+[\r\n]*|\s*[\r\n]+|\s+(?!\S)|\s+"""
59
+ )
60
+
61
+ def build_xonelm_tokenizer(
62
+ corpus: Optional[Union[str, List[str]]] = None,
63
+ vocab_size: int = 32000,
64
+ save_path: Optional[str] = None,
65
+ ) -> PreTrainedTokenizerFast:
66
+ bpe_model = BPE(unk_token="<unk>")
67
+ tokenizer_raw = Tokenizer(bpe_model)
68
+
69
+ tokenizer_raw.normalizer = NormalizerSequence([NFKC()])
70
+
71
+ tokenizer_raw.pre_tokenizer = PreTokenizerSequence([
72
+ Split(pattern=Regex(LLM_SPLIT_REGEX), behavior="isolated", invert=False),
73
+ Digits(individual_digits=True),
74
+ ByteLevel(add_prefix_space=False, use_regex=False),
75
+ ])
76
+
77
+ tokenizer_raw.decoder = ByteLevelDecoder()
78
+
79
+ trainer = BpeTrainer(
80
+ vocab_size=vocab_size,
81
+ special_tokens=ALL_TOKENS,
82
+ initial_alphabet=ByteLevel.alphabet(),
83
+ show_progress=False,
84
+ )
85
+
86
+ if corpus is not None:
87
+ if isinstance(corpus, str) and os.path.isfile(corpus):
88
+ tokenizer_raw.train([corpus], trainer)
89
+ elif isinstance(corpus, list) and len(corpus) > 0 and os.path.isfile(corpus[0]):
90
+ tokenizer_raw.train(corpus, trainer)
91
+ else:
92
+ iterator = [corpus] if isinstance(corpus, str) else corpus
93
+ tokenizer_raw.train_from_iterator(iterator, trainer)
94
+ else:
95
+ tokenizer_raw.train_from_iterator(["Hello world 123 \\alpha \\beta == async await"], trainer)
96
+
97
+ if save_path is not None:
98
+ tokenizer_raw.save(save_path)
99
+
100
+ hf_tokenizer = PreTrainedTokenizerFast(
101
+ tokenizer_object=tokenizer_raw,
102
+ bos_token="<s>",
103
+ eos_token="</s>",
104
+ pad_token="<pad>",
105
+ unk_token="<unk>",
106
+ additional_special_tokens=ALL_TOKENS,
107
+ )
108
+ return hf_tokenizer
109
+
110
+ @dataclass
111
+ class SpecialTokenConfig:
112
+ pad_token_id: int = 0
113
+ bos_token_id: int = 1
114
+ eos_token_id: int = 2
115
+ unk_token_id: int = 3
116
+ eod_token_id: int = 4
117
+ im_start_id: Optional[int] = None
118
+ im_end_id: Optional[int] = None
119
+ separator_token_id: Optional[int] = None
120
+
121
+ class MultiTurnConversationFormatter:
122
+ def __init__(self, tokenizer: Any, token_config: Optional[SpecialTokenConfig] = None):
123
+ self.tokenizer = tokenizer
124
+ self.config = token_config or SpecialTokenConfig()
125
+
126
+ def _get_id(token_str: str) -> Optional[int]:
127
+ if hasattr(tokenizer, "token_to_id"):
128
+ return tokenizer.token_to_id(token_str)
129
+ elif hasattr(tokenizer, "convert_tokens_to_ids"):
130
+ res = tokenizer.convert_tokens_to_ids(token_str)
131
+ return res if isinstance(res, int) and res >= 0 else None
132
+ return None
133
+
134
+ if self.config.im_start_id is None:
135
+ self.config.im_start_id = _get_id("<|im_start|>")
136
+ if self.config.im_end_id is None:
137
+ self.config.im_end_id = _get_id("<|im_end|>")
138
+ if self.config.eod_token_id is None:
139
+ self.config.eod_token_id = _get_id("[EOD]")
140
+
141
+ def format_conversation(
142
+ self, messages: List[Dict[str, str]], max_len: Optional[int] = None
143
+ ) -> Dict[str, List[int]]:
144
+ input_ids = []
145
+ labels = []
146
+
147
+ def _encode_text(t: str) -> List[int]:
148
+ if hasattr(self.tokenizer, "encode"):
149
+ res = self.tokenizer.encode(t)
150
+ return res.ids if hasattr(res, "ids") else res
151
+ elif callable(self.tokenizer):
152
+ return self.tokenizer(t)["input_ids"]
153
+ return []
154
+
155
+ for msg in messages:
156
+ role = msg["role"]
157
+ content = msg["content"].strip()
158
+
159
+ header_text = f"<|im_start|>{role}\n"
160
+ body_text = f"{content}<|im_end|>\n"
161
+
162
+ header_ids = _encode_text(header_text)
163
+ body_ids = _encode_text(body_text)
164
+
165
+ turn_input_ids = header_ids + body_ids
166
+ input_ids.extend(turn_input_ids)
167
+
168
+ if role == "assistant":
169
+ turn_labels = [-100] * len(header_ids) + body_ids
170
+ labels.extend(turn_labels)
171
+ else:
172
+ labels.extend([-100] * len(turn_input_ids))
173
+
174
+ if self.config.eod_token_id is not None:
175
+ input_ids.append(self.config.eod_token_id)
176
+ labels.append(self.config.eod_token_id)
177
+
178
+ if max_len is not None:
179
+ input_ids = input_ids[:max_len]
180
+ labels = labels[:max_len]
181
+
182
+ return {"input_ids": input_ids, "labels": labels}
train_example.py ADDED
@@ -0,0 +1,91 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import time
2
+ import torch
3
+ import torch.nn.functional as F
4
+ from modeling_xonelm import XoneLM, HardwareContext, create_universal_document_boundary_mask
5
+ from luminav import LuminaV
6
+ from tokenizer import build_xonelm_tokenizer
7
+
8
+ def run_train_demo():
9
+ device = HardwareContext.get_optimal_device()
10
+ autocast_dtype = HardwareContext.get_optimal_autocast_dtype(device)
11
+
12
+ print("Compute Device :", device)
13
+ print("Autocast Dtype :", autocast_dtype)
14
+
15
+ tokenizer = build_xonelm_tokenizer()
16
+ vocab_size = len(tokenizer)
17
+
18
+ model = XoneLM(
19
+ vocab_size=vocab_size,
20
+ dim=512,
21
+ num_layers=12,
22
+ num_heads=8,
23
+ kv_latent_dim=64,
24
+ hub_size=512,
25
+ num_specialized_hubs=12,
26
+ num_terminals=32,
27
+ slots_per_terminal=16,
28
+ chunk_size=1024,
29
+ ).to(device)
30
+
31
+ total_params = sum(p.numel() for p in model.parameters())
32
+ print(f"Total Parameters: {total_params / 1e6:.2f}M")
33
+
34
+ optimizer = LuminaV(
35
+ model.parameters(),
36
+ lr=8e-4,
37
+ betas=(0.9, 0.999),
38
+ eps=1e-8,
39
+ weight_decay=8e-2,
40
+ tau=0.8,
41
+ buffer=2,
42
+ cautious=True,
43
+ execution="auto",
44
+ )
45
+
46
+ use_scaler = (device.type == "cuda" and autocast_dtype == torch.float16)
47
+ scaler = torch.amp.GradScaler("cuda", enabled=True) if use_scaler else None
48
+
49
+ batch_size = 2
50
+ seq_len = 512
51
+ num_steps = 5
52
+
53
+ model.train()
54
+ optimizer.zero_grad()
55
+ start_time = time.time()
56
+
57
+ for step in range(num_steps):
58
+ x = torch.randint(0, vocab_size, (batch_size, seq_len), device=device)
59
+ y = torch.randint(0, vocab_size, (batch_size, seq_len), device=device)
60
+
61
+ doc_mask = create_universal_document_boundary_mask(
62
+ x_tokens=x,
63
+ hub_size=model.hub_size,
64
+ past_k_len=model.hub_size,
65
+ eod_token_id=4,
66
+ is_dense_with_hub=True,
67
+ )
68
+
69
+ with HardwareContext.get_autocast_context(device):
70
+ output = model(x, labels=y, attn_mask=doc_mask)
71
+ loss = output.loss
72
+
73
+ if scaler is not None:
74
+ scaler.scale(loss).backward()
75
+ scaler.unscale_(optimizer)
76
+ torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
77
+ scaler.step(optimizer)
78
+ scaler.update()
79
+ else:
80
+ loss.backward()
81
+ torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
82
+ optimizer.step()
83
+
84
+ optimizer.zero_grad()
85
+ print(f"Step [{step+1}/{num_steps}] | Loss: {loss.item():.4f} | Z-Loss: {output.z_loss.item():.4f}")
86
+
87
+ elapsed = time.time() - start_time
88
+ print(f"Demo training completed in {elapsed:.2f}s!")
89
+
90
+ if __name__ == "__main__":
91
+ run_train_demo()