climbmix-d26-10tpp-noright / configuration_nanochat_gpt.py
Eugleo's picture
exp-088 noright: d26 10-TPP base model on ClimbMix with the noright filter (arm d26-r10-a737eac7df78)
970f21f verified
Raw History Blame Contribute Delete
6.93 kB
"""Configuration for the nanochat-GPT architecture (HuggingFace export).
Derived from karpathy/nanochat (MIT License, Copyright (c) 2025 Andrej
Karpathy). This file is uploaded to the model repo and loaded with
trust_remote_code=True.
Two families of checkpoints share this configuration:
- the "clean" architecture (the d26 L-baseline): every speedrun mechanism
ablated, full dense attention. All mechanism fields below default to that
configuration, so config.json files written before these fields existed
keep loading with identical behavior.
- the full nanochat architecture (the 200-tokens-per-parameter seed-variance
models): value embeddings, x0 re-injection, per-layer residual scaling,
smear, backout, QK sharpening, and an "SSSL" sliding-window pattern all
active. The exporter (convert.py) fills these fields from the training
meta json.
"""
from transformers import PretrainedConfig
class NanochatGPTConfig(PretrainedConfig):
model_type = "nanochat_gpt"
def __init__(
self,
vocab_size=32768,
hidden_size=1664,
num_hidden_layers=26,
num_attention_heads=13,
num_key_value_heads=None,
intermediate_size=None,
max_position_embeddings=2048,
rope_theta=100000.0,
logit_softcap=15.0,
bos_token_id=32759,
eos_token_id=32759,
tie_word_embeddings=False,
# --- speedrun mechanisms (defaults = the clean architecture: all off).
# window_pattern: sliding-window attention pattern tiled across layers,
# "L"=full context (window = max_position_embeddings), "S"=short window
# (quarter context, rounded up to a 128 multiple). The final layer is
# always L. "L" alone means every layer sees the full context.
window_pattern="L",
# value_embedding_layers: layer indices with a value-embedding table
# (ResFormer-style value residual) and its per-head sigmoid gate.
value_embedding_layers=None,
# ve_gate_channels: how many leading channels of the (normed) hidden
# state feed each value-embedding gate.
ve_gate_channels=12,
# use_resid_lambdas: learned per-layer scalar on the residual stream.
use_resid_lambdas=False,
# use_x0_lambdas: learned per-layer scalar re-injecting the initial
# (post-embedding-norm, post-smear) representation at every layer.
use_x0_lambdas=False,
# use_smear: mix the previous token's embedding into the current one
# through a learned gate (cheap bigram-like information).
use_smear=False,
# smear_gate_channels: leading channels of the embedding feeding the
# smear gate.
smear_gate_channels=24,
# backout_layer: subtract backout_lambda * (that layer's output) before
# the final norm. None = no backout.
backout_layer=None,
# qk_sharpen_scale: multiply queries and keys by this after QK norm
# (nanochat uses 1.2). None = no sharpening.
qk_sharpen_scale=None,
**kwargs,
):
self.vocab_size = vocab_size
self.hidden_size = hidden_size
self.num_hidden_layers = num_hidden_layers
self.num_attention_heads = num_attention_heads
self.num_key_value_heads = num_key_value_heads if num_key_value_heads is not None else num_attention_heads
self.intermediate_size = intermediate_size if intermediate_size is not None else 4 * hidden_size
self.max_position_embeddings = max_position_embeddings
self.rope_theta = rope_theta
self.logit_softcap = logit_softcap
assert window_pattern and all(c in "SL" for c in window_pattern.upper()), (
f"invalid window_pattern {window_pattern!r}: use only S and L"
)
self.window_pattern = window_pattern.upper()
value_embedding_layers = list(value_embedding_layers) if value_embedding_layers else []
assert value_embedding_layers == sorted(set(value_embedding_layers)), (
f"value_embedding_layers must be sorted and unique: {value_embedding_layers}"
)
assert all(0 <= i < num_hidden_layers for i in value_embedding_layers), (
f"value_embedding_layers out of range for {num_hidden_layers} layers: {value_embedding_layers}"
)
self.value_embedding_layers = value_embedding_layers
assert 0 < ve_gate_channels <= hidden_size, ve_gate_channels
self.ve_gate_channels = ve_gate_channels
self.use_resid_lambdas = use_resid_lambdas
self.use_x0_lambdas = use_x0_lambdas
self.use_smear = use_smear
assert 0 < smear_gate_channels <= hidden_size, smear_gate_channels
self.smear_gate_channels = smear_gate_channels
assert backout_layer is None or 0 <= backout_layer < num_hidden_layers, backout_layer
self.backout_layer = backout_layer
assert qk_sharpen_scale is None or qk_sharpen_scale > 0, qk_sharpen_scale
self.qk_sharpen_scale = qk_sharpen_scale
# --- engine-facing aliases. vLLM's transformers backend reads these
# STANDARD keys; our own modeling code never does. ---
# vLLM bypasses NanochatGPTForCausalLM.forward (it builds its own
# lm_head + logits processor) and applies final-logit soft-capping
# from this gemma-2-convention key — same formula as ours.
self.final_logit_softcapping = logit_softcap
# Per-layer attention windows: vLLM builds its attention instances
# from layer_types + sliding_window. Emitted ONLY when a short window
# exists, so clean-architecture config.json files are unchanged.
# Semantics mapping (pinned in tests): our window w = "self + w
# previous positions" (w+1 keys); HF/vLLM sliding_window n = "the
# last n keys including self" — so n = w + 1. The window list here
# must stay identical to modeling's compute_window_sizes (asserted
# at model init).
long_window = max_position_embeddings
short_window = -(-long_window // 4 // 128) * 128
pattern = self.window_pattern
sizes = [
{"L": long_window, "S": short_window}[pattern[i % len(pattern)]]
for i in range(num_hidden_layers)
]
sizes[-1] = long_window
if any(w < long_window for w in sizes):
self.sliding_window = short_window + 1
self.layer_types = [
"sliding_attention" if w < long_window else "full_attention"
for w in sizes
]
super().__init__(
bos_token_id=bos_token_id,
eos_token_id=eos_token_id,
tie_word_embeddings=tie_word_embeddings,
**kwargs,
)
@property
def head_dim(self):
assert self.hidden_size % self.num_attention_heads == 0
return self.hidden_size // self.num_attention_heads