shikhar007 commited on
Commit
105f101
·
verified ·
1 Parent(s): 3c8ab1e

Add train_gpt_gram_ns.py

Browse files
Files changed (1) hide show
  1. train_gpt_gram_ns.py +2006 -0
train_gpt_gram_ns.py ADDED
@@ -0,0 +1,2006 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+ import copy
3
+ import glob
4
+ import io
5
+ import lzma
6
+ import math
7
+ import os
8
+ import random
9
+ import subprocess
10
+ import sys
11
+ import time
12
+ import uuid
13
+ import zlib
14
+ from pathlib import Path
15
+ try:
16
+ import zstandard
17
+ _COMPRESSOR = "zstd"
18
+ except ImportError:
19
+ _COMPRESSOR = "zlib"
20
+ import numpy as np
21
+ import sentencepiece as spm
22
+ import torch
23
+ import torch.distributed as dist
24
+ import torch.nn.functional as F
25
+ from torch import Tensor, nn
26
+ from torch.nn.parallel import DistributedDataParallel as DDP
27
+ # FlashAttention fallback chain: FA3 (H100) -> FA2 (Ampere+) -> PyTorch SDPA
28
+ try:
29
+ from flash_attn_interface import flash_attn_func as flash_attn_3_func
30
+ _ATTN_BACKEND = "fa3"
31
+ except ImportError:
32
+ try:
33
+ from flash_attn.flash_attn_interface import flash_attn_func as flash_attn_3_func
34
+ _ATTN_BACKEND = "fa3"
35
+ except ImportError:
36
+ try:
37
+ from flash_attn import flash_attn_func as flash_attn_3_func
38
+ _ATTN_BACKEND = "fa2"
39
+ except ImportError:
40
+ def flash_attn_3_func(q, k, v, causal=True):
41
+ # q,k,v: (B, T, H, D) -> transpose to (B, H, T, D) for SDPA
42
+ q_t = q.transpose(1, 2)
43
+ k_t = k.transpose(1, 2)
44
+ v_t = v.transpose(1, 2)
45
+ y = F.scaled_dot_product_attention(
46
+ q_t, k_t, v_t, attn_mask=None, is_causal=causal,
47
+ enable_gqa=(q.size(2) != k.size(2)),
48
+ )
49
+ return y.transpose(1, 2) # back to (B, T, H, D)
50
+ _ATTN_BACKEND = "sdpa"
51
+ class Hyperparameters:
52
+ data_path = os.environ.get("DATA_PATH", "./data/datasets/fineweb10B_sp1024")
53
+ train_files = os.path.join(data_path, "fineweb_train_*.bin")
54
+ val_files = os.path.join(data_path, "fineweb_val_*.bin")
55
+ tokenizer_path = os.environ.get("TOKENIZER_PATH", "./data/tokenizers/fineweb_1024_bpe.model")
56
+ run_id = os.environ.get("RUN_ID", str(uuid.uuid4()))
57
+ seed = int(os.environ.get("SEED", 1337))
58
+ val_batch_size = int(os.environ.get("VAL_BATCH_SIZE", 524_288))
59
+ val_loss_every = int(os.environ.get("VAL_LOSS_EVERY", 4000))
60
+ train_log_every = int(os.environ.get("TRAIN_LOG_EVERY", 500))
61
+ iterations = int(os.environ.get("ITERATIONS", 20000))
62
+ warmdown_iters = int(os.environ.get("WARMDOWN_ITERS", 3500))
63
+ warmup_steps = int(os.environ.get("WARMUP_STEPS", 20))
64
+ train_batch_tokens = int(os.environ.get("TRAIN_BATCH_TOKENS", 786_432))
65
+ train_seq_len = int(os.environ.get("TRAIN_SEQ_LEN", 2048))
66
+ eval_seq_len = int(os.environ.get("EVAL_SEQ_LEN", 2048))
67
+ max_wallclock_seconds = float(os.environ.get("MAX_WALLCLOCK_SECONDS", 600.0))
68
+ qk_gain_init = float(os.environ.get("QK_GAIN_INIT", 1.5))
69
+ vocab_size = int(os.environ.get("VOCAB_SIZE", 1024))
70
+ num_layers = int(os.environ.get("NUM_LAYERS", 11))
71
+ num_kv_heads = int(os.environ.get("NUM_KV_HEADS", 4))
72
+ model_dim = int(os.environ.get("MODEL_DIM", 512))
73
+ num_heads = int(os.environ.get("NUM_HEADS", 8))
74
+ mlp_mult = float(os.environ.get("MLP_MULT", 3.0))
75
+ tie_embeddings = bool(int(os.environ.get("TIE_EMBEDDINGS", "1")))
76
+ rope_base = float(os.environ.get("ROPE_BASE", 10000.0))
77
+ logit_softcap = float(os.environ.get("LOGIT_SOFTCAP", 30.0))
78
+ embed_lr = float(os.environ.get("EMBED_LR", 0.6))
79
+ head_lr = float(os.environ.get("HEAD_LR", 0.008))
80
+ tied_embed_lr = float(os.environ.get("TIED_EMBED_LR", 0.035))
81
+ tied_embed_init_std = float(os.environ.get("TIED_EMBED_INIT_STD", 0.005))
82
+ matrix_lr = float(os.environ.get("MATRIX_LR", 0.025))
83
+ scalar_lr = float(os.environ.get("SCALAR_LR", 0.025))
84
+ muon_momentum = float(os.environ.get("MUON_MOMENTUM", 0.99))
85
+ muon_backend_steps = int(os.environ.get("MUON_BACKEND_STEPS", 5))
86
+ muon_momentum_warmup_start = float(os.environ.get("MUON_MOMENTUM_WARMUP_START", 0.92))
87
+ muon_momentum_warmup_steps = int(os.environ.get("MUON_MOMENTUM_WARMUP_STEPS", 1500))
88
+ beta1 = float(os.environ.get("BETA1", 0.9))
89
+ beta2 = float(os.environ.get("BETA2", 0.95))
90
+ adam_eps = float(os.environ.get("ADAM_EPS", 1e-8))
91
+ grad_clip_norm = float(os.environ.get("GRAD_CLIP_NORM", 0.3))
92
+ eval_stride = int(os.environ.get("EVAL_STRIDE", 64))
93
+ mtp_num_heads = int(os.environ.get("MTP_NUM_HEADS", 0))
94
+ mtp_loss_weight = float(os.environ.get("MTP_LOSS_WEIGHT", 0.2))
95
+ muon_beta2 = float(os.environ.get("MUON_BETA2", 0.95))
96
+ swa_enabled = bool(int(os.environ.get("SWA_ENABLED", "1")))
97
+ swa_every = int(os.environ.get("SWA_EVERY", 50))
98
+ lawa_enabled = bool(int(os.environ.get("LAWA_ENABLED", "0")))
99
+ lawa_k = int(os.environ.get("LAWA_K", 10))
100
+ lawa_freq = int(os.environ.get("LAWA_FREQ", 100))
101
+ muon_wd = float(os.environ.get("MUON_WD", 0.04))
102
+ adam_wd = float(os.environ.get("ADAM_WD", 0.04))
103
+ qat_enabled = bool(int(os.environ.get("QAT_ENABLED", "0")))
104
+ bigram_vocab_size = int(os.environ.get("BIGRAM_VOCAB_SIZE", 2048))
105
+ bigram_dim = int(os.environ.get("BIGRAM_DIM", 128))
106
+ xsa_last_n = int(os.environ.get("XSA_LAST_N", 4))
107
+ rope_dims = int(os.environ.get("ROPE_DIMS", 16))
108
+ ln_scale = bool(int(os.environ.get("LN_SCALE", "1")))
109
+ dtg_enabled = bool(int(os.environ.get("DTG_ENABLED", "0")))
110
+ late_qat_threshold = float(os.environ.get("LATE_QAT_THRESHOLD", 0.15))
111
+ ve_enabled = bool(int(os.environ.get("VE_ENABLED", "1")))
112
+ ve_dim = int(os.environ.get("VE_DIM", 128))
113
+ ve_layers = os.environ.get("VE_LAYERS", "9,10")
114
+ gated_attention = bool(int(os.environ.get("GATED_ATTENTION", "0")))
115
+ value_residual = bool(int(os.environ.get("VALUE_RESIDUAL", "0")))
116
+ ttt_enabled = bool(int(os.environ.get("TTT_ENABLED", "0")))
117
+ ttt_lr = float(os.environ.get("TTT_LR", 0.002))
118
+ ttt_epochs = int(os.environ.get("TTT_EPOCHS", 3))
119
+ ttt_chunk_tokens = int(os.environ.get("TTT_CHUNK_TOKENS", 32768))
120
+ ttt_freeze_blocks = int(os.environ.get("TTT_FREEZE_BLOCKS", 2))
121
+ ttt_momentum = float(os.environ.get("TTT_MOMENTUM", 0.9))
122
+ ttt_batch_seqs = int(os.environ.get("TTT_BATCH_SEQS", 32))
123
+ ttt_grad_clip = float(os.environ.get("TTT_GRAD_CLIP", 1.0))
124
+
125
+ # --- Gram Newton-Schulz orthogonalization ---
126
+ # Reformulates Newton-Schulz to iterate on the smaller n×n Gram matrix R = X @ X^T
127
+ # instead of the full n×m matrix X. All inner-loop matmuls are n×n (symmetric),
128
+ # and the expensive n×m matmul only happens at restarts and the final step.
129
+ # Reference: https://github.com/Dao-AILab/gram-newton-schulz
130
+ #
131
+ # Per-step coefficients from Polar Express (arxiv 2505.16932) with 1.05x safety factor,
132
+ # matching the Dao-AILab reference implementation.
133
+
134
+ _POLAR_EXPRESS_SAFETY = 1.05
135
+ _POLAR_EXPRESS_RAW = [
136
+ (8.28721201814563, -23.595886519098837, 17.300387312530933),
137
+ (4.107059111542203, -2.9478499167379106, 0.5448431082926601),
138
+ (3.9486908534822946, -2.908902115962949, 0.5518191394370137),
139
+ (3.3184196573706015, -2.488488024314874, 0.51004894012372),
140
+ (2.300652019954817, -1.6689039845747493, 0.4188073119525673),
141
+ ]
142
+ NS_COEFFICIENTS = [
143
+ (a / _POLAR_EXPRESS_SAFETY,
144
+ b / _POLAR_EXPRESS_SAFETY ** 3,
145
+ c / _POLAR_EXPRESS_SAFETY ** 5)
146
+ for a, b, c in _POLAR_EXPRESS_RAW
147
+ ]
148
+
149
+
150
+ def _standard_newtonschulz(X: Tensor, coefficients: list[tuple[float, float, float]]) -> Tensor:
151
+ """Standard NS iteration on the full matrix. Used for square matrices where
152
+ the Gram reformulation offers no FLOP savings (n == m)."""
153
+ for a, b, c in coefficients:
154
+ A = X @ X.mT
155
+ B = torch.baddbmm(A, A, A, beta=b, alpha=c) # b*A + c*A@A
156
+ X = torch.baddbmm(X, B, X, beta=a) # a*X + B@X
157
+ return X
158
+
159
+
160
+ def _gram_newtonschulz(X: Tensor, coefficients: list[tuple[float, float, float]],
161
+ restart_at: frozenset[int]) -> Tensor:
162
+ """Gram NS iteration on the smaller n×n Gram matrix R = X @ X^T.
163
+ Only touches the full n×m matrix X at restarts and the final step.
164
+ Used for rectangular matrices where n < m."""
165
+ n = X.size(-2)
166
+ batch = X.size(0)
167
+ num_steps = len(coefficients)
168
+
169
+ R = X @ X.mT # (B, n, n) Gram matrix
170
+ I = torch.eye(n, device=X.device, dtype=X.dtype).unsqueeze(0).expand(batch, -1, -1).contiguous()
171
+ Q = None
172
+
173
+ for i, (a, b, c) in enumerate(coefficients):
174
+ # Restart: fold Q into X, recompute R, reset Q
175
+ if i in restart_at and i != 0:
176
+ X = Q @ X
177
+ R = X @ X.mT
178
+ Q = None
179
+
180
+ # Z = b*R + c*R^2
181
+ Z = torch.baddbmm(R, R, R, beta=b, alpha=c)
182
+
183
+ # Q update: first iteration (or after restart) initializes from identity
184
+ if Q is None:
185
+ Q = Z + a * I # = aI + bR + cR^2
186
+ else:
187
+ Q = torch.baddbmm(Q, Q, Z, beta=a) # a*Q + Q@Z
188
+
189
+ # R update: skip on last iteration and before restart iterations
190
+ # (R won't be used again, or will be recomputed from scratch)
191
+ if i < num_steps - 1 and (i + 1) not in restart_at:
192
+ RZ = torch.baddbmm(R, R, Z, beta=a) # a*R + R@Z
193
+ R = torch.baddbmm(RZ, Z, RZ, beta=a) # a*RZ + Z@RZ
194
+
195
+ X = Q @ X # final: apply accumulated orthogonal factor
196
+ return X
197
+
198
+
199
+ def zeropower_via_newtonschulz5(G: Tensor, steps: int = 5, eps: float = 1e-7) -> Tensor:
200
+ """Batched Newton-Schulz orthogonalization. G: (B,M,N) or (M,N).
201
+
202
+ Uses the Gram reformulation for rectangular matrices (n < m) and
203
+ standard NS for square matrices (n == m). Per-step coefficients
204
+ from Polar Express with safety factor.
205
+ """
206
+ was_2d = G.ndim == 2
207
+ if was_2d:
208
+ G = G.unsqueeze(0)
209
+
210
+ X = G.bfloat16()
211
+ transposed = X.size(-2) > X.size(-1)
212
+ if transposed:
213
+ X = X.mT
214
+ # X is now (B, n, m) with n <= m
215
+
216
+ X = X / (X.norm(dim=(-2, -1), keepdim=True) + eps)
217
+
218
+ coefficients = NS_COEFFICIENTS[:steps]
219
+
220
+ if X.size(-2) == X.size(-1):
221
+ # Square: Gram reformulation has no FLOP savings, use standard NS
222
+ X = _standard_newtonschulz(X, coefficients)
223
+ else:
224
+ # Rectangular: Gram NS iterates on the smaller n×n Gram matrix
225
+ X = _gram_newtonschulz(X, coefficients, restart_at=frozenset({2}))
226
+
227
+ if transposed:
228
+ X = X.mT.contiguous()
229
+ if was_2d:
230
+ X = X.squeeze(0)
231
+ return X
232
+
233
+ # --- Parallel Muon optimizer ---
234
+
235
+ class Muon(torch.optim.Optimizer):
236
+ """Parallel Muon: post-backward reduce-scatter -> local NS5 -> all-gather.
237
+
238
+ No DDP for bank params. After backward, this optimizer:
239
+ 1. Launches async reduce-scatter for all banks (biggest first)
240
+ 2. Returns control so Adam can step on small params while RS is in-flight
241
+ 3. Waits for each RS, runs local NS5 on the shard, launches async all-gather
242
+ 4. Each all-gather overlaps with next bank's NS5
243
+ """
244
+ def __init__(self, params, lr: float, momentum: float, backend_steps: int,
245
+ nesterov: bool = True, weight_decay: float = 0.0):
246
+ super().__init__(
247
+ params,
248
+ dict(lr=lr, momentum=momentum, backend_steps=backend_steps,
249
+ nesterov=nesterov, weight_decay=weight_decay),
250
+ )
251
+ self._built = False
252
+
253
+ def _build(self):
254
+ self._distributed = dist.is_available() and dist.is_initialized()
255
+ self._world_size = dist.get_world_size() if self._distributed else 1
256
+ self._rank = dist.get_rank() if self._distributed else 0
257
+ ws = self._world_size
258
+
259
+ self._bank_meta = []
260
+ for group in self.param_groups:
261
+ for p in group["params"]:
262
+ B = p.shape[0]
263
+ padded_B = ((B + ws - 1) // ws) * ws
264
+ shard_B = padded_B // ws
265
+ tail = p.shape[1:]
266
+ dev = p.device
267
+ self._bank_meta.append({
268
+ 'p': p,
269
+ 'B': B,
270
+ 'padded_grad': torch.zeros(padded_B, *tail, device=dev, dtype=torch.bfloat16),
271
+ 'shard': torch.zeros(shard_B, *tail, device=dev, dtype=torch.bfloat16),
272
+ 'shard_mom': torch.zeros(shard_B, *tail, device=dev, dtype=torch.bfloat16),
273
+ 'full_update': torch.zeros(padded_B, *tail, device=dev, dtype=torch.bfloat16),
274
+ 'scale': max(1, p.shape[-2] / p.shape[-1]) ** 0.5,
275
+ })
276
+ # Sort by size descending -- launch biggest reduce-scatters first
277
+ self._bank_meta.sort(key=lambda m: -m['p'].numel())
278
+ self._built = True
279
+
280
+ def launch_reduce_scatters(self):
281
+ """Phase 1: launch async reduce-scatter for all banks. Call right after backward."""
282
+ if not self._built:
283
+ self._build()
284
+ if not self._distributed:
285
+ return
286
+ self._rs_futures = []
287
+ for m in self._bank_meta:
288
+ p = m['p']
289
+ if p.grad is None:
290
+ self._rs_futures.append(None)
291
+ continue
292
+ pg = m['padded_grad']
293
+ pg[:m['B']].copy_(p.grad.bfloat16())
294
+ if pg.shape[0] > m['B']:
295
+ pg[m['B']:].zero_()
296
+ fut = dist.reduce_scatter_tensor(m['shard'], pg, op=dist.ReduceOp.AVG, async_op=True)
297
+ self._rs_futures.append(fut)
298
+
299
+ @torch.no_grad()
300
+ def step(self, closure=None):
301
+ """Phase 3: wait for RS, local NS5, all-gather. Call AFTER Adam steps."""
302
+ loss = None
303
+ if closure is not None:
304
+ with torch.enable_grad():
305
+ loss = closure()
306
+
307
+ if not self._built:
308
+ self._build()
309
+
310
+ for group in self.param_groups:
311
+ lr = group["lr"]
312
+ momentum = group["momentum"]
313
+ backend_steps = group["backend_steps"]
314
+ nesterov = group["nesterov"]
315
+ wd = group.get("weight_decay", 0.0)
316
+
317
+ prev_ag_handle = None
318
+ prev_m = None
319
+
320
+ sharded = self._distributed and hasattr(self, '_rs_futures')
321
+
322
+ for i, m in enumerate(self._bank_meta):
323
+ p = m['p']
324
+ if p.grad is None:
325
+ continue
326
+
327
+ if prev_ag_handle is not None:
328
+ prev_ag_handle.wait()
329
+ pp = prev_m['p']
330
+ upd = prev_m['full_update'][:prev_m['B']]
331
+ if wd > 0.0:
332
+ pp.data.mul_(1.0 - lr * wd)
333
+ pp.add_(upd.to(dtype=pp.dtype), alpha=-lr * prev_m['scale'])
334
+
335
+ if sharded and self._rs_futures[i] is not None:
336
+ self._rs_futures[i].wait()
337
+ g = m['shard']
338
+ buf = m['shard_mom']
339
+ else:
340
+ g = p.grad.bfloat16()
341
+ state = self.state[p]
342
+ if "momentum_buffer" not in state:
343
+ state["momentum_buffer"] = torch.zeros_like(g)
344
+ buf = state["momentum_buffer"]
345
+
346
+ buf.mul_(momentum).add_(g)
347
+ if nesterov:
348
+ update = g.add(buf, alpha=momentum)
349
+ else:
350
+ update = buf
351
+
352
+ update = zeropower_via_newtonschulz5(update, steps=backend_steps)
353
+
354
+ if sharded:
355
+ prev_ag_handle = dist.all_gather_into_tensor(
356
+ m['full_update'], update, async_op=True)
357
+ prev_m = m
358
+ else:
359
+ if wd > 0.0:
360
+ p.data.mul_(1.0 - lr * wd)
361
+ p.add_(update.to(dtype=p.dtype), alpha=-lr * m['scale'])
362
+
363
+ if prev_ag_handle is not None:
364
+ prev_ag_handle.wait()
365
+ pp = prev_m['p']
366
+ upd = prev_m['full_update'][:prev_m['B']]
367
+ if wd > 0.0:
368
+ pp.data.mul_(1.0 - lr * wd)
369
+ pp.add_(upd.to(dtype=pp.dtype), alpha=-lr * prev_m['scale'])
370
+
371
+ if hasattr(self, '_rs_futures'):
372
+ del self._rs_futures
373
+
374
+ return loss
375
+
376
+ # --- Tokenizer evaluation helpers ---
377
+
378
+ def build_sentencepiece_luts(
379
+ sp: spm.SentencePieceProcessor, vocab_size: int, device: torch.device
380
+ ) -> tuple[Tensor, Tensor, Tensor]:
381
+ sp_vocab_size = int(sp.vocab_size())
382
+ table_size = max(sp_vocab_size, vocab_size)
383
+ base_bytes_np = np.zeros((table_size,), dtype=np.int16)
384
+ has_leading_space_np = np.zeros((table_size,), dtype=np.bool_)
385
+ is_boundary_token_np = np.ones((table_size,), dtype=np.bool_)
386
+ for token_id in range(sp_vocab_size):
387
+ if sp.is_control(token_id) or sp.is_unknown(token_id) or sp.is_unused(token_id):
388
+ continue
389
+ is_boundary_token_np[token_id] = False
390
+ if sp.is_byte(token_id):
391
+ base_bytes_np[token_id] = 1
392
+ continue
393
+ piece = sp.id_to_piece(token_id)
394
+ if piece.startswith("\u2581"):
395
+ has_leading_space_np[token_id] = True
396
+ piece = piece[1:]
397
+ base_bytes_np[token_id] = len(piece.encode("utf-8"))
398
+ return (
399
+ torch.tensor(base_bytes_np, dtype=torch.int16, device=device),
400
+ torch.tensor(has_leading_space_np, dtype=torch.bool, device=device),
401
+ torch.tensor(is_boundary_token_np, dtype=torch.bool, device=device),
402
+ )
403
+ def load_validation_tokens(pattern: str, seq_len: int) -> Tensor:
404
+ files = [Path(p) for p in sorted(glob.glob(pattern))]
405
+ if not files:
406
+ raise FileNotFoundError(f"No files found for pattern: {pattern}")
407
+ tokens = torch.cat([load_data_shard(file) for file in files]).contiguous()
408
+ usable = ((tokens.numel() - 1) // seq_len) * seq_len
409
+ if usable <= 0:
410
+ raise ValueError(f"Validation split is too short for TRAIN_SEQ_LEN={seq_len}")
411
+ return tokens[: usable + 1]
412
+ def eval_val(
413
+ args: Hyperparameters,
414
+ model: nn.Module,
415
+ rank: int,
416
+ world_size: int,
417
+ device: torch.device,
418
+ grad_accum_steps: int,
419
+ val_tokens: Tensor,
420
+ base_bytes_lut: Tensor,
421
+ has_leading_space_lut: Tensor,
422
+ is_boundary_token_lut: Tensor,
423
+ eval_seq_len: int | None = None,
424
+ ) -> tuple[float, float]:
425
+ seq_len = eval_seq_len or args.train_seq_len
426
+ local_batch_tokens = args.val_batch_size // (world_size * grad_accum_steps)
427
+ if local_batch_tokens < seq_len:
428
+ raise ValueError(
429
+ "VAL_BATCH_SIZE must provide at least one sequence per rank; "
430
+ f"got VAL_BATCH_SIZE={args.val_batch_size}, WORLD_SIZE={world_size}, "
431
+ f"GRAD_ACCUM_STEPS={grad_accum_steps}, seq_len={seq_len}"
432
+ )
433
+ local_batch_seqs = local_batch_tokens // seq_len
434
+ total_seqs = (val_tokens.numel() - 1) // seq_len
435
+ seq_start = (total_seqs * rank) // world_size
436
+ seq_end = (total_seqs * (rank + 1)) // world_size
437
+ val_loss_sum = torch.zeros((), device=device, dtype=torch.float64)
438
+ val_token_count = torch.zeros((), device=device, dtype=torch.float64)
439
+ val_byte_count = torch.zeros((), device=device, dtype=torch.float64)
440
+ model.eval()
441
+ with torch.inference_mode():
442
+ for batch_seq_start in range(seq_start, seq_end, local_batch_seqs):
443
+ batch_seq_end = min(batch_seq_start + local_batch_seqs, seq_end)
444
+ raw_start = batch_seq_start * seq_len
445
+ raw_end = batch_seq_end * seq_len + 1
446
+ local = val_tokens[raw_start:raw_end].to(device=device, dtype=torch.int64, non_blocking=True)
447
+ x = local[:-1].reshape(-1, seq_len)
448
+ y = local[1:].reshape(-1, seq_len)
449
+ with torch.autocast(device_type="cuda", dtype=torch.bfloat16, enabled=True):
450
+ batch_loss = model(x, y).detach()
451
+ batch_token_count = float(y.numel())
452
+ val_loss_sum += batch_loss.to(torch.float64) * batch_token_count
453
+ val_token_count += batch_token_count
454
+ prev_ids = x.reshape(-1)
455
+ tgt_ids = y.reshape(-1)
456
+ token_bytes = base_bytes_lut[tgt_ids].to(dtype=torch.int16)
457
+ token_bytes += (has_leading_space_lut[tgt_ids] & ~is_boundary_token_lut[prev_ids]).to(dtype=torch.int16)
458
+ val_byte_count += token_bytes.to(torch.float64).sum()
459
+ if dist.is_available() and dist.is_initialized():
460
+ dist.all_reduce(val_loss_sum, op=dist.ReduceOp.SUM)
461
+ dist.all_reduce(val_token_count, op=dist.ReduceOp.SUM)
462
+ dist.all_reduce(val_byte_count, op=dist.ReduceOp.SUM)
463
+ val_loss = val_loss_sum / val_token_count
464
+ bits_per_token = val_loss.item() / math.log(2.0)
465
+ tokens_per_byte = val_token_count.item() / val_byte_count.item()
466
+ model.train()
467
+ return float(val_loss.item()), float(bits_per_token * tokens_per_byte)
468
+
469
+ # --- Quantization helpers ---
470
+
471
+ CONTROL_TENSOR_NAME_PATTERNS = tuple(
472
+ pattern
473
+ for pattern in os.environ.get(
474
+ "CONTROL_TENSOR_NAME_PATTERNS",
475
+ "attn_scale,attn_scales,mlp_scale,mlp_scales,resid_mix,resid_mixes,q_gain,skip_weight,skip_weights,smear,dtg_gate,ve_layer_scales,ve_shared.scale,attn_gate,vr_lambda",
476
+ ).split(",")
477
+ if pattern
478
+ )
479
+ INT8_KEEP_FLOAT_FP32_NAME_PATTERNS = tuple(
480
+ pattern
481
+ for pattern in os.environ.get(
482
+ "INT8_KEEP_FLOAT_FP32_NAME_PATTERNS",
483
+ ",".join(CONTROL_TENSOR_NAME_PATTERNS),
484
+ ).split(",")
485
+ if pattern
486
+ )
487
+ INT8_KEEP_FLOAT_MAX_NUMEL = 65_536
488
+ INT8_KEEP_FLOAT_STORE_DTYPE = torch.float16
489
+ INT8_PER_ROW_SCALE_DTYPE = torch.float16
490
+ INT8_CLIP_PERCENTILE = 99.99984
491
+ INT8_CLIP_Q = INT8_CLIP_PERCENTILE / 100.0
492
+ def tensor_nbytes(t: Tensor) -> int:
493
+ return int(t.numel()) * int(t.element_size())
494
+ def keep_float_tensor(name: str, t: Tensor, passthrough_orig_dtypes: dict[str, str]) -> Tensor:
495
+ if any(pattern in name for pattern in INT8_KEEP_FLOAT_FP32_NAME_PATTERNS):
496
+ return t.float().contiguous()
497
+ if t.dtype in {torch.float32, torch.bfloat16}:
498
+ passthrough_orig_dtypes[name] = str(t.dtype).removeprefix("torch.")
499
+ return t.to(dtype=INT8_KEEP_FLOAT_STORE_DTYPE).contiguous()
500
+ return t
501
+ def quantize_float_tensor(t: Tensor) -> tuple[Tensor, Tensor]:
502
+ t32 = t.float()
503
+ if t32.ndim == 2:
504
+ clip_abs = (
505
+ torch.quantile(t32.abs(), INT8_CLIP_Q, dim=1)
506
+ if t32.numel()
507
+ else torch.empty((t32.shape[0],), dtype=torch.float32)
508
+ )
509
+ clipped = torch.maximum(torch.minimum(t32, clip_abs[:, None]), -clip_abs[:, None])
510
+ scale = (clip_abs / 127.0).clamp_min(1.0 / 127.0)
511
+ q = torch.clamp(torch.round(clipped / scale[:, None]), -127, 127).to(torch.int8).contiguous()
512
+ return q, scale.to(dtype=INT8_PER_ROW_SCALE_DTYPE).contiguous()
513
+ clip_abs = float(torch.quantile(t32.abs().flatten(), INT8_CLIP_Q).item()) if t32.numel() else 0.0
514
+ scale = torch.tensor(clip_abs / 127.0 if clip_abs > 0 else 1.0, dtype=torch.float32)
515
+ q = torch.clamp(torch.round(torch.clamp(t32, -clip_abs, clip_abs) / scale), -127, 127).to(torch.int8).contiguous()
516
+ return q, scale
517
+ def quantize_state_dict_int8(state_dict: dict[str, Tensor]):
518
+ quantized: dict[str, Tensor] = {}
519
+ scales: dict[str, Tensor] = {}
520
+ dtypes: dict[str, str] = {}
521
+ passthrough: dict[str, Tensor] = {}
522
+ passthrough_orig_dtypes: dict[str, str] = {}
523
+ qmeta: dict[str, dict[str, object]] = {}
524
+ stats = dict.fromkeys(
525
+ ("param_count", "num_tensors", "num_float_tensors", "num_nonfloat_tensors", "baseline_tensor_bytes", "int8_payload_bytes"),
526
+ 0,
527
+ )
528
+ for name, tensor in state_dict.items():
529
+ t = tensor.detach().to("cpu").contiguous()
530
+ stats["param_count"] += int(t.numel())
531
+ stats["num_tensors"] += 1
532
+ stats["baseline_tensor_bytes"] += tensor_nbytes(t)
533
+ if not t.is_floating_point():
534
+ stats["num_nonfloat_tensors"] += 1
535
+ passthrough[name] = t
536
+ stats["int8_payload_bytes"] += tensor_nbytes(t)
537
+ continue
538
+ if t.numel() <= INT8_KEEP_FLOAT_MAX_NUMEL:
539
+ kept = keep_float_tensor(name, t, passthrough_orig_dtypes)
540
+ passthrough[name] = kept
541
+ stats["int8_payload_bytes"] += tensor_nbytes(kept)
542
+ continue
543
+ stats["num_float_tensors"] += 1
544
+ q, s = quantize_float_tensor(t)
545
+ if s.ndim > 0:
546
+ qmeta[name] = {"scheme": "per_row", "axis": 0}
547
+ quantized[name] = q
548
+ scales[name] = s
549
+ dtypes[name] = str(t.dtype).removeprefix("torch.")
550
+ stats["int8_payload_bytes"] += tensor_nbytes(q) + tensor_nbytes(s)
551
+ obj: dict[str, object] = {
552
+ "__quant_format__": "int8_clean_per_row_v1",
553
+ "quantized": quantized,
554
+ "scales": scales,
555
+ "dtypes": dtypes,
556
+ "passthrough": passthrough,
557
+ }
558
+ if qmeta:
559
+ obj["qmeta"] = qmeta
560
+ if passthrough_orig_dtypes:
561
+ obj["passthrough_orig_dtypes"] = passthrough_orig_dtypes
562
+ return obj, stats
563
+ def dequantize_state_dict_int8(obj: dict[str, object]) -> dict[str, Tensor]:
564
+ out: dict[str, Tensor] = {}
565
+ qmeta = obj.get("qmeta", {})
566
+ passthrough_orig_dtypes = obj.get("passthrough_orig_dtypes", {})
567
+ for name, q in obj["quantized"].items():
568
+ dtype = getattr(torch, obj["dtypes"][name])
569
+ s = obj["scales"][name]
570
+ if qmeta.get(name, {}).get("scheme") == "per_row" or s.ndim > 0:
571
+ s = s.to(dtype=torch.float32)
572
+ out[name] = (q.float() * s.view(q.shape[0], *([1] * (q.ndim - 1)))).to(dtype=dtype).contiguous()
573
+ else:
574
+ scale = float(s.item())
575
+ out[name] = (q.float() * scale).to(dtype=dtype).contiguous()
576
+ for name, t in obj["passthrough"].items():
577
+ out_t = t.detach().to("cpu").contiguous()
578
+ orig_dtype = passthrough_orig_dtypes.get(name)
579
+ if isinstance(orig_dtype, str):
580
+ out_t = out_t.to(dtype=getattr(torch, orig_dtype)).contiguous()
581
+ out[name] = out_t
582
+ return out
583
+
584
+ # --- Data loading ---
585
+
586
+ def load_data_shard(file: Path) -> Tensor:
587
+ header_bytes = 256 * np.dtype("<i4").itemsize
588
+ token_bytes = np.dtype("<u2").itemsize
589
+ header = np.fromfile(file, dtype="<i4", count=256)
590
+ if header.size != 256 or int(header[0]) != 20240520 or int(header[1]) != 1:
591
+ raise ValueError(f"Unexpected shard header for {file}")
592
+ num_tokens = int(header[2])
593
+ expected_size = header_bytes + num_tokens * token_bytes
594
+ if file.stat().st_size != expected_size:
595
+ raise ValueError(f"Shard size mismatch for {file}: expected {expected_size} bytes")
596
+ tokens_np = np.fromfile(file, dtype="<u2", count=num_tokens, offset=header_bytes)
597
+ if tokens_np.size != num_tokens:
598
+ raise ValueError(f"Short read for {file}")
599
+ return torch.from_numpy(tokens_np.astype(np.uint16, copy=False))
600
+ class TokenStream:
601
+ def __init__(self, pattern: str):
602
+ self.files = [Path(p) for p in sorted(glob.glob(pattern))]
603
+ if not self.files:
604
+ raise FileNotFoundError(f"No files found for pattern: {pattern}")
605
+ self.file_idx = 0
606
+ self.tokens = load_data_shard(self.files[0])
607
+ self.pos = 0
608
+ def _advance_file(self) -> None:
609
+ self.file_idx = (self.file_idx + 1) % len(self.files)
610
+ self.tokens = load_data_shard(self.files[self.file_idx])
611
+ self.pos = 0
612
+ def take(self, n: int) -> Tensor:
613
+ chunks: list[Tensor] = []
614
+ remaining = n
615
+ while remaining > 0:
616
+ avail = self.tokens.numel() - self.pos
617
+ if avail <= 0:
618
+ self._advance_file()
619
+ continue
620
+ k = min(remaining, avail)
621
+ chunks.append(self.tokens[self.pos : self.pos + k])
622
+ self.pos += k
623
+ remaining -= k
624
+ return chunks[0] if len(chunks) == 1 else torch.cat(chunks)
625
+ class DistributedTokenLoader:
626
+ def __init__(self, pattern: str, rank: int, world_size: int, device: torch.device):
627
+ self.rank = rank
628
+ self.world_size = world_size
629
+ self.device = device
630
+ self.stream = TokenStream(pattern)
631
+ def next_batch(self, global_tokens: int, seq_len: int, grad_accum_steps: int) -> tuple[Tensor, Tensor]:
632
+ local_tokens = global_tokens // (self.world_size * grad_accum_steps)
633
+ per_rank_span = local_tokens + 1
634
+ chunk = self.stream.take(per_rank_span * self.world_size)
635
+ start = self.rank * per_rank_span
636
+ local = chunk[start : start + per_rank_span].to(dtype=torch.int64)
637
+ x = local[:-1].reshape(-1, seq_len)
638
+ y = local[1:].reshape(-1, seq_len)
639
+ return x.to(self.device, non_blocking=True), y.to(self.device, non_blocking=True)
640
+
641
+ # --- Transformer modules ---
642
+
643
+ class RMSNorm(nn.Module):
644
+ def __init__(self, eps: float | None = None):
645
+ super().__init__()
646
+ self.eps = eps
647
+ def forward(self, x: Tensor) -> Tensor:
648
+ return F.rms_norm(x, (x.size(-1),), eps=self.eps)
649
+ class CastedLinear(nn.Linear):
650
+ _qat_enabled: bool = False
651
+ def forward(self, x: Tensor) -> Tensor:
652
+ w = self.weight.to(x.dtype)
653
+ if CastedLinear._qat_enabled and self.training and w.ndim == 2:
654
+ with torch.no_grad():
655
+ w32 = self.weight.float()
656
+ row_max = w32.abs().amax(dim=1)
657
+ scale = (row_max / 31.0).clamp_min(1.0 / 31.0)
658
+ w_q = (torch.clamp(torch.round(w32 / scale[:, None]), -32, 31) * scale[:, None]).to(x.dtype)
659
+ w = w + (w_q - w).detach()
660
+ bias = self.bias.to(x.dtype) if self.bias is not None else None
661
+ return F.linear(x, w, bias)
662
+ def restore_low_dim_params_to_fp32(module: nn.Module) -> None:
663
+ with torch.no_grad():
664
+ for name, param in module.named_parameters():
665
+ if (param.ndim < 2 or any(pattern in name for pattern in CONTROL_TENSOR_NAME_PATTERNS)) and param.dtype != torch.float32:
666
+ param.data = param.data.float()
667
+ class Rotary(nn.Module):
668
+ def __init__(self, dim: int, base: float = 10000.0, train_seq_len: int = 1024, rope_dims: int = 0):
669
+ super().__init__()
670
+ self.dim = dim
671
+ self.base = base
672
+ self.train_seq_len = train_seq_len
673
+ self.rope_dims = rope_dims if rope_dims > 0 else dim
674
+ inv_freq = 1.0 / (base ** (torch.arange(0, self.rope_dims, 2, dtype=torch.float32) / self.rope_dims))
675
+ self.register_buffer("inv_freq", inv_freq, persistent=False)
676
+ self._seq_len_cached = 0
677
+ self._cos_cached: Tensor | None = None
678
+ self._sin_cached: Tensor | None = None
679
+ def forward(self, seq_len: int, device: torch.device, dtype: torch.dtype) -> tuple[Tensor, Tensor]:
680
+ if (
681
+ self._cos_cached is None
682
+ or self._sin_cached is None
683
+ or self._seq_len_cached != seq_len
684
+ or self._cos_cached.device != device
685
+ ):
686
+ rd = self.rope_dims
687
+ if seq_len > self.train_seq_len:
688
+ scale = seq_len / self.train_seq_len
689
+ new_base = self.base * (scale ** (rd / (rd - 2)))
690
+ inv_freq = 1.0 / (new_base ** (torch.arange(0, rd, 2, dtype=torch.float32, device=device) / rd))
691
+ else:
692
+ inv_freq = self.inv_freq.to(device)
693
+ t = torch.arange(seq_len, device=device, dtype=inv_freq.dtype)
694
+ freqs = torch.outer(t, inv_freq)
695
+ self._cos_cached = freqs.cos()[None, :, None, :]
696
+ self._sin_cached = freqs.sin()[None, :, None, :]
697
+ self._seq_len_cached = seq_len
698
+ return self._cos_cached.to(dtype=dtype), self._sin_cached.to(dtype=dtype)
699
+ def apply_rotary_emb(x: Tensor, cos: Tensor, sin: Tensor, rope_dims: int = 0) -> Tensor:
700
+ if rope_dims > 0 and rope_dims < x.size(-1):
701
+ x_rope, x_pass = x[..., :rope_dims], x[..., rope_dims:]
702
+ half = rope_dims // 2
703
+ x1, x2 = x_rope[..., :half], x_rope[..., half:]
704
+ x_rope = torch.cat((x1 * cos + x2 * sin, x1 * (-sin) + x2 * cos), dim=-1)
705
+ return torch.cat((x_rope, x_pass), dim=-1)
706
+ half = x.size(-1) // 2
707
+ x1, x2 = x[..., :half], x[..., half:]
708
+ return torch.cat((x1 * cos + x2 * sin, x1 * (-sin) + x2 * cos), dim=-1)
709
+
710
+ class CausalSelfAttention(nn.Module):
711
+ def __init__(
712
+ self,
713
+ dim: int,
714
+ num_heads: int,
715
+ num_kv_heads: int,
716
+ rope_base: float,
717
+ qk_gain_init: float,
718
+ gated_attention: bool = False,
719
+ value_residual: bool = False,
720
+ ):
721
+ super().__init__()
722
+ if dim % num_heads != 0:
723
+ raise ValueError("model_dim must be divisible by num_heads")
724
+ if num_heads % num_kv_heads != 0:
725
+ raise ValueError("num_heads must be divisible by num_kv_heads")
726
+ self.num_heads = num_heads
727
+ self.num_kv_heads = num_kv_heads
728
+ self.head_dim = dim // num_heads
729
+ if self.head_dim % 2 != 0:
730
+ raise ValueError("head_dim must be even for RoPE")
731
+ # No CastedLinear -- weights come from banks
732
+ self.q_gain = nn.Parameter(torch.full((num_heads,), qk_gain_init, dtype=torch.float32))
733
+ self.rope_dims = 0 # set by GPT.__init__ for partial RoPE
734
+ self.rotary = Rotary(self.head_dim, base=rope_base, train_seq_len=1024)
735
+ self.use_xsa = False # set by GPT.__init__ for deep layers only
736
+ # Gated attention and value residual (non-banked small params)
737
+ self.gated_attention = gated_attention
738
+ if gated_attention:
739
+ self.attn_gate = nn.Linear(dim, num_heads, bias=True)
740
+ nn.init.zeros_(self.attn_gate.weight)
741
+ nn.init.constant_(self.attn_gate.bias, 4.0)
742
+ self.value_residual = value_residual
743
+ if value_residual:
744
+ self.vr_lambda = nn.Parameter(torch.tensor([0.5, 0.5], dtype=torch.float32))
745
+ def _xsa_efficient(self, y: Tensor, v: Tensor) -> Tensor:
746
+ """Efficient XSA: subtract self-value projection via GQA-aware reshape (no repeat_interleave).
747
+ y: [B, T, H, D], v: [B, T, Hkv, D]. H must be divisible by Hkv."""
748
+ B, T, H, D = y.shape
749
+ Hkv = v.size(-2)
750
+ group = H // Hkv
751
+ y_g = y.reshape(B, T, Hkv, group, D) # [B, T, Hkv, group, D]
752
+ vn = F.normalize(v, dim=-1).unsqueeze(-2) # [B, T, Hkv, 1, D] -- broadcast ready
753
+ proj = (y_g * vn).sum(dim=-1, keepdim=True) * vn
754
+ return (y_g - proj).reshape(B, T, H, D)
755
+ def forward(self, x: Tensor, q_w: Tensor, k_w: Tensor, v_w: Tensor, out_w: Tensor, v_embed: Tensor | None = None, v0: Tensor | None = None) -> tuple[Tensor, Tensor | None]:
756
+ bsz, seqlen, dim = x.shape
757
+ q = F.linear(x, q_w.to(x.dtype)).reshape(bsz, seqlen, self.num_heads, self.head_dim)
758
+ k = F.linear(x, k_w.to(x.dtype)).reshape(bsz, seqlen, self.num_kv_heads, self.head_dim)
759
+ v = F.linear(x, v_w.to(x.dtype))
760
+ if v_embed is not None:
761
+ v = v + v_embed
762
+ v = v.reshape(bsz, seqlen, self.num_kv_heads, self.head_dim)
763
+ raw_v = v if self.value_residual else None
764
+ if self.value_residual and v0 is not None:
765
+ lam = self.vr_lambda.to(dtype=v.dtype)
766
+ v = lam[0] * v0 + lam[1] * v
767
+ q = F.rms_norm(q, (q.size(-1),))
768
+ k = F.rms_norm(k, (k.size(-1),))
769
+ cos, sin = self.rotary(seqlen, x.device, q.dtype)
770
+ q = apply_rotary_emb(q, cos, sin, self.rope_dims)
771
+ k = apply_rotary_emb(k, cos, sin, self.rope_dims)
772
+ q = q * self.q_gain.to(dtype=q.dtype)[None, None, :, None]
773
+ y = flash_attn_3_func(q, k, v, causal=True)
774
+ if self.use_xsa:
775
+ y = self._xsa_efficient(y, v)
776
+ if self.gated_attention:
777
+ # gate shape: (bsz, seqlen, num_heads) -> (bsz, seqlen, num_heads, 1) for B,T,H,D layout
778
+ gate = torch.sigmoid(self.attn_gate(x)).unsqueeze(-1)
779
+ y = y * gate
780
+ y = y.reshape(bsz, seqlen, dim)
781
+ return F.linear(y, out_w.to(x.dtype)), raw_v
782
+
783
+ class SmearGate(nn.Module):
784
+ def __init__(self, dim: int):
785
+ super().__init__()
786
+ self.gate = nn.Parameter(torch.zeros(dim, dtype=torch.float32))
787
+ def forward(self, x: Tensor) -> Tensor:
788
+ g = torch.sigmoid(self.gate.to(dtype=x.dtype))[None, None, :]
789
+ x_prev = torch.cat([torch.zeros_like(x[:, :1]), x[:, :-1]], dim=1)
790
+ return (1 - g) * x + g * x_prev
791
+
792
+ class BigramHashEmbedding(nn.Module):
793
+ def __init__(self, bigram_vocab_size: int, bigram_dim: int, model_dim: int):
794
+ super().__init__()
795
+ self.bigram_vocab_size = bigram_vocab_size
796
+ self.embed = nn.Embedding(bigram_vocab_size, bigram_dim)
797
+ nn.init.zeros_(self.embed.weight)
798
+ self.proj = CastedLinear(bigram_dim, model_dim, bias=False) if bigram_dim != model_dim else None
799
+ if self.proj is not None:
800
+ nn.init.zeros_(self.proj.weight)
801
+ self.scale = nn.Parameter(torch.tensor(0.05, dtype=torch.float32))
802
+ def bigram_hash(self, tokens: Tensor) -> Tensor:
803
+ t = tokens.to(torch.int32)
804
+ mod = self.bigram_vocab_size - 1
805
+ out = torch.empty_like(t)
806
+ out[..., 0] = mod
807
+ out[..., 1:] = torch.bitwise_xor(36313 * t[..., 1:], 27191 * t[..., :-1]) % mod
808
+ return out.long()
809
+ def forward(self, token_ids: Tensor) -> Tensor:
810
+ h = self.embed(self.bigram_hash(token_ids))
811
+ if self.proj is not None:
812
+ h = self.proj(h)
813
+ return h * self.scale.to(dtype=h.dtype)
814
+
815
+ class ValueEmbedding(nn.Module):
816
+ """Reinject token identity into attention values at specific layers.
817
+ Each table maps vocab tokens to a low-dim embedding, projected to model_dim."""
818
+ def __init__(self, vocab_size: int, ve_dim: int, model_dim: int):
819
+ super().__init__()
820
+ self.embed = nn.Embedding(vocab_size, ve_dim)
821
+ nn.init.normal_(self.embed.weight, std=0.01)
822
+ self.proj = CastedLinear(ve_dim, model_dim, bias=False) if ve_dim != model_dim else None
823
+ if self.proj is not None:
824
+ nn.init.zeros_(self.proj.weight)
825
+ self.scale = nn.Parameter(torch.tensor(0.1, dtype=torch.float32))
826
+ def forward(self, token_ids: Tensor) -> Tensor:
827
+ h = self.embed(token_ids)
828
+ if self.proj is not None:
829
+ h = self.proj(h)
830
+ return h * self.scale.to(dtype=h.dtype)
831
+
832
+ class MLP(nn.Module):
833
+ def __init__(self, dim: int, mlp_mult: int):
834
+ super().__init__()
835
+ # No CastedLinear -- weights come from banks
836
+ def forward(self, x: Tensor, up_w: Tensor, down_w: Tensor) -> Tensor:
837
+ x = F.leaky_relu(F.linear(x, up_w.to(x.dtype)), negative_slope=0.5)
838
+ return F.linear(x.square(), down_w.to(x.dtype))
839
+
840
+ class Block(nn.Module):
841
+ def __init__(
842
+ self,
843
+ dim: int,
844
+ num_heads: int,
845
+ num_kv_heads: int,
846
+ mlp_mult: int,
847
+ rope_base: float,
848
+ qk_gain_init: float,
849
+ layer_idx: int = 0,
850
+ ln_scale: bool = False,
851
+ dtg: bool = False,
852
+ gated_attention: bool = False,
853
+ value_residual: bool = False,
854
+ ):
855
+ super().__init__()
856
+ self.attn_norm = RMSNorm()
857
+ self.mlp_norm = RMSNorm()
858
+ self.attn = CausalSelfAttention(dim, num_heads, num_kv_heads, rope_base, qk_gain_init,
859
+ gated_attention=gated_attention, value_residual=value_residual)
860
+ self.mlp = MLP(dim, mlp_mult)
861
+ self.attn_scale = nn.Parameter(torch.ones(dim, dtype=torch.float32))
862
+ self.mlp_scale = nn.Parameter(torch.ones(dim, dtype=torch.float32))
863
+ self.resid_mix = nn.Parameter(torch.stack((torch.ones(dim), torch.zeros(dim))).float())
864
+ self.ln_scale_factor = 1.0 / math.sqrt(layer_idx + 1) if ln_scale else 1.0
865
+ if dtg:
866
+ self.dtg_gate = nn.Linear(dim, 1, bias=True)
867
+ nn.init.zeros_(self.dtg_gate.weight)
868
+ nn.init.constant_(self.dtg_gate.bias, 2.0)
869
+ else:
870
+ self.dtg_gate = None
871
+ def forward(self, x: Tensor, x0: Tensor, q_w: Tensor, k_w: Tensor, v_w: Tensor, out_w: Tensor, up_w: Tensor, down_w: Tensor, v_embed: Tensor | None = None, v0: Tensor | None = None) -> tuple[Tensor, Tensor | None]:
872
+ mix = self.resid_mix.to(dtype=x.dtype)
873
+ x_in = mix[0][None, None, :] * x + mix[1][None, None, :] * x0
874
+ attn_out, raw_v = self.attn(self.attn_norm(x_in) * self.ln_scale_factor, q_w, k_w, v_w, out_w, v_embed=v_embed, v0=v0)
875
+ x_out = x_in + self.attn_scale.to(dtype=x_in.dtype)[None, None, :] * attn_out
876
+ x_out = x_out + self.mlp_scale.to(dtype=x_out.dtype)[None, None, :] * self.mlp(self.mlp_norm(x_out) * self.ln_scale_factor, up_w, down_w)
877
+ if self.dtg_gate is not None:
878
+ gate = torch.sigmoid(self.dtg_gate(x_in.detach()))
879
+ x_out = x_in + gate * (x_out - x_in)
880
+ return x_out, raw_v
881
+
882
+ class GPT(nn.Module):
883
+ def __init__(
884
+ self,
885
+ vocab_size: int,
886
+ num_layers: int,
887
+ model_dim: int,
888
+ num_heads: int,
889
+ num_kv_heads: int,
890
+ mlp_mult: int,
891
+ tie_embeddings: bool,
892
+ tied_embed_init_std: float,
893
+ logit_softcap: float,
894
+ rope_base: float,
895
+ qk_gain_init: float,
896
+ mtp_num_heads: int = 0,
897
+ mtp_loss_weight: float = 0.1,
898
+ bigram_vocab_size: int = 0,
899
+ bigram_dim: int = 128,
900
+ xsa_last_n: int = 0,
901
+ rope_dims: int = 0,
902
+ ln_scale: bool = False,
903
+ dtg: bool = False,
904
+ ve_enabled: bool = False,
905
+ ve_dim: int = 128,
906
+ ve_layers: str = "9,10",
907
+ gated_attention: bool = False,
908
+ value_residual: bool = False,
909
+ ):
910
+ super().__init__()
911
+ self._ve_target_dim = num_kv_heads * (model_dim // num_heads) # kv_dim for value projection
912
+ if logit_softcap <= 0.0:
913
+ raise ValueError(f"logit_softcap must be positive, got {logit_softcap}")
914
+ self.tie_embeddings = tie_embeddings
915
+ self.tied_embed_init_std = tied_embed_init_std
916
+ self.logit_softcap = logit_softcap
917
+ self.value_residual = value_residual
918
+ self.mtp_num_heads = mtp_num_heads
919
+ self.mtp_loss_weight = mtp_loss_weight
920
+ self.tok_emb = nn.Embedding(vocab_size, model_dim)
921
+ self.bigram = BigramHashEmbedding(bigram_vocab_size, bigram_dim, model_dim) if bigram_vocab_size > 0 else None
922
+ self.smear = SmearGate(model_dim)
923
+ self.num_encoder_layers = num_layers // 2
924
+ self.num_decoder_layers = num_layers - self.num_encoder_layers
925
+ self.num_skip_weights = min(self.num_encoder_layers, self.num_decoder_layers)
926
+ self.skip_weights = nn.Parameter(torch.ones(self.num_skip_weights, model_dim, dtype=torch.float32))
927
+ # Parameter banks: contiguous 3D tensors for batched optimizer
928
+ head_dim = model_dim // num_heads
929
+ kv_dim = num_kv_heads * head_dim
930
+ mlp_dim = int(mlp_mult * model_dim)
931
+ self.num_layers = num_layers
932
+ self.qo_bank = nn.Parameter(torch.empty(2 * num_layers, model_dim, model_dim))
933
+ self.kv_bank = nn.Parameter(torch.empty(2 * num_layers, kv_dim, model_dim))
934
+ self.mlp_up_bank = nn.Parameter(torch.empty(num_layers, mlp_dim, model_dim))
935
+ self.mlp_down_bank = nn.Parameter(torch.empty(num_layers, model_dim, mlp_dim))
936
+ self.blocks = nn.ModuleList(
937
+ [
938
+ Block(
939
+ model_dim,
940
+ num_heads,
941
+ num_kv_heads,
942
+ mlp_mult,
943
+ rope_base,
944
+ qk_gain_init,
945
+ layer_idx=i,
946
+ ln_scale=ln_scale,
947
+ dtg=dtg,
948
+ gated_attention=gated_attention,
949
+ value_residual=value_residual,
950
+ )
951
+ for i in range(num_layers)
952
+ ]
953
+ )
954
+ if rope_dims > 0:
955
+ head_dim = model_dim // num_heads
956
+ for block in self.blocks:
957
+ block.attn.rope_dims = rope_dims
958
+ block.attn.rotary = Rotary(head_dim, base=rope_base, train_seq_len=1024, rope_dims=rope_dims)
959
+ self.ve_layer_indices = [int(x) for x in ve_layers.split(",") if x.strip()] if ve_enabled else []
960
+ kv_dim_ve = self._ve_target_dim
961
+ if self.ve_layer_indices:
962
+ self.ve_shared = ValueEmbedding(vocab_size, ve_dim, kv_dim_ve)
963
+ self.ve_layer_scales = nn.ParameterList(
964
+ [nn.Parameter(torch.ones(1, dtype=torch.float32)) for _ in self.ve_layer_indices]
965
+ )
966
+ else:
967
+ self.ve_shared = None
968
+ self.ve_layer_scales = nn.ParameterList()
969
+ self.value_embeds = nn.ModuleList() # keep empty for compat
970
+ self.final_norm = RMSNorm()
971
+ self.lm_head = None if tie_embeddings else CastedLinear(model_dim, vocab_size, bias=False)
972
+ if self.lm_head is not None:
973
+ self.lm_head._zero_init = True
974
+ self.mtp_heads = nn.ModuleList(
975
+ [CastedLinear(model_dim, vocab_size, bias=False) for _ in range(mtp_num_heads)]
976
+ )
977
+ for head in self.mtp_heads:
978
+ head._zero_init = True
979
+ if xsa_last_n > 0:
980
+ for i in range(max(0, num_layers - xsa_last_n), num_layers):
981
+ self.blocks[i].attn.use_xsa = True
982
+ self._init_weights()
983
+ def _init_weights(self) -> None:
984
+ if self.tie_embeddings:
985
+ nn.init.normal_(self.tok_emb.weight, mean=0.0, std=self.tied_embed_init_std)
986
+ n = self.num_layers
987
+ proj_scale = 1.0 / math.sqrt(2 * n)
988
+ # Init banks: orthogonal, with proj layers scaled down and out/down zero-init
989
+ for i in range(n):
990
+ nn.init.orthogonal_(self.qo_bank.data[i], gain=1.0) # Q
991
+ nn.init.zeros_(self.qo_bank.data[n + i]) # Out (zero init)
992
+ nn.init.orthogonal_(self.kv_bank.data[i], gain=1.0) # K
993
+ nn.init.orthogonal_(self.kv_bank.data[n + i], gain=1.0) # V
994
+ nn.init.orthogonal_(self.mlp_up_bank.data[i], gain=1.0) # MLP up
995
+ nn.init.zeros_(self.mlp_down_bank.data[i]) # MLP down (zero init)
996
+ # Scale proj layers (out_proj and mlp_down are "proj" layers)
997
+ self.qo_bank.data[n + i].mul_(proj_scale)
998
+ self.mlp_down_bank.data[i].mul_(proj_scale)
999
+ # Init remaining nn.Linear modules (bigram proj, mtp heads, lm_head)
1000
+ for name, module in self.named_modules():
1001
+ if isinstance(module, nn.Linear):
1002
+ if getattr(module, "_zero_init", False):
1003
+ nn.init.zeros_(module.weight)
1004
+ elif module.weight.ndim == 2 and module.weight.shape[0] >= 64 and module.weight.shape[1] >= 64:
1005
+ nn.init.orthogonal_(module.weight, gain=1.0)
1006
+ def _get_ve(self, layer_idx: int, input_ids: Tensor, ve_cache: dict | None = None) -> Tensor | None:
1007
+ """Get value embedding for a specific layer using shared table + per-layer scale."""
1008
+ if self.ve_shared is None or layer_idx not in self.ve_layer_indices:
1009
+ return None
1010
+ if ve_cache is not None and 've' not in ve_cache:
1011
+ ve_cache['ve'] = self.ve_shared(input_ids)
1012
+ ve_base = ve_cache['ve'] if ve_cache is not None else self.ve_shared(input_ids)
1013
+ ve_idx = self.ve_layer_indices.index(layer_idx)
1014
+ return ve_base * self.ve_layer_scales[ve_idx].to(dtype=ve_base.dtype)
1015
+ def forward(self, input_ids: Tensor, target_ids: Tensor) -> Tensor:
1016
+ n = self.num_layers
1017
+ x = self.tok_emb(input_ids)
1018
+ if self.bigram is not None:
1019
+ x = x + self.bigram(input_ids)
1020
+ x = F.rms_norm(x, (x.size(-1),))
1021
+ x = self.smear(x)
1022
+ x0 = x
1023
+ v0 = None
1024
+ skips: list[Tensor] = []
1025
+ ve_cache: dict = {}
1026
+ for i in range(self.num_encoder_layers):
1027
+ ve = self._get_ve(i, input_ids, ve_cache)
1028
+ x, raw_v = self.blocks[i](x, x0,
1029
+ self.qo_bank[i], self.kv_bank[i], self.kv_bank[n + i],
1030
+ self.qo_bank[n + i], self.mlp_up_bank[i], self.mlp_down_bank[i],
1031
+ v_embed=ve, v0=v0)
1032
+ if v0 is None and raw_v is not None:
1033
+ v0 = raw_v
1034
+ skips.append(x)
1035
+ for i in range(self.num_decoder_layers):
1036
+ bi = self.num_encoder_layers + i
1037
+ if skips:
1038
+ x = x + self.skip_weights[i].to(dtype=x.dtype)[None, None, :] * skips.pop()
1039
+ ve = self._get_ve(bi, input_ids, ve_cache)
1040
+ x, _ = self.blocks[bi](x, x0,
1041
+ self.qo_bank[bi], self.kv_bank[bi], self.kv_bank[n + bi],
1042
+ self.qo_bank[n + bi], self.mlp_up_bank[bi], self.mlp_down_bank[bi],
1043
+ v_embed=ve, v0=v0)
1044
+ x = self.final_norm(x)
1045
+ x_flat = x.reshape(-1, x.size(-1))
1046
+ targets = target_ids.reshape(-1)
1047
+ if self.tie_embeddings:
1048
+ logits_proj = F.linear(x_flat, self.tok_emb.weight)
1049
+ else:
1050
+ if self.lm_head is None:
1051
+ raise RuntimeError("lm_head is required when tie_embeddings=False")
1052
+ logits_proj = self.lm_head(x_flat)
1053
+ logits = self.logit_softcap * torch.tanh(logits_proj / self.logit_softcap)
1054
+ main_loss = F.cross_entropy(logits.float(), targets, reduction="mean")
1055
+ if self.training and self.mtp_num_heads > 0 and self.mtp_loss_weight > 0.0:
1056
+ _, seqlen, dim = x.shape
1057
+ mtp_loss_sum = x.new_zeros(())
1058
+ mtp_loss_count = 0
1059
+ for k, mtp_head in enumerate(self.mtp_heads):
1060
+ valid_t = seqlen - (k + 1)
1061
+ if valid_t <= 0:
1062
+ continue
1063
+ mtp_hidden = x[:, :valid_t, :].reshape(-1, dim)
1064
+ mtp_targets = target_ids[:, k + 1 :].reshape(-1)
1065
+ mtp_logits_proj = mtp_head(mtp_hidden)
1066
+ mtp_logits = self.logit_softcap * torch.tanh(mtp_logits_proj / self.logit_softcap)
1067
+ mtp_loss_sum = mtp_loss_sum + F.cross_entropy(mtp_logits.float(), mtp_targets, reduction="mean")
1068
+ mtp_loss_count += 1
1069
+ if mtp_loss_count > 0:
1070
+ main_loss = main_loss + self.mtp_loss_weight * (mtp_loss_sum / mtp_loss_count)
1071
+ return main_loss
1072
+ def forward_logits(self, input_ids: Tensor) -> Tensor:
1073
+ """Return logits (bsz, seq_len, vocab) without computing loss."""
1074
+ n = self.num_layers
1075
+ x = self.tok_emb(input_ids)
1076
+ if self.bigram is not None:
1077
+ x = x + self.bigram(input_ids)
1078
+ x = F.rms_norm(x, (x.size(-1),))
1079
+ x = self.smear(x)
1080
+ x0 = x
1081
+ v0 = None
1082
+ skips: list[Tensor] = []
1083
+ ve_cache: dict = {}
1084
+ for i in range(self.num_encoder_layers):
1085
+ ve = self._get_ve(i, input_ids, ve_cache)
1086
+ x, raw_v = self.blocks[i](x, x0,
1087
+ self.qo_bank[i], self.kv_bank[i], self.kv_bank[n + i],
1088
+ self.qo_bank[n + i], self.mlp_up_bank[i], self.mlp_down_bank[i],
1089
+ v_embed=ve, v0=v0)
1090
+ if v0 is None and raw_v is not None:
1091
+ v0 = raw_v
1092
+ skips.append(x)
1093
+ for i in range(self.num_decoder_layers):
1094
+ bi = self.num_encoder_layers + i
1095
+ if skips:
1096
+ x = x + self.skip_weights[i].to(dtype=x.dtype)[None, None, :] * skips.pop()
1097
+ ve = self._get_ve(bi, input_ids, ve_cache)
1098
+ x, _ = self.blocks[bi](x, x0,
1099
+ self.qo_bank[bi], self.kv_bank[bi], self.kv_bank[n + bi],
1100
+ self.qo_bank[n + bi], self.mlp_up_bank[bi], self.mlp_down_bank[bi],
1101
+ v_embed=ve, v0=v0)
1102
+ x = self.final_norm(x)
1103
+ if self.tie_embeddings:
1104
+ logits_proj = F.linear(x, self.tok_emb.weight)
1105
+ else:
1106
+ logits_proj = self.lm_head(x)
1107
+ return self.logit_softcap * torch.tanh(logits_proj / self.logit_softcap)
1108
+
1109
+ # --- Sliding window evaluation ---
1110
+
1111
+ def eval_val_sliding(
1112
+ args: Hyperparameters,
1113
+ base_model: nn.Module,
1114
+ rank: int,
1115
+ world_size: int,
1116
+ device: torch.device,
1117
+ val_tokens: Tensor,
1118
+ base_bytes_lut: Tensor,
1119
+ has_leading_space_lut: Tensor,
1120
+ is_boundary_token_lut: Tensor,
1121
+ stride: int,
1122
+ batch_seqs: int = 32,
1123
+ eval_seq_len: int | None = None,
1124
+ ) -> tuple[float, float]:
1125
+ """Sliding window evaluation: each token scored with maximum context."""
1126
+ seq_len = eval_seq_len or args.train_seq_len
1127
+ total_tokens = val_tokens.numel() - 1
1128
+ window_starts = [ws for ws in range(0, total_tokens, stride)
1129
+ if min(ws + seq_len, total_tokens) - ws >= 1]
1130
+ total_windows = len(window_starts)
1131
+ my_s = (total_windows * rank) // world_size
1132
+ my_e = (total_windows * (rank + 1)) // world_size
1133
+ my_windows = window_starts[my_s:my_e]
1134
+ loss_sum = torch.zeros((), device=device, dtype=torch.float64)
1135
+ token_count = torch.zeros((), device=device, dtype=torch.float64)
1136
+ byte_count = torch.zeros((), device=device, dtype=torch.float64)
1137
+ base_model.eval()
1138
+ compiled_logits = torch.compile(base_model.forward_logits, dynamic=False, fullgraph=True)
1139
+ with torch.inference_mode():
1140
+ for bi in range(0, len(my_windows), batch_seqs):
1141
+ batch_ws = my_windows[bi:bi + batch_seqs]
1142
+ bsz = len(batch_ws)
1143
+ x_batch = torch.zeros(bsz, seq_len, dtype=torch.int64, device=device)
1144
+ y_batch = torch.zeros(bsz, seq_len, dtype=torch.int64, device=device)
1145
+ wlens: list[int] = []
1146
+ for i, ws in enumerate(batch_ws):
1147
+ end = min(ws + seq_len, total_tokens)
1148
+ wlen = end - ws
1149
+ wlens.append(wlen)
1150
+ chunk = val_tokens[ws:end + 1].to(dtype=torch.int64, device=device)
1151
+ x_batch[i, :wlen] = chunk[:-1]
1152
+ y_batch[i, :wlen] = chunk[1:]
1153
+ with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
1154
+ logits = compiled_logits(x_batch)
1155
+ nll = F.cross_entropy(
1156
+ logits.reshape(-1, logits.size(-1)).float(),
1157
+ y_batch.reshape(-1),
1158
+ reduction="none",
1159
+ ).reshape(bsz, seq_len)
1160
+ for i, ws in enumerate(batch_ws):
1161
+ wlen = wlens[i]
1162
+ s = 0 if ws == 0 else max(wlen - stride, 0)
1163
+ scored_nll = nll[i, s:wlen].to(torch.float64)
1164
+ loss_sum += scored_nll.sum()
1165
+ token_count += float(wlen - s)
1166
+ tgt = y_batch[i, s:wlen]
1167
+ prev = x_batch[i, s:wlen]
1168
+ tb = base_bytes_lut[tgt].to(torch.float64)
1169
+ tb += (has_leading_space_lut[tgt] & ~is_boundary_token_lut[prev]).to(torch.float64)
1170
+ byte_count += tb.sum()
1171
+ if dist.is_available() and dist.is_initialized():
1172
+ dist.all_reduce(loss_sum, op=dist.ReduceOp.SUM)
1173
+ dist.all_reduce(token_count, op=dist.ReduceOp.SUM)
1174
+ dist.all_reduce(byte_count, op=dist.ReduceOp.SUM)
1175
+ val_loss = (loss_sum / token_count).item()
1176
+ bits_per_token = val_loss / math.log(2.0)
1177
+ tokens_per_byte = token_count.item() / byte_count.item()
1178
+ base_model.train()
1179
+ return val_loss, bits_per_token * tokens_per_byte
1180
+
1181
+
1182
+ def eval_val_sliding_ttt(
1183
+ args: Hyperparameters, base_model: nn.Module, rank: int, world_size: int,
1184
+ device: torch.device, val_tokens: Tensor, base_bytes_lut: Tensor,
1185
+ has_leading_space_lut: Tensor, is_boundary_token_lut: Tensor,
1186
+ stride: int, batch_seqs: int = 32, log0=print,
1187
+ ) -> tuple[float, float]:
1188
+ """Legal score-first TTT (PR #461 recipe): score each chunk with sliding windows,
1189
+ then train on it. Every token scored BEFORE any update that could use it."""
1190
+ seq_len = args.train_seq_len
1191
+ total_tokens = val_tokens.numel() - 1
1192
+ ttt_chunk = args.ttt_chunk_tokens
1193
+
1194
+ # Pre-compute all window starts
1195
+ window_starts = [ws for ws in range(0, total_tokens, stride)
1196
+ if min(ws + seq_len, total_tokens) - ws >= stride or ws == 0]
1197
+
1198
+ # Assign each window to a chunk based on the first token it scores
1199
+ num_chunks = (total_tokens + ttt_chunk - 1) // ttt_chunk
1200
+ chunk_windows: list[list[int]] = [[] for _ in range(num_chunks)]
1201
+ for ws in window_starts:
1202
+ end = min(ws + seq_len, total_tokens)
1203
+ wlen = end - ws
1204
+ s = 0 if ws == 0 else max(wlen - stride, 0)
1205
+ scored_start = ws + s
1206
+ ci = min(scored_start // ttt_chunk, num_chunks - 1)
1207
+ chunk_windows[ci].append(ws)
1208
+
1209
+ log0(f"ttt_sliding:start chunks={num_chunks} chunk_tokens={ttt_chunk} "
1210
+ f"total_windows={len(window_starts)} stride={stride} "
1211
+ f"ttt_lr={args.ttt_lr} ttt_epochs={args.ttt_epochs} "
1212
+ f"freeze_blocks={args.ttt_freeze_blocks}")
1213
+
1214
+ loss_sum = torch.zeros((), device=device, dtype=torch.float64)
1215
+ token_count = torch.zeros((), device=device, dtype=torch.float64)
1216
+ byte_count = torch.zeros((), device=device, dtype=torch.float64)
1217
+
1218
+ # Freeze first N blocks
1219
+ frozen_block_ids = set(range(min(args.ttt_freeze_blocks, len(base_model.blocks))))
1220
+ ttt_params = []
1221
+ for name, p in base_model.named_parameters():
1222
+ freeze = False
1223
+ for bi in frozen_block_ids:
1224
+ if f"blocks.{bi}." in name:
1225
+ freeze = True
1226
+ break
1227
+ if freeze:
1228
+ p.requires_grad_(False)
1229
+ else:
1230
+ p.requires_grad_(True)
1231
+ ttt_params.append(p)
1232
+
1233
+ log0(f"ttt_sliding:params unfrozen={sum(p.numel() for p in ttt_params)} "
1234
+ f"frozen={sum(p.numel() for p in base_model.parameters() if not p.requires_grad)}")
1235
+
1236
+ optimizer = torch.optim.SGD(ttt_params, lr=args.ttt_lr, momentum=args.ttt_momentum)
1237
+ t0 = time.perf_counter()
1238
+
1239
+ for ci in range(num_chunks):
1240
+ windows = chunk_windows[ci]
1241
+ if not windows:
1242
+ continue
1243
+ chunk_start = ci * ttt_chunk
1244
+ chunk_end = min((ci + 1) * ttt_chunk, total_tokens)
1245
+
1246
+ # --- Phase 1: SCORE this chunk's windows (inference_mode) ---
1247
+ my_s = (len(windows) * rank) // world_size
1248
+ my_e = (len(windows) * (rank + 1)) // world_size
1249
+ my_windows = windows[my_s:my_e]
1250
+
1251
+ base_model.eval()
1252
+ with torch.inference_mode():
1253
+ for bi in range(0, len(my_windows), batch_seqs):
1254
+ batch_ws = my_windows[bi:bi + batch_seqs]
1255
+ bsz = len(batch_ws)
1256
+ x_batch = torch.zeros(bsz, seq_len, dtype=torch.int64, device=device)
1257
+ y_batch = torch.zeros(bsz, seq_len, dtype=torch.int64, device=device)
1258
+ wlens: list[int] = []
1259
+ for i, ws in enumerate(batch_ws):
1260
+ end = min(ws + seq_len, total_tokens)
1261
+ wlen = end - ws
1262
+ wlens.append(wlen)
1263
+ chunk_tok = val_tokens[ws:end + 1].to(dtype=torch.int64, device=device)
1264
+ x_batch[i, :wlen] = chunk_tok[:-1]
1265
+ y_batch[i, :wlen] = chunk_tok[1:]
1266
+ with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
1267
+ logits = base_model.forward_logits(x_batch)
1268
+ nll = F.cross_entropy(
1269
+ logits.reshape(-1, logits.size(-1)).float(),
1270
+ y_batch.reshape(-1), reduction="none",
1271
+ ).reshape(bsz, seq_len)
1272
+ for i, ws in enumerate(batch_ws):
1273
+ wlen = wlens[i]
1274
+ s = 0 if ws == 0 else max(wlen - stride, 0)
1275
+ scored_nll = nll[i, s:wlen].to(torch.float64)
1276
+ loss_sum += scored_nll.sum()
1277
+ token_count += float(wlen - s)
1278
+ tgt, prev = y_batch[i, s:wlen], x_batch[i, s:wlen]
1279
+ tb = base_bytes_lut[tgt].to(torch.float64)
1280
+ tb += (has_leading_space_lut[tgt] & ~is_boundary_token_lut[prev]).to(torch.float64)
1281
+ byte_count += tb.sum()
1282
+
1283
+ # --- Phase 2: TRAIN on this chunk (already scored = legal) ---
1284
+ is_last_chunk = (ci == num_chunks - 1)
1285
+ if not is_last_chunk and args.ttt_epochs > 0:
1286
+ base_model.train()
1287
+ chunk_seqs = (chunk_end - chunk_start) // seq_len
1288
+ if chunk_seqs > 0:
1289
+ cos_lr = args.ttt_lr * 0.5 * (1.0 + math.cos(math.pi * ci / max(num_chunks - 1, 1)))
1290
+ for pg in optimizer.param_groups:
1291
+ pg['lr'] = cos_lr
1292
+ my_seq_s = (chunk_seqs * rank) // world_size
1293
+ my_seq_e = (chunk_seqs * (rank + 1)) // world_size
1294
+ my_chunk_seqs = my_seq_e - my_seq_s
1295
+ for _ep in range(args.ttt_epochs):
1296
+ for bs in range(0, my_chunk_seqs, args.ttt_batch_seqs):
1297
+ be = min(bs + args.ttt_batch_seqs, my_chunk_seqs)
1298
+ actual_bs = my_seq_s + bs
1299
+ start_tok = chunk_start + actual_bs * seq_len
1300
+ end_tok = chunk_start + (my_seq_s + be) * seq_len + 1
1301
+ if end_tok > val_tokens.numel():
1302
+ continue
1303
+ local = val_tokens[start_tok:end_tok].to(device=device, dtype=torch.int64)
1304
+ x = local[:-1].reshape(-1, seq_len)
1305
+ y = local[1:].reshape(-1, seq_len)
1306
+ optimizer.zero_grad(set_to_none=True)
1307
+ with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
1308
+ loss = base_model(x, y)
1309
+ loss.backward()
1310
+ if world_size > 1:
1311
+ for p in ttt_params:
1312
+ if p.grad is not None:
1313
+ dist.all_reduce(p.grad, op=dist.ReduceOp.AVG)
1314
+ torch.nn.utils.clip_grad_norm_(ttt_params, args.ttt_grad_clip)
1315
+ optimizer.step()
1316
+
1317
+ if rank == 0 and (ci % 10 == 0 or ci == num_chunks - 1):
1318
+ elapsed = time.perf_counter() - t0
1319
+ rl = loss_sum.item() / max(token_count.item(), 1)
1320
+ rbpb = rl / math.log(2.0) * (token_count.item() / max(byte_count.item(), 1)) if token_count.item() > 0 else 0.0
1321
+ log0(f" ttt_chunk [{ci+1}/{num_chunks}] bpb={rbpb:.6f} time={elapsed:.1f}s")
1322
+
1323
+ if dist.is_available() and dist.is_initialized():
1324
+ dist.all_reduce(loss_sum, op=dist.ReduceOp.SUM)
1325
+ dist.all_reduce(token_count, op=dist.ReduceOp.SUM)
1326
+ dist.all_reduce(byte_count, op=dist.ReduceOp.SUM)
1327
+
1328
+ val_loss = (loss_sum / token_count).item()
1329
+ val_bpb = val_loss / math.log(2.0) * (token_count.item() / byte_count.item())
1330
+
1331
+ for p in base_model.parameters():
1332
+ p.requires_grad_(True)
1333
+ base_model.eval()
1334
+
1335
+ log0(f"ttt_sliding:done val_loss={val_loss:.6f} val_bpb={val_bpb:.6f} "
1336
+ f"elapsed={time.perf_counter() - t0:.1f}s")
1337
+ return val_loss, val_bpb
1338
+
1339
+
1340
+ # --- GPTQ-lite int6 quantization ---
1341
+
1342
+ def _classify_param(name: str) -> str:
1343
+ if "tok_emb" in name or "lm_head" in name:
1344
+ return "embed"
1345
+ if ".mlp." in name:
1346
+ return "mlp"
1347
+ if ".attn." in name or (".proj." in name and ".mlp." not in name):
1348
+ return "attn"
1349
+ return "other"
1350
+ def quantize_int6_per_row(t: Tensor, clip_range: int = 31) -> tuple[Tensor, Tensor]:
1351
+ t32 = t.float()
1352
+ if t32.ndim == 2:
1353
+ best_q, best_s, best_err = None, None, float('inf')
1354
+ for pct in [0.9990, 0.9995, 0.9999, 0.99999, 1.0]:
1355
+ if pct < 1.0:
1356
+ row_clip = torch.quantile(t32.abs(), pct, dim=1)
1357
+ else:
1358
+ row_clip = t32.abs().amax(dim=1)
1359
+ s = (row_clip / clip_range).clamp_min(1.0 / clip_range).to(torch.float16)
1360
+ q = torch.clamp(torch.round(t32 / s.float()[:, None]), -clip_range, clip_range).to(torch.int8)
1361
+ recon = q.float() * s.float()[:, None]
1362
+ err = (t32 - recon).pow(2).mean().item()
1363
+ if err < best_err:
1364
+ best_q, best_s, best_err = q, s, err
1365
+ return best_q, best_s
1366
+ amax = t32.abs().max().item()
1367
+ scale = torch.tensor(amax / clip_range if amax > 0 else 1.0, dtype=torch.float16)
1368
+ q = torch.clamp(torch.round(t32 / scale.float()), -clip_range, clip_range).to(torch.int8)
1369
+ return q, scale
1370
+
1371
+ def _unbank_state_dict(sd: dict[str, Tensor], num_layers: int) -> dict[str, Tensor]:
1372
+ """Convert 3D bank tensors into individual 2D tensors with standard names."""
1373
+ out: dict[str, Tensor] = {}
1374
+ n = num_layers
1375
+ for name, tensor in sd.items():
1376
+ if name == "qo_bank":
1377
+ for i in range(n):
1378
+ out[f"blocks.{i}.attn.c_q.weight"] = tensor[i]
1379
+ out[f"blocks.{i}.attn.proj.weight"] = tensor[n + i]
1380
+ elif name == "kv_bank":
1381
+ for i in range(n):
1382
+ out[f"blocks.{i}.attn.c_k.weight"] = tensor[i]
1383
+ out[f"blocks.{i}.attn.c_v.weight"] = tensor[n + i]
1384
+ elif name == "mlp_up_bank":
1385
+ for i in range(n):
1386
+ out[f"blocks.{i}.mlp.fc.weight"] = tensor[i]
1387
+ elif name == "mlp_down_bank":
1388
+ for i in range(n):
1389
+ out[f"blocks.{i}.mlp.proj.weight"] = tensor[i]
1390
+ else:
1391
+ out[name] = tensor
1392
+ return out
1393
+
1394
+ def _rebank_state_dict(sd: dict[str, Tensor], num_layers: int, template_sd: dict[str, Tensor]) -> dict[str, Tensor]:
1395
+ """Convert individual 2D tensors back into 3D bank tensors."""
1396
+ out: dict[str, Tensor] = {}
1397
+ n = num_layers
1398
+ # Reconstruct banks from individual weight keys
1399
+ qo_slices = [None] * (2 * n)
1400
+ kv_slices = [None] * (2 * n)
1401
+ up_slices = [None] * n
1402
+ down_slices = [None] * n
1403
+ consumed = set()
1404
+ for i in range(n):
1405
+ qk = f"blocks.{i}.attn.c_q.weight"
1406
+ if qk in sd:
1407
+ qo_slices[i] = sd[qk]
1408
+ consumed.add(qk)
1409
+ ok = f"blocks.{i}.attn.proj.weight"
1410
+ if ok in sd:
1411
+ qo_slices[n + i] = sd[ok]
1412
+ consumed.add(ok)
1413
+ kk = f"blocks.{i}.attn.c_k.weight"
1414
+ if kk in sd:
1415
+ kv_slices[i] = sd[kk]
1416
+ consumed.add(kk)
1417
+ vk = f"blocks.{i}.attn.c_v.weight"
1418
+ if vk in sd:
1419
+ kv_slices[n + i] = sd[vk]
1420
+ consumed.add(vk)
1421
+ fk = f"blocks.{i}.mlp.fc.weight"
1422
+ if fk in sd:
1423
+ up_slices[i] = sd[fk]
1424
+ consumed.add(fk)
1425
+ dk = f"blocks.{i}.mlp.proj.weight"
1426
+ if dk in sd:
1427
+ down_slices[i] = sd[dk]
1428
+ consumed.add(dk)
1429
+ out["qo_bank"] = torch.stack(qo_slices).to(dtype=template_sd["qo_bank"].dtype)
1430
+ out["kv_bank"] = torch.stack(kv_slices).to(dtype=template_sd["kv_bank"].dtype)
1431
+ out["mlp_up_bank"] = torch.stack(up_slices).to(dtype=template_sd["mlp_up_bank"].dtype)
1432
+ out["mlp_down_bank"] = torch.stack(down_slices).to(dtype=template_sd["mlp_down_bank"].dtype)
1433
+ for name, tensor in sd.items():
1434
+ if name not in consumed:
1435
+ out[name] = tensor
1436
+ return out
1437
+
1438
+ def mixed_quantize_int6(state_dict: dict[str, Tensor], int6_cats: set[str]):
1439
+ num_layers_total = max(
1440
+ (int(k.split(".")[1]) for k in state_dict if k.startswith("blocks.")),
1441
+ default=0,
1442
+ ) + 1
1443
+ late_k_layers = set(range(num_layers_total - 2, num_layers_total))
1444
+ result: dict[str, Tensor] = {}
1445
+ meta: dict[str, object] = {}
1446
+ for name, tensor in state_dict.items():
1447
+ t = tensor.detach().cpu().contiguous()
1448
+ cat = _classify_param(name)
1449
+ if not t.is_floating_point() or t.numel() <= 65536:
1450
+ result[name] = t.to(torch.float16) if t.is_floating_point() else t
1451
+ meta[name] = "passthrough"
1452
+ continue
1453
+ if any(p in name for p in CONTROL_TENSOR_NAME_PATTERNS):
1454
+ result[name] = t.float()
1455
+ meta[name] = "passthrough_ctrl"
1456
+ continue
1457
+ if cat in int6_cats and t.ndim >= 1:
1458
+ q, s = quantize_int6_per_row(t)
1459
+ result[name + ".q"] = q
1460
+ result[name + ".scale"] = s
1461
+ meta[name] = {"type": "int6"}
1462
+ else:
1463
+ q, s = quantize_float_tensor(t)
1464
+ result[name + ".q"] = q
1465
+ result[name + ".scale"] = s
1466
+ meta[name] = {"type": "int8"}
1467
+ return result, meta
1468
+ def dequantize_mixed_int6(result: dict[str, Tensor], meta: dict[str, object],
1469
+ template_sd: dict[str, Tensor]) -> dict[str, Tensor]:
1470
+ out: dict[str, Tensor] = {}
1471
+ for name, orig in template_sd.items():
1472
+ info = meta.get(name)
1473
+ if info is None:
1474
+ continue
1475
+ orig_dtype = orig.dtype
1476
+ if info in ("passthrough", "passthrough_ctrl", "passthrough_fp16"):
1477
+ t = result[name]
1478
+ if t.dtype == torch.float16 and orig_dtype in (torch.float32, torch.bfloat16):
1479
+ t = t.to(orig_dtype)
1480
+ out[name] = t
1481
+ continue
1482
+ q, s = result[name + ".q"], result[name + ".scale"]
1483
+ if s.ndim > 0:
1484
+ out[name] = (q.float() * s.float().view(q.shape[0], *([1] * (q.ndim - 1)))).to(orig_dtype)
1485
+ else:
1486
+ out[name] = (q.float() * float(s.item())).to(orig_dtype)
1487
+ return out
1488
+
1489
+ # --- Training ---
1490
+
1491
+ def main() -> None:
1492
+ code = Path(__file__).read_text(encoding="utf-8")
1493
+ args = Hyperparameters()
1494
+ # zeropower_via_newtonschulz5 runs eagerly with bmm -- do NOT compile
1495
+ distributed = "RANK" in os.environ and "WORLD_SIZE" in os.environ
1496
+ rank = int(os.environ.get("RANK", "0"))
1497
+ world_size = int(os.environ.get("WORLD_SIZE", "1"))
1498
+ local_rank = int(os.environ.get("LOCAL_RANK", "0"))
1499
+ if world_size <= 0:
1500
+ raise ValueError(f"WORLD_SIZE must be positive, got {world_size}")
1501
+ if 8 % world_size != 0:
1502
+ raise ValueError(f"WORLD_SIZE={world_size} must divide 8 so grad_accum_steps stays integral")
1503
+ grad_accum_steps = 8 // world_size
1504
+ grad_scale = 1.0 / grad_accum_steps
1505
+ if not torch.cuda.is_available():
1506
+ raise RuntimeError("CUDA is required")
1507
+ device = torch.device("cuda", local_rank)
1508
+ torch.cuda.set_device(device)
1509
+ if distributed:
1510
+ dist.init_process_group(backend="nccl", device_id=device)
1511
+ dist.barrier()
1512
+ master_process = rank == 0
1513
+ torch.backends.cuda.matmul.allow_tf32 = True
1514
+ torch.backends.cudnn.allow_tf32 = True
1515
+ from torch.backends.cuda import enable_cudnn_sdp, enable_flash_sdp, enable_math_sdp, enable_mem_efficient_sdp
1516
+ enable_cudnn_sdp(False)
1517
+ enable_flash_sdp(True)
1518
+ enable_mem_efficient_sdp(False)
1519
+ enable_math_sdp(False)
1520
+ logfile = None
1521
+ if master_process:
1522
+ os.makedirs("logs", exist_ok=True)
1523
+ logfile = f"logs/{args.run_id}.txt"
1524
+ print(logfile)
1525
+ def log0(msg: str, console: bool = True) -> None:
1526
+ if not master_process:
1527
+ return
1528
+ if console:
1529
+ print(msg)
1530
+ if logfile is not None:
1531
+ with open(logfile, "a", encoding="utf-8") as f:
1532
+ print(msg, file=f)
1533
+ log0(code, console=False)
1534
+ log0("=" * 100, console=False)
1535
+ log0(f"Running Python {sys.version}", console=False)
1536
+ log0(f"Running PyTorch {torch.__version__}", console=False)
1537
+ log0(
1538
+ subprocess.run(["nvidia-smi"], stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, check=False).stdout,
1539
+ console=False,
1540
+ )
1541
+ log0("=" * 100, console=False)
1542
+ random.seed(args.seed)
1543
+ np.random.seed(args.seed)
1544
+ torch.manual_seed(args.seed)
1545
+ torch.cuda.manual_seed_all(args.seed)
1546
+ if not args.tokenizer_path.endswith(".model"):
1547
+ raise ValueError(f"Script only setup for SentencePiece .model file: {args.tokenizer_path}")
1548
+ sp = spm.SentencePieceProcessor(model_file=args.tokenizer_path)
1549
+ if int(sp.vocab_size()) != args.vocab_size:
1550
+ raise ValueError(
1551
+ f"VOCAB_SIZE={args.vocab_size} does not match tokenizer vocab_size={int(sp.vocab_size())}"
1552
+ )
1553
+ dataset_dir = Path(args.data_path).resolve()
1554
+ actual_train_files = len(list(dataset_dir.glob("fineweb_train_*.bin")))
1555
+ effective_eval_seq_len = args.eval_seq_len if args.eval_seq_len > 0 else args.train_seq_len
1556
+ val_seq_len = max(args.train_seq_len, effective_eval_seq_len)
1557
+ val_tokens = load_validation_tokens(args.val_files, val_seq_len)
1558
+ base_bytes_lut, has_leading_space_lut, is_boundary_token_lut = build_sentencepiece_luts(
1559
+ sp, args.vocab_size, device
1560
+ )
1561
+ log0(f"val_bpb:enabled tokenizer_kind=sentencepiece tokenizer_path={args.tokenizer_path}")
1562
+ log0(f"train_loader:dataset:{dataset_dir.name} train_shards:{actual_train_files}")
1563
+ log0(f"val_loader:shards pattern={args.val_files} tokens:{val_tokens.numel() - 1}")
1564
+ CastedLinear._qat_enabled = args.qat_enabled
1565
+ base_model = GPT(
1566
+ vocab_size=args.vocab_size,
1567
+ num_layers=args.num_layers,
1568
+ model_dim=args.model_dim,
1569
+ num_heads=args.num_heads,
1570
+ num_kv_heads=args.num_kv_heads,
1571
+ mlp_mult=args.mlp_mult,
1572
+ tie_embeddings=args.tie_embeddings,
1573
+ tied_embed_init_std=args.tied_embed_init_std,
1574
+ logit_softcap=args.logit_softcap,
1575
+ rope_base=args.rope_base,
1576
+ qk_gain_init=args.qk_gain_init,
1577
+ mtp_num_heads=args.mtp_num_heads,
1578
+ mtp_loss_weight=args.mtp_loss_weight,
1579
+ bigram_vocab_size=args.bigram_vocab_size,
1580
+ bigram_dim=args.bigram_dim,
1581
+ xsa_last_n=args.xsa_last_n,
1582
+ rope_dims=args.rope_dims,
1583
+ ln_scale=args.ln_scale,
1584
+ dtg=args.dtg_enabled,
1585
+ ve_enabled=args.ve_enabled,
1586
+ ve_dim=args.ve_dim,
1587
+ ve_layers=args.ve_layers,
1588
+ gated_attention=args.gated_attention,
1589
+ value_residual=args.value_residual,
1590
+ ).to(device).bfloat16()
1591
+ # Banks stay FP32 (like CastedLinear weights), cast to BF16 in forward
1592
+ base_model.qo_bank.data = base_model.qo_bank.data.float()
1593
+ base_model.kv_bank.data = base_model.kv_bank.data.float()
1594
+ base_model.mlp_up_bank.data = base_model.mlp_up_bank.data.float()
1595
+ base_model.mlp_down_bank.data = base_model.mlp_down_bank.data.float()
1596
+ for module in base_model.modules():
1597
+ if isinstance(module, CastedLinear):
1598
+ module.float()
1599
+ restore_low_dim_params_to_fp32(base_model)
1600
+ # No DDP -- Parallel Muon handles bank grad communication via reduce-scatter,
1601
+ # and non-bank grads are manually all-reduced before Adam steps.
1602
+ compiled_model = torch.compile(base_model, dynamic=False, fullgraph=True)
1603
+ model = compiled_model
1604
+
1605
+ # Optimizer split:
1606
+ # - 4 parameter banks -> Muon (batched Newton-Schulz)
1607
+ # - token embedding -> Adam
1608
+ # - scalars/control tensors -> Adam
1609
+ # - bigram proj, mtp heads, VE proj -> Adam (small matrix params not worth banking)
1610
+ matrix_params = [
1611
+ base_model.qo_bank, base_model.kv_bank,
1612
+ base_model.mlp_up_bank, base_model.mlp_down_bank,
1613
+ ]
1614
+ block_named_params = list(base_model.blocks.named_parameters())
1615
+ scalar_params = [
1616
+ p
1617
+ for name, p in block_named_params
1618
+ if p.ndim < 2 or any(pattern in name for pattern in CONTROL_TENSOR_NAME_PATTERNS)
1619
+ ]
1620
+ if base_model.skip_weights.numel() > 0:
1621
+ scalar_params.append(base_model.skip_weights)
1622
+ scalar_params.append(base_model.smear.gate)
1623
+ if base_model.bigram is not None:
1624
+ scalar_params.append(base_model.bigram.scale)
1625
+ token_lr = args.tied_embed_lr if args.tie_embeddings else args.embed_lr
1626
+ tok_params = [{"params": [base_model.tok_emb.weight], "lr": token_lr, "base_lr": token_lr}]
1627
+ if base_model.bigram is not None:
1628
+ tok_params.append({"params": [base_model.bigram.embed.weight], "lr": token_lr, "base_lr": token_lr})
1629
+ if base_model.bigram.proj is not None:
1630
+ scalar_params.append(base_model.bigram.proj.weight)
1631
+ if base_model.ve_shared is not None:
1632
+ tok_params.append({"params": [base_model.ve_shared.embed.weight], "lr": token_lr, "base_lr": token_lr})
1633
+ if base_model.ve_shared.proj is not None:
1634
+ scalar_params.append(base_model.ve_shared.proj.weight)
1635
+ scalar_params.append(base_model.ve_shared.scale)
1636
+ for s in base_model.ve_layer_scales:
1637
+ scalar_params.append(s)
1638
+ optimizer_tok = torch.optim.AdamW(
1639
+ tok_params,
1640
+ betas=(args.beta1, args.beta2),
1641
+ eps=args.adam_eps,
1642
+ weight_decay=args.adam_wd,
1643
+ fused=True,
1644
+ )
1645
+ optimizer_muon = Muon(
1646
+ matrix_params,
1647
+ lr=args.matrix_lr,
1648
+ momentum=args.muon_momentum,
1649
+ backend_steps=args.muon_backend_steps,
1650
+ weight_decay=args.muon_wd,
1651
+ )
1652
+ for group in optimizer_muon.param_groups:
1653
+ group["base_lr"] = args.matrix_lr
1654
+ optimizer_scalar = torch.optim.AdamW(
1655
+ [{"params": scalar_params, "lr": args.scalar_lr, "base_lr": args.scalar_lr}],
1656
+ betas=(args.beta1, args.beta2),
1657
+ eps=args.adam_eps,
1658
+ weight_decay=args.adam_wd,
1659
+ fused=True,
1660
+ )
1661
+ # Non-bank params that need manual all-reduce (replicated across GPUs)
1662
+ replicated_params = list(optimizer_tok.param_groups[0]["params"])
1663
+ for pg in optimizer_tok.param_groups[1:]:
1664
+ replicated_params.extend(pg["params"])
1665
+ replicated_params.extend(scalar_params)
1666
+
1667
+ optimizer_head = None
1668
+ if base_model.lm_head is not None:
1669
+ optimizer_head = torch.optim.Adam(
1670
+ [{"params": [base_model.lm_head.weight], "lr": args.head_lr, "base_lr": args.head_lr}],
1671
+ betas=(args.beta1, args.beta2),
1672
+ eps=args.adam_eps,
1673
+ fused=True,
1674
+ )
1675
+ replicated_params.append(base_model.lm_head.weight)
1676
+ optimizers: list[torch.optim.Optimizer] = [optimizer_tok, optimizer_muon, optimizer_scalar]
1677
+ if optimizer_head is not None:
1678
+ optimizers.append(optimizer_head)
1679
+ n_params = sum(p.numel() for p in base_model.parameters())
1680
+ mtp_params = sum(p.numel() for p in base_model.mtp_heads.parameters())
1681
+ log0(f"model_params:{n_params}")
1682
+ log0(f"mtp_num_heads:{args.mtp_num_heads} mtp_loss_weight:{args.mtp_loss_weight} mtp_params:{mtp_params}")
1683
+ xsa_layers = [i for i, b in enumerate(base_model.blocks) if b.attn.use_xsa]
1684
+ log0(f"XSA:last_{args.xsa_last_n} active_layers:{xsa_layers}")
1685
+ log0(f"world_size:{world_size} grad_accum_steps:{grad_accum_steps}")
1686
+ log0("sdp_backends:cudnn=False flash=True mem_efficient=False math=False")
1687
+ log0(f"attention_mode:gqa num_heads:{args.num_heads} num_kv_heads:{args.num_kv_heads}")
1688
+ log0(
1689
+ f"tie_embeddings:{args.tie_embeddings} embed_lr:{token_lr} "
1690
+ f"head_lr:{args.head_lr if base_model.lm_head is not None else 0.0} "
1691
+ f"matrix_lr:{args.matrix_lr} scalar_lr:{args.scalar_lr}"
1692
+ )
1693
+ log0(
1694
+ f"train_batch_tokens:{args.train_batch_tokens} train_seq_len:{args.train_seq_len} "
1695
+ f"iterations:{args.iterations} warmup_steps:{args.warmup_steps} "
1696
+ f"max_wallclock_seconds:{args.max_wallclock_seconds:.3f}"
1697
+ )
1698
+ log0(f"seed:{args.seed}")
1699
+ train_loader = DistributedTokenLoader(args.train_files, rank, world_size, device)
1700
+ def zero_grad_all() -> None:
1701
+ for opt in optimizers:
1702
+ opt.zero_grad(set_to_none=True)
1703
+ max_wallclock_ms = 1000.0 * args.max_wallclock_seconds if args.max_wallclock_seconds > 0 else None
1704
+ def lr_mul(step: int, elapsed_ms: float) -> float:
1705
+ if args.warmdown_iters <= 0:
1706
+ return 1.0
1707
+ if max_wallclock_ms is None:
1708
+ warmdown_start = max(args.iterations - args.warmdown_iters, 0)
1709
+ return max((args.iterations - step) / max(args.warmdown_iters, 1), 0.0) if warmdown_start <= step < args.iterations else 1.0
1710
+ step_ms = elapsed_ms / max(step, 1)
1711
+ warmdown_ms = args.warmdown_iters * step_ms
1712
+ remaining_ms = max(max_wallclock_ms - elapsed_ms, 0.0)
1713
+ return remaining_ms / max(warmdown_ms, 1e-9) if remaining_ms <= warmdown_ms else 1.0
1714
+ if args.warmup_steps > 0:
1715
+ initial_model_state = {name: tensor.detach().cpu().clone() for name, tensor in base_model.state_dict().items()}
1716
+ initial_optimizer_states = [copy.deepcopy(opt.state_dict()) for opt in optimizers]
1717
+ model.train()
1718
+ for warmup_step in range(args.warmup_steps):
1719
+ zero_grad_all()
1720
+ for micro_step in range(grad_accum_steps):
1721
+ x, y = train_loader.next_batch(args.train_batch_tokens, args.train_seq_len, grad_accum_steps)
1722
+ with torch.autocast(device_type="cuda", dtype=torch.bfloat16, enabled=True):
1723
+ warmup_loss = model(x, y)
1724
+ (warmup_loss * grad_scale).backward()
1725
+ # All-reduce all grads for warmup (simple, not optimized)
1726
+ if distributed:
1727
+ for p in base_model.parameters():
1728
+ if p.grad is not None:
1729
+ dist.all_reduce(p.grad, op=dist.ReduceOp.AVG)
1730
+ for opt in optimizers:
1731
+ opt.step()
1732
+ zero_grad_all()
1733
+ if args.warmup_steps <= 20 or (warmup_step + 1) % 10 == 0 or warmup_step + 1 == args.warmup_steps:
1734
+ log0(f"warmup_step:{warmup_step + 1}/{args.warmup_steps}")
1735
+ base_model.load_state_dict(initial_model_state, strict=True)
1736
+ for opt, state in zip(optimizers, initial_optimizer_states, strict=True):
1737
+ opt.load_state_dict(state)
1738
+ zero_grad_all()
1739
+ train_loader = DistributedTokenLoader(args.train_files, rank, world_size, device)
1740
+ swa_state: dict[str, Tensor] | None = None
1741
+ swa_count = 0
1742
+ from collections import deque
1743
+ lawa_queue: deque[dict[str, Tensor]] = deque(maxlen=args.lawa_k)
1744
+ ema_state = {name: t.detach().float().clone() for name, t in base_model.state_dict().items()}
1745
+ ema_decay = 0.997
1746
+ training_time_ms = 0.0
1747
+ stop_after_step: int | None = None
1748
+ torch.cuda.synchronize()
1749
+ t0 = time.perf_counter()
1750
+ step = 0
1751
+ while True:
1752
+ last_step = step == args.iterations or (stop_after_step is not None and step >= stop_after_step)
1753
+ should_validate = last_step or (args.val_loss_every > 0 and step % args.val_loss_every == 0)
1754
+ if should_validate:
1755
+ torch.cuda.synchronize()
1756
+ training_time_ms += 1000.0 * (time.perf_counter() - t0)
1757
+ val_loss, val_bpb = eval_val(
1758
+ args,
1759
+ model,
1760
+ rank,
1761
+ world_size,
1762
+ device,
1763
+ grad_accum_steps,
1764
+ val_tokens,
1765
+ base_bytes_lut,
1766
+ has_leading_space_lut,
1767
+ is_boundary_token_lut,
1768
+ )
1769
+ log0(
1770
+ f"step:{step}/{args.iterations} val_loss:{val_loss:.4f} val_bpb:{val_bpb:.4f} "
1771
+ f"train_time:{training_time_ms:.0f}ms step_avg:{training_time_ms / max(step, 1):.2f}ms"
1772
+ )
1773
+ torch.cuda.synchronize()
1774
+ t0 = time.perf_counter()
1775
+ if last_step:
1776
+ if stop_after_step is not None and step < args.iterations:
1777
+ log0(
1778
+ f"stopping_early: wallclock_cap train_time:{training_time_ms:.0f}ms "
1779
+ f"step:{step}/{args.iterations}"
1780
+ )
1781
+ break
1782
+ elapsed_ms = training_time_ms + 1000.0 * (time.perf_counter() - t0)
1783
+ scale = lr_mul(step, elapsed_ms)
1784
+ if args.late_qat_threshold > 0 and scale < args.late_qat_threshold and not CastedLinear._qat_enabled:
1785
+ CastedLinear._qat_enabled = True
1786
+ log0(f"late_qat:enabled step:{step} scale:{scale:.4f}")
1787
+ zero_grad_all()
1788
+ train_loss = torch.zeros((), device=device)
1789
+ for micro_step in range(grad_accum_steps):
1790
+ x, y = train_loader.next_batch(args.train_batch_tokens, args.train_seq_len, grad_accum_steps)
1791
+ with torch.autocast(device_type="cuda", dtype=torch.bfloat16, enabled=True):
1792
+ loss = model(x, y)
1793
+ train_loss += loss.detach()
1794
+ (loss * grad_scale).backward()
1795
+ train_loss /= grad_accum_steps
1796
+ frac = min(step / args.muon_momentum_warmup_steps, 1.0) if args.muon_momentum_warmup_steps > 0 else 1.0
1797
+ muon_momentum = (1 - frac) * args.muon_momentum_warmup_start + frac * args.muon_momentum
1798
+ for group in optimizer_muon.param_groups:
1799
+ group["momentum"] = muon_momentum
1800
+ for opt in optimizers:
1801
+ for group in opt.param_groups:
1802
+ group["lr"] = group["base_lr"] * scale
1803
+ if args.grad_clip_norm > 0:
1804
+ torch.nn.utils.clip_grad_norm_(base_model.parameters(), args.grad_clip_norm)
1805
+ # === 3-phase overlapped optimizer step ===
1806
+ # Phase 1: Launch async reduce-scatter for banks (biggest first)
1807
+ optimizer_muon.launch_reduce_scatters()
1808
+ # Phase 2: All-reduce non-bank grads + step Adam (while bank RS is in-flight)
1809
+ if distributed:
1810
+ for p in replicated_params:
1811
+ if p.grad is not None:
1812
+ dist.all_reduce(p.grad, op=dist.ReduceOp.AVG)
1813
+ optimizer_tok.step()
1814
+ optimizer_scalar.step()
1815
+ if optimizer_head is not None:
1816
+ optimizer_head.step()
1817
+ # Phase 3: Wait for RS, local NS5, all-gather (banks processed last)
1818
+ optimizer_muon.step()
1819
+ zero_grad_all()
1820
+ # EMA update
1821
+ with torch.no_grad():
1822
+ for name, t in base_model.state_dict().items():
1823
+ ema_state[name].mul_(ema_decay).add_(t.detach().float(), alpha=1.0 - ema_decay)
1824
+ step += 1
1825
+ approx_training_time_ms = training_time_ms + 1000.0 * (time.perf_counter() - t0)
1826
+ if args.swa_enabled and scale < 0.2 and step % args.swa_every == 0:
1827
+ if swa_state is None:
1828
+ swa_state = {name: t.detach().cpu().clone() for name, t in base_model.state_dict().items()}
1829
+ swa_count = 1
1830
+ log0(f"swa:start step:{step}")
1831
+ else:
1832
+ for name, t in base_model.state_dict().items():
1833
+ swa_state[name] += t.detach().cpu()
1834
+ swa_count += 1
1835
+ if args.lawa_enabled and step % args.lawa_freq == 0:
1836
+ lawa_queue.append({name: t.detach().cpu().clone() for name, t in base_model.state_dict().items()})
1837
+ should_log_train = (
1838
+ args.train_log_every > 0
1839
+ and (step <= 10 or step % args.train_log_every == 0 or stop_after_step is not None)
1840
+ )
1841
+ if should_log_train:
1842
+ log0(
1843
+ f"step:{step}/{args.iterations} train_loss:{train_loss.item():.4f} "
1844
+ f"train_time:{approx_training_time_ms:.0f}ms step_avg:{approx_training_time_ms / step:.2f}ms"
1845
+ )
1846
+ reached_cap = max_wallclock_ms is not None and approx_training_time_ms >= max_wallclock_ms
1847
+ if distributed and max_wallclock_ms is not None:
1848
+ reached_cap_tensor = torch.tensor(int(reached_cap), device=device)
1849
+ dist.all_reduce(reached_cap_tensor, op=dist.ReduceOp.MAX)
1850
+ reached_cap = bool(reached_cap_tensor.item())
1851
+ if stop_after_step is None and reached_cap:
1852
+ stop_after_step = step
1853
+ log0(
1854
+ f"peak memory allocated: {torch.cuda.max_memory_allocated() // 1024 // 1024} MiB "
1855
+ f"reserved: {torch.cuda.max_memory_reserved() // 1024 // 1024} MiB"
1856
+ )
1857
+ # Apply weight averaging
1858
+ if args.lawa_enabled and len(lawa_queue) > 1:
1859
+ log0(f"lawa:applying LAWA averaging k={len(lawa_queue)}")
1860
+ current_state = base_model.state_dict()
1861
+ avg_state = {name: torch.zeros(t.shape, dtype=torch.float32, device='cpu') for name, t in current_state.items()}
1862
+ for snap in lawa_queue:
1863
+ for name in avg_state:
1864
+ avg_state[name] += snap[name].float()
1865
+ for name in avg_state:
1866
+ avg_state[name] /= len(lawa_queue)
1867
+ avg_state[name] = avg_state[name].to(dtype=current_state[name].dtype)
1868
+ base_model.load_state_dict(avg_state, strict=True)
1869
+ else:
1870
+ log0("ema:applying EMA weights")
1871
+ current_state = base_model.state_dict()
1872
+ avg_state = {name: t.to(dtype=current_state[name].dtype) for name, t in ema_state.items()}
1873
+ base_model.load_state_dict(avg_state, strict=True)
1874
+ torch.cuda.synchronize()
1875
+ t_diag = time.perf_counter()
1876
+ diag_val_loss, diag_val_bpb = eval_val(
1877
+ args, compiled_model, rank, world_size, device, grad_accum_steps,
1878
+ val_tokens, base_bytes_lut, has_leading_space_lut, is_boundary_token_lut,
1879
+ )
1880
+ torch.cuda.synchronize()
1881
+ log0(
1882
+ f"DIAGNOSTIC post_ema val_loss:{diag_val_loss:.4f} val_bpb:{diag_val_bpb:.4f} "
1883
+ f"eval_time:{1000.0 * (time.perf_counter() - t_diag):.0f}ms"
1884
+ )
1885
+ full_state_dict = base_model.state_dict()
1886
+ export_sd = {k: v for k, v in full_state_dict.items() if "mtp_heads" not in k}
1887
+ excluded_mtp = sum(int(t.numel()) for k, t in full_state_dict.items() if "mtp_heads" in k)
1888
+ if excluded_mtp > 0:
1889
+ log0(f"export_excluding_mtp_params:{excluded_mtp}")
1890
+ if master_process:
1891
+ torch.save(export_sd, "final_model.pt")
1892
+ model_bytes = os.path.getsize("final_model.pt")
1893
+ code_bytes = len(code.encode("utf-8"))
1894
+ log0(f"Serialized model: {model_bytes} bytes")
1895
+ log0(f"Code size: {code_bytes} bytes")
1896
+ # Unbank 3D tensors into individual 2D tensors for quantization
1897
+ sd_cpu = {k: v.detach().cpu() for k, v in export_sd.items()}
1898
+ unbanked_sd = _unbank_state_dict(sd_cpu, args.num_layers)
1899
+ quant_result, quant_meta = mixed_quantize_int6(unbanked_sd, {"mlp", "attn"})
1900
+ quant_buf = io.BytesIO()
1901
+ torch.save({"w": quant_result, "m": quant_meta}, quant_buf)
1902
+ quant_raw = quant_buf.getvalue()
1903
+ quant_blob = lzma.compress(quant_raw, preset=6)
1904
+ if master_process:
1905
+ with open("final_model.int6.ptz", "wb") as f:
1906
+ f.write(quant_blob)
1907
+ quant_file_bytes = len(quant_blob)
1908
+ code_bytes = len(code.encode("utf-8"))
1909
+ log0(f"Serialized model int6+lzma: {quant_file_bytes} bytes")
1910
+ log0(f"Total submission size int6+lzma: {quant_file_bytes + code_bytes} bytes")
1911
+ if distributed:
1912
+ dist.barrier()
1913
+ with open("final_model.int6.ptz", "rb") as f:
1914
+ quant_blob_disk = f.read()
1915
+ quant_state = torch.load(
1916
+ io.BytesIO(lzma.decompress(quant_blob_disk)),
1917
+ map_location="cpu",
1918
+ )
1919
+ deq_unbanked = dequantize_mixed_int6(quant_state["w"], quant_state["m"], unbanked_sd)
1920
+ # Re-bank the dequantized tensors
1921
+ deq_state = _rebank_state_dict(deq_unbanked, args.num_layers, sd_cpu)
1922
+ eval_model = GPT(
1923
+ vocab_size=args.vocab_size, num_layers=args.num_layers, model_dim=args.model_dim,
1924
+ num_heads=args.num_heads, num_kv_heads=args.num_kv_heads, mlp_mult=args.mlp_mult,
1925
+ tie_embeddings=args.tie_embeddings, tied_embed_init_std=args.tied_embed_init_std,
1926
+ logit_softcap=args.logit_softcap, rope_base=args.rope_base, qk_gain_init=args.qk_gain_init,
1927
+ mtp_num_heads=0, mtp_loss_weight=0.0,
1928
+ bigram_vocab_size=args.bigram_vocab_size, bigram_dim=args.bigram_dim,
1929
+ xsa_last_n=args.xsa_last_n,
1930
+ rope_dims=args.rope_dims, ln_scale=args.ln_scale, dtg=args.dtg_enabled,
1931
+ ve_enabled=args.ve_enabled, ve_dim=args.ve_dim, ve_layers=args.ve_layers,
1932
+ gated_attention=args.gated_attention, value_residual=args.value_residual,
1933
+ ).to(device).bfloat16()
1934
+ eval_model.qo_bank.data = eval_model.qo_bank.data.float()
1935
+ eval_model.kv_bank.data = eval_model.kv_bank.data.float()
1936
+ eval_model.mlp_up_bank.data = eval_model.mlp_up_bank.data.float()
1937
+ eval_model.mlp_down_bank.data = eval_model.mlp_down_bank.data.float()
1938
+ for m in eval_model.modules():
1939
+ if isinstance(m, CastedLinear):
1940
+ m.float()
1941
+ restore_low_dim_params_to_fp32(eval_model)
1942
+ eval_model.load_state_dict(deq_state, strict=True)
1943
+ compiled_eval = torch.compile(eval_model, dynamic=False, fullgraph=True)
1944
+ torch.cuda.synchronize()
1945
+ t_qeval = time.perf_counter()
1946
+ q_val_loss, q_val_bpb = eval_val(
1947
+ args, compiled_eval, rank, world_size, device, grad_accum_steps,
1948
+ val_tokens, base_bytes_lut, has_leading_space_lut, is_boundary_token_lut,
1949
+ eval_seq_len=effective_eval_seq_len,
1950
+ )
1951
+ torch.cuda.synchronize()
1952
+ log0(
1953
+ f"final_int6_roundtrip val_loss:{q_val_loss:.4f} val_bpb:{q_val_bpb:.4f} "
1954
+ f"eval_time:{1000.0 * (time.perf_counter() - t_qeval):.0f}ms"
1955
+ )
1956
+ log0(f"final_int6_roundtrip_exact val_loss:{q_val_loss:.8f} val_bpb:{q_val_bpb:.8f}")
1957
+ sw_seq_len = effective_eval_seq_len
1958
+ if args.eval_stride > 0 and args.eval_stride < sw_seq_len:
1959
+ torch.cuda.synchronize()
1960
+ t_slide = time.perf_counter()
1961
+ sw_val_loss, sw_val_bpb = eval_val_sliding(
1962
+ args, eval_model, rank, world_size, device,
1963
+ val_tokens, base_bytes_lut, has_leading_space_lut, is_boundary_token_lut,
1964
+ stride=args.eval_stride,
1965
+ eval_seq_len=sw_seq_len,
1966
+ )
1967
+ torch.cuda.synchronize()
1968
+ log0(
1969
+ f"final_int6_sliding_window val_loss:{sw_val_loss:.4f} val_bpb:{sw_val_bpb:.4f} "
1970
+ f"stride:{args.eval_stride} eval_time:{1000.0 * (time.perf_counter() - t_slide):.0f}ms"
1971
+ )
1972
+ log0(f"final_int6_sliding_window_exact val_loss:{sw_val_loss:.8f} val_bpb:{sw_val_bpb:.8f}")
1973
+ log0(f"final_int8_zlib_roundtrip_exact val_loss:{sw_val_loss:.8f} val_bpb:{sw_val_bpb:.8f}")
1974
+ if args.eval_stride != 64 and 64 < sw_seq_len:
1975
+ torch.cuda.synchronize()
1976
+ t_slide64 = time.perf_counter()
1977
+ sw64_val_loss, sw64_val_bpb = eval_val_sliding(
1978
+ args, eval_model, rank, world_size, device,
1979
+ val_tokens, base_bytes_lut, has_leading_space_lut, is_boundary_token_lut,
1980
+ stride=64,
1981
+ eval_seq_len=sw_seq_len,
1982
+ )
1983
+ torch.cuda.synchronize()
1984
+ log0(
1985
+ f"final_int6_sliding_window_s64 val_loss:{sw64_val_loss:.4f} val_bpb:{sw64_val_bpb:.4f} "
1986
+ f"stride:64 eval_time:{1000.0 * (time.perf_counter() - t_slide64):.0f}ms"
1987
+ )
1988
+ log0(f"final_int6_sliding_window_s64_exact val_loss:{sw64_val_loss:.8f} val_bpb:{sw64_val_bpb:.8f}")
1989
+ log0(f"final_int8_zlib_roundtrip_exact val_loss:{sw64_val_loss:.8f} val_bpb:{sw64_val_bpb:.8f}")
1990
+ # Legal score-first TTT (PR #461 recipe)
1991
+ if args.ttt_enabled:
1992
+ torch.cuda.synchronize()
1993
+ t_ttt = time.perf_counter()
1994
+ ttt_loss, ttt_bpb = eval_val_sliding_ttt(
1995
+ args, eval_model, rank, world_size, device,
1996
+ val_tokens, base_bytes_lut, has_leading_space_lut, is_boundary_token_lut,
1997
+ stride=args.eval_stride, log0=log0,
1998
+ )
1999
+ torch.cuda.synchronize()
2000
+ log0(f"legal_ttt val_loss:{ttt_loss:.4f} val_bpb:{ttt_bpb:.4f} "
2001
+ f"eval_time:{1000.0 * (time.perf_counter() - t_ttt):.0f}ms")
2002
+ log0(f"legal_ttt_exact val_loss:{ttt_loss:.8f} val_bpb:{ttt_bpb:.8f}")
2003
+ if distributed:
2004
+ dist.destroy_process_group()
2005
+ if __name__ == "__main__":
2006
+ main()