thisisisheanesu commited on
Commit
5629a1a
·
verified ·
1 Parent(s): 267b4eb

MORENA release: morena-1.5b-instruct

Browse files
Files changed (7) hide show
  1. README.md +113 -0
  2. SHA256SUMS +5 -0
  3. config.json +44 -0
  4. load_example.py +62 -0
  5. model.safetensors +3 -0
  6. modeling_morena.py +1081 -0
  7. tokenizer.json +0 -0
README.md ADDED
@@ -0,0 +1,113 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: other
3
+ language: [sn, sw, ha, yo, ig, zu, xh, rw, tn, af, nr, pcm, en]
4
+ tags: [african-languages, from-scratch, instruct, research-preview, isheanesu-misi, vambo-ai]
5
+ inference: false
6
+ datasets:
7
+ - HuggingFaceFW/fineweb-edu
8
+ - HuggingFaceFW/fineweb-2
9
+ - codeparrot/codeparrot-clean
10
+ - codeparrot/github-code-clean
11
+ - open-web-math/open-web-math
12
+ - EleutherAI/proof-pile-2
13
+ - castorini/wura
14
+ - allenai/MADLAD-400
15
+ - castorini/afriberta-corpus
16
+ - vamboai/fikira
17
+ - ise-uiuc/Magicoder-OSS-Instruct-75K
18
+ - m-a-p/CodeFeedback-Filtered-Instruction
19
+ - thisisisheanesu/morena-sft-corpus
20
+ base_model: thisisisheanesu/morena-1.5b-base
21
+ ---
22
+
23
+ # MORENA 1.5B instruct
24
+
25
+ MORENA, Sesotho and Setswana for a king, a lord or a chief, is a 1.5B-parameter decoder trained
26
+ from scratch for twelve Latin-script African languages plus English, French and code. This is the
27
+ instruction-tuned model described in the whitepaper *MORENA: Built, not adapted*. **Private research
28
+ preview. Not for user-facing deployment.**
29
+
30
+ ## Headline numbers
31
+
32
+ | | MORENA 1.5B instruct | reference |
33
+ |---|---|---|
34
+ | African bits per byte, mean of 12 (lower is better) | **1.441** | Lugha-Llama-8B 1.423, gemma-3-12b-it 2.159, Llama-3.2-1B 2.498 |
35
+ | Translation, FLORES+ chrF++, English into 5 African languages, 3-shot | **45.8** | NLLB-600M 45.6, NLLB-1.3B 47.2, Lugha-Llama-8B 36.8, 1B-class general models 9 to 14 |
36
+ | Translation, 5 African languages into English | **48.7** | NLLB-600M 55.1 |
37
+ | Tool calling, correct tool / valid JSON (constrained, prefilled) | **98.1% / 100%** | |
38
+ | Retrieval QA, open-book accuracy / grounding, African mean | 0.325 / 0.734 | 0.25 chance |
39
+ | Safety, good of scorable, 3,341 prompts, 13 languages, 11 categories | **90.2%** | |
40
+ | Degenerate output | 1.3% | |
41
+ | Benign requests answered well | 58.9% | |
42
+
43
+ Safety by category (share of all attempts handled well, 24 prompts per category per language):
44
+ self-harm 92.5, child safety 91.2, drugs 86.2, violence 88.9, election 89.2, privacy 93.8, fraud 94.3,
45
+ weapons 93.3, hate 95.4, medication 95.4. Weakest languages: Igbo 75%, Yoruba 82%, Setswana 83%,
46
+ isiXhosa 87%. Every number is a model judging a model; no native speaker has yet rated an answer.
47
+
48
+ ## What it is not good at
49
+
50
+ Retrieval-augmented QA is at chance in African languages even though the model demonstrably reads
51
+ the passage (grounding 0.73). Grounded generation is fully faithful to given facts in about 31% of
52
+ attempts. The model over-refuses: 41% of ordinary benign requests still get a refusal. Tool calling
53
+ is measured with the tool marker prefilled; the model does not reliably decide on its own that a tool
54
+ is needed. Reading comprehension (belebele) 0.302 against 0.250 chance.
55
+
56
+ ## Chat format
57
+
58
+ Single reserved tokens mark turns: `<reserved_0>` opens a user turn and `<reserved_1>` an assistant
59
+ turn (token ids 3 and 4). `load_example.py` in this repo shows a full prompt. Do not use
60
+ `<|user|>`-style strings; they are not in the vocabulary and produce degenerate output.
61
+
62
+ ## Files
63
+
64
+ `model.safetensors` (bf16), `config.json`, `tokenizer.json`, `modeling_morena.py` (reference
65
+ implementation, plain PyTorch, no transformers dependency), `load_example.py`, `SHA256SUMS`.
66
+ A GGUF build for llama.cpp is in `thisisisheanesu/morena-1.5b-instruct-gguf`.
67
+
68
+ ## Training
69
+
70
+ 251.7B tokens of pretraining (45% English, 23% code, 10% native African text, 22% machine-translated
71
+ African text, 7.5% French), 63B tokens of mid-training, 4,500 steps of supervised fine-tuning with
72
+ loss masking and a 500-step safety anneal. Release lineage 12,834 A100 GPU-hours; the research
73
+ programme that produced it about 22,450. Architecture: 28 layers x 2048, GQA 16/4, SwiGLU 6144,
74
+ RoPE theta 500,000, 4,096 context, tied embeddings. Optimiser: Muon for non-embedding weights,
75
+ AdamW for the rest, WSD schedule.
76
+
77
+ ## The MORENA family
78
+
79
+ | model | params | African bpb (all 12, lower is better) | role |
80
+ |---|---|---|---|
81
+ | MORENA 1.5B base | 1.485B | 1.408 | pretrained and mid-trained; fine-tuning starting point |
82
+ | MORENA 1.5B instruct | 1.485B | 1.441 | chat, translation, tool calling; the model described in the whitepaper |
83
+ | MORENA 0.5B mini | 503M | 1.520 | pruned and distilled from the 1.5B base |
84
+ | MORENA 0.5B mini instruct | 503M | 1.540 | chat fine-tune of the mini |
85
+ | MORENA 0.2B nano | 209M | 1.583 | cheap trunk for ASR rescoring, keyboards, normalisation |
86
+
87
+ Every outside model we measured (25 in total, from 125M to 12B, including the 8B African specialist
88
+ Lugha-Llama-8B at 1.423 and gemma-3-12b-it at 2.159) sits behind all five on African bits per byte.
89
+ Twelve languages: Shona, Swahili, Hausa, Yoruba, Igbo, isiZulu, isiXhosa, Kinyarwanda, Setswana,
90
+ Afrikaans, isiNdebele, Naija Pidgin; plus English, French and code. Tokenizer: 65,536 entries trained
91
+ on the target mix, fertility 0.239 tokens per byte on African text against 0.246 on English.
92
+
93
+ ## Author and citation
94
+
95
+ Isheanesu Misi, Vambo AI. Trained on CINECA Leonardo (EuroHPC allocation aih4a_vamboai) with support
96
+ from the AI Hub for Sustainable Development.
97
+
98
+ ```
99
+ @techreport{misi2026morena,
100
+ title = {MORENA: Built, not adapted. A 1.5B-parameter language model trained from scratch for twelve African languages},
101
+ author = {Misi, Isheanesu},
102
+ institution = {Vambo AI},
103
+ year = {2026},
104
+ month = {December}
105
+ }
106
+ ```
107
+
108
+ ## Licence and status
109
+
110
+ Private research preview. Weights are released under a custom licence pending resolution of one
111
+ data-licensing question: 22% of the pretraining corpus derives from an NLLB model under a
112
+ NonCommercial licence. Until that is settled the weights are not for commercial use or public
113
+ redistribution. **Not for user-facing deployment**: see the evaluation notes above.
SHA256SUMS ADDED
@@ -0,0 +1,5 @@
 
 
 
 
 
 
1
+ e46cbea0588ff434cc5ffec0f06d444579b1f55a47e93f11ce43658a04926391 config.json
2
+ b0730bd0f6afea78ccb9151c340f0216eacf6335245d24957edae91316bae285 load_example.py
3
+ 1b2693d823e709a4614f091025448febc19ca0e32aa6a431c28f090e3b486de9 modeling_morena.py
4
+ 2ab0b7f3b88ddb22a7fc91bd674e10481f1c8eefa404595557083edf6bc8029e model.safetensors
5
+ 97e5dc822407e4ecb6bac02b0fb8a883e465344953e2fac3becf7b5aacae4de3 tokenizer.json
config.json ADDED
@@ -0,0 +1,44 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "_comment": "Morena-1.5B target: 28L x 2048d, GQA 16/4 (head_dim 128), SwiGLU 6144, tied 64k embeddings, RoPE 5e5, seq 4096 -> 1.48B params (1.35B non-embedding). GLOBAL BATCH = global_batch_seqs 1024 x 4096 = 4.19M tokens/step, independent of node count: train.py sets grad_accum = round(1024 / (micro_batch * n_gpus)). 64 GPUs: accum 4 (exact). 48 GPUs: 5.33 -> 5 (3.9M). 40 GPUs: 6.4 -> 6 (3.9M). 12 GPUs: 21.3 -> 21 (4.1M). Keep the realized global batch constant across resumes (log line [data] prints it). 250B tokens = ~60k steps. fsdp_shard_size 4 = HSDP: parameters sharded inside each 4-GPU node (NVLink all-gathers), gradients all-reduced across nodes -- per-GPU state for 1.48B fp32 master+grad+Muon momentum+Adam (embed only) ~ 19 GB/4 = 5 GB, fits 64 GB with act_ckpt on. Set fsdp_shard_size 0 for full FSDP across all GPUs (lower memory, more inter-node traffic). LR 1e-3 shared (Moonlight scaling), warmup 2k, WSD decay over the last ~10% -- launch the anneal explicitly with --anneal from any stable checkpoint.",
3
+ "model": {
4
+ "vocab_size": 65536,
5
+ "n_layer": 28,
6
+ "d_model": 2048,
7
+ "n_head": 16,
8
+ "n_kv_head": 4,
9
+ "d_ff": 6144,
10
+ "rope_theta": 500000.0,
11
+ "norm_eps": 1e-05,
12
+ "tie_embeddings": true,
13
+ "init_std": 0.02
14
+ },
15
+ "train": {
16
+ "seq_len": 4096,
17
+ "micro_batch": 4,
18
+ "grad_accum": 4,
19
+ "lr": 0.001,
20
+ "min_lr_ratio": 0.1,
21
+ "weight_decay": 0.1,
22
+ "muon_momentum": 0.95,
23
+ "muon_ns_steps": 5,
24
+ "muon_ns_mode": "roundrobin",
25
+ "adam_betas": [
26
+ 0.9,
27
+ 0.95
28
+ ],
29
+ "adam_eps": 1e-08,
30
+ "grad_clip": 1.0,
31
+ "warmup_steps": 2000,
32
+ "total_steps": 60000,
33
+ "decay_steps": 6000,
34
+ "decay_start": -1,
35
+ "decay_shape": "1-sqrt",
36
+ "attn": "auto",
37
+ "act_ckpt": true,
38
+ "compile": false,
39
+ "seed": 1234,
40
+ "global_batch_seqs": 1024,
41
+ "fsdp_shard_size": 4,
42
+ "optimizer": "muon"
43
+ }
44
+ }
load_example.py ADDED
@@ -0,0 +1,62 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Load the Morena preview weights (morena1p5b-sft16) and sample from them.
2
+
3
+ Morena is not a transformers architecture; it is `Transformer` from modeling_morena.py, and that
4
+ class has NO .generate() and a training-shaped forward:
5
+
6
+ forward(idx, positions, cu_seqlens, max_seqlen, mask)
7
+
8
+ For single-sequence eval pass cu_seqlens=None and mask=None, which is plain causal SDPA. The
9
+ sampler below is the one the eval harness uses (no KV cache; fine for a couple hundred tokens).
10
+
11
+ pip install torch safetensors tokenizers
12
+ python load_example.py
13
+ """
14
+ import json, torch
15
+ from safetensors.torch import load_file
16
+ from tokenizers import Tokenizer
17
+ import modeling_morena as M
18
+
19
+ cfg = json.load(open("config.json"))
20
+ model = M.Transformer(M.ModelConfig(**cfg["model"]), "sdpa") # "sdpa" = plain causal attention
21
+ model.load_state_dict(load_file("model.safetensors"), strict=True)
22
+ model = model.to(torch.bfloat16).cuda().eval()
23
+ tok = Tokenizer.from_file("tokenizer.json")
24
+
25
+ EOS = tok.token_to_id("<eos>")
26
+ # ---------------------------------------------------------------------------
27
+ # CHAT MARKERS -- DO NOT "MODERNISE" OR "CORRECT" THESE.
28
+ #
29
+ # The marker set is a property of the CHECKPOINT, not of the project. This folder ships
30
+ # morena1p5b-sft16, which was fine-tuned on the SINGLE reserved tokens <reserved_0> (user, id 3)
31
+ # and <reserved_1> (assistant, id 4). They are real single tokens, not the literal angle-bracket
32
+ # strings they look like.
33
+ #
34
+ # The legacy multi-token strings <|user|>/<|assistant|> belong to sft3 and earlier. Substituting
35
+ # them here does NOT fail loudly -- it makes a healthy model emit degenerate text, and it silently
36
+ # invalidated an entire safety evaluation on 2026-09-08. If you are looking at this line because
37
+ # these markers "look wrong", they are not: check which checkpoint the folder ships first.
38
+ # ---------------------------------------------------------------------------
39
+ USER, ASSISTANT = "<reserved_0>", "<reserved_1>"
40
+
41
+
42
+ @torch.no_grad()
43
+ def generate(prompt, max_new_tokens=100, temperature=0.7, top_p=0.9):
44
+ ids = [EOS] + tok.encode(prompt, add_special_tokens=False).ids # docs start with <eos>
45
+ x = torch.tensor([ids], device="cuda")
46
+ out = []
47
+ for _ in range(max_new_tokens):
48
+ pos = torch.arange(x.shape[1], device=x.device).unsqueeze(0)
49
+ logits = model(x, pos, None, x.shape[1], None)[:, -1, :].float()
50
+ probs = torch.softmax(logits / max(temperature, 1e-5), dim=-1)
51
+ sp, si = torch.sort(probs, descending=True, dim=-1)
52
+ sp[sp.cumsum(-1) - sp > top_p] = 0.0
53
+ nxt = si.gather(-1, torch.multinomial(sp / sp.sum(-1, keepdim=True), 1))
54
+ t = int(nxt[0, 0])
55
+ if t == EOS:
56
+ break
57
+ out.append(t)
58
+ x = torch.cat([x, nxt], dim=1)
59
+ return tok.decode(out, skip_special_tokens=True)
60
+
61
+
62
+ print(generate(f"{USER}\nNdeipi guta guru reZimbabwe?\n{ASSISTANT}\n"))
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:2ab0b7f3b88ddb22a7fc91bd674e10481f1c8eefa404595557083edf6bc8029e
3
+ size 2969826328
modeling_morena.py ADDED
@@ -0,0 +1,1081 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python
2
+ """Morena production trainer — single file, dependency-light.
3
+
4
+ Dense Llama-style decoder (RMSNorm, RoPE, GQA, SwiGLU, tied embeddings) trained with
5
+ FSDP2 (per-parameter sharding, bf16 compute / fp32 master+reduce), Muon (2-D hidden
6
+ weights) + AdamW (embeddings, norms), WSD schedule, document-packed sequences with
7
+ cross-document masking (FlashAttention-2 varlen when available), a deterministic
8
+ mixture sampler, rotating + milestone checkpoints (torch.distributed.checkpoint),
9
+ bit-exact resume and a walltime-aware clean exit for 24h SLURM chunks.
10
+
11
+ Usage (single node): torchrun --nproc_per_node 4 train.py --config configs/proxy200m.json \
12
+ --mix configs/mix_330b.json --data-root /scratch/morena/data --out runs/proxy
13
+ See TRAIN_README.md. Requires torch >= 2.4 (FSDP2 / DTensor); flash-attn >= 2.5 optional.
14
+ """
15
+ from __future__ import annotations
16
+ import argparse, json, math, os, random, re, shutil, signal, sys, threading, time, queue, zlib
17
+ from dataclasses import dataclass, asdict, field
18
+ from typing import Optional
19
+
20
+ import numpy as np
21
+ import torch
22
+ import torch.nn as nn
23
+ import torch.nn.functional as F
24
+ import torch.distributed as dist
25
+
26
+ # ----------------------------------------------------------------------------------------
27
+ # FSDP2 / DTensor imports (torch 2.4: _composable path; torch >= 2.6: public path)
28
+ # ----------------------------------------------------------------------------------------
29
+ try:
30
+ from torch.distributed.fsdp import fully_shard, MixedPrecisionPolicy # torch >= 2.6
31
+ except ImportError: # torch 2.4 / 2.5
32
+ from torch.distributed._composable.fsdp import fully_shard, MixedPrecisionPolicy
33
+ try:
34
+ from torch.distributed.tensor import DTensor
35
+ except ImportError:
36
+ from torch.distributed._tensor import DTensor
37
+ from torch.distributed.device_mesh import init_device_mesh
38
+ import torch.distributed.checkpoint as dcp
39
+ from torch.distributed.checkpoint.state_dict import get_model_state_dict, set_model_state_dict
40
+
41
+ # ----------------------------------------------------------------------------------------
42
+ # Attention backend selection
43
+ # ----------------------------------------------------------------------------------------
44
+ _FA_VARLEN = None
45
+ try:
46
+ from flash_attn import flash_attn_varlen_func as _FA_VARLEN # type: ignore
47
+ except Exception:
48
+ _FA_VARLEN = None
49
+
50
+
51
+ def log0(*a, **k):
52
+ if int(os.environ.get("RANK", "0")) == 0:
53
+ print(*a, **k, flush=True)
54
+
55
+
56
+ # ----------------------------------------------------------------------------------------
57
+ # Config
58
+ # ----------------------------------------------------------------------------------------
59
+ @dataclass
60
+ class ModelConfig:
61
+ vocab_size: int = 65536
62
+ n_layer: int = 24
63
+ d_model: int = 2048
64
+ n_head: int = 16
65
+ n_kv_head: int = 4
66
+ d_ff: int = 5632
67
+ rope_theta: float = 500000.0
68
+ norm_eps: float = 1e-5
69
+ tie_embeddings: bool = True
70
+ init_std: float = 0.02
71
+
72
+
73
+ @dataclass
74
+ class TrainConfig:
75
+ seq_len: int = 4096
76
+ micro_batch: int = 4 # sequences per GPU per micro-step
77
+ grad_accum: int = 1 # used only if global_batch_seqs == 0
78
+ global_batch_seqs: int = 0 # if > 0: grad_accum = round(global_batch_seqs / (micro_batch * world)) -> node-count invariant
79
+ fsdp_shard_size: int = 0 # 0 = shard over all GPUs (FSDP); N = HSDP: shard within groups of N (e.g. 4 = one node), replicate across
80
+ optimizer: str = "muon" # muon (hidden 2-D) + adamw (rest) | adamw (everything; fallback)
81
+ lr: float = 1e-3 # shared Muon/AdamW LR (Moonlight RMS-0.2 scaling makes this valid)
82
+ adam_lr: Optional[float] = None # override for the AdamW group (embeddings/norms)
83
+ min_lr_ratio: float = 0.1
84
+ weight_decay: float = 0.1
85
+ muon_momentum: float = 0.95
86
+ muon_ns_steps: int = 5
87
+ muon_ns_mode: str = "roundrobin" # roundrobin | redundant
88
+ adam_betas: tuple = (0.9, 0.95)
89
+ adam_eps: float = 1e-8
90
+ grad_clip: float = 1.0
91
+ warmup_steps: int = 2000
92
+ total_steps: int = 100000 # steps in stable phase end by default == total - decay
93
+ decay_steps: int = 10000 # length of the WSD decay phase
94
+ decay_start: int = -1 # -1 => total_steps - decay_steps; set explicitly for anneal
95
+ decay_shape: str = "linear" # linear | 1-sqrt | cosine
96
+ attn: str = "auto" # auto | flash | sdpa_mask | sdpa
97
+ # Train only on positions whose mask byte is 1 (build_sft.py writes them). Default OFF so every
98
+ # existing run and every pretraining shard behaves exactly as before; SFT configs opt in.
99
+ loss_mask: bool = False
100
+ act_ckpt: bool = False
101
+ compile: bool = False
102
+ seed: int = 1234
103
+
104
+
105
+ def load_json(p):
106
+ with open(p) as f:
107
+ return json.load(f)
108
+
109
+
110
+ # ----------------------------------------------------------------------------------------
111
+ # Model
112
+ # ----------------------------------------------------------------------------------------
113
+ class RMSNorm(nn.Module):
114
+ def __init__(self, d, eps):
115
+ super().__init__()
116
+ self.eps = eps
117
+ self.weight = nn.Parameter(torch.ones(d))
118
+
119
+ def forward(self, x):
120
+ xf = x.float()
121
+ y = xf * torch.rsqrt(xf.pow(2).mean(-1, keepdim=True) + self.eps)
122
+ return (y * self.weight.float()).to(x.dtype)
123
+
124
+
125
+ def rope_cos_sin(positions: torch.Tensor, head_dim: int, theta: float, dtype):
126
+ # positions: (N,) int64 -> cos/sin (N, head_dim/2)
127
+ inv = 1.0 / (theta ** (torch.arange(0, head_dim, 2, device=positions.device).float() / head_dim))
128
+ freqs = positions.float()[:, None] * inv[None, :]
129
+ return freqs.cos().to(dtype), freqs.sin().to(dtype)
130
+
131
+
132
+ def apply_rope(x, cos, sin):
133
+ # x: (N, H, D); cos/sin: (N, D/2)
134
+ x1, x2 = x[..., : x.shape[-1] // 2], x[..., x.shape[-1] // 2:]
135
+ c, s = cos[:, None, :], sin[:, None, :]
136
+ return torch.cat([x1 * c - x2 * s, x1 * s + x2 * c], dim=-1)
137
+
138
+
139
+ class Attention(nn.Module):
140
+ def __init__(self, cfg: ModelConfig, attn_mode: str):
141
+ super().__init__()
142
+ self.n_head, self.n_kv = cfg.n_head, cfg.n_kv_head
143
+ self.hd = cfg.d_model // cfg.n_head
144
+ self.wq = nn.Linear(cfg.d_model, cfg.n_head * self.hd, bias=False)
145
+ self.wk = nn.Linear(cfg.d_model, cfg.n_kv_head * self.hd, bias=False)
146
+ self.wv = nn.Linear(cfg.d_model, cfg.n_kv_head * self.hd, bias=False)
147
+ self.wo = nn.Linear(cfg.n_head * self.hd, cfg.d_model, bias=False)
148
+ self.attn_mode = attn_mode
149
+
150
+ def forward(self, x, cos, sin, cu_seqlens, max_seqlen, mask):
151
+ B, T, C = x.shape
152
+ N = B * T
153
+ q = self.wq(x).view(N, self.n_head, self.hd)
154
+ k = self.wk(x).view(N, self.n_kv, self.hd)
155
+ v = self.wv(x).view(N, self.n_kv, self.hd)
156
+ q, k = apply_rope(q, cos, sin), apply_rope(k, cos, sin)
157
+ if self.attn_mode == "flash":
158
+ o = _FA_VARLEN(q, k, v, cu_seqlens, cu_seqlens, max_seqlen, max_seqlen, causal=True)
159
+ o = o.view(B, T, C)
160
+ else:
161
+ q = q.view(B, T, self.n_head, self.hd).transpose(1, 2)
162
+ k = k.view(B, T, self.n_kv, self.hd).transpose(1, 2)
163
+ v = v.view(B, T, self.n_kv, self.hd).transpose(1, 2)
164
+ rep = self.n_head // self.n_kv
165
+ if rep > 1:
166
+ k = k.repeat_interleave(rep, dim=1)
167
+ v = v.repeat_interleave(rep, dim=1)
168
+ if self.attn_mode == "sdpa_mask":
169
+ o = F.scaled_dot_product_attention(q, k, v, attn_mask=mask) # mask: (B,1,T,T) bool
170
+ else:
171
+ o = F.scaled_dot_product_attention(q, k, v, is_causal=True)
172
+ o = o.transpose(1, 2).reshape(B, T, C)
173
+ return self.wo(o)
174
+
175
+
176
+ class MLP(nn.Module):
177
+ def __init__(self, cfg: ModelConfig):
178
+ super().__init__()
179
+ self.w1 = nn.Linear(cfg.d_model, cfg.d_ff, bias=False) # gate
180
+ self.w3 = nn.Linear(cfg.d_model, cfg.d_ff, bias=False) # up
181
+ self.w2 = nn.Linear(cfg.d_ff, cfg.d_model, bias=False) # down
182
+
183
+ def forward(self, x):
184
+ return self.w2(F.silu(self.w1(x)) * self.w3(x))
185
+
186
+
187
+ class Block(nn.Module):
188
+ def __init__(self, cfg: ModelConfig, attn_mode: str):
189
+ super().__init__()
190
+ self.attn_norm = RMSNorm(cfg.d_model, cfg.norm_eps)
191
+ self.attn = Attention(cfg, attn_mode)
192
+ self.mlp_norm = RMSNorm(cfg.d_model, cfg.norm_eps)
193
+ self.mlp = MLP(cfg)
194
+
195
+ def forward(self, x, cos, sin, cu, mx, mask):
196
+ x = x + self.attn(self.attn_norm(x), cos, sin, cu, mx, mask)
197
+ return x + self.mlp(self.mlp_norm(x))
198
+
199
+
200
+ class Transformer(nn.Module):
201
+ def __init__(self, cfg: ModelConfig, attn_mode: str):
202
+ super().__init__()
203
+ self.cfg, self.attn_mode = cfg, attn_mode
204
+ self.embed = nn.Embedding(cfg.vocab_size, cfg.d_model)
205
+ self.layers = nn.ModuleList([Block(cfg, attn_mode) for _ in range(cfg.n_layer)])
206
+ self.norm = RMSNorm(cfg.d_model, cfg.norm_eps)
207
+ if not cfg.tie_embeddings:
208
+ self.lm_head = nn.Linear(cfg.d_model, cfg.vocab_size, bias=False)
209
+ self.act_ckpt = False
210
+ self.apply(self._init)
211
+ for n, p in self.named_parameters(): # GPT-2-style residual-output scaling
212
+ if n.endswith("wo.weight") or n.endswith("w2.weight"):
213
+ nn.init.normal_(p, std=cfg.init_std / math.sqrt(2 * cfg.n_layer))
214
+
215
+ def _init(self, m):
216
+ if isinstance(m, (nn.Linear, nn.Embedding)):
217
+ nn.init.normal_(m.weight, std=self.cfg.init_std)
218
+
219
+ def forward(self, idx, positions, cu_seqlens, max_seqlen, mask):
220
+ B, T = idx.shape
221
+ x = self.embed(idx)
222
+ cos, sin = rope_cos_sin(positions.view(-1), self.cfg.d_model // self.cfg.n_head,
223
+ self.cfg.rope_theta, x.dtype)
224
+ for blk in self.layers:
225
+ if self.act_ckpt and self.training:
226
+ x = torch.utils.checkpoint.checkpoint(blk, x, cos, sin, cu_seqlens, max_seqlen, mask,
227
+ use_reentrant=False)
228
+ else:
229
+ x = blk(x, cos, sin, cu_seqlens, max_seqlen, mask)
230
+ x = self.norm(x)
231
+ w = self.embed.weight if self.cfg.tie_embeddings else self.lm_head.weight
232
+ return F.linear(x, w)
233
+
234
+ def n_params(self, non_embed=False):
235
+ n = sum(p.numel() for p in self.parameters())
236
+ return n - (self.embed.weight.numel() if non_embed else 0)
237
+
238
+
239
+ def choose_attn(requested: str) -> str:
240
+ if requested == "auto":
241
+ if _FA_VARLEN is not None:
242
+ return "flash"
243
+ log0("=" * 88)
244
+ log0("!! WARNING: flash-attn varlen NOT available -> falling back to plain SDPA causal attention.")
245
+ log0("!! Packed sequences will attend ACROSS document boundaries (no doc mask).")
246
+ log0("!! Use --attn sdpa_mask for correct (slower, memory-hungry) masking without flash-attn.")
247
+ log0("=" * 88)
248
+ return "sdpa"
249
+ if requested == "flash" and _FA_VARLEN is None:
250
+ raise RuntimeError("--attn flash requested but flash_attn is not importable")
251
+ return requested
252
+
253
+
254
+ # ----------------------------------------------------------------------------------------
255
+ # Data: memmap shards per source + deterministic mixture sampler
256
+ # ----------------------------------------------------------------------------------------
257
+ class Source:
258
+ """A directory of raw token shards with index.json {dtype, eos_id, shards:[{file,n_tokens}]}.
259
+ Windows of seq_len+1 tokens are enumerated per shard (stride seq_len), never crossing shards."""
260
+
261
+ def __init__(self, name, path, seq_len):
262
+ self.name, self.path, self.L = name, path, seq_len
263
+ idx = load_json(os.path.join(path, "index.json"))
264
+ self.dtype = np.dtype(idx["dtype"])
265
+ self.eos_id = int(idx["eos_id"])
266
+ # Optional parallel loss mask: one uint8 per token, 1 = train on this position. Written by
267
+ # build_sft.py so that SFT trains on the response and not on the prompt it was handed. A
268
+ # source without mask files behaves exactly as before (mask of all ones), so pretraining
269
+ # shards and older SFT builds keep working untouched.
270
+ self.mask_files = {s_["file"]: s_.get("mask") for s_ in idx["shards"]}
271
+ self.has_mask = any(self.mask_files.values())
272
+ self._mmm = {}
273
+ self.shards, self.win_cum = [], [0]
274
+ for s in idx["shards"]:
275
+ n = int(s["n_tokens"])
276
+ w = max(0, (n - 1) // seq_len)
277
+ self.shards.append((os.path.join(path, s["file"]), n, w))
278
+ self.win_cum.append(self.win_cum[-1] + w)
279
+ self.n_windows = self.win_cum[-1]
280
+ self.n_tokens = sum(s[1] for s in self.shards)
281
+ self._mm = {}
282
+ if self.n_windows == 0:
283
+ raise ValueError(f"source {name} at {path} has no full windows of {seq_len + 1} tokens")
284
+
285
+ def _mmap(self, i):
286
+ if i not in self._mm:
287
+ self._mm[i] = np.memmap(self.shards[i][0], dtype=self.dtype, mode="r")
288
+ return self._mm[i]
289
+
290
+ def _mmap_mask(self, i):
291
+ if i not in self._mmm:
292
+ fn = self.mask_files.get(os.path.basename(self.shards[i][0]))
293
+ self._mmm[i] = (np.memmap(os.path.join(self.path, fn), dtype=np.uint8, mode="r")
294
+ if fn else None)
295
+ return self._mmm[i]
296
+
297
+ def window(self, w):
298
+ i = int(np.searchsorted(self.win_cum, w, side="right") - 1)
299
+ off = (w - self.win_cum[i]) * self.L
300
+ a = self._mmap(i)[off: off + self.L + 1]
301
+ return np.asarray(a, dtype=np.int64)
302
+
303
+ def window_mask(self, w):
304
+ """(L+1,) uint8 aligned with window(w). All ones when this source has no mask stream."""
305
+ i = int(np.searchsorted(self.win_cum, w, side="right") - 1)
306
+ mm = self._mmap_mask(i)
307
+ if mm is None:
308
+ return np.ones(self.L + 1, dtype=np.uint8)
309
+ off = (w - self.win_cum[i]) * self.L
310
+ return np.asarray(mm[off: off + self.L + 1], dtype=np.uint8)
311
+
312
+
313
+ def _coprime_multiplier(n, seed):
314
+ rng = random.Random(seed)
315
+ while True:
316
+ a = rng.randrange(1, n) if n > 1 else 1
317
+ if math.gcd(a, n) == 1:
318
+ return a
319
+
320
+
321
+ class MixtureSampler:
322
+ """Deterministic: for global step s, every rank derives the same per-sample source list from
323
+ hash(seed, s); per-source window positions come from monotone counters (checkpointed, and
324
+ reconstructible by replay) mapped through an epoch-keyed affine permutation. Resume is exact."""
325
+
326
+ def __init__(self, sources: dict, weights: dict, global_batch: int, seed: int):
327
+ self.names = sorted(sources)
328
+ self.sources = sources
329
+ w = np.array([float(weights[n]) for n in self.names])
330
+ self.probs = w / w.sum()
331
+ self.B, self.seed = global_batch, seed
332
+ self.counters = {n: 0 for n in self.names}
333
+ self._perm_cache = {}
334
+
335
+ def state_dict(self):
336
+ return {"counters": dict(self.counters)}
337
+
338
+ def load_state_dict(self, sd):
339
+ self.counters = {n: int(sd["counters"].get(n, 0)) for n in self.names}
340
+
341
+ def _perm(self, name, epoch):
342
+ key = (name, epoch)
343
+ if key not in self._perm_cache:
344
+ n = self.sources[name].n_windows
345
+ h = zlib.crc32(f"{self.seed}|{name}|{epoch}".encode()) # stable across processes
346
+ a = _coprime_multiplier(n, h)
347
+ b = random.Random(h ^ 0x9E3779B9).randrange(n)
348
+ self._perm_cache[key] = (a, b, n)
349
+ return self._perm_cache[key]
350
+
351
+ def _pos(self, name, counter):
352
+ a, b, n = self._perm(name, counter // self.sources[name].n_windows)
353
+ return (a * (counter % n) + b) % n
354
+
355
+ def step_assignments(self, step):
356
+ """Return list of (source_name, window_idx) for ALL global_batch samples of `step`,
357
+ and advance counters. Must be called exactly once per step, in order, on every rank."""
358
+ rng = np.random.default_rng([self.seed, step])
359
+ srcs = rng.choice(len(self.names), size=self.B, p=self.probs)
360
+ out = []
361
+ for si in srcs:
362
+ name = self.names[int(si)]
363
+ c = self.counters[name]
364
+ out.append((name, self._pos(name, c)))
365
+ self.counters[name] = c + 1
366
+ return out
367
+
368
+ def epochs(self):
369
+ return {n: self.counters[n] / self.sources[n].n_windows for n in self.names}
370
+
371
+
372
+ _LOADER_ERR = object() # sentinel: the prefetch thread failed (see Loader._run / Loader.next)
373
+
374
+
375
+ class Loader:
376
+ """Background-threaded prefetch of this rank's slice of each step's assignments."""
377
+
378
+ def __init__(self, sampler: MixtureSampler, rank, world, micro_batch, grad_accum, seq_len,
379
+ start_step, prefetch=4):
380
+ self.s, self.rank, self.world = sampler, rank, world
381
+ self.mb, self.ga, self.L = micro_batch, grad_accum, seq_len
382
+ self.per_rank = micro_batch * grad_accum
383
+ self.q = queue.Queue(maxsize=prefetch)
384
+ self.step = start_step
385
+ self.stop = False
386
+ self.err = None
387
+ self.t = threading.Thread(target=self._run, daemon=True)
388
+ self.t.start()
389
+
390
+ def _run(self):
391
+ while not self.stop:
392
+ step = self.step
393
+ try:
394
+ assign = self.s.step_assignments(step)
395
+ mine = assign[self.rank * self.per_rank:(self.rank + 1) * self.per_rank]
396
+ micro = []
397
+ for m in range(self.ga):
398
+ part = mine[m * self.mb:(m + 1) * self.mb]
399
+ toks = np.stack([self.s.sources[n].window(w) for n, w in part])
400
+ msk = np.stack([self.s.sources[n].window_mask(w) for n, w in part])
401
+ micro.append((torch.from_numpy(toks), [self.s.sources[n].eos_id for n, _ in part][0],
402
+ torch.from_numpy(msk)))
403
+ except BaseException as e: # an unreadable shard must not hang the whole job forever
404
+ self.err = e
405
+ self.q.put((_LOADER_ERR, None, None))
406
+ return
407
+ self.q.put((step, micro, self.s.state_dict()))
408
+ self.step += 1
409
+
410
+ def next(self):
411
+ item = self.q.get()
412
+ if item[0] is _LOADER_ERR:
413
+ raise RuntimeError(f"data loader thread died at step {self.step}") from self.err
414
+ return item
415
+
416
+
417
+ def build_batch(tokens: torch.Tensor, eos_id: int, attn_mode: str, device):
418
+ """tokens: (B, L+1) int64. Returns inputs, targets, positions, cu_seqlens, max_seqlen, mask."""
419
+ x, y = tokens[:, :-1], tokens[:, 1:]
420
+ B, T = x.shape
421
+ if attn_mode == "sdpa":
422
+ pos = torch.arange(T).repeat(B, 1)
423
+ return (x.to(device, non_blocking=True), y.to(device, non_blocking=True),
424
+ pos.to(device, non_blocking=True), None, T, None)
425
+ # document boundaries: a new doc starts right after each EOS token (EOS belongs to the previous doc)
426
+ is_eos = (x == eos_id)
427
+ starts = torch.zeros(B, T, dtype=torch.bool)
428
+ starts[:, 0] = True
429
+ starts[:, 1:] = is_eos[:, :-1]
430
+ doc_id = torch.cumsum(starts.long(), dim=1) - 1 # (B,T) doc index within row
431
+ # positions restart at every document
432
+ idx = torch.arange(T).repeat(B, 1)
433
+ start_pos = torch.where(starts, idx, torch.zeros_like(idx))
434
+ start_pos = torch.cummax(start_pos, dim=1).values
435
+ pos = idx - start_pos
436
+ if attn_mode == "flash":
437
+ flat_starts = starts.clone()
438
+ flat_starts[:, 0] = True
439
+ s = flat_starts.view(-1).nonzero().squeeze(1)
440
+ cu = torch.cat([s, torch.tensor([B * T])]).to(torch.int32)
441
+ max_len = int((cu[1:] - cu[:-1]).max())
442
+ return (x.to(device, non_blocking=True), y.to(device, non_blocking=True),
443
+ pos.to(device, non_blocking=True), cu.to(device), max_len, None)
444
+ # sdpa_mask: block-diagonal causal mask (B,1,T,T)
445
+ same = doc_id[:, :, None] == doc_id[:, None, :]
446
+ causal = torch.tril(torch.ones(T, T, dtype=torch.bool))
447
+ mask = (same & causal)[:, None]
448
+ return (x.to(device, non_blocking=True), y.to(device, non_blocking=True),
449
+ pos.to(device, non_blocking=True), None, T, mask.to(device, non_blocking=True))
450
+
451
+
452
+ # ----------------------------------------------------------------------------------------
453
+ # Muon (distributed-safe over FSDP2 DTensors)
454
+ # ----------------------------------------------------------------------------------------
455
+ def newton_schulz5(G: torch.Tensor, steps: int = 5, eps: float = 1e-7):
456
+ """Quintic Newton-Schulz iteration (Keller Jordan coefficients) -> approx. orthogonal matrix."""
457
+ a, b, c = (3.4445, -4.7750, 2.0315)
458
+ X = G.to(torch.bfloat16)
459
+ transposed = X.size(0) > X.size(1)
460
+ if transposed:
461
+ X = X.T
462
+ X = X / (X.norm() + eps)
463
+ for _ in range(steps):
464
+ A = X @ X.T
465
+ B = b * A + c * (A @ A)
466
+ X = a * X + B @ X
467
+ return X.T if transposed else X
468
+
469
+
470
+ class Muon(torch.optim.Optimizer):
471
+ """Muon for 2-D hidden weights. Works for plain tensors and for FSDP2 DTensors (Shard(0)).
472
+
473
+ DESIGN CHOICE (documented): we use FSDP2 (fully_shard, per-parameter DTensor sharding) rather
474
+ than FSDP1 + use_orig_params. With FSDP2 every parameter and its gradient is a DTensor whose
475
+ 2-D shape is preserved and whose local shard is a contiguous row-block, so the
476
+ orthogonalization is computed on the FULL gathered matrix: momentum is kept sharded
477
+ (memory = 1 extra sharded copy), the Nesterov update is all-gathered (`full_tensor()`),
478
+ Newton-Schulz runs on the full 2-D matrix, and each rank applies its own row-slice.
479
+ `ns_mode='roundrobin'` assigns each matrix to one owner rank which computes NS and broadcasts
480
+ the result (compute / world); `'redundant'` makes every rank compute NS (no broadcast, ~10%
481
+ extra compute at 12 GPUs for 1.5B). Both are bit-identical across ranks.
482
+ Scaling follows Moonlight: update *= 0.2 * sqrt(max(rows, cols)) so Muon and AdamW share LR/WD.
483
+ """
484
+
485
+ def __init__(self, params, lr=1e-3, momentum=0.95, nesterov=True, ns_steps=5, weight_decay=0.1,
486
+ ns_mode="roundrobin"):
487
+ super().__init__(params, dict(lr=lr, momentum=momentum, nesterov=nesterov, ns_steps=ns_steps,
488
+ weight_decay=weight_decay))
489
+ self.ns_mode = ns_mode
490
+ self._row_meta = {} # id(p) -> (offset, nrows) of local shard
491
+
492
+ def _local_rows(self, p: DTensor):
493
+ key = id(p)
494
+ if key not in self._row_meta:
495
+ mesh = p.device_mesh
496
+ pg = mesh.get_group(mesh_dim=mesh.ndim - 1) # the shard dim (last); replicate dim holds identical copies
497
+ local_n = p.to_local().shape[0]
498
+ sizes = [0] * dist.get_world_size(pg)
499
+ dist.all_gather_object(sizes, local_n, group=pg)
500
+ r = dist.get_rank(pg)
501
+ self._row_meta[key] = (sum(sizes[:r]), local_n, pg)
502
+ return self._row_meta[key]
503
+
504
+ @torch.no_grad()
505
+ def step(self, closure=None):
506
+ for group in self.param_groups:
507
+ lr, mu, wd = group["lr"], group["momentum"], group["weight_decay"]
508
+ params = [p for p in group["params"] if p.grad is not None]
509
+ for i, p in enumerate(params):
510
+ g = p.grad
511
+ st = self.state[p]
512
+ if "momentum_buffer" not in st:
513
+ st["momentum_buffer"] = torch.zeros_like(p)
514
+ buf = st["momentum_buffer"]
515
+ buf.mul_(mu).add_(g)
516
+ u = g.add(buf, alpha=mu) if group["nesterov"] else buf
517
+ is_dt = isinstance(p, DTensor)
518
+ if is_dt:
519
+ off, n, pg = self._local_rows(p)
520
+ full = u.full_tensor()
521
+ owner = i % dist.get_world_size(pg)
522
+ if self.ns_mode == "roundrobin":
523
+ if dist.get_rank(pg) == owner:
524
+ O = newton_schulz5(full, group["ns_steps"])
525
+ else:
526
+ O = torch.empty_like(full, dtype=torch.bfloat16)
527
+ dist.broadcast(O, src=dist.get_global_rank(pg, owner), group=pg)
528
+ else:
529
+ O = newton_schulz5(full, group["ns_steps"])
530
+ O_local = O[off: off + n]
531
+ p_local = p.to_local()
532
+ else:
533
+ O_local = newton_schulz5(u, group["ns_steps"])
534
+ p_local = p
535
+ scale = 0.2 * math.sqrt(max(p.shape[0], p.shape[1]))
536
+ p_local.mul_(1 - lr * wd).add_(O_local.to(p_local.dtype), alpha=-lr * scale)
537
+
538
+
539
+ # ----------------------------------------------------------------------------------------
540
+ # Schedule
541
+ # ----------------------------------------------------------------------------------------
542
+ def lr_at(step, tc: TrainConfig):
543
+ base, minr = tc.lr, tc.min_lr_ratio
544
+ if step < tc.warmup_steps:
545
+ return base * (step + 1) / tc.warmup_steps
546
+ ds = tc.decay_start if tc.decay_start >= 0 else tc.total_steps - tc.decay_steps
547
+ if step < ds:
548
+ return base
549
+ p = min(1.0, (step - ds) / max(1, tc.decay_steps))
550
+ if tc.decay_shape == "1-sqrt":
551
+ f = 1 - math.sqrt(p)
552
+ elif tc.decay_shape == "cosine":
553
+ f = 0.5 * (1 + math.cos(math.pi * p))
554
+ else:
555
+ f = 1 - p
556
+ return base * (minr + (1 - minr) * f)
557
+
558
+
559
+ # ----------------------------------------------------------------------------------------
560
+ # Checkpointing (DCP: each rank writes its shards; resharding-safe across world sizes)
561
+ # ----------------------------------------------------------------------------------------
562
+ def ckpt_dir(out, step):
563
+ return os.path.join(out, "ckpt", f"step_{step:08d}")
564
+
565
+
566
+ def _param_fqns(model):
567
+ return {p: n for n, p in model.named_parameters()}
568
+
569
+
570
+ def _init_opt_state(opt, fqns):
571
+ """Make every optimizer state tensor exist before loading, without running a fake step.
572
+ Muon: momentum_buffer (sharded like the param). AdamW: exp_avg, exp_avg_sq (sharded), step (CPU scalar)."""
573
+ for group in opt.param_groups:
574
+ for p in group["params"]:
575
+ st = opt.state[p]
576
+ if isinstance(opt, Muon):
577
+ st.setdefault("momentum_buffer", torch.zeros_like(p))
578
+ else:
579
+ st.setdefault("step", torch.tensor(0.0, dtype=torch.float32))
580
+ st.setdefault("exp_avg", torch.zeros_like(p))
581
+ st.setdefault("exp_avg_sq", torch.zeros_like(p))
582
+
583
+
584
+ def optimizer_state_for_dcp(opt, fqns):
585
+ """Flat {fqn.key: tensor} view of the optimizer state (tensors are the live buffers -> DCP loads in place)."""
586
+ _init_opt_state(opt, fqns)
587
+ out = {}
588
+ for p, st in opt.state.items():
589
+ for k, v in st.items():
590
+ if isinstance(v, torch.Tensor):
591
+ out[f"{fqns[p]}.{k}"] = v
592
+ return out
593
+
594
+
595
+ def state_fingerprint(model, opts):
596
+ """Sum of L2 norms of all model params and optimizer state tensors (identical on all ranks)."""
597
+ tot = torch.zeros(2, dtype=torch.float64)
598
+ for p in model.parameters():
599
+ t = p.full_tensor() if isinstance(p, DTensor) else p
600
+ tot[0] += t.detach().double().norm().cpu()
601
+ for o in opts:
602
+ for st in o.state.values():
603
+ for v in st.values():
604
+ if isinstance(v, torch.Tensor):
605
+ t = v.full_tensor() if isinstance(v, DTensor) else v
606
+ tot[1] += t.detach().double().norm().cpu()
607
+ return [round(float(x), 3) for x in tot]
608
+
609
+
610
+ def save_checkpoint(out, step, model, opts, sampler_state, extra, keep_last, milestone_every, rank):
611
+ d = ckpt_dir(out, step)
612
+ tmp = d + ".tmp"
613
+ if rank == 0:
614
+ shutil.rmtree(tmp, ignore_errors=True)
615
+ os.makedirs(tmp, exist_ok=True)
616
+ if dist.is_initialized():
617
+ dist.barrier()
618
+ fqns = _param_fqns(model)
619
+ sd = {"model": get_model_state_dict(model)}
620
+ for i, o in enumerate(opts):
621
+ sd[f"opt{i}"] = optimizer_state_for_dcp(o, fqns)
622
+ dcp.save(sd, checkpoint_id=tmp)
623
+ if rank == 0:
624
+ meta = {"step": step, "sampler": sampler_state, "fingerprint": extra.pop("fingerprint", None), "rng": {
625
+ "torch": torch.get_rng_state().tolist(),
626
+ "cuda": torch.cuda.get_rng_state().tolist() if torch.cuda.is_available() else []}, **extra}
627
+ with open(os.path.join(tmp, "meta.json"), "w") as f:
628
+ json.dump(meta, f)
629
+ if os.path.exists(d): # re-saving a step (e.g. after a manual resume from an older ckpt)
630
+ shutil.rmtree(d, ignore_errors=True)
631
+ os.replace(tmp, d)
632
+ with open(os.path.join(out, "ckpt", "latest.txt.tmp"), "w") as f:
633
+ f.write(str(step))
634
+ os.replace(os.path.join(out, "ckpt", "latest.txt.tmp"), os.path.join(out, "ckpt", "latest.txt"))
635
+ # rotate: keep last `keep_last` non-milestone checkpoints; milestones are kept forever
636
+ steps = sorted(int(m.group(1)) for n in os.listdir(os.path.join(out, "ckpt"))
637
+ if (m := re.fullmatch(r"step_(\d+)", n)))
638
+ rot = [s for s in steps if not (milestone_every and s % milestone_every == 0 and s > 0)]
639
+ for s in rot[:-keep_last]:
640
+ shutil.rmtree(ckpt_dir(out, s), ignore_errors=True)
641
+ if dist.is_initialized():
642
+ dist.barrier()
643
+
644
+
645
+ def load_checkpoint(path, model, opts):
646
+ parts = os.environ.get("MORENA_RESUME_PARTS", "model,opt,rng").split(",") # debugging knob
647
+ fqns = _param_fqns(model)
648
+ sd = {}
649
+ if "model" in parts:
650
+ sd["model"] = get_model_state_dict(model)
651
+ if "opt" in parts:
652
+ for i, o in enumerate(opts):
653
+ sd[f"opt{i}"] = optimizer_state_for_dcp(o, fqns)
654
+ live = {k: dict(v) for k, v in sd.items() if k.startswith("opt")}
655
+ dcp.load(sd, checkpoint_id=path) # loads IN PLACE into the live model / optimizer tensors ...
656
+ if "model" in parts:
657
+ set_model_state_dict(model, sd["model"])
658
+ for k, d_ in live.items(): # ... but copy explicitly in case the planner returned new tensors
659
+ for name, t in d_.items():
660
+ loaded = sd[k][name]
661
+ if loaded is not t:
662
+ t.copy_(loaded)
663
+ with open(os.path.join(path, "meta.json")) as f:
664
+ return json.load(f)
665
+
666
+
667
+ def find_latest(out):
668
+ p = os.path.join(out, "ckpt", "latest.txt")
669
+ if not os.path.exists(p):
670
+ return None
671
+ with open(p) as f:
672
+ step = int(f.read().strip())
673
+ d = ckpt_dir(out, step)
674
+ return d if os.path.exists(os.path.join(d, "meta.json")) else None
675
+
676
+
677
+ # ----------------------------------------------------------------------------------------
678
+ # Main
679
+ # ----------------------------------------------------------------------------------------
680
+ def parse():
681
+ ap = argparse.ArgumentParser()
682
+ ap.add_argument("--config", required=True, help="model+train JSON (configs/*.json)")
683
+ ap.add_argument("--mix", required=True, help="mixture JSON (configs/mix_*.json)")
684
+ ap.add_argument("--data-root", required=True, help="dir containing one shard dir per source")
685
+ ap.add_argument("--out", required=True, help="run dir (checkpoints, logs)")
686
+ ap.add_argument("--resume", default="auto", help="auto | none | /path/to/ckpt/step_XXXXXXXX")
687
+ ap.add_argument("--override", default="", help='JSON string of config overrides, e.g. {"train":{"lr":2e-3}}')
688
+ ap.add_argument("--override-file", default="", help="JSON file of config overrides (configs/fallback_*.json); applied before --override")
689
+ ap.add_argument("--anneal", type=int, default=0,
690
+ help="start WSD decay NOW (from the resumed step) lasting this many steps; sets total_steps accordingly")
691
+ ap.add_argument("--loss-mask", action="store_true",
692
+ help="train only on positions whose mask byte is 1 (SFT)")
693
+ ap.add_argument("--ckpt-minutes", type=float, default=30)
694
+ ap.add_argument("--ckpt-keep", type=int, default=3)
695
+ ap.add_argument("--milestone-every", type=int, default=10000, help="steps; these checkpoints are never rotated")
696
+ ap.add_argument("--walltime", default=os.environ.get("MORENA_WALLTIME", ""),
697
+ help="HH:MM:SS budget from process start (or env MORENA_WALLTIME)")
698
+ ap.add_argument("--deadline-unix", type=float, default=float(os.environ.get("MORENA_DEADLINE", "0") or 0),
699
+ help="absolute unix time the job will be killed (env MORENA_DEADLINE); overrides --walltime")
700
+ ap.add_argument("--exit-margin-min", type=float, default=20)
701
+ ap.add_argument("--log-every", type=int, default=1)
702
+ ap.add_argument("--wandb", default="", help="W&B project name; offline mode unless WANDB_MODE set")
703
+ ap.add_argument("--max-steps", type=int, default=0, help="stop after this many steps in THIS process (testing)")
704
+ ap.add_argument("--no-fsdp", action="store_true", help="single-GPU debugging without sharding")
705
+ return ap.parse_args()
706
+
707
+
708
+ def hms_to_sec(s):
709
+ parts = [int(x) for x in s.split(":")]
710
+ while len(parts) < 3:
711
+ parts.insert(0, 0)
712
+ return parts[0] * 3600 + parts[1] * 60 + parts[2]
713
+
714
+
715
+ def main():
716
+ args = parse()
717
+ t_start = time.time()
718
+ cfg = load_json(args.config)
719
+ for ov in ([load_json(args.override_file)] if args.override_file else []) + ([json.loads(args.override)] if args.override else []):
720
+ for k in ov:
721
+ if k.startswith("_"):
722
+ continue
723
+ cfg.setdefault(k, {}).update(ov[k])
724
+ log0(f"[config] override applied: { {k: v for k, v in ov.items() if not k.startswith('_')} }")
725
+ mc = ModelConfig(**cfg["model"])
726
+ tc = TrainConfig(**cfg["train"])
727
+ # CLI wins over the config file, so a config can enable masking and a run can force it off
728
+ # (or on) without editing JSON. Only applies when the flag was actually passed.
729
+ if args.loss_mask:
730
+ tc.loss_mask = True
731
+ tc.adam_betas = tuple(tc.adam_betas)
732
+
733
+ # --- distributed init ---
734
+ dist.init_process_group("nccl" if torch.cuda.is_available() else "gloo")
735
+ rank, world = dist.get_rank(), dist.get_world_size()
736
+ local_rank = int(os.environ.get("LOCAL_RANK", 0))
737
+ device = torch.device("cuda", local_rank) if torch.cuda.is_available() else torch.device("cpu")
738
+ if device.type == "cuda":
739
+ torch.cuda.set_device(device)
740
+ torch.manual_seed(tc.seed + rank)
741
+ torch.backends.cuda.matmul.allow_tf32 = True
742
+ torch.backends.cudnn.allow_tf32 = True
743
+
744
+ # --- deadline ---
745
+ deadline = None
746
+ if args.deadline_unix > 0:
747
+ deadline = args.deadline_unix
748
+ elif args.walltime:
749
+ deadline = t_start + hms_to_sec(args.walltime)
750
+ if deadline:
751
+ log0(f"[time] deadline in {(deadline - time.time()) / 60:.1f} min; will exit {args.exit_margin_min} min early")
752
+
753
+ os.makedirs(os.path.join(args.out, "ckpt"), exist_ok=True)
754
+ attn_mode = choose_attn(tc.attn)
755
+ log0(f"[attn] backend = {attn_mode} (flash_attn importable: {_FA_VARLEN is not None})")
756
+
757
+ # --- model ---
758
+ with torch.device("meta"):
759
+ model = Transformer(mc, attn_mode)
760
+ n_params, n_nonembed = model.n_params(), model.n_params(non_embed=True)
761
+ log0(f"[model] {n_params / 1e6:.1f}M params ({n_nonembed / 1e6:.1f}M non-embedding) {asdict(mc)}")
762
+ model.act_ckpt = tc.act_ckpt
763
+
764
+ # materialize on device (full init on every rank, then shard; fine up to a few B params)
765
+ model.to_empty(device=device)
766
+ torch.manual_seed(tc.seed) # identical init on all ranks
767
+ model.apply(model._init)
768
+ for n, p in model.named_parameters():
769
+ if n.endswith("wo.weight") or n.endswith("w2.weight"):
770
+ nn.init.normal_(p, std=mc.init_std / math.sqrt(2 * mc.n_layer))
771
+ for m in model.modules():
772
+ if isinstance(m, RMSNorm):
773
+ nn.init.ones_(m.weight)
774
+
775
+ if not args.no_fsdp:
776
+ shard = tc.fsdp_shard_size if tc.fsdp_shard_size > 0 else world
777
+ assert world % shard == 0, f"world {world} not divisible by fsdp_shard_size {shard}"
778
+ if shard == world:
779
+ mesh = init_device_mesh(device.type, (world,), mesh_dim_names=("dp_shard",))
780
+ else: # HSDP: all-gathers stay inside a node (NVLink), only gradient all-reduce crosses the fabric
781
+ mesh = init_device_mesh(device.type, (world // shard, shard), mesh_dim_names=("dp_replicate", "dp_shard"))
782
+ log0(f"[fsdp] mesh {tuple(mesh.shape)} dims {mesh.mesh_dim_names}")
783
+ mp = MixedPrecisionPolicy(param_dtype=torch.bfloat16, reduce_dtype=torch.float32)
784
+ for blk in model.layers:
785
+ fully_shard(blk, mesh=mesh, mp_policy=mp)
786
+ fully_shard(model, mesh=mesh, mp_policy=mp)
787
+ if tc.compile:
788
+ for blk in model.layers:
789
+ blk.compile()
790
+
791
+ # --- optimizers: Muon for 2-D hidden weights, AdamW for embeddings + norms ---
792
+ muon_params, adam_params, adam_nodecay = [], [], []
793
+ for n, p in model.named_parameters():
794
+ if tc.optimizer == "muon" and p.ndim == 2 and "embed" not in n and "lm_head" not in n:
795
+ muon_params.append(p)
796
+ elif p.ndim >= 2:
797
+ adam_params.append(p)
798
+ else:
799
+ adam_nodecay.append(p)
800
+ adam_lr = tc.adam_lr if tc.adam_lr else tc.lr
801
+ opt_muon = Muon(muon_params, lr=tc.lr, momentum=tc.muon_momentum, ns_steps=tc.muon_ns_steps,
802
+ weight_decay=tc.weight_decay, ns_mode=tc.muon_ns_mode) if muon_params else None
803
+ opt_adam = torch.optim.AdamW([{"params": adam_params, "weight_decay": tc.weight_decay},
804
+ {"params": adam_nodecay, "weight_decay": 0.0}],
805
+ lr=adam_lr, betas=tc.adam_betas, eps=tc.adam_eps, fused=False)
806
+ opts = [opt_muon, opt_adam] if muon_params else [opt_adam]
807
+ log0(f"[optim] {tc.optimizer}: muon: {sum(p.numel() for p in muon_params) / 1e6:.1f}M adamw: "
808
+ f"{(sum(p.numel() for p in adam_params) + sum(p.numel() for p in adam_nodecay)) / 1e6:.1f}M ns_mode={tc.muon_ns_mode}")
809
+
810
+ # --- data ---
811
+ # Mixture / launch manifest: every listed source is OPTIONAL unless "required": true -- the run can start
812
+ # with whatever shards are on SCRATCH and pick up more sources at a later link (weights are relative and
813
+ # renormalized over the sources present; the sampler is keyed on the global step, so this is reproducible).
814
+ mix = load_json(args.mix)
815
+ sources, weights, missing = {}, {}, []
816
+ for name, spec in mix["sources"].items():
817
+ if float(spec.get("weight", 0)) <= 0:
818
+ continue
819
+ path = os.path.join(args.data_root, spec.get("path", name))
820
+ if not os.path.exists(os.path.join(path, "index.json")):
821
+ if spec.get("required", False):
822
+ raise FileNotFoundError(f"required source {name}: {path}/index.json")
823
+ missing.append(name)
824
+ continue
825
+ sources[name] = Source(name, path, tc.seq_len)
826
+ weights[name] = float(spec["weight"])
827
+ if missing:
828
+ log0(f"[data] sources listed but NOT on disk (skipped, weights renormalized): {missing}")
829
+ extra = sorted(d for d in os.listdir(args.data_root) if os.path.exists(os.path.join(args.data_root, d, "index.json"))
830
+ and d not in {spec.get("path", n) for n, spec in mix["sources"].items()})
831
+ if extra:
832
+ log0(f"[data] shard dirs on disk but not in the mix (ignored): {extra}")
833
+ if not sources:
834
+ raise RuntimeError("no sources available")
835
+ if tc.global_batch_seqs > 0:
836
+ ga = max(1, round(tc.global_batch_seqs / (tc.micro_batch * world)))
837
+ got = ga * tc.micro_batch * world
838
+ if got != tc.global_batch_seqs:
839
+ log0(f"!! global_batch_seqs {tc.global_batch_seqs} not reachable with micro {tc.micro_batch} x world {world}; using {got} "
840
+ f"(changing the global batch changes the sampler stream -- keep it constant across resumes)")
841
+ tc.grad_accum = ga
842
+ eos_ids = {s.eos_id for s in sources.values()}
843
+ assert len(eos_ids) == 1, f"all sources must share one eos_id, got {eos_ids}"
844
+ global_batch = tc.micro_batch * tc.grad_accum * world
845
+ tokens_per_step = global_batch * tc.seq_len
846
+ sampler = MixtureSampler(sources, weights, global_batch, tc.seed)
847
+ if rank == 0:
848
+ with open(os.path.join(args.out, f"mix_realized_{int(time.time())}.json"), "w") as f:
849
+ json.dump({"mix_file": os.path.abspath(args.mix), "world": world, "global_batch_seqs": global_batch,
850
+ "grad_accum": tc.grad_accum, "sources": {n: {"prob": float(sampler.probs[i]), "tokens": sources[n].n_tokens,
851
+ "windows": sources[n].n_windows, "path": sources[n].path} for i, n in enumerate(sampler.names)},
852
+ "missing": missing, "ignored_on_disk": extra}, f, indent=1)
853
+ log0(f"[data] {len(sources)} sources, global batch {global_batch} seqs = {tokens_per_step / 1e6:.2f}M tokens/step; "
854
+ + ", ".join(f"{n}:{sources[n].n_tokens / 1e9:.2f}B tok/{sampler.probs[i]:.3f}" for i, n in enumerate(sampler.names)))
855
+
856
+ # --- resume ---
857
+ step, resumed = 0, None
858
+ if args.resume == "auto":
859
+ resumed = find_latest(args.out)
860
+ elif args.resume != "none":
861
+ resumed = args.resume
862
+ if resumed:
863
+ meta = load_checkpoint(resumed, model, opts)
864
+ step = int(meta["step"])
865
+ sampler.load_state_dict(meta["sampler"])
866
+ if "rng" in os.environ.get("MORENA_RESUME_PARTS", "model,opt,rng").split(","):
867
+ torch.set_rng_state(torch.tensor(meta["rng"]["torch"], dtype=torch.uint8))
868
+ if device.type == "cuda" and meta["rng"].get("cuda"):
869
+ torch.cuda.set_rng_state(torch.tensor(meta["rng"]["cuda"], dtype=torch.uint8))
870
+ if meta.get("decay_start", -1) >= 0 and tc.decay_start < 0 and not args.anneal:
871
+ tc.decay_start, tc.decay_steps, tc.total_steps = meta["decay_start"], meta["decay_steps"], meta["total_steps"]
872
+ fp = state_fingerprint(model, opts)
873
+ log0(f"[resume] from {resumed} at step {step}; epochs {sampler.epochs()}")
874
+ log0(f"[resume] fingerprint loaded {fp} vs saved {meta.get('fingerprint')} "
875
+ f"{'MATCH' if fp == meta.get('fingerprint') else '!! MISMATCH'}")
876
+ if args.anneal:
877
+ tc.decay_start, tc.decay_steps = step, args.anneal
878
+ tc.total_steps = step + args.anneal
879
+ log0(f"[anneal] WSD decay from step {step} for {args.anneal} steps -> total {tc.total_steps}")
880
+ if step >= tc.total_steps:
881
+ log0("[done] already at total_steps; nothing to do")
882
+ _mark_done(args.out, rank)
883
+ dist.destroy_process_group()
884
+ return
885
+
886
+ loader = Loader(sampler, rank, world, tc.micro_batch, tc.grad_accum, tc.seq_len, step)
887
+
888
+ # --- logging ---
889
+ logf = open(os.path.join(args.out, f"log_rank{rank}.jsonl" if rank else "log.jsonl"), "a") if rank == 0 else None
890
+ wb = None
891
+ if args.wandb and rank == 0:
892
+ os.environ.setdefault("WANDB_MODE", "offline")
893
+ import wandb
894
+ wb = wandb.init(project=args.wandb, dir=args.out, resume="allow", id=os.path.basename(os.path.abspath(args.out)),
895
+ config={"model": asdict(mc), "train": asdict(tc), "mix": mix})
896
+ peak_flops = 312e12 if device.type == "cuda" else 1e12
897
+ hd = mc.d_model // mc.n_head
898
+ flops_per_token = 6 * n_params + 12 * mc.n_layer * mc.d_model * tc.seq_len # fwd+bwd incl. attention
899
+ if mc.tie_embeddings:
900
+ flops_per_token += 0 # output projection already counted via tied embed params
901
+
902
+ # --- signals (SLURM sends SIGTERM/SIGUSR1 before kill) ---
903
+ stop_flag = {"v": False}
904
+
905
+ def _sig(signum, frame):
906
+ stop_flag["v"] = True
907
+ log0(f"[signal] {signum} received -> will checkpoint and exit")
908
+ for s in (signal.SIGTERM, signal.SIGUSR1):
909
+ signal.signal(s, _sig)
910
+
911
+ def should_stop_for_time():
912
+ if deadline is None:
913
+ return False
914
+ return time.time() > deadline - args.exit_margin_min * 60
915
+
916
+ # --- train loop ---
917
+ model.train()
918
+ last_ckpt_t = time.time()
919
+ t_step = time.time()
920
+ steps_this_proc = 0
921
+ finished = False
922
+ stop_t = torch.zeros(1, device=device)
923
+ log0(f"[train] start step {step}/{tc.total_steps} attn={attn_mode} world={world}")
924
+ timing = os.environ.get("MORENA_TIMING") == "1"
925
+ tsec = {"data": 0.0, "fwdbwd": 0.0, "clip": 0.0, "muon": 0.0, "adam": 0.0, "n": 0}
926
+
927
+ def _tick():
928
+ if timing and device.type == "cuda":
929
+ torch.cuda.synchronize()
930
+ return time.time()
931
+ while step < tc.total_steps:
932
+ lr = lr_at(step, tc)
933
+ for o in opts:
934
+ for g in o.param_groups:
935
+ g["lr"] = lr if o is opt_muon else lr * (adam_lr / tc.lr)
936
+ if os.environ.get("MORENA_PROFILE") == "1" and steps_this_proc == 8:
937
+ prof = torch.profiler.profile(activities=[torch.profiler.ProfilerActivity.CPU, torch.profiler.ProfilerActivity.CUDA])
938
+ prof.__enter__()
939
+ elif os.environ.get("MORENA_PROFILE") == "1" and steps_this_proc == 11:
940
+ prof.__exit__(None, None, None)
941
+ log0(prof.key_averages().table(sort_by="cuda_time_total", row_limit=18))
942
+ # also dump param dtypes/strides/contiguity of a few params
943
+ for n_, p_ in list(model.named_parameters())[:6]:
944
+ lt = p_.to_local() if isinstance(p_, DTensor) else p_
945
+ log0(f"[param] {n_} {type(p_).__name__} {p_.dtype} local {tuple(lt.shape)} stride {lt.stride()} contig {lt.is_contiguous()} req_grad {p_.requires_grad}")
946
+ _t0 = _tick()
947
+ got_step, micro, sampler_state = loader.next()
948
+ assert got_step == step, (got_step, step)
949
+ _t1 = _tick()
950
+ loss_acc = torch.zeros(1, device=device)
951
+
952
+ # LOSS MASKING. The denominator has to be the count of unmasked target tokens across the
953
+ # WHOLE global batch, not per micro-batch: micro-batches hold different numbers of unmasked
954
+ # tokens, so averaging per-micro means silently reweights them. And FSDP averages gradients
955
+ # across ranks, so a rank scaling by its own local count would make the result depend on how
956
+ # the batch happened to shard. We therefore sum the mask over every micro-batch and every
957
+ # rank first, then scale by world_size to undo FSDP's mean. Cheap: one scalar all-reduce.
958
+ use_mask = tc.loss_mask and len(micro[0]) > 2
959
+ if use_mask:
960
+ den = torch.zeros((), device=device, dtype=torch.float32)
961
+ for _t, _e, _m in micro:
962
+ den += _m[:, 1:].to(device, non_blocking=True).sum()
963
+ if dist.is_initialized():
964
+ dist.all_reduce(den, op=dist.ReduceOp.SUM)
965
+ den = den.clamp(min=1.0)
966
+ wsz = float(dist.get_world_size()) if dist.is_initialized() else 1.0
967
+
968
+ for mi, item in enumerate(micro):
969
+ tokens, eos_id = item[0], item[1]
970
+ tmask = item[2] if len(item) > 2 else None
971
+ x, y, pos, cu, mx, mask = build_batch(tokens, eos_id, attn_mode, device)
972
+ if not args.no_fsdp and hasattr(model, "set_requires_gradient_sync"):
973
+ model.set_requires_gradient_sync(mi == len(micro) - 1)
974
+ with torch.autocast(device.type, dtype=torch.bfloat16, enabled=(device.type == "cuda")):
975
+ logits = model(x, pos, cu, mx, mask)
976
+ if use_mask:
977
+ m = tmask[:, 1:].to(device, non_blocking=True).reshape(-1).float()
978
+ tok = F.cross_entropy(logits.float().view(-1, logits.shape[-1]), y.view(-1),
979
+ reduction="none")
980
+ num = (tok * m).sum()
981
+ (num * (wsz / den)).backward()
982
+ # log the true global masked mean; the AVG all-reduce below undoes the wsz factor
983
+ loss_acc += num.detach() * (wsz / den)
984
+ else:
985
+ loss = F.cross_entropy(logits.float().view(-1, logits.shape[-1]), y.view(-1), reduction="mean")
986
+ (loss / len(micro)).backward()
987
+ loss_acc += loss.detach() / len(micro)
988
+ _t2 = _tick()
989
+ gn = torch.nn.utils.clip_grad_norm_(model.parameters(), tc.grad_clip)
990
+ if isinstance(gn, DTensor):
991
+ gn = gn.full_tensor()
992
+ _t3 = _tick()
993
+ if muon_params:
994
+ opt_muon.step()
995
+ _t4 = _tick()
996
+ opt_adam.step()
997
+ _t5 = _tick()
998
+ for o in opts:
999
+ o.zero_grad(set_to_none=True)
1000
+ if timing:
1001
+ tsec["data"] += _t1 - _t0; tsec["fwdbwd"] += _t2 - _t1; tsec["clip"] += _t3 - _t2
1002
+ tsec["muon"] += _t4 - _t3; tsec["adam"] += _t5 - _t4; tsec["n"] += 1
1003
+ if tsec["n"] % 10 == 0:
1004
+ log0("[timing] " + " ".join(f"{k} {v / tsec['n']:.3f}s" for k, v in tsec.items() if k != "n"))
1005
+ for k in tsec: tsec[k] = 0 if k == "n" else 0.0
1006
+ step += 1
1007
+ steps_this_proc += 1
1008
+
1009
+ # --- logging ---
1010
+ if step % args.log_every == 0 or step == tc.total_steps:
1011
+ if world > 1:
1012
+ dist.all_reduce(loss_acc, op=dist.ReduceOp.AVG)
1013
+ if device.type == "cuda":
1014
+ torch.cuda.synchronize()
1015
+ now = time.time()
1016
+ dt = now - t_step
1017
+ t_step = now
1018
+ tps = tokens_per_step * args.log_every / dt
1019
+ mfu = flops_per_token * tps / (world * peak_flops)
1020
+ rec = {"step": step, "loss": round(loss_acc.item(), 5), "lr": lr, "gnorm": round(float(gn), 4),
1021
+ "tok_s": round(tps), "tok_s_gpu": round(tps / world), "mfu": round(mfu, 4),
1022
+ "step_s": round(dt / args.log_every, 3), "tokens": step * tokens_per_step, "t": round(now - t_start)}
1023
+ if rank == 0:
1024
+ logf.write(json.dumps(rec) + "\n"); logf.flush()
1025
+ if wb:
1026
+ wb.log(rec, step=step)
1027
+ if step % (args.log_every * 10) == 0 or steps_this_proc <= 5:
1028
+ print(f"[step {step}] loss {rec['loss']:.4f} lr {lr:.2e} gn {rec['gnorm']:.3f} "
1029
+ f"{tps / 1e3:.1f}k tok/s mfu {mfu * 100:.1f}% {rec['step_s']:.2f}s/step "
1030
+ f"mem {torch.cuda.max_memory_allocated() / 2**30 if device.type == 'cuda' else 0:.1f}GB", flush=True)
1031
+
1032
+ # --- checkpoint / exit decisions (agreed across ranks via all_reduce of a flag) ---
1033
+ time_ckpt = (time.time() - last_ckpt_t) > args.ckpt_minutes * 60
1034
+ milestone = args.milestone_every and step % args.milestone_every == 0
1035
+ stop_time = should_stop_for_time() or stop_flag["v"]
1036
+ stop_max = args.max_steps and steps_this_proc >= args.max_steps
1037
+ finished = step >= tc.total_steps
1038
+ stop_t[0] = float(stop_time or stop_max or finished)
1039
+ flag_t = torch.tensor([float(time_ckpt or milestone)], device=device)
1040
+ if world > 1:
1041
+ dist.all_reduce(stop_t, op=dist.ReduceOp.MAX); dist.all_reduce(flag_t, op=dist.ReduceOp.MAX)
1042
+ do_stop = stop_t.item() > 0
1043
+ if flag_t.item() > 0 or do_stop:
1044
+ t0 = time.time()
1045
+ ds = tc.decay_start if tc.decay_start >= 0 else -1
1046
+ fp = state_fingerprint(model, opts)
1047
+ save_checkpoint(args.out, step, model, opts, sampler_state, {"fingerprint": fp,
1048
+ "decay_start": ds, "decay_steps": tc.decay_steps, "total_steps": tc.total_steps,
1049
+ "attn": attn_mode, "world": world, "tokens": step * tokens_per_step, "lr": lr},
1050
+ args.ckpt_keep, args.milestone_every, rank)
1051
+ last_ckpt_t = time.time()
1052
+ log0(f"[ckpt] step {step} saved in {last_ckpt_t - t0:.1f}s -> {ckpt_dir(args.out, step)}"
1053
+ + (" (milestone)" if milestone else ""))
1054
+ if rank == 0:
1055
+ logf.write(json.dumps({"event": "checkpoint", "step": step, "reason":
1056
+ "finished" if finished else "time_budget" if stop_time else "max_steps" if stop_max else
1057
+ "milestone" if milestone else "interval"}) + "\n"); logf.flush()
1058
+ if do_stop:
1059
+ break
1060
+
1061
+ loader.stop = True
1062
+ if finished:
1063
+ _mark_done(args.out, rank)
1064
+ log0(f"[done] training finished at step {step}")
1065
+ else:
1066
+ log0(f"[exit] clean exit at step {step} (not finished) — resubmit/resume to continue")
1067
+ if wb:
1068
+ wb.finish()
1069
+ dist.barrier()
1070
+ dist.destroy_process_group()
1071
+ sys.exit(0)
1072
+
1073
+
1074
+ def _mark_done(out, rank):
1075
+ if rank == 0:
1076
+ with open(os.path.join(out, "DONE"), "w") as f:
1077
+ f.write(time.strftime("%Y-%m-%d %H:%M:%S\n"))
1078
+
1079
+
1080
+ if __name__ == "__main__":
1081
+ main()
tokenizer.json ADDED
The diff for this file is too large to render. See raw diff