MORENA release: morena-1.5b-instruct
Browse files- README.md +113 -0
- SHA256SUMS +5 -0
- config.json +44 -0
- load_example.py +62 -0
- model.safetensors +3 -0
- modeling_morena.py +1081 -0
- 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
|
|
|