akoumpa commited on
Commit
b3986d2
·
verified ·
1 Parent(s): c0bd3a0

Add Moonlight-V4-1B variant config, tokenizer, model card and training helpers

Browse files
README.md ADDED
@@ -0,0 +1,186 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: mit
3
+ library_name: transformers
4
+ pipeline_tag: text-generation
5
+ tags:
6
+ - deepseek_v4
7
+ - mixture-of-experts
8
+ - moe
9
+ - hybrid-attention
10
+ - compressed-sparse-attention
11
+ - hyper-connections
12
+ - mhc
13
+ - architecture-config
14
+ - untrained
15
+ - from-scratch
16
+ - nemo-automodel
17
+ - moonlight
18
+ ---
19
+
20
+ # Moonlight-V4-1B-h16d256
21
+
22
+ **An untrained, ~1B-parameter DeepSeek-V4-architecture configuration** (no weights): DeepSeek-V4-architecture 1B config with 16 x 256 attention heads (CSA ratio 4 + lightning indexer).
23
+ It is the small end of a Moonlight-style down-scaling of DeepSeek-V4, sized to pre-train on two 48 GB GPUs, with
24
+ attention dimensions chosen so that the TileLang sparse-attention kernel fits GPUs with 99 KB of shared memory
25
+ (Ada / consumer class). This is the V4-faithful variant: 7 Compressed Sparse Attention (CSA) layers with ratio-4 overlapped compression and a lightning indexer, 6 Heavily Compressed Attention (HCA) layers (ratio 128) and 2 pure sliding-window layers.
26
+
27
+ The sibling repo replaces the ratio-4 CSA layers by ratio-8 compression without an indexer, which is cheaper to train at short context and avoids the indexer kernel entirely. Sibling: [`akoumpa/Moonlight-V4-1B-h16d256-r8`](https://huggingface.co/akoumpa/Moonlight-V4-1B-h16d256-r8).
28
+
29
+ | | total | non-embedding | activated / token | activated non-embedding |
30
+ | --- | ---: | ---: | ---: | ---: |
31
+ | parameters | **999,670,879** (999.7M) | 734.9M | 539.6M | 274.8M |
32
+
33
+ ## Lineage
34
+
35
+ | model | architecture | size | status |
36
+ | --- | --- | --- | --- |
37
+ | [Moonlight-16B-A3B](https://huggingface.co/moonshotai/Moonlight-16B-A3B) (Moonshot) | DeepSeek-V3 | 16B / 3B active | released, trained with Muon |
38
+ | Moonlight-V4-16B-A3B | DeepSeek-V4 at Moonlight's width/depth/experts | 16.5B / 3.0B active | proposed config |
39
+ | Moonlight-V4-1B (8 heads x 512) | same, shrunk | 1.0B / 0.54B active | config; eager attention only |
40
+ | **Moonlight-V4-1B-h16d256** | same, attention re-shaped for the kernels | 999.7M / 539.6M active | **this repo** |
41
+
42
+ DeepSeek-V4 (Flash: 284B/13B, Pro: 1.6T/49B) keeps `head_dim` = 512 and 64 or 128 query heads; Moonlight's 16 heads
43
+ were kept for the 16B analogue, and here 16 heads x 256 dims give the same 4096-wide query space as the 8 x 512
44
+ 1B config with half the shared-KV width.
45
+
46
+ ## Architecture
47
+
48
+ | component | setting |
49
+ | --- | --- |
50
+ | hidden size / layers | 1024 / 15 |
51
+ | attention schedule (`compress_ratios`) | `[0, 0, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4]`: 2 sliding-window, 7 compressed (ratio 4, overlapped, with indexer), 6 HCA (ratio 128) |
52
+ | attention | shared-KV MQA (`num_key_value_heads` 1): 16 query heads x `head_dim` 256 (last 64 dims RoPE), the same 256-dim entry is key and value |
53
+ | query path | `q_lora_rank` 256 -> 16 x 256; per-head RMSNorm before RoPE |
54
+ | output projection | grouped low-rank: `o_groups` 2 (8 heads per group) x `o_lora_rank` 1024 -> hidden |
55
+ | sliding window / attention sinks | 128 tokens on every layer / one learnable sink logit per head |
56
+ | Lightning indexer | 64 heads x 128 dims, `index_topk` 1024 (a power of two, as the indexer kernels require); at <= 4K tokens this is >= the 1024 pooled entries, i.e. dense |
57
+ | MoE (every layer) | 32 routed experts x 384 (SwiGLU, clamp 10.0), top-6, 1 shared expert; `sqrtsoftplus` scoring, aux-loss-free bias (`noaux_tc`), `routed_scaling_factor` 2.436 |
58
+ | hash-routed layers | first 1 MoE layer(s) route by a fixed token-id table (`tid2eid`) |
59
+ | residual stream | manifold-constrained hyper-connections, `hc_mult` 4, 20 Sinkhorn iterations |
60
+ | positions | 4096 tokens, plain RoPE (theta 10000 on sliding layers, 160000 on compressed layers, no YaRN) |
61
+ | MTP | none (`num_nextn_predict_layers` 0) |
62
+ | vocabulary | 129280 (DeepSeek-V4 tokenizer, BOS 0, EOS 1) |
63
+ | norm eps / init | 1e-06 / 0.02 |
64
+
65
+ Per-layer parameters: attention 7.9M (sliding) / 11.6M (ratio 4) / 8.4M (HCA); MoE 39.0M
66
+ total, 8.3M activated (each expert 1.2M); mHC mixers 196,662. Embedding and
67
+ head are 132.4M each. KV cache per sequence at 4K / 32K tokens (FP8 non-RoPE dims, bf16 RoPE dims): 3.3 MiB / 22.1 MiB.
68
+ Core-attention FLOPs per generated token at 4K context: 0.27 GF, against 0.81 GF of linear layers.
69
+ `python count_params.py config.json` reproduces these numbers (the script also reproduces the published sizes of
70
+ Moonlight, DeepSeek-V3, V4-Flash and V4-Pro).
71
+
72
+ Design rules: `head_dim` and `index_topk` are powers of two and `index_n_heads` divides 128 (kernel requirements);
73
+ `o_lora_rank` 1024 keeps DeepSeek-V4's per-group output projection shape; `routed_scaling_factor` follows Moonlight's
74
+ recipe (expected 1/||p||_2 of renormalised top-k scores) applied to sqrt-softplus with 32 experts / top-6.
75
+
76
+ ## Files
77
+
78
+ | file | purpose |
79
+ | --- | --- |
80
+ | `config.json` | Hugging Face config (same key set as `deepseek-ai/DeepSeek-V4-Flash`) |
81
+ | `inference_config.json` | the same model in the key format of DeepSeek's reference `inference/model.py` |
82
+ | `tokenizer.json`, `tokenizer_config.json` | DeepSeek-V4 tokenizer (MIT), copied from `deepseek-ai/DeepSeek-V4-Flash` |
83
+ | `count_params.py` | parameter / KV-cache / FLOP calculator for DeepSeek-V3- and V4-style configs |
84
+ | `training/pretrain.yaml`, `training/train.py` | NeMo Automodel recipe (FSDP2, 2 GPUs) and launcher that seeds the hash table and mHC mixers |
85
+ | `training/prepare_data.py`, `training/finite_nanogpt.py`, `training/init_utils.py` | data shards from parquet text, bounded validation dataset, from-scratch initialisers |
86
+
87
+ ## Loading
88
+
89
+ ### transformers (>= 5.8, native `deepseek_v4`)
90
+
91
+ ```python
92
+ from transformers import AutoTokenizer, DeepseekV4Config, DeepseekV4ForCausalLM
93
+ cfg = DeepseekV4Config.from_pretrained("akoumpa/Moonlight-V4-1B-h16d256") # legacy keys (compress_ratios, num_hash_layers, ...) are folded in
94
+ tok = AutoTokenizer.from_pretrained("akoumpa/Moonlight-V4-1B-h16d256")
95
+ model = DeepseekV4ForCausalLM(cfg) # random init, 999.7M parameters
96
+ ```
97
+
98
+ The transformers implementation is **inference-oriented**: with no KV cache it appends compressed entries to the
99
+ key axis without a causal mask and gathers per-query top-k entries into an `S x k` key axis, so do not train with it.
100
+ Use it for architecture inspection and, once you have weights, for generation.
101
+
102
+ ### NeMo Automodel (native training implementation)
103
+
104
+ ```python
105
+ from nemo_automodel.components.models.common import BackendConfig
106
+ from nemo_automodel.components.models.deepseek_v4.config import DeepseekV4Config
107
+ from nemo_automodel.components.models.deepseek_v4.model import DeepseekV4ForCausalLM
108
+ cfg = DeepseekV4Config.from_pretrained("akoumpa/Moonlight-V4-1B-h16d256")
109
+ backend = BackendConfig(attn="eager", linear="torch", rms_norm="torch_fp32", rope_fusion=False,
110
+ dispatcher="torch", experts="torch_mm", enable_hf_state_dict_adapter=False)
111
+ model = DeepseekV4ForCausalLM(cfg, backend=backend)
112
+ model.initialize_weights(dtype=torch.bfloat16)
113
+ ```
114
+
115
+ Or in a recipe: `NeMoAutoModelForCausalLM.from_config` with `config: DeepseekV4Config.from_pretrained("akoumpa/Moonlight-V4-1B-h16d256")`
116
+ (see `training/pretrain.yaml`).
117
+
118
+ ### DeepSeek reference code
119
+
120
+ `inference_config.json` drops into the `inference/` folder of the DeepSeek-V4 release (`ModelArgs` keys, `n_mtp_layers` 0).
121
+
122
+ ## Training from scratch
123
+
124
+ The released implementations only ever load checkpoints, so two things must be initialised by hand; `training/init_utils.py`
125
+ does both and `training/train.py` calls it after the recipe's own weight init on a fresh start:
126
+
127
+ 1. **Hash-routing table.** `tid2eid` is created as zeros (every token to expert 0). `fill_hash_tables` writes a balanced
128
+ token-id hash: 32 experts, 6 distinct experts per token, equal load over the vocabulary.
129
+ 2. **mHC mixers.** NeMo Automodel leaves the `fn` / `base` / `scale` tensors uninitialised; `init_hyper_connections` applies
130
+ transformers' rule (normal(0, 0.02) projection, zero bias, unit gates).
131
+
132
+ Recipe notes that cost time to find (all encoded in `training/`):
133
+
134
+ - Automodel wraps DeepSeek-V4's fp32 tensors (attention sinks, compressor position biases, mHC mixers, `lm_head`) as their
135
+ own FSDP2 units whose forward returns the parameter; if they reshard after forward, attention reads a freed tensor.
136
+ `train.py` calls `set_reshard_after_forward(False)` on those units.
137
+ - Use the logits-based `MaskedCrossEntropy`: the fused linear cross-entropy rejects the fp32 `lm_head` x bf16 hidden states.
138
+ - For iterable datasets the recipe passes no batch size to the DataLoader; set `dataloader.batch_size` explicitly.
139
+ - `NanogptDataset` is an infinite stream; validation uses `FiniteNanogptDataset`.
140
+ - The indexer's top-k has no gradient path (Automodel freezes its parameters). With `index_topk` 1024 every query sees all causal compressed entries up to 4K tokens, i.e. DeepSeek's own dense warm-up regime; sparse training at longer contexts needs an indexer distillation loss.
141
+ - To train beyond 4K, add DeepSeek-V4's YaRN block (`rope_scaling`: factor 16, `original_max_position_embeddings` 65536)
142
+ and raise `max_position_embeddings` and `index_topk`.
143
+
144
+ Measured on two RTX 5880 Ada (48 GB, sm_89) GPUs, bf16, `torch_mm` experts, chunked cross-entropy, single-GPU forward+backward
145
+ (TileLang sparse attention with Sinkhorn and indexer on torch; the eager path for the same model reaches 4.2k tok/s at
146
+ B=2 x 2048 and runs out of memory at B=4):
147
+
148
+ | micro-batch | fwd+bwd time, throughput, peak memory |
149
+ | --- | --- |
150
+ | 2 x 2048 | 740 ms, 5.5k tok/s, 15.1 GiB |
151
+ | 4 x 2048 | 1227 ms, 6.7k tok/s, 27.9 GiB |
152
+ | 2 x 4096 | 1421 ms, 5.8k tok/s, 33.1 GiB |
153
+
154
+ Forward-time breakdown at B=2 x 2048: attention 59 ms, indexer 82 ms, MoE 39 ms, mHC mixers 35 ms, compressor 12 ms (247 ms forward). With the recipe's FSDP2 data parallelism over 2 GPUs and
155
+ AdamW, a 32-sequence x 2048-token global batch is a reasonable starting point (`training/pretrain.yaml`).
156
+
157
+ ## TileLang kernels
158
+
159
+ Automodel's vendored Miles/TileLang kernels (sparse attention, indexer) and DeepSeek's TileKernels Sinkhorn were written
160
+ for Hopper's 227 KB of shared memory. On a 99 KB-per-block GPU:
161
+
162
+ | kernel | shape rule | shared memory | Ada (99 KB) |
163
+ | --- | --- | --- | --- |
164
+ | sparse attention fwd/bwd | `head_dim` power of two; heads padded to 16, chunked by 16, multiples of 64 above 64 | 147 KB at `head_dim` 512, fits at **256** | runs with this config; parity with the torch reference verified |
165
+ | lightning indexer fwd | `index_n_heads` <= 64, multiple of 8, divides 128; `index_topk` power of two (bwd) | 224 KB (`block_N` 256 x 128 fp32) | does not fit; run the indexer on torch |
166
+ | Sinkhorn (mHC) | needs the `tile_kernels` package | small | needs TileKernels installed or the torch fallback |
167
+
168
+ Automodel currently switches all three together (`backend.attn: tilelang`); the measurements above used per-kernel
169
+ selection (only sparse attention on TileLang). On Hopper GPUs the default `head_dim` 512 kernels fit and the 16 x 256
170
+ choice is optional.
171
+
172
+ ## Limitations and notes
173
+
174
+ - No trained weights are provided; all numbers are architecture-derived or short from-scratch measurements.
175
+ - The name follows the Moonlight / DeepSeek-V4 lineage for clarity; this repository is not affiliated with Moonshot AI
176
+ or DeepSeek.
177
+ - The tokenizer files are DeepSeek's (MIT licence, `deepseek-ai/DeepSeek-V4-Flash`).
178
+
179
+ ## References
180
+
181
+ - Liu et al., *Muon is Scalable for LLM Training* (Moonlight), [arXiv:2502.16982](https://arxiv.org/abs/2502.16982)
182
+ - DeepSeek-AI, *DeepSeek-V4: Towards Highly Efficient Million-Token Context Intelligence*, [arXiv:2606.19348](https://arxiv.org/abs/2606.19348)
183
+ - Xie et al., *Manifold-Constrained Hyper-Connections* (mHC), 2026
184
+ - Roller et al., *Hash Layers for Large Sparse Models*, NeurIPS 2021
185
+ - DeepSeek-AI, *DeepSeek-V3.2* (DeepSeek Sparse Attention, lightning indexer), 2025
186
+ - Kernels: [Miles](https://github.com/yueming-yuan/miles) (sparse attention / indexer, vendored in NeMo Automodel), [TileKernels](https://github.com/deepseek-ai/TileKernels) (Sinkhorn), [TileLang](https://github.com/tile-ai/tilelang)
config.json ADDED
@@ -0,0 +1,65 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "DeepseekV4ForCausalLM"
4
+ ],
5
+ "attention_bias": false,
6
+ "attention_dropout": 0.0,
7
+ "bos_token_id": 0,
8
+ "eos_token_id": 1,
9
+ "hc_eps": 1e-06,
10
+ "hc_mult": 4,
11
+ "hc_sinkhorn_iters": 20,
12
+ "head_dim": 256,
13
+ "hidden_act": "silu",
14
+ "hidden_size": 1024,
15
+ "index_head_dim": 128,
16
+ "index_n_heads": 64,
17
+ "index_topk": 1024,
18
+ "initializer_range": 0.02,
19
+ "max_position_embeddings": 4096,
20
+ "model_type": "deepseek_v4",
21
+ "moe_intermediate_size": 384,
22
+ "n_routed_experts": 32,
23
+ "n_shared_experts": 1,
24
+ "norm_topk_prob": true,
25
+ "num_attention_heads": 16,
26
+ "num_experts_per_tok": 6,
27
+ "num_hidden_layers": 15,
28
+ "num_hash_layers": 1,
29
+ "num_key_value_heads": 1,
30
+ "num_nextn_predict_layers": 0,
31
+ "o_groups": 2,
32
+ "o_lora_rank": 1024,
33
+ "q_lora_rank": 256,
34
+ "qk_rope_head_dim": 64,
35
+ "rms_norm_eps": 1e-06,
36
+ "rope_scaling": null,
37
+ "rope_theta": 10000,
38
+ "routed_scaling_factor": 2.436,
39
+ "scoring_func": "sqrtsoftplus",
40
+ "sliding_window": 128,
41
+ "swiglu_limit": 10.0,
42
+ "tie_word_embeddings": false,
43
+ "topk_method": "noaux_tc",
44
+ "torch_dtype": "bfloat16",
45
+ "use_cache": true,
46
+ "vocab_size": 129280,
47
+ "compress_rope_theta": 160000,
48
+ "compress_ratios": [
49
+ 0,
50
+ 0,
51
+ 4,
52
+ 128,
53
+ 4,
54
+ 128,
55
+ 4,
56
+ 128,
57
+ 4,
58
+ 128,
59
+ 4,
60
+ 128,
61
+ 4,
62
+ 128,
63
+ 4
64
+ ]
65
+ }
count_params.py ADDED
@@ -0,0 +1,316 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """
3
+ Parameter / KV-cache / attention-FLOP calculator for
4
+ * DeepSeek-V3-style configs (Moonlight-16B-A3B, DeepSeek-V3)
5
+ * DeepSeek-V4-style configs (DeepSeek-V4-Flash / -Pro, Moonlight-V4)
6
+
7
+ Formulas follow the tensor shapes of the released checkpoints:
8
+ * Moonlight: modeling_deepseek.py (DeepseekV3) shapes read from model.safetensors headers
9
+ * DeepSeek-V4: inference/model.py (ModelArgs / Block / Attention / Compressor / Indexer / MoE)
10
+
11
+ Published numbers used for validation (all reproduced by this script):
12
+ Moonlight-16B-A3B : 15.29B total / 2.24B activated non-embedding; 16B / 3B incl. embedding
13
+ DeepSeek-V4-Flash : 284B total / 13B activated
14
+ DeepSeek-V4-Pro : 1.6T total / 49B activated
15
+
16
+ Usage:
17
+ python3 count_params.py # validation + Moonlight-V4 report
18
+ python3 count_params.py path/to/config.json [...] # report for arbitrary HF-style configs
19
+ """
20
+ import json
21
+ import math
22
+ import random
23
+ import sys
24
+ from pathlib import Path
25
+
26
+ HERE = Path(__file__).resolve().parent
27
+
28
+ # ----------------------------------------------------------------------------------------------
29
+ # DeepSeek-V3 architecture (Moonlight, DeepSeek-V3)
30
+ # ----------------------------------------------------------------------------------------------
31
+ def count_v3(c):
32
+ V, D, L = c["vocab_size"], c["hidden_size"], c["num_hidden_layers"]
33
+ H = c["num_attention_heads"]
34
+ qk_nope, qk_rope, v_hd = c["qk_nope_head_dim"], c["qk_rope_head_dim"], c["v_head_dim"]
35
+ kv_r, q_r = c["kv_lora_rank"], c.get("q_lora_rank")
36
+ E, K, I = c["n_routed_experts"], c["num_experts_per_tok"], c["moe_intermediate_size"]
37
+ n_sh, I_dense = c["n_shared_experts"], c["intermediate_size"]
38
+ k_dense, moe_freq = c["first_k_dense_replace"], c.get("moe_layer_freq", 1)
39
+ n_mtp = c.get("num_nextn_predict_layers", 0)
40
+
41
+ # attention (MLA)
42
+ if q_r:
43
+ attn = D * q_r + q_r + q_r * H * (qk_nope + qk_rope) # q_a_proj, q_a_layernorm, q_b_proj
44
+ else:
45
+ attn = D * H * (qk_nope + qk_rope) # q_proj
46
+ attn += D * (kv_r + qk_rope) + kv_r # kv_a_proj_with_mqa, kv_a_layernorm
47
+ attn += kv_r * H * (qk_nope + v_hd) # kv_b_proj
48
+ attn += H * v_hd * D # o_proj
49
+ norms = 2 * D
50
+
51
+ dense_mlp = 3 * D * I_dense
52
+ gate = E * D + E # gate.weight + e_score_correction_bias
53
+ expert = 3 * D * I
54
+ moe_total = gate + E * expert + n_sh * expert
55
+ moe_active = gate + K * expert + n_sh * expert
56
+
57
+ def is_moe(l):
58
+ return l >= k_dense and (l - k_dense) % moe_freq == 0 if moe_freq > 1 else l >= k_dense
59
+
60
+ n_moe = sum(1 for l in range(L) if is_moe(l))
61
+ n_dense = L - n_moe
62
+ body_total = L * (attn + norms) + n_dense * dense_mlp + n_moe * moe_total + D
63
+ body_active = L * (attn + norms) + n_dense * dense_mlp + n_moe * moe_active + D
64
+ emb = V * D
65
+ head = V * D
66
+ # MTP module (DeepSeek-V3 style): eh_proj (2D x D) + enorm + hnorm + shared_head.norm + one full layer
67
+ mtp_total = n_mtp * (2 * D * D + 3 * D + attn + norms + moe_total)
68
+ mtp_active = n_mtp * (2 * D * D + 3 * D + attn + norms + moe_active)
69
+ return dict(
70
+ family="deepseek_v3", layers=L, n_dense=n_dense, n_moe=n_moe,
71
+ attn_per_layer=attn, dense_mlp=dense_mlp, moe_total_per_layer=moe_total,
72
+ moe_active_per_layer=moe_active, expert=expert,
73
+ embed=emb, head=head, body_total=body_total, body_active=body_active,
74
+ total_nonembed=body_total, active_nonembed=body_active,
75
+ total_incl_embed=body_total + emb + head, active_incl_embed=body_active + emb + head,
76
+ mtp_total=mtp_total, mtp_active=mtp_active,
77
+ )
78
+
79
+
80
+ # ----------------------------------------------------------------------------------------------
81
+ # DeepSeek-V4 architecture (Flash, Pro, Moonlight-V4)
82
+ # ----------------------------------------------------------------------------------------------
83
+ def v4_attention_params(c, ratio):
84
+ """Parameters of one Attention module (incl. compressor / indexer) for a layer with the given compress ratio."""
85
+ D, H, hd = c["hidden_size"], c["num_attention_heads"], c["head_dim"]
86
+ Qr, Or, G = c["q_lora_rank"], c["o_lora_rank"], c["o_groups"]
87
+ IH, Ihd = c["index_n_heads"], c["index_head_dim"]
88
+ p = {}
89
+ p["attn_sink"] = H
90
+ p["wq_a"] = D * Qr
91
+ p["q_norm"] = Qr
92
+ p["wq_b"] = Qr * H * hd
93
+ p["wkv"] = D * hd
94
+ p["kv_norm"] = hd
95
+ p["wo_a"] = (H * hd // G) * (G * Or) # == H*hd*Or, applied group-wise as [G, Or, H*hd/G]
96
+ p["wo_b"] = G * Or * D
97
+ if ratio:
98
+ coff = 2 if ratio == 4 else 1 # overlapped compression only for the CSA ratio (4)
99
+ p["compressor"] = ratio * coff * hd + 2 * D * coff * hd + hd # ape + wkv + wgate + norm
100
+ if ratio == 4: # CSA: lightning indexer on top of compressed keys
101
+ idx = Qr * IH * Ihd + D * IH # wq_b + weights_proj
102
+ idx += 4 * 2 * Ihd + 2 * D * 2 * Ihd + Ihd # indexer.compressor (ratio 4, overlap)
103
+ p["indexer"] = idx
104
+ return p
105
+
106
+
107
+ def v4_hc_params(c):
108
+ hc, D = c["hc_mult"], c["hidden_size"]
109
+ mix = (2 + hc) * hc
110
+ return 2 * (mix * hc * D + mix + 3) # (fn + base + scale) for attn and for ffn
111
+
112
+
113
+ def count_v4(c):
114
+ V, D, L = c["vocab_size"], c["hidden_size"], c["num_hidden_layers"]
115
+ E, K, I = c["n_routed_experts"], c["num_experts_per_tok"], c["moe_intermediate_size"]
116
+ n_sh, n_hash = c["n_shared_experts"], c["num_hash_layers"]
117
+ hc = c["hc_mult"]
118
+ n_mtp = c.get("num_nextn_predict_layers", 0)
119
+ ratios = list(c["compress_ratios"])
120
+ assert len(ratios) >= L + n_mtp, f"compress_ratios needs >= {L + n_mtp} entries, got {len(ratios)}"
121
+
122
+ expert = 3 * D * I
123
+ gate_w, gate_b = E * D, E
124
+ moe_total = lambda hash_layer: gate_w + (0 if hash_layer else gate_b) + E * expert + n_sh * expert
125
+ moe_active = lambda hash_layer: gate_w + (0 if hash_layer else gate_b) + K * expert + n_sh * expert
126
+ hc_per_layer = v4_hc_params(c)
127
+ norms = 2 * D
128
+
129
+ per_layer = []
130
+ for l in range(L):
131
+ ap = v4_attention_params(c, ratios[l])
132
+ attn = sum(ap.values())
133
+ is_hash = l < n_hash
134
+ per_layer.append(dict(
135
+ layer=l, ratio=ratios[l], hash=is_hash, attn=attn, attn_parts=ap, hc=hc_per_layer,
136
+ total=attn + hc_per_layer + norms + moe_total(is_hash),
137
+ active=attn + hc_per_layer + norms + moe_active(is_hash),
138
+ ))
139
+ head_hc = hc * hc * D + hc + 1
140
+ body_total = sum(x["total"] for x in per_layer) + D + head_hc
141
+ body_active = sum(x["active"] for x in per_layer) + D + head_hc
142
+ emb, head = V * D, V * D
143
+
144
+ mtp_total = mtp_active = 0
145
+ for i in range(n_mtp):
146
+ ap = sum(v4_attention_params(c, ratios[L + i]).values())
147
+ extra = 2 * D * D + 3 * D + head_hc # e_proj, h_proj, enorm, hnorm, norm, hc_head_*
148
+ mtp_total += ap + hc_per_layer + norms + moe_total(False) + extra
149
+ mtp_active += ap + hc_per_layer + norms + moe_active(False) + extra
150
+
151
+ kinds = {0: "SWA", 4: "CSA", 128: "HCA"}
152
+ layer_kinds = {k: sum(1 for l in range(L) if kinds.get(ratios[l]) == k) for k in kinds.values()}
153
+ return dict(
154
+ family="deepseek_v4", layers=L, layer_kinds=layer_kinds, n_hash=n_hash,
155
+ per_layer=per_layer, expert=expert, hc_per_layer=hc_per_layer,
156
+ moe_total_per_layer=moe_total(False), moe_active_per_layer=moe_active(False),
157
+ embed=emb, head=head, body_total=body_total, body_active=body_active,
158
+ total_nonembed=body_total, active_nonembed=body_active,
159
+ total_incl_embed=body_total + emb + head, active_incl_embed=body_active + emb + head,
160
+ mtp_total=mtp_total, mtp_active=mtp_active,
161
+ tid2eid_entries=n_hash * V * K, # non-trainable hash-routing lookup tables
162
+ )
163
+
164
+
165
+ def count(c):
166
+ return count_v4(c) if c.get("model_type") == "deepseek_v4" or "compress_ratios" in c else count_v3(c)
167
+
168
+
169
+ # ----------------------------------------------------------------------------------------------
170
+ # KV cache and attention FLOPs per token as a function of context length
171
+ # ----------------------------------------------------------------------------------------------
172
+ def kv_cache_bytes_v3(c, L, kv_bytes=2):
173
+ """MLA cache: one (kv_lora_rank + rope) latent per token per layer."""
174
+ per_tok = (c["kv_lora_rank"] + c["qk_rope_head_dim"]) * kv_bytes
175
+ return c["num_hidden_layers"] * per_tok * L
176
+
177
+
178
+ def kv_cache_bytes_v4(c, L, mixed=True):
179
+ """V4 cache. Per entry: head_dim dims; mixed storage = fp8 for non-rope dims + bf16 for the 64 rope dims
180
+ (paper sec. 2.3.4). CSA layers additionally cache indexer keys (index_head_dim, FP4)."""
181
+ hd, rd = c["head_dim"], c["qk_rope_head_dim"]
182
+ win = c["sliding_window"]
183
+ entry = (hd - rd) * 1 + rd * 2 if mixed else hd * 2
184
+ idx_entry = c["index_head_dim"] * (0.5 if mixed else 2)
185
+ total = 0
186
+ for r in c["compress_ratios"][: c["num_hidden_layers"]]:
187
+ n = min(win, L)
188
+ if r:
189
+ n += L // r
190
+ total += n * entry
191
+ if r == 4:
192
+ total += (L // r) * idx_entry
193
+ return total
194
+
195
+
196
+ def attn_flops_v3(c, L):
197
+ """Core attention FLOPs for ONE query token attending to L cached tokens, all layers.
198
+ Naive (non-absorbed) MLA: QK over (nope+rope) dims and PV over v dims, per head."""
199
+ H = c["num_attention_heads"]
200
+ qk = c["qk_nope_head_dim"] + c["qk_rope_head_dim"]
201
+ return c["num_hidden_layers"] * 2 * H * (qk + c["v_head_dim"]) * L
202
+
203
+
204
+ def attn_flops_v4(c, L):
205
+ """Core attention (+ lightning indexer) FLOPs for ONE query token at context L, all layers.
206
+ KV entries serve as both key and value: 2*H*hd per entry for QK and 2*H*hd for PV."""
207
+ H, hd, win = c["num_attention_heads"], c["head_dim"], c["sliding_window"]
208
+ IH, Ihd, topk = c["index_n_heads"], c["index_head_dim"], c["index_topk"]
209
+ total = 0
210
+ for r in c["compress_ratios"][: c["num_hidden_layers"]]:
211
+ n = min(win, L)
212
+ if r == 4:
213
+ n += min(L // r, topk)
214
+ total += 2 * IH * Ihd * (L // r) # indexer scores over all compressed keys
215
+ elif r:
216
+ n += L // r
217
+ total += 4 * H * hd * n
218
+ return total
219
+
220
+
221
+ # ----------------------------------------------------------------------------------------------
222
+ # Routed scaling factor, Moonlight's recipe (appendix C, fig. 6) generalised to any scoring function
223
+ # ----------------------------------------------------------------------------------------------
224
+ def gate_scaling_factor(num_experts, topk, score="sigmoid", iters=200_000, seed=0):
225
+ rng = random.Random(seed)
226
+ if score == "sigmoid":
227
+ f = lambda x: 1.0 / (1.0 + math.exp(-x))
228
+ elif score == "sqrtsoftplus":
229
+ f = lambda x: math.sqrt(math.log1p(math.exp(x)))
230
+ else:
231
+ raise ValueError(score)
232
+ acc = 0.0
233
+ for _ in range(iters):
234
+ p = sorted((f(rng.gauss(0, 1)) for _ in range(num_experts)), reverse=True)[:topk]
235
+ s = sum(p)
236
+ acc += 1.0 / math.sqrt(sum((x / s) ** 2 for x in p))
237
+ return acc / iters
238
+
239
+
240
+ # ----------------------------------------------------------------------------------------------
241
+ def fmt(n):
242
+ if n >= 1e12:
243
+ return f"{n/1e12:.3f}T"
244
+ if n >= 1e9:
245
+ return f"{n/1e9:.3f}B"
246
+ if n >= 1e6:
247
+ return f"{n/1e6:.2f}M"
248
+ return f"{n:,}"
249
+
250
+
251
+ def gb(x):
252
+ return f"{x/2**30:.3f} GiB" if x >= 2**30 else f"{x/2**20:.1f} MiB"
253
+
254
+
255
+ def report(name, c):
256
+ r = count(c)
257
+ print(f"== {name} ({r['family']}) ==")
258
+ print(f" layers={r['layers']}" + (f" kinds={r['layer_kinds']} hash_layers={r['n_hash']}" if r["family"] == "deepseek_v4"
259
+ else f" dense={r['n_dense']} moe={r['n_moe']}"))
260
+ print(f" embed={fmt(r['embed'])} head={fmt(r['head'])} expert={fmt(r['expert'])}")
261
+ print(f" total non-embedding: {fmt(r['total_nonembed'])} incl. embed+head: {fmt(r['total_incl_embed'])}")
262
+ print(f" active non-embedding: {fmt(r['active_nonembed'])} incl. embed+head: {fmt(r['active_incl_embed'])}")
263
+ if r["mtp_total"]:
264
+ print(f" MTP module(s): total {fmt(r['mtp_total'])}, active {fmt(r['mtp_active'])}"
265
+ f" -> grand total incl. MTP {fmt(r['total_incl_embed'] + r['mtp_total'])}")
266
+ if r["family"] == "deepseek_v4":
267
+ for kind in ("SWA", "CSA", "HCA"):
268
+ ex = next((x for x in r["per_layer"] if {0: "SWA", 4: "CSA", 128: "HCA"}[x["ratio"]] == kind), None)
269
+ if ex:
270
+ parts = ", ".join(f"{k}={fmt(v)}" for k, v in ex["attn_parts"].items())
271
+ print(f" attention/{kind} layer: {fmt(ex['attn'])} [{parts}]")
272
+ print(f" mHC per layer: {fmt(r['hc_per_layer'])} tid2eid entries (non-trainable): {fmt(r['tid2eid_entries'])}")
273
+ else:
274
+ print(f" attention per layer: {fmt(r['attn_per_layer'])} dense MLP: {fmt(r['dense_mlp'])}")
275
+ print(f" MoE per layer: total {fmt(r['moe_total_per_layer'])}, active {fmt(r['moe_active_per_layer'])}")
276
+ return r
277
+
278
+
279
+ def main(argv):
280
+ if argv:
281
+ for p in argv:
282
+ report(Path(p).stem, json.load(open(p)))
283
+ return
284
+ ref = HERE / "reference"
285
+ moonlight = json.load(open(ref / "moonlight_16b_a3b_config.json"))
286
+ flash = json.load(open(ref / "deepseek_v4_flash_config.json"))
287
+ pro = json.load(open(ref / "deepseek_v4_pro_config.json"))
288
+ mv4 = json.load(open(HERE / "config.json"))
289
+
290
+ print("#" * 100 + "\n# Validation against published numbers\n" + "#" * 100)
291
+ r_ml = report("Moonlight-16B-A3B (published: 15.29B/2.24B non-embed, 16B/3B incl. embed)", moonlight)
292
+ r_fl = report("DeepSeek-V4-Flash (published: 284B total, 13B active)", flash)
293
+ r_pr = report("DeepSeek-V4-Pro (published: 1.6T total, 49B active)", pro)
294
+ print("\n" + "#" * 100 + "\n# Moonlight-V4 (this repo)\n" + "#" * 100)
295
+ r_m4 = report("Moonlight-V4-16B-A3B", mv4)
296
+
297
+ print("\n#### Routed scaling factor via Moonlight's recipe (E[1/||p||_2] over renormalised top-k scores)")
298
+ for (E, K, s) in [(64, 6, "sigmoid"), (64, 6, "sqrtsoftplus"), (256, 6, "sqrtsoftplus"), (384, 6, "sqrtsoftplus"), (256, 8, "sigmoid")]:
299
+ print(f" experts={E:4d} topk={K} score={s:12s}: {gate_scaling_factor(E, K, s, iters=20000):.3f}")
300
+
301
+ print("\n#### KV cache per sequence (Moonlight: bf16 MLA latent; V4: fp8 non-rope + bf16 rope + fp4 indexer keys)")
302
+ for L in (8192, 65536, 1048576):
303
+ print(f" L={L:>8}: Moonlight {gb(kv_cache_bytes_v3(moonlight, L)):>12} | Moonlight-V4 {gb(kv_cache_bytes_v4(mv4, L)):>12}"
304
+ f" (bf16-only {gb(kv_cache_bytes_v4(mv4, L, mixed=False)):>12}) | V4-Flash {gb(kv_cache_bytes_v4(flash, L)):>12}")
305
+
306
+ print("\n#### Per-token FLOPs at context L (2*active params incl. head + core attention [+ indexer])")
307
+ for L in (8192, 65536, 1048576):
308
+ lin_ml = 2 * (r_ml["active_nonembed"] + r_ml["head"])
309
+ lin_m4 = 2 * (r_m4["active_nonembed"] + r_m4["head"])
310
+ a_ml, a_m4 = attn_flops_v3(moonlight, L), attn_flops_v4(mv4, L)
311
+ print(f" L={L:>8}: Moonlight linear {lin_ml/1e9:6.2f} GF + attn {a_ml/1e9:8.2f} GF = {(lin_ml+a_ml)/1e9:8.2f} GF"
312
+ f" | Moonlight-V4 linear {lin_m4/1e9:6.2f} GF + attn {a_m4/1e9:6.2f} GF = {(lin_m4+a_m4)/1e9:6.2f} GF")
313
+
314
+
315
+ if __name__ == "__main__":
316
+ main(sys.argv[1:])
inference_config.json ADDED
@@ -0,0 +1,50 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "vocab_size": 129280,
3
+ "dim": 1024,
4
+ "moe_inter_dim": 384,
5
+ "n_layers": 15,
6
+ "n_hash_layers": 1,
7
+ "n_mtp_layers": 0,
8
+ "n_heads": 16,
9
+ "n_routed_experts": 32,
10
+ "n_shared_experts": 1,
11
+ "n_activated_experts": 6,
12
+ "score_func": "sqrtsoftplus",
13
+ "route_scale": 2.436,
14
+ "swiglu_limit": 10.0,
15
+ "q_lora_rank": 256,
16
+ "head_dim": 256,
17
+ "rope_head_dim": 64,
18
+ "o_groups": 2,
19
+ "o_lora_rank": 1024,
20
+ "window_size": 128,
21
+ "original_seq_len": 0,
22
+ "rope_theta": 10000,
23
+ "rope_factor": 1,
24
+ "beta_fast": 32,
25
+ "beta_slow": 1,
26
+ "index_n_heads": 64,
27
+ "index_head_dim": 128,
28
+ "index_topk": 1024,
29
+ "hc_mult": 4,
30
+ "hc_sinkhorn_iters": 20,
31
+ "dtype": "bf16",
32
+ "compress_rope_theta": 160000,
33
+ "compress_ratios": [
34
+ 0,
35
+ 0,
36
+ 4,
37
+ 128,
38
+ 4,
39
+ 128,
40
+ 4,
41
+ 128,
42
+ 4,
43
+ 128,
44
+ 4,
45
+ 128,
46
+ 4,
47
+ 128,
48
+ 4
49
+ ]
50
+ }
tokenizer.json ADDED
The diff for this file is too large to render. See raw diff
 
tokenizer_config.json ADDED
@@ -0,0 +1,34 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "add_bos_token": false,
3
+ "add_eos_token": false,
4
+ "bos_token": {
5
+ "__type": "AddedToken",
6
+ "content": "<|begin▁of▁sentence|>",
7
+ "lstrip": false,
8
+ "normalized": true,
9
+ "rstrip": false,
10
+ "single_word": false
11
+ },
12
+ "clean_up_tokenization_spaces": false,
13
+ "eos_token": {
14
+ "__type": "AddedToken",
15
+ "content": "<|end▁of▁sentence|>",
16
+ "lstrip": false,
17
+ "normalized": true,
18
+ "rstrip": false,
19
+ "single_word": false
20
+ },
21
+ "legacy": true,
22
+ "model_max_length": 1048576,
23
+ "pad_token": {
24
+ "__type": "AddedToken",
25
+ "content": "<|end▁of▁sentence|>",
26
+ "lstrip": false,
27
+ "normalized": true,
28
+ "rstrip": false,
29
+ "single_word": false
30
+ },
31
+ "sp_model_kwargs": {},
32
+ "unk_token": null,
33
+ "tokenizer_class": "PreTrainedTokenizerFast"
34
+ }
training/finite_nanogpt.py ADDED
@@ -0,0 +1,22 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Bounded variant of NeMo Automodel's NanogptDataset for validation.
2
+
3
+ ``NanogptDataset`` is an infinite stream (its file iterator restarts the shard list forever, which is what
4
+ training wants), so a validation epoch over it never terminates. This subclass stops after ``max_samples``
5
+ samples per iterator (i.e. per data-parallel rank / dataloader worker).
6
+ """
7
+ from nemo_automodel.components.datasets.llm.nanogpt_dataset import NanogptDataset
8
+
9
+
10
+ class FiniteNanogptDataset(NanogptDataset):
11
+ def __init__(self, *args, max_samples: int = 64, **kwargs):
12
+ super().__init__(*args, **kwargs)
13
+ self.max_samples = int(max_samples)
14
+
15
+ def __iter__(self):
16
+ for i, sample in enumerate(super().__iter__()):
17
+ if i >= self.max_samples:
18
+ return
19
+ yield sample
20
+
21
+ def __len__(self) -> int: # type: ignore[override]
22
+ return self.max_samples
training/init_utils.py ADDED
@@ -0,0 +1,90 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Helpers for training DeepSeek-V4-architecture models from scratch.
2
+
3
+ DeepSeek-V4 routes the first ``num_hash_layers`` MoE layers with a fixed token-id -> expert-id table
4
+ (``tid2eid``, Roller et al. 2021 "Hash Layers"). Both NeMo Automodel (``DeepseekV4HashGate``) and
5
+ transformers (``DeepseekV4HashRouter``) zero-initialise that table, which sends every token to expert 0
6
+ when a model is built from a config instead of a checkpoint. Call :func:`fill_hash_tables` after the model
7
+ weights have been initialised (and after FSDP wrapping; the table is a plain buffer).
8
+ """
9
+ from __future__ import annotations
10
+
11
+ import torch
12
+
13
+
14
+ def balanced_hash_tid2eid(vocab_size: int, n_experts: int, topk: int, seed: int = 0,
15
+ device: torch.device | str | None = None) -> torch.Tensor:
16
+ """Return a ``[vocab_size, topk]`` int64 table mapping each token id to ``topk`` distinct experts.
17
+
18
+ Token ids are randomly permuted (seeded) and mapped to a base expert; the remaining ``topk - 1`` experts are
19
+ spaced ``n_experts // topk`` apart, so every expert receives the same number of (token, slot) assignments
20
+ up to rounding. This is a uniform-over-token-ids hash; the released V4 tables are not published, so a
21
+ frequency-balanced variant (weighting by corpus token counts) is left to the user.
22
+ """
23
+ assert n_experts >= topk, "need at least topk experts"
24
+ stride = n_experts // topk
25
+ g = torch.Generator().manual_seed(seed)
26
+ base = torch.randperm(vocab_size, generator=g) % n_experts
27
+ offsets = torch.arange(topk) * stride
28
+ table = (base.unsqueeze(1) + offsets.unsqueeze(0)) % n_experts
29
+ return table.to(dtype=torch.int64, device=device)
30
+
31
+
32
+ def fill_hash_tables(model: torch.nn.Module, seed: int = 0) -> int:
33
+ """Fill every ``tid2eid`` buffer in ``model`` (Automodel or transformers V4) with a balanced hash.
34
+
35
+ Returns the number of tables filled. Idempotent for a given seed.
36
+ """
37
+ n = 0
38
+ for module in model.modules():
39
+ table = getattr(module, "tid2eid", None)
40
+ if not isinstance(table, torch.Tensor):
41
+ continue
42
+ n_experts = getattr(module, "n_experts", None) or getattr(module, "num_experts", None)
43
+ if n_experts is None:
44
+ raise AttributeError(f"{type(module).__name__} has tid2eid but no n_experts/num_experts attribute")
45
+ vocab_size, topk = table.shape
46
+ with torch.no_grad():
47
+ table.copy_(balanced_hash_tid2eid(vocab_size, int(n_experts), topk, seed + n, table.device))
48
+ n += 1
49
+ return n
50
+
51
+
52
+ def check_hash_tables(model: torch.nn.Module) -> dict:
53
+ """Return per-table expert-load statistics (max/mean assignment count) to confirm balance."""
54
+ stats = {}
55
+ for name, buf in model.named_buffers():
56
+ if name.endswith("tid2eid"):
57
+ counts = torch.bincount(buf.flatten().cpu())
58
+ stats[name] = {"min": int(counts.min()), "max": int(counts.max()), "mean": float(counts.float().mean())}
59
+ return stats
60
+
61
+
62
+ def init_hyper_connections(model: torch.nn.Module, std: float = 0.02) -> int:
63
+ """Initialise the mHC mixer parameters (``fn`` / ``base`` / ``scale`` and the head's ``hc_fn`` /
64
+ ``hc_base`` / ``hc_scale``) for from-scratch training.
65
+
66
+ NeMo Automodel constructs these with ``torch.empty`` and its ``init_weights`` leaves them untouched
67
+ (it only ever loads them from a checkpoint). transformers' ``_init_weights`` uses
68
+ normal(0, initializer_range) for the projection, zeros for the static bias and ones for the gates, and that
69
+ is what this function applies. Works on plain tensors and on FSDP2 DTensors (operates on the local shard).
70
+ Returns the number of tensors initialised.
71
+ """
72
+ n = 0
73
+
74
+ def _local(p):
75
+ return p.to_local() if hasattr(p, "to_local") else p
76
+
77
+ with torch.no_grad():
78
+ for name, p in model.named_parameters():
79
+ leaf = name.rsplit(".", 1)[-1]
80
+ parent = name.rsplit(".", 2)[-2] if name.count(".") >= 1 else ""
81
+ if parent in ("attn_hc", "ffn_hc") and leaf in ("fn", "base", "scale") or leaf in ("hc_fn", "hc_base", "hc_scale"):
82
+ t = _local(p)
83
+ if leaf in ("fn", "hc_fn"):
84
+ torch.nn.init.normal_(t, mean=0.0, std=std)
85
+ elif leaf in ("base", "hc_base"):
86
+ torch.nn.init.zeros_(t)
87
+ else:
88
+ torch.nn.init.ones_(t)
89
+ n += 1
90
+ return n
training/prepare_data.py ADDED
@@ -0,0 +1,107 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """Tokenise local FineWeb parquet files into NeMo Automodel ``NanogptDataset`` shards.
3
+
4
+ Shard format (nemo_automodel/components/datasets/llm/nanogpt_dataset.py, "new" format):
5
+ int32[256] header: [MAGIC=278895051, VERSION=1, num_tokens, bytes_per_token], then uint32 tokens.
6
+ A sibling ``.bos.idx`` file holds int32 positions of BOS tokens (used when ``align_to_bos=True``).
7
+ Every document is prefixed with the tokenizer's BOS token (id 0 for the DeepSeek-V4 tokenizer), no EOS is added,
8
+ matching Automodel's tools/nanogpt_data_processor.py.
9
+
10
+ Example (200M train tokens + 4M validation tokens, ~0.8 GB on disk):
11
+ python3 prepare_data.py --parquet "/path/to/fineweb/*.parquet" \
12
+ --tokenizer akoumpa/Moonlight-V4-1B-h16d256 --out /path/to/data \
13
+ --train-tokens 200M --val-tokens 4M --shard-tokens 50M
14
+ """
15
+ import argparse
16
+ import glob
17
+ import os
18
+ import time
19
+
20
+ import numpy as np
21
+ import pyarrow.parquet as pq
22
+ from transformers import AutoTokenizer
23
+
24
+ MAGIC, VERSION, HEADER_SIZE = 278895051, 1, 256
25
+
26
+
27
+ def parse_tokens(s: str) -> int:
28
+ s = s.strip().upper()
29
+ mult = {"K": 10**3, "M": 10**6, "B": 10**9}.get(s[-1], 1)
30
+ return int(float(s[:-1]) * mult) if s[-1] in "KMB" else int(s)
31
+
32
+
33
+ class ShardWriter:
34
+ def __init__(self, path: str, bos_id: int):
35
+ self.path, self.bos_id, self.n = path, bos_id, 0
36
+ self.fp = open(path, "wb")
37
+ self.idx = open(path[:-4] + ".bos.idx", "wb")
38
+ self.fp.write(np.zeros(HEADER_SIZE, dtype=np.int32).tobytes())
39
+
40
+ def write(self, toks: np.ndarray):
41
+ pos = self.n
42
+ self.fp.write(toks.astype(np.uint32).tobytes())
43
+ self.idx.write((pos + np.flatnonzero(toks == self.bos_id)).astype(np.int32).tobytes())
44
+ self.n += toks.size
45
+
46
+ def close(self):
47
+ header = np.array([MAGIC, VERSION, self.n, 4] + [0] * (HEADER_SIZE - 4), dtype=np.int32)
48
+ self.fp.seek(0); self.fp.write(header.tobytes()); self.fp.close(); self.idx.close()
49
+
50
+
51
+ def doc_stream(files, text_col="text", batch_rows=1024):
52
+ for f in files:
53
+ pf = pq.ParquetFile(f)
54
+ for batch in pf.iter_batches(batch_size=batch_rows, columns=[text_col]):
55
+ yield batch.column(text_col).to_pylist()
56
+
57
+
58
+ def main():
59
+ ap = argparse.ArgumentParser()
60
+ ap.add_argument("--parquet", required=True, help="glob of parquet files with a 'text' column")
61
+ ap.add_argument("--tokenizer", required=True)
62
+ ap.add_argument("--out", required=True)
63
+ ap.add_argument("--train-tokens", default="200M", type=parse_tokens)
64
+ ap.add_argument("--val-tokens", default="4M", type=parse_tokens)
65
+ ap.add_argument("--shard-tokens", default="50M", type=parse_tokens)
66
+ ap.add_argument("--max-doc-tokens", default=32768, type=int)
67
+ args = ap.parse_args()
68
+
69
+ files = sorted(glob.glob(os.path.expanduser(args.parquet)))
70
+ assert files, f"no parquet files match {args.parquet}"
71
+ out = os.path.expanduser(args.out)
72
+ os.makedirs(out, exist_ok=True)
73
+ tok = AutoTokenizer.from_pretrained(args.tokenizer)
74
+ bos = tok.bos_token_id
75
+ assert bos is not None
76
+
77
+ # plan: validation shard first, then train shards
78
+ plan = [("val", args.val_tokens, args.val_tokens)] + [("train", args.train_tokens, args.shard_tokens)]
79
+ stream = doc_stream(files)
80
+ t0, total, n_docs = time.time(), 0, 0
81
+ for split, budget, shard_size in plan:
82
+ written, shard_id, writer = 0, 0, None
83
+ while written < budget:
84
+ if writer is None:
85
+ writer = ShardWriter(os.path.join(out, f"fineweb_{split}_{shard_id:04d}.bin"), bos)
86
+ try:
87
+ texts = next(stream)
88
+ except StopIteration:
89
+ break
90
+ ids = tok(texts, add_special_tokens=False)["input_ids"]
91
+ buf = np.concatenate([np.array([bos] + d[: args.max_doc_tokens - 1], dtype=np.uint32) for d in ids])
92
+ writer.write(buf)
93
+ written += buf.size; total += buf.size; n_docs += len(texts)
94
+ if writer.n >= shard_size:
95
+ writer.close(); writer = None; shard_id += 1
96
+ if n_docs % (1024 * 20) == 0:
97
+ print(f"[{split}] {written/1e6:8.1f}M / {budget/1e6:.0f}M tokens, {n_docs} docs, {total/(time.time()-t0)/1e6:.2f} Mtok/s", flush=True)
98
+ if writer is not None:
99
+ writer.close()
100
+ print(f"[{split}] done: {written/1e6:.1f}M tokens in {shard_id + 1} shard(s)")
101
+ print(f"total {total/1e6:.1f}M tokens, {n_docs} docs, {time.time()-t0:.0f}s -> {out}")
102
+ print(open(os.path.join(out, "README.txt"), "w").write(
103
+ f"tokenizer={args.tokenizer} bos={bos} train_tokens={args.train_tokens} val_tokens={args.val_tokens} source={args.parquet}\n") and "")
104
+
105
+
106
+ if __name__ == "__main__":
107
+ main()
training/pretrain.yaml ADDED
@@ -0,0 +1,113 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # akoumpa/Moonlight-V4-1B-h16d256: from-scratch pre-training recipe for NeMo Automodel (native deepseek_v4 model, FSDP2 over 2 GPUs).
2
+ #
3
+ # torchrun --nproc-per-node 2 train.py --config pretrain.yaml # any key can be overridden: --step_scheduler.max_steps 100
4
+ #
5
+ # train.py seeds the hash-routing table and the mHC mixers on a fresh start (see README, "Training from scratch").
6
+ # Replace the /path/to/... placeholders; data shards come from prepare_data.py.
7
+
8
+ recipe: TrainFinetuneRecipeForNextTokenPrediction
9
+
10
+ step_scheduler:
11
+ global_batch_size: 32 # sequences per optimizer step (32 x 2048 = 65k tokens)
12
+ local_batch_size: 2 # sequences per GPU per micro-step; keep equal to dataloader.batch_size below (4 OOMs on 48 GB)
13
+ ckpt_every_steps: 250
14
+ val_every_steps: 100
15
+ gc_every_steps: 10
16
+ num_epochs: 1
17
+ max_steps: 3000 # ~200M tokens = one pass over the prepared shards; ~10.6 s/step measured -> ~9 h
18
+
19
+ dist_env:
20
+ backend: nccl
21
+ timeout_minutes: 15
22
+
23
+ rng:
24
+ _target_: nemo_automodel.components.training.rng.StatefulRNG
25
+ seed: 1234
26
+ ranked: true
27
+
28
+ model:
29
+ _target_: nemo_automodel.NeMoAutoModelForCausalLM.from_config
30
+ config:
31
+ _target_: nemo_automodel.components.models.deepseek_v4.config.DeepseekV4Config.from_pretrained
32
+ pretrained_model_name_or_path: akoumpa/Moonlight-V4-1B-h16d256
33
+ backend:
34
+ _target_: nemo_automodel.components.models.common.BackendConfig
35
+ attn: eager # see README section "TileLang kernels" before switching to 'tilelang'
36
+ linear: torch
37
+ rms_norm: torch_fp32
38
+ rope_fusion: false
39
+ dispatcher: torch
40
+ experts: torch_mm
41
+ enable_hf_state_dict_adapter: false
42
+
43
+ distributed:
44
+ strategy: fsdp2
45
+ dp_size: 2
46
+ tp_size: 1
47
+ cp_size: 1
48
+ pp_size: 1
49
+ ep_size: 1
50
+ sequence_parallel: false
51
+ activation_checkpointing: false
52
+ moe:
53
+ # Automodel wraps DSV4's fp32 parameter holders (attention sinks, compressor ape, mHC mixers) as their own
54
+ # FSDP2 units; resharding them right after their forward frees the tensor before attention uses it.
55
+ # The shipped DeepSeek-V4 fine-tune recipes disable resharding for the same reason.
56
+ reshard_after_forward: false
57
+ wrap_outer_model: false
58
+
59
+ checkpoint:
60
+ enabled: true
61
+ checkpoint_dir: /path/to/checkpoints/Moonlight-V4-1B-h16d256
62
+ model_save_format: torch_save
63
+ save_consolidated: false
64
+
65
+ loss_fn:
66
+ # Logits-based loss (as in Automodel's own DeepSeek-V4 recipes). FusedLinearCrossEntropy cannot be used here:
67
+ # the model keeps lm_head in fp32 while hidden states are bf16 and the fused Triton kernel requires one dtype.
68
+ _target_: nemo_automodel.components.loss.masked_ce.MaskedCrossEntropy
69
+
70
+ dataset:
71
+ _target_: nemo_automodel.components.datasets.llm.nanogpt_dataset.NanogptDataset
72
+ file_pattern: /path/to/data/fineweb_train_*.bin
73
+ seq_len: 2048
74
+ shuffle_files: true
75
+ align_to_bos: false
76
+ bos_token: 0
77
+
78
+ dataloader:
79
+ _target_: torchdata.stateful_dataloader.StatefulDataLoader
80
+ # The recipe passes no batch size to the DataLoader for IterableDatasets (it only uses local_batch_size for the
81
+ # accumulation count), so without this line every micro-batch is a single sequence.
82
+ batch_size: 2
83
+ shuffle: false
84
+ collate_fn: nemo_automodel.components.datasets.utils.default_collater
85
+
86
+ validation_dataset:
87
+ # NanogptDataset is an infinite stream; FiniteNanogptDataset (finite_nanogpt.py, importable because train_1b.py
88
+ # puts this directory on sys.path) stops after max_samples sequences per rank.
89
+ _target_: finite_nanogpt.FiniteNanogptDataset
90
+ file_pattern: /path/to/data/fineweb_val_0000.bin
91
+ max_samples: 64 # 64 x 2048 tokens per rank per validation
92
+ seq_len: 2048
93
+ shuffle_files: false
94
+ align_to_bos: false
95
+ bos_token: 0
96
+
97
+ validation_dataloader:
98
+ _target_: torchdata.stateful_dataloader.StatefulDataLoader
99
+ batch_size: 2
100
+ shuffle: false
101
+ collate_fn: nemo_automodel.components.datasets.utils.default_collater
102
+
103
+ optimizer:
104
+ _target_: torch.optim.AdamW
105
+ lr: 4.0e-4
106
+ betas: [0.9, 0.95]
107
+ eps: 1.0e-8
108
+ weight_decay: 0.1
109
+
110
+ lr_scheduler:
111
+ lr_decay_style: cosine
112
+ lr_warmup_steps: 200
113
+ min_lr: 4.0e-5
training/train.py ADDED
@@ -0,0 +1,82 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """Launcher for pretrain.yaml (NeMo Automodel recipe).
3
+
4
+ Thin wrapper around NeMo Automodel's TrainFinetuneRecipeForNextTokenPrediction that, on a fresh start
5
+ (no checkpoint to resume from), seeds two things Automodel's from-config path leaves at zero /
6
+ uninitialised because it only ever loads them from released checkpoints:
7
+ * the hash-routing table ``tid2eid`` of the first ``num_hash_layers`` MoE layers (balanced hash), and
8
+ * the mHC mixer parameters (``attn_hc``/``ffn_hc``/``hc_head``: normal(0, 0.02) projection, zero bias, unit gates).
9
+
10
+ Usage: torchrun --nproc-per-node 2 train.py --config pretrain_1b.yaml [--dotted.key value ...]
11
+ """
12
+ import os
13
+ import pathlib
14
+ import sys
15
+
16
+ sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
17
+ from init_utils import check_hash_tables, fill_hash_tables, init_hyper_connections # noqa: E402
18
+
19
+ import torch # noqa: E402
20
+ from torch.distributed.fsdp import FSDPModule # noqa: E402
21
+
22
+ import nemo_automodel # noqa: E402
23
+ from nemo_automodel.components.config._arg_parser import parse_args_and_load_config # noqa: E402
24
+ from nemo_automodel.recipes.llm.train_ft import TrainFinetuneRecipeForNextTokenPrediction # noqa: E402
25
+
26
+
27
+ def will_resume(cfg) -> bool:
28
+ """Mirror the recipe's auto-resume rule: explicit restore_from, or a checkpoint_dir with a training log."""
29
+ if cfg.get("checkpoint.restore_from", None):
30
+ return True
31
+ ckpt_dir = cfg.get("checkpoint.checkpoint_dir", None)
32
+ return bool(ckpt_dir) and (pathlib.Path(ckpt_dir) / "training.jsonl").exists()
33
+
34
+
35
+ def keep_fp32_units_unsharded(model) -> int:
36
+ """Automodel wraps DSV4's fp32 tensors (attention sinks, compressor position biases, mHC mixers, lm_head)
37
+ as separate FSDP2 units whose forward *returns the parameter itself*. If such a unit reshards right after
38
+ its forward, the returned tensor is freed before attention consumes it (``setStorage ... storage of size 0``).
39
+ Keep those (tiny) units unsharded between forward and backward."""
40
+ n = n_resharding = 0
41
+ for m in model.modules():
42
+ if isinstance(m, FSDPModule) and {p.dtype for p in m.parameters()} == {torch.float32}:
43
+ try: # diagnostic only: was this unit going to reshard after forward?
44
+ group = m._get_fsdp_state()._fsdp_param_group
45
+ n_resharding += int(group is not None and group.post_forward_mesh_info is not None)
46
+ except Exception:
47
+ pass
48
+ m.set_reshard_after_forward(False, recurse=False)
49
+ n += 1
50
+ if int(os.environ.get("RANK", "0")) == 0:
51
+ print(f"[train] fp32 FSDP units: {n} total, {n_resharding} were set to reshard after forward")
52
+ return n
53
+
54
+
55
+ def main():
56
+ cfg = parse_args_and_load_config(os.path.join(os.path.dirname(os.path.abspath(__file__)), "pretrain.yaml"))
57
+ resume = will_resume(cfg)
58
+ recipe = TrainFinetuneRecipeForNextTokenPrediction(cfg)
59
+ recipe.setup()
60
+ rank = int(os.environ.get("RANK", "0"))
61
+ if rank == 0:
62
+ print(f"[train] nemo_automodel from {os.path.dirname(nemo_automodel.__file__)}")
63
+ n_units = sum(keep_fp32_units_unsharded(part) for part in recipe.model_parts)
64
+ if rank == 0:
65
+ print(f"[train] fp32 FSDP units kept unsharded after forward: {n_units}")
66
+ if resume:
67
+ if rank == 0:
68
+ print("[train] resuming from checkpoint: keeping hash tables / mHC mixers from the checkpoint")
69
+ else:
70
+ seed = int(cfg.get("rng.seed", 0) or 0)
71
+ std = float(getattr(recipe.model_parts[0].config, "initializer_range", 0.02))
72
+ for part in recipe.model_parts:
73
+ n_tab = fill_hash_tables(part, seed=seed)
74
+ n_hc = init_hyper_connections(part, std=std)
75
+ if rank == 0:
76
+ print(f"[train] fresh start: filled {n_tab} hash table(s), initialised {n_hc} mHC tensors; "
77
+ f"table balance {check_hash_tables(part)}")
78
+ recipe.run_train_validation_loop()
79
+
80
+
81
+ if __name__ == "__main__":
82
+ main()