anpaurehf commited on
Commit
23e91c8
·
verified ·
1 Parent(s): bbbd4d1

Upload train_inverter_v6.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. train_inverter_v6.py +695 -0
train_inverter_v6.py ADDED
@@ -0,0 +1,695 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python
2
+ import argparse
3
+ import json
4
+ import math
5
+ import os
6
+ import time
7
+ from dataclasses import dataclass
8
+ from pathlib import Path
9
+
10
+ import numpy as np
11
+ import torch
12
+ import torch.nn.functional as F
13
+ from transformers import AutoTokenizer
14
+
15
+ from v6_model import EncoderOnlyModel
16
+
17
+
18
+ @dataclass
19
+ class TrainState:
20
+ tokens_seen: int = 0
21
+ example_index: int = 0
22
+ example_token_offset: int = 0
23
+ step: int = 0
24
+
25
+
26
+ def _write_json_atomic(path, payload):
27
+ tmp = f"{path}.tmp"
28
+ with open(tmp, "w") as f:
29
+ json.dump(payload, f, indent=2, sort_keys=True)
30
+ os.replace(tmp, path)
31
+
32
+
33
+ def _torch_save_atomic(path, payload):
34
+ tmp = f"{path}.tmp"
35
+ torch.save(payload, tmp)
36
+ os.replace(tmp, path)
37
+
38
+
39
+ def _unwrap_model(model):
40
+ return model._orig_mod if hasattr(model, "_orig_mod") else model
41
+
42
+
43
+ class ExpertStream:
44
+ def __init__(
45
+ self,
46
+ idx_path,
47
+ tokens_path,
48
+ doc_starts_path,
49
+ dataset_name,
50
+ dataset_revision,
51
+ tokenizer,
52
+ seq_len,
53
+ min_seq_len,
54
+ stride,
55
+ max_tokens,
56
+ batch_size,
57
+ state: TrainState,
58
+ ):
59
+ self.idx_path = idx_path
60
+ self.tokens_path = tokens_path
61
+ self.doc_starts_path = doc_starts_path
62
+ self.dataset_name = dataset_name
63
+ self.dataset_revision = dataset_revision
64
+ self.tokenizer = tokenizer
65
+ self.seq_len = seq_len
66
+ self.min_seq_len = min(min_seq_len, seq_len)
67
+ self.stride = stride if stride > 0 else max(1, seq_len // 2)
68
+ self.max_tokens = max_tokens
69
+ self.batch_size = batch_size
70
+ self.state = state
71
+
72
+ self.idx_mmap = np.load(self.idx_path, mmap_mode="r")
73
+ if self.idx_mmap.shape != (self.max_tokens, 24, 4):
74
+ raise ValueError(f"Unexpected idx shape {self.idx_mmap.shape}.")
75
+ self.tokens_mmap = None
76
+ self.doc_starts = None
77
+ if self.tokens_path:
78
+ self.tokens_mmap = np.load(self.tokens_path, mmap_mode="r")
79
+ if self.tokens_mmap.shape != (self.max_tokens,):
80
+ raise ValueError(f"Unexpected token shape {self.tokens_mmap.shape}.")
81
+ if self.doc_starts_path and os.path.exists(self.doc_starts_path):
82
+ doc_starts_mmap = np.load(self.doc_starts_path, mmap_mode="r")
83
+ if doc_starts_mmap.shape != (self.max_tokens,):
84
+ raise ValueError(
85
+ f"Unexpected doc-start shape {doc_starts_mmap.shape}."
86
+ )
87
+ starts = np.flatnonzero(doc_starts_mmap[: self.max_tokens])
88
+ if len(starts) == 0 or starts[0] != 0:
89
+ starts = np.concatenate(
90
+ [np.array([0], dtype=np.int64), starts[starts > 0]]
91
+ )
92
+ self.doc_starts = np.concatenate(
93
+ [
94
+ starts.astype(np.int64, copy=False),
95
+ np.array([self.max_tokens], dtype=np.int64),
96
+ ]
97
+ )
98
+
99
+ def _yield_batch(self, batch_tokens, batch_idx, batch_mask, state):
100
+ return {
101
+ "input_ids": torch.tensor(np.asarray(batch_tokens, dtype=np.int64)),
102
+ "expert_idx": torch.tensor(np.asarray(batch_idx, dtype=np.int64)),
103
+ "attention_mask": torch.tensor(np.asarray(batch_mask, dtype=np.bool_)),
104
+ "state": state,
105
+ }
106
+
107
+ def _pad_chunk(self, chunk, idx_chunk, current_len):
108
+ pad_len = self.seq_len - current_len
109
+ if pad_len:
110
+ chunk = np.pad(
111
+ chunk,
112
+ (0, pad_len),
113
+ mode="constant",
114
+ constant_values=self.tokenizer.pad_token_id,
115
+ )
116
+ idx_chunk = np.pad(
117
+ idx_chunk,
118
+ ((0, pad_len), (0, 0), (0, 0)),
119
+ mode="constant",
120
+ constant_values=0,
121
+ )
122
+ mask = [1] * current_len + [0] * pad_len
123
+ return chunk, idx_chunk, mask
124
+
125
+ def _pick_window_len(self, doc_start: int, start_pos: int, max_len: int) -> int:
126
+ if max_len <= self.min_seq_len:
127
+ return max_len
128
+ span = max_len - self.min_seq_len + 1
129
+ window_idx = (start_pos - doc_start) // self.stride
130
+ mix = (1103515245 * (window_idx + 1) + 12345 * (doc_start + 1)) & 0x7FFFFFFF
131
+ return self.min_seq_len + (mix % span)
132
+
133
+ def _local_doc_bounds(self, token_pos):
134
+ if self.doc_starts is None:
135
+ return 0, self.max_tokens, 0
136
+ doc_idx = max(0, int(np.searchsorted(self.doc_starts, token_pos, side="right")) - 1)
137
+ doc_start = int(self.doc_starts[doc_idx])
138
+ doc_end = int(self.doc_starts[doc_idx + 1])
139
+ return doc_start, doc_end, doc_idx
140
+
141
+ def _iter_local(self):
142
+ next_start = self.state.tokens_seen
143
+ batch_tokens = []
144
+ batch_idx = []
145
+ batch_mask = []
146
+
147
+ while next_start < self.max_tokens:
148
+ doc_start, doc_end, doc_idx = self._local_doc_bounds(next_start)
149
+ if next_start < doc_start:
150
+ next_start = doc_start
151
+ if next_start >= doc_end:
152
+ next_start = doc_end
153
+ continue
154
+
155
+ current_max_len = min(
156
+ self.seq_len,
157
+ doc_end - next_start,
158
+ self.max_tokens - next_start,
159
+ )
160
+ if current_max_len <= 0:
161
+ break
162
+
163
+ current_len = self._pick_window_len(doc_start, next_start, current_max_len)
164
+ chunk = self.tokens_mmap[next_start : next_start + current_len].astype(
165
+ np.int64, copy=False
166
+ )
167
+ idx_chunk = self.idx_mmap[next_start : next_start + current_len]
168
+ chunk, idx_chunk, mask = self._pad_chunk(chunk, idx_chunk, current_len)
169
+
170
+ new_offset = min(doc_end - doc_start, next_start + self.stride - doc_start)
171
+ state = TrainState(
172
+ tokens_seen=min(doc_end, next_start + self.stride),
173
+ example_index=doc_idx,
174
+ example_token_offset=new_offset,
175
+ step=self.state.step,
176
+ )
177
+ next_start = state.tokens_seen
178
+
179
+ batch_tokens.append(chunk)
180
+ batch_idx.append(idx_chunk)
181
+ batch_mask.append(mask)
182
+
183
+ if len(batch_tokens) >= self.batch_size:
184
+ yield self._yield_batch(batch_tokens, batch_idx, batch_mask, state)
185
+ batch_tokens, batch_idx, batch_mask = [], [], []
186
+
187
+ if batch_tokens:
188
+ yield self._yield_batch(
189
+ batch_tokens,
190
+ batch_idx,
191
+ batch_mask,
192
+ TrainState(
193
+ tokens_seen=next_start,
194
+ example_index=0,
195
+ example_token_offset=0,
196
+ step=self.state.step,
197
+ ),
198
+ )
199
+
200
+ def __iter__(self):
201
+ if self.tokens_mmap is not None:
202
+ yield from self._iter_local()
203
+ return
204
+
205
+ from datasets import load_dataset
206
+
207
+ ds = load_dataset(
208
+ self.dataset_name,
209
+ split="train",
210
+ streaming=True,
211
+ revision=self.dataset_revision,
212
+ )
213
+ tokens_seen = self.state.tokens_seen
214
+ example_index = self.state.example_index
215
+ example_token_offset = self.state.example_token_offset
216
+
217
+ batch_tokens = []
218
+ batch_idx = []
219
+ batch_mask = []
220
+
221
+ for idx, example in enumerate(ds):
222
+ if idx < example_index:
223
+ continue
224
+ if tokens_seen >= self.max_tokens:
225
+ break
226
+
227
+ token_ids = self.tokenizer.encode(example["text"], add_special_tokens=False)
228
+ doc_len = len(token_ids)
229
+ if not token_ids:
230
+ example_index = idx + 1
231
+ example_token_offset = 0
232
+ continue
233
+
234
+ if idx == example_index:
235
+ pos = example_token_offset
236
+ doc_abs_start = tokens_seen - example_token_offset
237
+ else:
238
+ pos = 0
239
+ doc_abs_start = tokens_seen
240
+
241
+ while pos < doc_len and tokens_seen < self.max_tokens:
242
+ abs_pos = doc_abs_start + pos
243
+ current_max_len = min(
244
+ self.seq_len,
245
+ doc_len - pos,
246
+ self.max_tokens - abs_pos,
247
+ )
248
+ if current_max_len <= 0:
249
+ break
250
+
251
+ current_len = self._pick_window_len(doc_abs_start, abs_pos, current_max_len)
252
+ chunk = np.asarray(token_ids[pos : pos + current_len], dtype=np.int64)
253
+ idx_chunk = self.idx_mmap[abs_pos : abs_pos + current_len]
254
+ chunk, idx_chunk, mask = self._pad_chunk(chunk, idx_chunk, current_len)
255
+
256
+ next_pos = min(doc_len, pos + self.stride)
257
+ next_abs = doc_abs_start + next_pos
258
+ state = TrainState(
259
+ tokens_seen=next_abs,
260
+ example_index=idx,
261
+ example_token_offset=next_pos,
262
+ step=self.state.step,
263
+ )
264
+
265
+ batch_tokens.append(chunk)
266
+ batch_idx.append(idx_chunk)
267
+ batch_mask.append(mask)
268
+
269
+ pos = next_pos
270
+ tokens_seen = next_abs
271
+ example_token_offset = next_pos
272
+
273
+ if len(batch_tokens) >= self.batch_size:
274
+ yield self._yield_batch(batch_tokens, batch_idx, batch_mask, state)
275
+ batch_tokens, batch_idx, batch_mask = [], [], []
276
+
277
+ example_index = idx + 1
278
+ example_token_offset = 0
279
+
280
+ if batch_tokens:
281
+ yield self._yield_batch(
282
+ batch_tokens,
283
+ batch_idx,
284
+ batch_mask,
285
+ TrainState(
286
+ tokens_seen=tokens_seen,
287
+ example_index=example_index,
288
+ example_token_offset=example_token_offset,
289
+ step=self.state.step,
290
+ ),
291
+ )
292
+
293
+
294
+ def _derive_companion_path(idx_path: str, suffix: str) -> str | None:
295
+ path = Path(idx_path)
296
+ if path.name.endswith("_idx.npy"):
297
+ candidate = path.with_name(f"{path.name[:-len('_idx.npy')]}_{suffix}.npy")
298
+ if candidate.exists():
299
+ return str(candidate)
300
+ return None
301
+
302
+
303
+ def _load_train_state(args, checkpoint_payload):
304
+ if checkpoint_payload and "train_state" in checkpoint_payload:
305
+ return TrainState(**checkpoint_payload["train_state"])
306
+ if args.resume and os.path.exists(args.state_path):
307
+ with open(args.state_path, "r") as f:
308
+ payload = json.load(f)
309
+ return TrainState(
310
+ tokens_seen=payload.get("tokens_seen", 0),
311
+ example_index=payload.get("example_index", 0),
312
+ example_token_offset=payload.get("example_token_offset", 0),
313
+ step=payload.get("step", 0),
314
+ )
315
+ return TrainState()
316
+
317
+
318
+ def _load_checkpoint(path, model, optimizers, schedulers, scaler, device):
319
+ if not os.path.exists(path):
320
+ return None
321
+ payload = torch.load(path, map_location=device)
322
+ _unwrap_model(model).load_state_dict(payload["model"])
323
+ optimizer_states = payload.get("optimizers")
324
+ if optimizer_states is not None:
325
+ for opt, opt_state in zip(optimizers, optimizer_states):
326
+ opt.load_state_dict(opt_state)
327
+ scheduler_states = payload.get("schedulers")
328
+ if scheduler_states is not None:
329
+ for sched, sched_state in zip(schedulers, scheduler_states):
330
+ sched.load_state_dict(sched_state)
331
+ scaler_state = payload.get("scaler")
332
+ if scaler_state and scaler.is_enabled():
333
+ scaler.load_state_dict(scaler_state)
334
+ return payload
335
+
336
+
337
+ def _save_checkpoint(path, model, args, state, step, optimizers, schedulers, scaler):
338
+ payload = {
339
+ "model": _unwrap_model(model).state_dict(),
340
+ "optimizers": [opt.state_dict() for opt in optimizers],
341
+ "schedulers": [sched.state_dict() for sched in schedulers],
342
+ "scaler": scaler.state_dict() if scaler.is_enabled() else None,
343
+ "config": vars(args),
344
+ "train_state": state.__dict__,
345
+ "step": step,
346
+ }
347
+ _torch_save_atomic(path, payload)
348
+
349
+
350
+ def split_muon_params(model):
351
+ muon_params = []
352
+ adam_params = []
353
+ for name, p in model.named_parameters():
354
+ if not p.requires_grad:
355
+ continue
356
+ is_matrix = p.ndim == 2
357
+ is_embedding_or_head = name.endswith("head.weight") or name.endswith("pos_emb.weight")
358
+ if is_matrix and not is_embedding_or_head:
359
+ muon_params.append(p)
360
+ else:
361
+ adam_params.append(p)
362
+ return muon_params, adam_params
363
+
364
+
365
+ def make_trapezoidal_lr(step_idx, max_steps, warmup_ratio, warmdown_ratio):
366
+ warmup_steps = max(1, int(warmup_ratio * max_steps)) if warmup_ratio > 0 else 0
367
+ warmdown_steps = max(1, int(warmdown_ratio * max_steps)) if warmdown_ratio > 0 else 0
368
+ if warmup_steps > 0 and step_idx < warmup_steps:
369
+ return float(step_idx + 1) / float(warmup_steps)
370
+ warmdown_start = max_steps - warmdown_steps
371
+ if step_idx < warmdown_start:
372
+ return 1.0
373
+ if warmdown_steps > 0 and step_idx < max_steps:
374
+ remaining = max_steps - step_idx
375
+ return max(0.0, float(remaining) / float(warmdown_steps))
376
+ return 0.0
377
+
378
+
379
+ def _set_tf32():
380
+ if not torch.cuda.is_available():
381
+ return
382
+ try:
383
+ torch.backends.cuda.matmul.fp32_precision = "tf32"
384
+ torch.backends.cudnn.conv.fp32_precision = "tf32"
385
+ except AttributeError:
386
+ torch.backends.cuda.matmul.allow_tf32 = True
387
+ torch.backends.cudnn.allow_tf32 = True
388
+ torch.set_float32_matmul_precision("high")
389
+
390
+
391
+ def _estimate_schedule_steps(max_tokens: int, batch_size: int, stride: int, grad_accum: int) -> int:
392
+ unique_tokens_per_step = max(1, batch_size * stride * grad_accum)
393
+ return max(1, math.ceil(max_tokens / unique_tokens_per_step))
394
+
395
+
396
+ def main():
397
+ parser = argparse.ArgumentParser(
398
+ description="V6 inverter trainer (ctx256 + RoPE + QK-Norm + aligned schedule)."
399
+ )
400
+ parser.add_argument("--idx", required=True)
401
+ parser.add_argument("--tokens", default=None)
402
+ parser.add_argument("--doc-starts", default=None)
403
+ parser.add_argument("--dataset", default="vietgpt/openwebtext_en")
404
+ parser.add_argument("--dataset-revision", default=None)
405
+ parser.add_argument("--model", default="openai/gpt-oss-20b")
406
+ parser.add_argument("--model-revision", default=None)
407
+ parser.add_argument("--seq-len", type=int, default=256)
408
+ parser.add_argument("--min-seq-len", type=int, default=128)
409
+ parser.add_argument("--stride", type=int, default=128)
410
+ parser.add_argument("--layers", type=int, default=24)
411
+ parser.add_argument("--max-tokens", type=int, default=200000000)
412
+ parser.add_argument("--batch-size", type=int, default=8)
413
+ parser.add_argument("--grad-accum", type=int, default=2)
414
+ parser.add_argument("--steps", type=int, default=50000)
415
+ parser.add_argument("--schedule-steps", type=int, default=0)
416
+ parser.add_argument("--save-every", type=int, default=1000)
417
+ parser.add_argument("--log-every", type=int, default=50)
418
+ parser.add_argument("--out", default="inverter_v6.pt")
419
+ parser.add_argument("--state-path", default="train_state_v6.json")
420
+ parser.add_argument("--resume", action="store_true")
421
+ parser.add_argument("--compile", action="store_true")
422
+ parser.add_argument(
423
+ "--attn-impl",
424
+ choices=["auto", "flash", "mem_efficient", "math"],
425
+ default="auto",
426
+ )
427
+ parser.add_argument("--position-type", choices=["learned", "rope"], default="rope")
428
+ parser.add_argument("--rope-theta", type=float, default=10000.0)
429
+ parser.add_argument("--qk-norm", action="store_true", default=True)
430
+ parser.add_argument("--no-qk-norm", action="store_false", dest="qk_norm")
431
+ parser.add_argument("--qk-norm-eps", type=float, default=1e-5)
432
+ parser.add_argument("--logit-softcap", type=float, default=0.0)
433
+ parser.add_argument("--layer-gating", action="store_true")
434
+ parser.add_argument("--d-model", type=int, default=768)
435
+ parser.add_argument("--n-head", type=int, default=12)
436
+ parser.add_argument("--d-ff", type=int, default=2048)
437
+ parser.add_argument("--n-layer", type=int, default=6)
438
+ parser.add_argument("--layer-hidden", type=int, default=64)
439
+ parser.add_argument("--layer-proj", type=int, default=64)
440
+ parser.add_argument("--dropout", type=float, default=0.1)
441
+ parser.add_argument("--adam-lr", type=float, default=1.75e-4)
442
+ parser.add_argument("--muon-lr-factor", type=float, default=4.0)
443
+ parser.add_argument("--weight-decay", type=float, default=0.1)
444
+ parser.add_argument("--warmup-ratio", type=float, default=0.02)
445
+ parser.add_argument("--warmdown-ratio", type=float, default=0.40)
446
+ parser.add_argument("--label-smoothing", type=float, default=0.05)
447
+ parser.add_argument("--clip-grad-norm", type=float, default=1.0)
448
+ parser.add_argument("--wandb", action="store_true")
449
+ parser.add_argument("--wandb-project", default="expert-inversion")
450
+ parser.add_argument("--wandb-entity", default=None)
451
+ parser.add_argument("--wandb-run-name", default=None)
452
+ args = parser.parse_args()
453
+
454
+ _set_tf32()
455
+ if torch.cuda.is_available() and args.attn_impl != "auto":
456
+ try:
457
+ torch.backends.cuda.enable_flash_sdp(args.attn_impl == "flash")
458
+ torch.backends.cuda.enable_mem_efficient_sdp(
459
+ args.attn_impl == "mem_efficient"
460
+ )
461
+ torch.backends.cuda.enable_math_sdp(args.attn_impl == "math")
462
+ except AttributeError:
463
+ pass
464
+
465
+ wandb_run = None
466
+ if args.wandb:
467
+ import wandb
468
+
469
+ wandb_run = wandb.init(
470
+ project=args.wandb_project,
471
+ entity=args.wandb_entity,
472
+ name=args.wandb_run_name,
473
+ config=vars(args),
474
+ )
475
+
476
+ tokenizer = AutoTokenizer.from_pretrained(
477
+ args.model,
478
+ revision=args.model_revision,
479
+ local_files_only=True,
480
+ )
481
+ if tokenizer.pad_token_id is None:
482
+ tokenizer.pad_token_id = tokenizer.eos_token_id
483
+
484
+ if args.tokens is None:
485
+ args.tokens = _derive_companion_path(args.idx, "tok")
486
+ if args.doc_starts is None:
487
+ args.doc_starts = _derive_companion_path(args.idx, "doc_start")
488
+ if args.tokens:
489
+ print(f"Using local token mmap: {args.tokens}")
490
+ if args.doc_starts:
491
+ print(f"Using local doc-start mmap: {args.doc_starts}")
492
+ else:
493
+ print("Local token mmap not found; falling back to dataset streaming.")
494
+
495
+ schedule_steps = (
496
+ args.schedule_steps
497
+ if args.schedule_steps > 0
498
+ else min(
499
+ args.steps,
500
+ _estimate_schedule_steps(
501
+ max_tokens=args.max_tokens,
502
+ batch_size=args.batch_size,
503
+ stride=args.stride,
504
+ grad_accum=args.grad_accum,
505
+ ),
506
+ )
507
+ )
508
+ unique_tokens_per_step = args.batch_size * args.stride * args.grad_accum
509
+ print(
510
+ f"Schedule horizon: {schedule_steps} steps "
511
+ f"(estimated unique tokens/step: {unique_tokens_per_step})"
512
+ )
513
+
514
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
515
+ model = EncoderOnlyModel(
516
+ vocab_size=len(tokenizer),
517
+ num_experts=32,
518
+ num_layers=args.layers,
519
+ topk=4,
520
+ d_model=args.d_model,
521
+ n_head=args.n_head,
522
+ d_ff=args.d_ff,
523
+ n_layer=args.n_layer,
524
+ dropout=args.dropout,
525
+ max_len=args.seq_len,
526
+ layer_gating=args.layer_gating,
527
+ logit_softcap=args.logit_softcap,
528
+ layer_hidden=args.layer_hidden,
529
+ layer_proj=args.layer_proj,
530
+ position_type=args.position_type,
531
+ rope_theta=args.rope_theta,
532
+ qk_norm=args.qk_norm,
533
+ qk_norm_eps=args.qk_norm_eps,
534
+ ).to(device)
535
+
536
+ muon_params, adam_params = split_muon_params(model)
537
+ if not muon_params:
538
+ raise RuntimeError("No Muon parameters found; check parameter names.")
539
+ if not hasattr(torch.optim, "Muon"):
540
+ raise RuntimeError("torch.optim.Muon not available in this environment.")
541
+
542
+ optimizer_adam = torch.optim.AdamW(
543
+ adam_params,
544
+ lr=args.adam_lr,
545
+ betas=(0.9, 0.95),
546
+ weight_decay=args.weight_decay,
547
+ )
548
+ optimizer_muon = torch.optim.Muon(
549
+ muon_params,
550
+ lr=args.adam_lr * args.muon_lr_factor,
551
+ weight_decay=args.weight_decay,
552
+ momentum=0.95,
553
+ nesterov=True,
554
+ adjust_lr_fn="match_rms_adamw",
555
+ )
556
+ optimizers = [optimizer_adam, optimizer_muon]
557
+
558
+ def lr_lambda(step_idx):
559
+ return make_trapezoidal_lr(
560
+ step_idx, schedule_steps, args.warmup_ratio, args.warmdown_ratio
561
+ )
562
+
563
+ schedulers = [
564
+ torch.optim.lr_scheduler.LambdaLR(optimizer_adam, lr_lambda=lr_lambda),
565
+ torch.optim.lr_scheduler.LambdaLR(optimizer_muon, lr_lambda=lr_lambda),
566
+ ]
567
+
568
+ scaler = torch.amp.GradScaler("cuda", enabled=device.type == "cuda")
569
+ checkpoint_payload = None
570
+ if args.resume:
571
+ checkpoint_payload = _load_checkpoint(
572
+ args.out,
573
+ model,
574
+ optimizers,
575
+ schedulers,
576
+ scaler,
577
+ device,
578
+ )
579
+ if checkpoint_payload is None:
580
+ print(f"Resume requested but checkpoint not found at {args.out}; starting fresh.")
581
+ state = _load_train_state(args, checkpoint_payload)
582
+
583
+ if args.compile and device.type == "cuda":
584
+ model = torch.compile(model, dynamic=False)
585
+
586
+ stream = ExpertStream(
587
+ idx_path=args.idx,
588
+ tokens_path=args.tokens,
589
+ doc_starts_path=args.doc_starts,
590
+ dataset_name=args.dataset,
591
+ dataset_revision=args.dataset_revision,
592
+ tokenizer=tokenizer,
593
+ seq_len=args.seq_len,
594
+ min_seq_len=args.min_seq_len,
595
+ stride=args.stride,
596
+ max_tokens=args.max_tokens,
597
+ batch_size=args.batch_size,
598
+ state=state,
599
+ )
600
+
601
+ model.train()
602
+ step = state.step
603
+ micro_step = 0
604
+ start_time = time.time()
605
+ for opt in optimizers:
606
+ opt.zero_grad(set_to_none=True)
607
+
608
+ for batch in stream:
609
+ if step >= args.steps:
610
+ break
611
+
612
+ micro_step += 1
613
+ input_ids = batch["input_ids"].to(device, non_blocking=True)
614
+ expert_idx = batch["expert_idx"][:, :, : args.layers].to(device, non_blocking=True)
615
+ attention_mask = batch["attention_mask"].to(device, non_blocking=True)
616
+
617
+ labels = input_ids.clone()
618
+ labels[~attention_mask] = -100
619
+
620
+ with torch.autocast(device_type=device.type, dtype=torch.bfloat16):
621
+ logits = model(expert_idx, attention_mask)
622
+ loss = F.cross_entropy(
623
+ logits.view(-1, logits.size(-1)),
624
+ labels.view(-1),
625
+ ignore_index=-100,
626
+ label_smoothing=args.label_smoothing,
627
+ )
628
+ scaled_loss = loss / args.grad_accum
629
+
630
+ scaler.scale(scaled_loss).backward()
631
+
632
+ if micro_step % args.grad_accum != 0:
633
+ continue
634
+
635
+ if args.clip_grad_norm > 0:
636
+ for opt in optimizers:
637
+ scaler.unscale_(opt)
638
+ params = [
639
+ p
640
+ for opt in optimizers
641
+ for group in opt.param_groups
642
+ for p in group["params"]
643
+ if p.grad is not None
644
+ ]
645
+ torch.nn.utils.clip_grad_norm_(params, args.clip_grad_norm)
646
+
647
+ for opt in optimizers:
648
+ scaler.step(opt)
649
+ scaler.update()
650
+ for opt in optimizers:
651
+ opt.zero_grad(set_to_none=True)
652
+ for sched in schedulers:
653
+ sched.step()
654
+
655
+ step += 1
656
+ state = batch["state"]
657
+ state.step = step
658
+ micro_step = 0
659
+
660
+ if step % args.log_every == 0:
661
+ elapsed = time.time() - start_time
662
+ lr_adam = schedulers[0].get_last_lr()[0]
663
+ lr_muon = schedulers[1].get_last_lr()[0]
664
+ loss_value = float(loss.item())
665
+ print(
666
+ f"step {step} loss {loss_value:.4f} lr_adam {lr_adam:.6e} "
667
+ f"lr_muon {lr_muon:.6e}"
668
+ )
669
+ if wandb_run:
670
+ wandb_run.log(
671
+ {
672
+ "train/loss": loss_value,
673
+ "train/lr_adam": lr_adam,
674
+ "train/lr_muon": lr_muon,
675
+ "train/step": step,
676
+ "train/time_elapsed_s": elapsed,
677
+ },
678
+ step=step,
679
+ )
680
+
681
+ if step % args.save_every == 0:
682
+ _save_checkpoint(
683
+ args.out, model, args, state, step, optimizers, schedulers, scaler
684
+ )
685
+ _write_json_atomic(args.state_path, state.__dict__)
686
+
687
+ _save_checkpoint(args.out, model, args, state, step, optimizers, schedulers, scaler)
688
+ _write_json_atomic(args.state_path, state.__dict__)
689
+
690
+ if wandb_run:
691
+ wandb_run.finish()
692
+
693
+
694
+ if __name__ == "__main__":
695
+ main()