Text Generation
Transformers
Safetensors
English
nanochat_gpt
chat
sft
research
nanochat
scaling-ladder
conversational
custom_code
Instructions to use jkminder/d26_973m_seed2_sft with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use jkminder/d26_973m_seed2_sft with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-generation", model="jkminder/d26_973m_seed2_sft", trust_remote_code=True) messages = [ {"role": "user", "content": "Who are you?"}, ] pipe(messages)# Load model directly from transformers import AutoModelForCausalLM model = AutoModelForCausalLM.from_pretrained("jkminder/d26_973m_seed2_sft", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- vLLM
How to use jkminder/d26_973m_seed2_sft with vLLM:
Install from pip and serve model
# Install vLLM from pip: pip install vllm # Start the vLLM server: vllm serve "jkminder/d26_973m_seed2_sft" # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:8000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "jkminder/d26_973m_seed2_sft", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }'Use Docker
docker model run hf.co/jkminder/d26_973m_seed2_sft
- SGLang
How to use jkminder/d26_973m_seed2_sft with SGLang:
Install from pip and serve model
# Install SGLang from pip: pip install sglang # Start the SGLang server: python3 -m sglang.launch_server \ --model-path "jkminder/d26_973m_seed2_sft" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "jkminder/d26_973m_seed2_sft", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }'Use Docker images
docker run --gpus all \ --shm-size 32g \ -p 30000:30000 \ -v ~/.cache/huggingface:/root/.cache/huggingface \ --env "HF_TOKEN=<secret>" \ --ipc=host \ lmsysorg/sglang:latest \ python3 -m sglang.launch_server \ --model-path "jkminder/d26_973m_seed2_sft" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "jkminder/d26_973m_seed2_sft", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }' - Docker Model Runner
How to use jkminder/d26_973m_seed2_sft with Docker Model Runner:
docker model run hf.co/jkminder/d26_973m_seed2_sft
Scaling Ladder d26_973m_seed2_sft main = ds0_r1
Browse files- LICENSE +28 -0
- README.md +117 -0
- chat_template.jinja +1 -0
- config.json +38 -0
- configuration_nanochat_gpt.py +147 -0
- generation_config.json +13 -0
- model.safetensors +3 -0
- modeling_nanochat_gpt.py +562 -0
- tokenizer.json +0 -0
- tokenizer_config.json +17 -0
- verify_results.json +26 -0
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/d26_973m_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 — d26 (973M total parameters), seed 2, chat-SFT
|
| 22 |
+
|
| 23 |
+
**Research artifact.** The chat-SFT of
|
| 24 |
+
[d26_973m_seed2](https://huggingface.co/jkminder/d26_973m_seed2) — size d26,
|
| 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.2662 | 0.7786 | 0.6220 | 0.4378 | eugleo | `affd94ff568e` |
|
| 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/d26_973m_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 |
+
[d26_973m_seed2](https://huggingface.co/jkminder/d26_973m_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": 1664,
|
| 19 |
+
"intermediate_size": 6656,
|
| 20 |
+
"logit_softcap": 15.0,
|
| 21 |
+
"max_position_embeddings": 2048,
|
| 22 |
+
"model_type": "nanochat_gpt",
|
| 23 |
+
"num_attention_heads": 13,
|
| 24 |
+
"num_hidden_layers": 26,
|
| 25 |
+
"num_key_value_heads": 13,
|
| 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:08261675355ae86c1c625db6c2bd4f93b31239311cf231676690ef9a295a540f
|
| 3 |
+
size 1945913144
|
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": "d26_clean_tpp200-2e546c39a771-sft-0fea1929",
|
| 4 |
+
"step": 467,
|
| 5 |
+
"export_dir": "ds0_r1",
|
| 6 |
+
"export_sha256": "07d88f7d4e23851133bd0077939e4aa613b90b3c971a34c3ba43f4e821ae7c2d",
|
| 7 |
+
"template_conversations_checked": 4,
|
| 8 |
+
"template_tokens_checked": 182,
|
| 9 |
+
"logit_max_abs_diff": 0.0,
|
| 10 |
+
"losses_original": [
|
| 11 |
+
3.513197183609009,
|
| 12 |
+
7.814452648162842,
|
| 13 |
+
4.3745927810668945,
|
| 14 |
+
0.6991087794303894
|
| 15 |
+
],
|
| 16 |
+
"losses_converted": [
|
| 17 |
+
3.513197183609009,
|
| 18 |
+
7.814452648162842,
|
| 19 |
+
4.3745927810668945,
|
| 20 |
+
0.6991087794303894
|
| 21 |
+
],
|
| 22 |
+
"greedy_replies": {
|
| 23 |
+
"Why is the sky blue?": "The sky appears blue due to a phenomenon known as Rayleigh scattering, which was first described by Lord Rayleigh in 1871. When sunlight enters the Earth's atmosphere, it encounters tiny molecules of gases such as nitrogen and oxygen. These molecules scatter the shorter, blue wavelengths of light more than the longer,"
|
| 24 |
+
},
|
| 25 |
+
"greedy_reply": "The sky appears blue due to a phenomenon known as Rayleigh scattering, which was first described by Lord Rayleigh in 1871. When sunlight enters the Earth's atmosphere, it encounters tiny molecules of gases such as nitrogen and oxygen. These molecules scatter the shorter, blue wavelengths of light more than the longer,"
|
| 26 |
+
}
|