jkminder commited on
Commit
6799a25
·
verified ·
1 Parent(s): 80e14d5

Scaling Ladder d20_477m_seed2_sft main = ds0_r1

Browse files
LICENSE ADDED
@@ -0,0 +1,28 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ MIT License
2
+
3
+ Copyright (c) 2025 Andrej Karpathy
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.
22
+
23
+ ---
24
+
25
+ Note: this license covers the modeling/configuration code in this repository,
26
+ which is derived from karpathy/nanochat. The model weights are a separate
27
+ artifact; see README.md for the weight license and training-data terms
28
+ (ClimbMix, CC BY-NC 4.0, research and development only).
README.md ADDED
@@ -0,0 +1,117 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: cc-by-nc-4.0
3
+ language:
4
+ - en
5
+ datasets:
6
+ - HuggingFaceTB/smol-smoltalk
7
+ - cais/mmlu
8
+ - openai/gsm8k
9
+ base_model: jkminder/d20_477m_seed2
10
+ base_model_relation: finetune
11
+ pipeline_tag: text-generation
12
+ library_name: transformers
13
+ tags:
14
+ - chat
15
+ - sft
16
+ - research
17
+ - nanochat
18
+ - scaling-ladder
19
+ ---
20
+
21
+ # Scaling Ladder — d20 (477M total parameters), seed 2, chat-SFT
22
+
23
+ **Research artifact.** The chat-SFT of
24
+ [d20_477m_seed2](https://huggingface.co/jkminder/d20_477m_seed2) — size d20,
25
+ pretraining seed 2 of the plain-architecture Scaling Ladder (base
26
+ models trained for 200 tokens per parameter). A small research model tuned
27
+ for basic chat: helpfulness is limited by its size, and it has **no safety
28
+ training**.
29
+
30
+ **This revision (`main`)** mirrors `ds0_r1`, this seed's standard chat-SFT.
31
+
32
+ ## Recipe
33
+
34
+ One pass of nanochat's chat-SFT mixture, applied to the base repository's
35
+ `main` revision (the 200-tokens-per-parameter model):
36
+ [smol-smoltalk](https://huggingface.co/datasets/HuggingFaceTB/smol-smoltalk)
37
+ (460K conversations) + MMLU auxiliary-train x3 + GSM8K main train x4 (with
38
+ one calculator tool-call rendered per solution), interleaved by a fixed
39
+ shuffle and then permuted by the revision's SFT data seed. The optimizer is
40
+ a **cold start** (`+sft.load_optimizer=0`): fresh optimizer state, not the
41
+ pretraining optimizer's. Learning rates start at 0.8x the pretraining
42
+ values, no warmup, linear decay to zero over the second half; 467
43
+ steps at this size. Per-revision training provenance (cluster, code commit)
44
+ is in the table below.
45
+
46
+ ## Revisions
47
+
48
+ Every revision is one chat-SFT run of the same base model:
49
+
50
+ - **`ds<k>`** — SFT data seed k: the permutation of the training-data order
51
+ (all runs share the data; only the order differs).
52
+ - **`r<j>`** — replicate j: an independent repeat at identical
53
+ configuration. The replicate index is never read by training, so repeats
54
+ differ only through run-to-run (GPU) nondeterminism.
55
+ - **`main`** mirrors `ds0_r1`, this seed's standard chat-SFT.
56
+
57
+ Seed-1 repositories carry a noise battery (replicates `ds0_r1..r8`, data
58
+ seeds `ds1..ds7` at `r1`) from a study of SFT run-to-run variance; the
59
+ other seeds have `ds0_r1` only. Runs are added as they finish, so a missing
60
+ revision only means it has not landed yet.
61
+
62
+ | revision | step | SFT val bpb | ARC-Easy | ARC-Challenge | MMLU | trained on | code commit |
63
+ |---|---|---|---|---|---|---|---|
64
+ | ds0_r1 | 467 | 0.3001 | 0.6667 | 0.4676 | 0.3702 | bulbasaur | `fe20a844ca92` |
65
+
66
+ Accuracies are fractions from each run's own chat_eval pass (full test
67
+ suites, greedy decoding: temperature 0, 1 sample, 512 max new tokens; the
68
+ same harness across all runs and sizes). "SFT val bpb" is the run's final
69
+ validation loss (bits per byte) on the mixture's held-out split. A "-"
70
+ means that run's eval has not landed yet.
71
+
72
+ ## Usage
73
+
74
+ The chat template is bundled; format conversations with
75
+ `apply_chat_template`:
76
+
77
+ ```python
78
+ from transformers import AutoModelForCausalLM, AutoTokenizer
79
+
80
+ repo = "jkminder/d20_477m_seed2_sft"
81
+ revision = "main" # or any revision above
82
+ tok = AutoTokenizer.from_pretrained(repo, revision=revision, trust_remote_code=True)
83
+ model = AutoModelForCausalLM.from_pretrained(
84
+ repo, revision=revision, trust_remote_code=True, dtype="bfloat16")
85
+
86
+ msgs = [{"role": "user", "content": "Why is the sky blue?"}]
87
+ ids = tok.apply_chat_template(msgs, add_generation_prompt=True, return_tensors="pt")
88
+ out = model.generate(ids, max_new_tokens=256)
89
+ print(tok.decode(out[0, ids.shape[1]:], skip_special_tokens=True))
90
+ ```
91
+
92
+ `trust_remote_code=True` is required: the architecture matches no stock
93
+ transformers class, so the modeling code ships in the repository
94
+ (`modeling_nanochat_gpt.py`, plain PyTorch). Generation stops on
95
+ `<|assistant_end|>`; sampling defaults (temperature 0.6, top_k 50) ship in
96
+ `generation_config.json`. The template renders nanochat's chat format
97
+ token-for-token (a leading system message is merged into the first user
98
+ message); conversion is verified per revision by chat-template, logit and
99
+ loss equivalence against the original training code
100
+ (`verify_results.json`, where present).
101
+
102
+ ## Architecture, tokenizer, training data
103
+
104
+ Identical to the base repository — a plain GPT (nanochat with all optional
105
+ architecture mechanisms disabled), nanochat BPE tokenizer (32,768 tokens),
106
+ base pretraining on ClimbMix; see
107
+ [d20_477m_seed2](https://huggingface.co/jkminder/d20_477m_seed2) for the full
108
+ description. Weights are bfloat16 safetensors, the training compute
109
+ precision.
110
+
111
+ ## License
112
+
113
+ - Model weights: **cc-by-nc-4.0** (the base model mirrors its ClimbMix
114
+ training data's research-only license, and this fine-tune mirrors the
115
+ base).
116
+ - Modeling/configuration code: MIT (derived from karpathy/nanochat; see the
117
+ bundled LICENSE file).
chat_template.jinja ADDED
@@ -0,0 +1 @@
 
 
1
+ {{ bos_token }}{% if messages[0]['role'] == 'system' %}{% if messages | length < 2 or messages[1]['role'] != 'user' %}{{ raise_exception('a system message must be followed by a user message') }}{% endif %}{% set loop_messages = messages[1:] %}{% else %}{% set loop_messages = messages %}{% endif %}{% for message in loop_messages %}{% if message['content'] is not string %}{{ raise_exception('only plain string contents are supported') }}{% endif %}{% if message['role'] == 'user' %}{{ '<|user_start|>' }}{% if loop.first and messages[0]['role'] == 'system' %}{{ messages[0]['content'] + '\n\n' }}{% endif %}{{ message['content'] }}{{ '<|user_end|>' }}{% elif message['role'] == 'assistant' %}{{ '<|assistant_start|>' }}{{ message['content'] }}{{ '<|assistant_end|>' }}{% else %}{{ raise_exception('only system (first), user, and assistant roles are supported') }}{% endif %}{% endfor %}{% if add_generation_prompt %}{{ '<|assistant_start|>' }}{% endif %}
config.json ADDED
@@ -0,0 +1,38 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "NanochatGPTForCausalLM"
4
+ ],
5
+ "auto_map": {
6
+ "AutoConfig": "configuration_nanochat_gpt.NanochatGPTConfig",
7
+ "AutoModel": "modeling_nanochat_gpt.NanochatGPTModel",
8
+ "AutoModelForCausalLM": "modeling_nanochat_gpt.NanochatGPTForCausalLM"
9
+ },
10
+ "backout_layer": null,
11
+ "bos_token_id": 32759,
12
+ "dtype": "bfloat16",
13
+ "eos_token_id": [
14
+ 32763,
15
+ 32759
16
+ ],
17
+ "final_logit_softcapping": 15.0,
18
+ "hidden_size": 1280,
19
+ "intermediate_size": 5120,
20
+ "logit_softcap": 15.0,
21
+ "max_position_embeddings": 2048,
22
+ "model_type": "nanochat_gpt",
23
+ "num_attention_heads": 10,
24
+ "num_hidden_layers": 20,
25
+ "num_key_value_heads": 10,
26
+ "qk_sharpen_scale": null,
27
+ "rope_theta": 100000.0,
28
+ "smear_gate_channels": 24,
29
+ "tie_word_embeddings": false,
30
+ "transformers_version": "5.14.1",
31
+ "use_resid_lambdas": false,
32
+ "use_smear": false,
33
+ "use_x0_lambdas": false,
34
+ "value_embedding_layers": [],
35
+ "ve_gate_channels": 12,
36
+ "vocab_size": 32768,
37
+ "window_pattern": "L"
38
+ }
configuration_nanochat_gpt.py ADDED
@@ -0,0 +1,147 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Configuration for the nanochat-GPT architecture (HuggingFace export).
2
+
3
+ Derived from karpathy/nanochat (MIT License, Copyright (c) 2025 Andrej
4
+ Karpathy). This file is uploaded to the model repo and loaded with
5
+ trust_remote_code=True.
6
+
7
+ Two families of checkpoints share this configuration:
8
+
9
+ - the "clean" architecture (the d26 L-baseline): every speedrun mechanism
10
+ ablated, full dense attention. All mechanism fields below default to that
11
+ configuration, so config.json files written before these fields existed
12
+ keep loading with identical behavior.
13
+ - the full nanochat architecture (the 200-tokens-per-parameter seed-variance
14
+ models): value embeddings, x0 re-injection, per-layer residual scaling,
15
+ smear, backout, QK sharpening, and an "SSSL" sliding-window pattern all
16
+ active. The exporter (convert.py) fills these fields from the training
17
+ meta json.
18
+ """
19
+
20
+ from transformers import PretrainedConfig
21
+
22
+
23
+ class NanochatGPTConfig(PretrainedConfig):
24
+ model_type = "nanochat_gpt"
25
+
26
+ def __init__(
27
+ self,
28
+ vocab_size=32768,
29
+ hidden_size=1664,
30
+ num_hidden_layers=26,
31
+ num_attention_heads=13,
32
+ num_key_value_heads=None,
33
+ intermediate_size=None,
34
+ max_position_embeddings=2048,
35
+ rope_theta=100000.0,
36
+ logit_softcap=15.0,
37
+ bos_token_id=32759,
38
+ eos_token_id=32759,
39
+ tie_word_embeddings=False,
40
+ # --- speedrun mechanisms (defaults = the clean architecture: all off).
41
+ # window_pattern: sliding-window attention pattern tiled across layers,
42
+ # "L"=full context (window = max_position_embeddings), "S"=short window
43
+ # (quarter context, rounded up to a 128 multiple). The final layer is
44
+ # always L. "L" alone means every layer sees the full context.
45
+ window_pattern="L",
46
+ # value_embedding_layers: layer indices with a value-embedding table
47
+ # (ResFormer-style value residual) and its per-head sigmoid gate.
48
+ value_embedding_layers=None,
49
+ # ve_gate_channels: how many leading channels of the (normed) hidden
50
+ # state feed each value-embedding gate.
51
+ ve_gate_channels=12,
52
+ # use_resid_lambdas: learned per-layer scalar on the residual stream.
53
+ use_resid_lambdas=False,
54
+ # use_x0_lambdas: learned per-layer scalar re-injecting the initial
55
+ # (post-embedding-norm, post-smear) representation at every layer.
56
+ use_x0_lambdas=False,
57
+ # use_smear: mix the previous token's embedding into the current one
58
+ # through a learned gate (cheap bigram-like information).
59
+ use_smear=False,
60
+ # smear_gate_channels: leading channels of the embedding feeding the
61
+ # smear gate.
62
+ smear_gate_channels=24,
63
+ # backout_layer: subtract backout_lambda * (that layer's output) before
64
+ # the final norm. None = no backout.
65
+ backout_layer=None,
66
+ # qk_sharpen_scale: multiply queries and keys by this after QK norm
67
+ # (nanochat uses 1.2). None = no sharpening.
68
+ qk_sharpen_scale=None,
69
+ **kwargs,
70
+ ):
71
+ self.vocab_size = vocab_size
72
+ self.hidden_size = hidden_size
73
+ self.num_hidden_layers = num_hidden_layers
74
+ self.num_attention_heads = num_attention_heads
75
+ self.num_key_value_heads = num_key_value_heads if num_key_value_heads is not None else num_attention_heads
76
+ self.intermediate_size = intermediate_size if intermediate_size is not None else 4 * hidden_size
77
+ self.max_position_embeddings = max_position_embeddings
78
+ self.rope_theta = rope_theta
79
+ self.logit_softcap = logit_softcap
80
+
81
+ assert window_pattern and all(c in "SL" for c in window_pattern.upper()), (
82
+ f"invalid window_pattern {window_pattern!r}: use only S and L"
83
+ )
84
+ self.window_pattern = window_pattern.upper()
85
+
86
+ value_embedding_layers = list(value_embedding_layers) if value_embedding_layers else []
87
+ assert value_embedding_layers == sorted(set(value_embedding_layers)), (
88
+ f"value_embedding_layers must be sorted and unique: {value_embedding_layers}"
89
+ )
90
+ assert all(0 <= i < num_hidden_layers for i in value_embedding_layers), (
91
+ f"value_embedding_layers out of range for {num_hidden_layers} layers: {value_embedding_layers}"
92
+ )
93
+ self.value_embedding_layers = value_embedding_layers
94
+ assert 0 < ve_gate_channels <= hidden_size, ve_gate_channels
95
+ self.ve_gate_channels = ve_gate_channels
96
+
97
+ self.use_resid_lambdas = use_resid_lambdas
98
+ self.use_x0_lambdas = use_x0_lambdas
99
+ self.use_smear = use_smear
100
+ assert 0 < smear_gate_channels <= hidden_size, smear_gate_channels
101
+ self.smear_gate_channels = smear_gate_channels
102
+
103
+ assert backout_layer is None or 0 <= backout_layer < num_hidden_layers, backout_layer
104
+ self.backout_layer = backout_layer
105
+ assert qk_sharpen_scale is None or qk_sharpen_scale > 0, qk_sharpen_scale
106
+ self.qk_sharpen_scale = qk_sharpen_scale
107
+
108
+ # --- engine-facing aliases. vLLM's transformers backend reads these
109
+ # STANDARD keys; our own modeling code never does. ---
110
+ # vLLM bypasses NanochatGPTForCausalLM.forward (it builds its own
111
+ # lm_head + logits processor) and applies final-logit soft-capping
112
+ # from this gemma-2-convention key — same formula as ours.
113
+ self.final_logit_softcapping = logit_softcap
114
+ # Per-layer attention windows: vLLM builds its attention instances
115
+ # from layer_types + sliding_window. Emitted ONLY when a short window
116
+ # exists, so clean-architecture config.json files are unchanged.
117
+ # Semantics mapping (pinned in tests): our window w = "self + w
118
+ # previous positions" (w+1 keys); HF/vLLM sliding_window n = "the
119
+ # last n keys including self" — so n = w + 1. The window list here
120
+ # must stay identical to modeling's compute_window_sizes (asserted
121
+ # at model init).
122
+ long_window = max_position_embeddings
123
+ short_window = -(-long_window // 4 // 128) * 128
124
+ pattern = self.window_pattern
125
+ sizes = [
126
+ {"L": long_window, "S": short_window}[pattern[i % len(pattern)]]
127
+ for i in range(num_hidden_layers)
128
+ ]
129
+ sizes[-1] = long_window
130
+ if any(w < long_window for w in sizes):
131
+ self.sliding_window = short_window + 1
132
+ self.layer_types = [
133
+ "sliding_attention" if w < long_window else "full_attention"
134
+ for w in sizes
135
+ ]
136
+
137
+ super().__init__(
138
+ bos_token_id=bos_token_id,
139
+ eos_token_id=eos_token_id,
140
+ tie_word_embeddings=tie_word_embeddings,
141
+ **kwargs,
142
+ )
143
+
144
+ @property
145
+ def head_dim(self):
146
+ assert self.hidden_size % self.num_attention_heads == 0
147
+ return self.hidden_size // self.num_attention_heads
generation_config.json ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "bos_token_id": 32759,
3
+ "do_sample": true,
4
+ "eos_token_id": [
5
+ 32763,
6
+ 32759
7
+ ],
8
+ "max_new_tokens": 256,
9
+ "pad_token_id": 32759,
10
+ "temperature": 0.6,
11
+ "top_k": 50,
12
+ "transformers_version": "5.14.1"
13
+ }
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:3e75e5a0f837f1f756db9d29a73b84020da3ee22a21c1df084402b8616ace675
3
+ size 954218088
modeling_nanochat_gpt.py ADDED
@@ -0,0 +1,562 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """nanochat-GPT for HuggingFace transformers (custom code, trust_remote_code).
2
+
3
+ Derived from karpathy/nanochat gpt.py (MIT License, Copyright (c) 2025 Andrej
4
+ Karpathy). The always-on pieces:
5
+ - decoder-only transformer, causal attention
6
+ - rotary position embeddings (base 100000, nanochat's half-split convention)
7
+ - RMSNorm with no learnable parameters (after embedding, pre-attn, pre-MLP, final)
8
+ - QK norm: queries/keys RMS-normalized AFTER rotary, no learnable weight
9
+ - MLP with relu(x)^2 activation, no gating
10
+ - no biases anywhere, untied input embedding / output head
11
+ - logit softcap: logits = softcap * tanh(logits / softcap), in float32
12
+
13
+ The speedrun mechanisms, each enabled by its config field (see
14
+ configuration_nanochat_gpt.py; all off = the clean d26-style architecture):
15
+ - sliding-window attention (window_pattern, "S"/"L" tiled across layers)
16
+ - value embeddings: per-layer token-embedding tables mixed into the attention
17
+ values through a learned per-head sigmoid gate (ResFormer-style)
18
+ - x0 re-injection and per-layer residual scaling (x0_lambdas, resid_lambdas)
19
+ - smear: gated mix of the previous token's embedding into the current one
20
+ - backout: subtract a scaled mid-layer residual before the final norm
21
+ - QK sharpening: fixed scale on queries and keys after QK norm
22
+
23
+ Numerical intent: weights are stored in bfloat16 and all matmuls run in
24
+ bfloat16 (this matches training, where fp32 master weights were cast to
25
+ bfloat16 for every forward). The logit softcap and the loss run in float32.
26
+ Every mechanism follows the reference (ppriors/utils/gpt.py) operation by
27
+ operation, in the same order and dtype flow, so logits reproduce the
28
+ original model bit for bit on the same kernel (verified in verify.py).
29
+ """
30
+
31
+ from typing import Optional
32
+
33
+ import torch
34
+ import torch.nn as nn
35
+ import torch.nn.functional as F
36
+
37
+ from transformers import PreTrainedModel
38
+ from transformers.generation import GenerationMixin
39
+ from transformers.cache_utils import Cache, DynamicCache
40
+ from transformers.modeling_outputs import BaseModelOutputWithPast, CausalLMOutputWithPast
41
+ from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS
42
+
43
+ from .configuration_nanochat_gpt import NanochatGPTConfig
44
+
45
+ # Attention implementations that take the REFERENCE path
46
+ # (sliding_window_sdpa — bit-identical to nanochat's SDPA kernel, what
47
+ # verify.py certifies against the training checkpoint). Anything else
48
+ # (e.g. the "vllm" implementation vLLM's transformers backend patches into
49
+ # config._attn_implementation) dispatches through HF's attention-interface
50
+ # registry; whether that engine's numbers match is what the equivalence gate
51
+ # (docs/archive/vllm-eval-acceptance.md) adjudicates. "eager" deliberately maps to
52
+ # the reference path too: this export has always run one attention code
53
+ # path, and silently switching kernels on an innocuous-looking config
54
+ # default would invalidate the verify.py certificate.
55
+ REFERENCE_ATTN_IMPLS = (None, "sdpa", "eager")
56
+
57
+
58
+ class NanochatDynamicCache(DynamicCache):
59
+ """DynamicCache plus the smear stash: the pre-smear embedding of the
60
+ newest position, consumed by the next single-token decode.
61
+
62
+ The stash must follow every batch-dimension shuffle of the k/v tensors.
63
+ Beam search permutes the cache between steps via reorder_cache(beam_idx);
64
+ a stash stored as a plain attribute on a stock DynamicCache stayed in the
65
+ OLD beam order, so each beam smeared with another beam's embedding —
66
+ silently wrong logits from the first reorder on (reviewer finding on
67
+ 649aa3e: beam(3) diverged from the 3rd generated token). Cropping
68
+ (assisted-decoding rollback) is refused: the stash holds only the newest
69
+ position, so after a shrink the right embedding is gone.
70
+ """
71
+
72
+ nanochat_prev_embedding = None # class default; instances stash their own
73
+
74
+ def reorder_cache(self, beam_idx):
75
+ super().reorder_cache(beam_idx)
76
+ prev = self.nanochat_prev_embedding
77
+ if prev is not None:
78
+ self.nanochat_prev_embedding = prev.index_select(0, beam_idx.to(prev.device))
79
+
80
+ def batch_repeat_interleave(self, repeats):
81
+ super().batch_repeat_interleave(repeats)
82
+ prev = self.nanochat_prev_embedding
83
+ if prev is not None:
84
+ self.nanochat_prev_embedding = prev.repeat_interleave(repeats, dim=0)
85
+
86
+ def batch_select_indices(self, indices):
87
+ super().batch_select_indices(indices)
88
+ prev = self.nanochat_prev_embedding
89
+ if prev is not None:
90
+ self.nanochat_prev_embedding = prev.index_select(0, indices.to(prev.device))
91
+
92
+ def crop(self, max_length):
93
+ assert self.nanochat_prev_embedding is None or \
94
+ max_length >= self.get_seq_length(), (
95
+ "cropping a smear model's KV cache is unsupported: the cache "
96
+ "stashes only the NEWEST position's pre-smear embedding, so a "
97
+ "shrunk cache would smear with a stale embedding (silently wrong "
98
+ "logits). Assisted decoding needs cropping; run without an "
99
+ "assistant model."
100
+ )
101
+ super().crop(max_length)
102
+
103
+
104
+ def rms_norm(x):
105
+ # RMSNorm without learnable parameters, computed by the framework kernel
106
+ # (same call as nanochat) so results match the original bit-for-bit.
107
+ return F.rms_norm(x, (x.size(-1),))
108
+
109
+
110
+ def apply_rotary_emb(x, cos, sin):
111
+ # nanochat convention: rotates by -theta relative to the textbook
112
+ # convention (only the relative q/k rotation matters, but q and k must
113
+ # both use this exact form to reproduce the checkpoint).
114
+ assert x.ndim == 4 # (B, T, H, D)
115
+ d = x.shape[3] // 2
116
+ x1, x2 = x[..., :d], x[..., d:]
117
+ y1 = x1 * cos + x2 * sin
118
+ y2 = x1 * (-sin) + x2 * cos
119
+ return torch.cat([y1, y2], 3)
120
+
121
+
122
+ def compute_rotary_cos_sin(positions, head_dim, base, device, dtype):
123
+ """cos/sin of shape (1, T, 1, head_dim/2), computed in fp32 then cast
124
+ (nanochat computes its rotary cache the same way)."""
125
+ channel_range = torch.arange(0, head_dim, 2, dtype=torch.float32, device=device)
126
+ inv_freq = 1.0 / (base ** (channel_range / head_dim))
127
+ t = positions.to(device=device, dtype=torch.float32)
128
+ freqs = torch.outer(t, inv_freq)
129
+ cos, sin = freqs.cos(), freqs.sin()
130
+ cos, sin = cos.to(dtype), sin.to(dtype)
131
+ return cos[None, :, None, :], sin[None, :, None, :]
132
+
133
+
134
+ def compute_window_sizes(config: NanochatGPTConfig):
135
+ """Per-layer left attention window, ported from nanochat GPT._compute_window_sizes.
136
+
137
+ The pattern string is tiled across layers; the final layer is always L.
138
+ L = the full trained context (max_position_embeddings); S = quarter
139
+ context, rounded up to a 128 multiple (nanochat rounds to the FA3 tile).
140
+ """
141
+ pattern = config.window_pattern.upper()
142
+ assert all(c in "SL" for c in pattern), f"Invalid window_pattern: {pattern}. Use only S and L."
143
+ long_window = config.max_position_embeddings
144
+ short_window = -(-long_window // 4 // 128) * 128
145
+ char_to_window = {"L": long_window, "S": short_window}
146
+ window_sizes = [char_to_window[pattern[i % len(pattern)]] for i in range(config.num_hidden_layers)]
147
+ window_sizes[-1] = long_window
148
+ return window_sizes
149
+
150
+
151
+ def sliding_window_sdpa(q, k, v, window, enable_gqa):
152
+ """SDPA with nanochat's left-window semantics (a row attends to itself and
153
+ the `window` previous positions). Ported from
154
+ ppriors/utils/flash_attention._sdpa_attention, the kernel the reference
155
+ model runs when Flash Attention 3 is unavailable (CPU verification).
156
+ q, k, v are (B, H, T, D); k/v already include any cached positions.
157
+ """
158
+ Tq = q.size(2)
159
+ Tk = k.size(2)
160
+
161
+ # Full context, same length
162
+ if (window < 0 or window >= Tq) and Tq == Tk:
163
+ return F.scaled_dot_product_attention(q, k, v, is_causal=True, enable_gqa=enable_gqa)
164
+
165
+ # Single token generation
166
+ if Tq == 1:
167
+ if window >= 0 and window < Tk:
168
+ # window is "left" tokens: include (window + 1) keys total
169
+ start = max(0, Tk - (window + 1))
170
+ k = k[:, :, start:, :]
171
+ v = v[:, :, start:, :]
172
+ return F.scaled_dot_product_attention(q, k, v, is_causal=False, enable_gqa=enable_gqa)
173
+
174
+ # Sliding window and/or chunked prefill on a cache: explicit bool mask
175
+ device = q.device
176
+ row_idx = (Tk - Tq) + torch.arange(Tq, device=device).unsqueeze(1)
177
+ col_idx = torch.arange(Tk, device=device).unsqueeze(0)
178
+ mask = col_idx <= row_idx
179
+ if window >= 0 and window < Tk:
180
+ mask = mask & ((row_idx - col_idx) <= window)
181
+ return F.scaled_dot_product_attention(q, k, v, attn_mask=mask, enable_gqa=enable_gqa)
182
+
183
+
184
+ class NanochatGPTAttention(nn.Module):
185
+ def __init__(self, config: NanochatGPTConfig, layer_idx: int):
186
+ super().__init__()
187
+ self.layer_idx = layer_idx
188
+ self.n_head = config.num_attention_heads
189
+ self.n_kv_head = config.num_key_value_heads
190
+ self.head_dim = config.head_dim
191
+ self.q_proj = nn.Linear(config.hidden_size, self.n_head * self.head_dim, bias=False)
192
+ self.k_proj = nn.Linear(config.hidden_size, self.n_kv_head * self.head_dim, bias=False)
193
+ self.v_proj = nn.Linear(config.hidden_size, self.n_kv_head * self.head_dim, bias=False)
194
+ self.o_proj = nn.Linear(config.hidden_size, config.hidden_size, bias=False)
195
+ self.ve_gate_channels = config.ve_gate_channels
196
+ self.ve_gate = (
197
+ nn.Linear(self.ve_gate_channels, self.n_kv_head, bias=False)
198
+ if layer_idx in config.value_embedding_layers else None
199
+ )
200
+ self.qk_sharpen_scale = config.qk_sharpen_scale
201
+ self.window = None # left attention window (int), set by NanochatGPTModel
202
+ # Attributes HF attention-interface implementations read off the module.
203
+ self.config = config
204
+ self.is_causal = True
205
+ self.num_key_value_groups = self.n_head // self.n_kv_head
206
+ self.scaling = self.head_dim**-0.5 # SDPA's default scale, made explicit
207
+
208
+ def forward(self, x, ve, cos_sin, past_key_values: Optional[Cache], cache_position, **kwargs):
209
+ B, T, C = x.size()
210
+ q = self.q_proj(x).view(B, T, self.n_head, self.head_dim)
211
+ k = self.k_proj(x).view(B, T, self.n_kv_head, self.head_dim)
212
+ v = self.v_proj(x).view(B, T, self.n_kv_head, self.head_dim)
213
+
214
+ # Value residual (ResFormer): mix in the value embedding with an
215
+ # input-dependent gate per kv head, before rotary/QK norm (which do
216
+ # not touch v anyway) — same point as the reference.
217
+ assert (ve is None) == (self.ve_gate is None), (
218
+ f"layer {self.layer_idx}: value embedding and gate must appear together"
219
+ )
220
+ if ve is not None:
221
+ ve = ve.view(B, T, self.n_kv_head, self.head_dim)
222
+ gate = 3 * torch.sigmoid(self.ve_gate(x[..., :self.ve_gate_channels])) # (B, T, n_kv_head), range (0, 3)
223
+ v = v + gate.unsqueeze(-1) * ve
224
+
225
+ cos, sin = cos_sin
226
+ q, k = apply_rotary_emb(q, cos, sin), apply_rotary_emb(k, cos, sin)
227
+ q, k = rms_norm(q), rms_norm(k) # QK norm, after rotary
228
+ if self.qk_sharpen_scale is not None:
229
+ q = q * self.qk_sharpen_scale # sharper attention, scale split between Q and K
230
+ k = k * self.qk_sharpen_scale
231
+
232
+ # SDPA layout (B, H, T, D)
233
+ q = q.transpose(1, 2)
234
+ k = k.transpose(1, 2)
235
+ v = v.transpose(1, 2)
236
+
237
+ if past_key_values is not None:
238
+ k, v = past_key_values.update(k, v, self.layer_idx)
239
+
240
+ assert self.window is not None, "window not set (NanochatGPTModel wires it)"
241
+ impl = getattr(self.config, "_attn_implementation", None)
242
+ if impl in REFERENCE_ATTN_IMPLS:
243
+ # The reference path: exactly the kernel verify.py certifies.
244
+ enable_gqa = self.n_kv_head != self.n_head
245
+ y = sliding_window_sdpa(q, k, v, self.window, enable_gqa)
246
+ y = y.transpose(1, 2).contiguous().view(B, T, -1)
247
+ else:
248
+ # Engine path (e.g. vLLM's "vllm" implementation): dispatch through
249
+ # HF's attention-interface registry. The engine owns KV caching and
250
+ # window/causality (vLLM: per-layer windows from config.layer_types
251
+ # + config.sliding_window); q/k/v here carry everything upstream of
252
+ # attention (rotary, QK norm, sharpening, value-embedding mix).
253
+ # Interface convention: q/k/v in (B, H, T, D), output (B, T, H, D).
254
+ attention_interface = ALL_ATTENTION_FUNCTIONS[impl]
255
+ y, _ = attention_interface(
256
+ self, q, k, v, None, scaling=self.scaling, **kwargs
257
+ )
258
+ y = y.reshape(B, T, -1).contiguous()
259
+ return self.o_proj(y)
260
+
261
+
262
+ class NanochatGPTMLP(nn.Module):
263
+ def __init__(self, config: NanochatGPTConfig):
264
+ super().__init__()
265
+ self.up_proj = nn.Linear(config.hidden_size, config.intermediate_size, bias=False)
266
+ self.down_proj = nn.Linear(config.intermediate_size, config.hidden_size, bias=False)
267
+
268
+ def forward(self, x):
269
+ return self.down_proj(F.relu(self.up_proj(x)).square())
270
+
271
+
272
+ class NanochatGPTBlock(nn.Module):
273
+ def __init__(self, config: NanochatGPTConfig, layer_idx: int):
274
+ super().__init__()
275
+ self.self_attn = NanochatGPTAttention(config, layer_idx)
276
+ self.mlp = NanochatGPTMLP(config)
277
+
278
+ def forward(self, x, ve, cos_sin, past_key_values, cache_position, **kwargs):
279
+ x = x + self.self_attn(rms_norm(x), ve, cos_sin, past_key_values, cache_position, **kwargs)
280
+ x = x + self.mlp(rms_norm(x))
281
+ return x
282
+
283
+
284
+ class NanochatGPTPreTrainedModel(PreTrainedModel):
285
+ config_class = NanochatGPTConfig
286
+ base_model_prefix = "model"
287
+ supports_gradient_checkpointing = False
288
+ _no_split_modules = ["NanochatGPTBlock"]
289
+ _supports_sdpa = True
290
+ _supports_cache_class = True
291
+ # Attention routes through HF's attention-interface registry when a
292
+ # non-reference implementation is patched in (REFERENCE_ATTN_IMPLS above),
293
+ # which is what vLLM's transformers backend requires
294
+ # (is_backend_compatible reads this flag).
295
+ _supports_attention_backend = True
296
+
297
+ def _init_weights(self, module):
298
+ # Export-only model: weights always come from a converted checkpoint.
299
+ if isinstance(module, nn.Linear):
300
+ module.weight.data.normal_(mean=0.0, std=0.02)
301
+ elif isinstance(module, nn.Embedding):
302
+ module.weight.data.normal_(mean=0.0, std=0.02)
303
+
304
+
305
+ class NanochatGPTModel(NanochatGPTPreTrainedModel):
306
+ def __init__(self, config: NanochatGPTConfig):
307
+ super().__init__(config)
308
+ self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size)
309
+ self.layers = nn.ModuleList(
310
+ [NanochatGPTBlock(config, layer_idx) for layer_idx in range(config.num_hidden_layers)]
311
+ )
312
+ self.window_sizes = compute_window_sizes(config)
313
+ for layer, window in zip(self.layers, self.window_sizes):
314
+ layer.self_attn.window = window
315
+ # The config's engine-facing layer_types/sliding_window (what vLLM
316
+ # builds its attention from) must describe the SAME windows this
317
+ # module enforces on the reference path — two derivations of one
318
+ # pattern, pinned against drift here.
319
+ layer_types = getattr(config, "layer_types", None)
320
+ if layer_types is not None:
321
+ long_window = config.max_position_embeddings
322
+ expected = ["sliding_attention" if w < long_window else "full_attention"
323
+ for w in self.window_sizes]
324
+ short = [w for w in self.window_sizes if w < long_window]
325
+ assert list(layer_types) == expected and \
326
+ all(w + 1 == config.sliding_window for w in short), (
327
+ "config.layer_types/sliding_window disagree with "
328
+ "compute_window_sizes — the engine would attend differently "
329
+ f"than the reference: {layer_types} vs {expected}, "
330
+ f"sliding_window={getattr(config, 'sliding_window', None)}"
331
+ )
332
+
333
+ # Mechanism parameters exist only when their mechanism is on, so the
334
+ # clean-architecture state dict (older exports) still loads strictly.
335
+ n_layer = config.num_hidden_layers
336
+ if config.use_resid_lambdas:
337
+ self.resid_lambdas = nn.Parameter(torch.ones(n_layer))
338
+ if config.use_x0_lambdas:
339
+ self.x0_lambdas = nn.Parameter(torch.zeros(n_layer))
340
+ if config.use_smear:
341
+ self.smear_gate = nn.Linear(config.smear_gate_channels, 1, bias=False)
342
+ self.smear_lambda = nn.Parameter(torch.zeros(1))
343
+ if config.backout_layer is not None:
344
+ self.backout_lambda = nn.Parameter(torch.zeros(1))
345
+ kv_dim = config.num_key_value_heads * config.head_dim
346
+ self.value_embeds = nn.ModuleDict(
347
+ {str(i): nn.Embedding(config.vocab_size, kv_dim) for i in config.value_embedding_layers}
348
+ )
349
+ self.post_init()
350
+
351
+ def _smear(self, x, past_key_values, cache_position):
352
+ """Mix the previous token's (pre-smear) embedding into each position.
353
+
354
+ Mirrors nanochat GPT.forward: full-sequence smear when every position
355
+ is present; with a KV cache, the pre-smear embedding of the newest
356
+ position is stashed on the cache object and consumed by the next
357
+ single-token decode step.
358
+ """
359
+ B, T, C = x.size()
360
+ ch = self.config.smear_gate_channels
361
+ if past_key_values is None:
362
+ # Full sequence available (no cache): position 0 has no predecessor.
363
+ assert T > 1, "smear on a full sequence needs T > 1"
364
+ gate = self.smear_lambda.to(x.dtype) * torch.sigmoid(self.smear_gate(x[:, 1:, :ch]))
365
+ return torch.cat([x[:, :1], x[:, 1:] + gate * x[:, :-1]], dim=1)
366
+ prev = getattr(past_key_values, "nanochat_prev_embedding", None)
367
+ past_key_values.nanochat_prev_embedding = x[:, -1:, :] # pre-smear, for the next step
368
+ if T > 1:
369
+ # Prefill: smear positions 1+, same as the full-sequence path.
370
+ gate = self.smear_lambda.to(x.dtype) * torch.sigmoid(self.smear_gate(x[:, 1:, :ch]))
371
+ return torch.cat([x[:, :1], x[:, 1:] + gate * x[:, :-1]], dim=1)
372
+ if int(cache_position[0]) == 0:
373
+ return x # single-token prefill at position 0: no predecessor exists
374
+ # Single-token decode: the previous step must have stashed its embedding.
375
+ # Refusing beats silently skipping the smear (wrong logits).
376
+ assert prev is not None, (
377
+ "single-token decode past position 0 without a stashed previous "
378
+ "embedding: the cache was not built by this model's forward"
379
+ )
380
+ gate = self.smear_lambda.to(x.dtype) * torch.sigmoid(self.smear_gate(x[:, :, :ch]))
381
+ return x + gate * prev
382
+
383
+ def forward(
384
+ self,
385
+ input_ids: Optional[torch.LongTensor] = None,
386
+ attention_mask: Optional[torch.Tensor] = None,
387
+ past_key_values: Optional[Cache] = None,
388
+ use_cache: Optional[bool] = None,
389
+ cache_position: Optional[torch.LongTensor] = None,
390
+ position_ids: Optional[torch.LongTensor] = None,
391
+ inputs_embeds: Optional[torch.Tensor] = None,
392
+ **kwargs,
393
+ ):
394
+ assert (input_ids is None) != (inputs_embeds is None), (
395
+ "pass exactly one of input_ids / inputs_embeds"
396
+ )
397
+ B, T = input_ids.size() if input_ids is not None else inputs_embeds.shape[:2]
398
+ if attention_mask is not None:
399
+ assert bool(torch.all(attention_mask == 1)), (
400
+ "NanochatGPT does not support padded batches; use batch size 1 "
401
+ "or unpadded sequences."
402
+ )
403
+
404
+ # Under an engine (vLLM's transformers backend) requests are packed
405
+ # into one flattened row: positions restart at each request boundary
406
+ # inside dim 1, and the engine owns attention. Mechanisms that mix
407
+ # information ACROSS positions in our own code (smear) or need token
408
+ # ids we were not given (value embeddings under an inputs_embeds-only
409
+ # call) would silently cross request boundaries or cannot run — refuse
410
+ # loudly instead.
411
+ engine_packed = "attention_instances" in kwargs
412
+ if engine_packed:
413
+ assert not self.config.use_smear, (
414
+ "smear models cannot run under an engine that packs requests "
415
+ "into one row: the previous-token embedding mix would cross "
416
+ "request boundaries (silently wrong logits). Run smear models "
417
+ "on the HF path."
418
+ )
419
+ if self.value_embeds:
420
+ assert input_ids is not None, (
421
+ "value-embedding models need input_ids (per-layer token-id "
422
+ "lookups); this call passed only inputs_embeds"
423
+ )
424
+
425
+ if use_cache and past_key_values is None:
426
+ past_key_values = NanochatDynamicCache()
427
+ if use_cache and self.config.use_smear and \
428
+ not isinstance(past_key_values, NanochatDynamicCache):
429
+ # generate() constructs a stock DynamicCache and passes it in.
430
+ # The smear stash must follow beam-search reorder (and refuse
431
+ # crop), so the EMPTY stock cache is grafted onto the stash-aware
432
+ # subclass in place — keeping all internal state and the object
433
+ # identity generate() holds. Any other cache cannot keep the
434
+ # stash in sync; refusing beats silently smearing with another
435
+ # batch row's embedding.
436
+ assert type(past_key_values) is DynamicCache and \
437
+ past_key_values.get_seq_length() == 0, (
438
+ "smear models support only the default dynamic KV cache: pass "
439
+ "past_key_values=None or a fresh DynamicCache, got "
440
+ f"{type(past_key_values).__name__} with "
441
+ f"{past_key_values.get_seq_length()} cached positions"
442
+ )
443
+ past_key_values.__class__ = NanochatDynamicCache
444
+ device = input_ids.device if input_ids is not None else inputs_embeds.device
445
+ if cache_position is None:
446
+ past_len = past_key_values.get_seq_length() if past_key_values is not None else 0
447
+ cache_position = torch.arange(past_len, past_len + T, device=device)
448
+ cache = past_key_values if use_cache else None
449
+
450
+ x = inputs_embeds if inputs_embeds is not None else self.embed_tokens(input_ids)
451
+ x = rms_norm(x)
452
+
453
+ if self.config.use_smear:
454
+ x = self._smear(x, cache, cache_position)
455
+
456
+ # Rotary positions: an engine passes explicit position_ids (packed
457
+ # rows restart positions per request); the HF path derives them from
458
+ # the cache. The rope table is built once per forward from 1-D
459
+ # positions and broadcast over the batch, so distinct per-row
460
+ # positions are refused rather than silently rotated wrong.
461
+ if position_ids is not None:
462
+ assert position_ids.ndim == 2, position_ids.shape
463
+ assert bool(torch.all(position_ids == position_ids[0:1])), (
464
+ "per-row position_ids differ; this model broadcasts one "
465
+ "rotary table over the batch"
466
+ )
467
+ rope_positions = position_ids[0]
468
+ else:
469
+ rope_positions = cache_position
470
+ cos_sin = compute_rotary_cos_sin(
471
+ rope_positions, self.config.head_dim, self.config.rope_theta, x.device, x.dtype
472
+ )
473
+
474
+ x0 = x # initial (post-smear) normalized embedding, for x0 re-injection
475
+ use_resid = self.config.use_resid_lambdas
476
+ use_x0 = self.config.use_x0_lambdas
477
+ backout_layer = self.config.backout_layer
478
+ x_backout = None
479
+ for i, layer in enumerate(self.layers):
480
+ # Same branch structure and expressions as the reference so the
481
+ # bf16 rounding sequence is identical.
482
+ if not use_resid and not use_x0:
483
+ pass
484
+ elif not use_x0:
485
+ x = self.resid_lambdas[i] * x
486
+ elif not use_resid:
487
+ x = x + self.x0_lambdas[i] * x0
488
+ else:
489
+ x = self.resid_lambdas[i] * x + self.x0_lambdas[i] * x0
490
+ ve = self.value_embeds[str(i)](input_ids).to(x.dtype) if str(i) in self.value_embeds else None
491
+ x = layer(x, ve, cos_sin, cache, cache_position, **kwargs)
492
+ if i == backout_layer:
493
+ x_backout = x
494
+ if backout_layer is not None:
495
+ assert x_backout is not None
496
+ x = x - self.backout_lambda.to(x.dtype) * x_backout
497
+ x = rms_norm(x)
498
+
499
+ return BaseModelOutputWithPast(
500
+ last_hidden_state=x,
501
+ past_key_values=past_key_values if use_cache else None,
502
+ )
503
+
504
+
505
+ class NanochatGPTForCausalLM(NanochatGPTPreTrainedModel, GenerationMixin):
506
+ _tied_weights_keys = []
507
+
508
+ def __init__(self, config: NanochatGPTConfig):
509
+ super().__init__(config)
510
+ self.model = NanochatGPTModel(config)
511
+ self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
512
+ self.post_init()
513
+
514
+ def get_input_embeddings(self):
515
+ return self.model.embed_tokens
516
+
517
+ def set_input_embeddings(self, value):
518
+ self.model.embed_tokens = value
519
+
520
+ def get_output_embeddings(self):
521
+ return self.lm_head
522
+
523
+ def forward(
524
+ self,
525
+ input_ids: Optional[torch.LongTensor] = None,
526
+ attention_mask: Optional[torch.Tensor] = None,
527
+ past_key_values: Optional[Cache] = None,
528
+ labels: Optional[torch.LongTensor] = None,
529
+ use_cache: Optional[bool] = None,
530
+ cache_position: Optional[torch.LongTensor] = None,
531
+ position_ids: Optional[torch.LongTensor] = None,
532
+ inputs_embeds: Optional[torch.Tensor] = None,
533
+ **kwargs,
534
+ ):
535
+ outputs = self.model(
536
+ input_ids=input_ids,
537
+ attention_mask=attention_mask,
538
+ past_key_values=past_key_values,
539
+ use_cache=use_cache,
540
+ cache_position=cache_position,
541
+ position_ids=position_ids,
542
+ inputs_embeds=inputs_embeds,
543
+ )
544
+ logits = self.lm_head(outputs.last_hidden_state)
545
+ logits = logits.float() # fp32 for softcap and loss, as in training
546
+ softcap = self.config.logit_softcap
547
+ if softcap is not None and softcap > 0:
548
+ logits = softcap * torch.tanh(logits / softcap)
549
+
550
+ loss = None
551
+ if labels is not None:
552
+ loss = F.cross_entropy(
553
+ logits[:, :-1].reshape(-1, logits.size(-1)),
554
+ labels[:, 1:].reshape(-1),
555
+ ignore_index=-100,
556
+ )
557
+
558
+ return CausalLMOutputWithPast(
559
+ loss=loss,
560
+ logits=logits,
561
+ past_key_values=outputs.past_key_values,
562
+ )
tokenizer.json ADDED
The diff for this file is too large to render. See raw diff
 
tokenizer_config.json ADDED
@@ -0,0 +1,17 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "backend": "tokenizers",
3
+ "bos_token": "<|bos|>",
4
+ "eos_token": "<|assistant_end|>",
5
+ "extra_special_tokens": [
6
+ "<|user_start|>",
7
+ "<|user_end|>",
8
+ "<|assistant_start|>",
9
+ "<|assistant_end|>",
10
+ "<|python_start|>",
11
+ "<|python_end|>",
12
+ "<|output_start|>",
13
+ "<|output_end|>"
14
+ ],
15
+ "model_max_length": 2048,
16
+ "tokenizer_class": "TokenizersBackend"
17
+ }
verify_results.json ADDED
@@ -0,0 +1,26 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "verification_passed": true,
3
+ "checkpoint_dir": "d20_clean_tpp200-ba383fd49540-sft-0fea1929",
4
+ "step": 467,
5
+ "export_dir": "ds0_r1",
6
+ "export_sha256": "43cc5fc2a34997c34eae6d57cd3b5a88880bb2aaf10fdaf993d06ebbbaacdd37",
7
+ "template_conversations_checked": 4,
8
+ "template_tokens_checked": 182,
9
+ "logit_max_abs_diff": 0.0,
10
+ "losses_original": [
11
+ 3.747992515563965,
12
+ 11.390064239501953,
13
+ 4.6220903396606445,
14
+ 0.7291238903999329
15
+ ],
16
+ "losses_converted": [
17
+ 3.747992515563965,
18
+ 11.390064239501953,
19
+ 4.6220903396606445,
20
+ 0.7291238903999329
21
+ ],
22
+ "greedy_replies": {
23
+ "Why is the sky blue?": "The sky appears blue due to the way our atmosphere scatters and absorbs light. The primary reason for this is the presence of tiny molecules of gases, such as nitrogen and oxygen, in the Earth's atmosphere. These molecules absorb and scatter shorter wavelengths of light, such as blue and violet, more than longer wavelengths, such"
24
+ },
25
+ "greedy_reply": "The sky appears blue due to the way our atmosphere scatters and absorbs light. The primary reason for this is the presence of tiny molecules of gases, such as nitrogen and oxygen, in the Earth's atmosphere. These molecules absorb and scatter shorter wavelengths of light, such as blue and violet, more than longer wavelengths, such"
26
+ }