"""nanochat-GPT for HuggingFace transformers (custom code, trust_remote_code). Derived from karpathy/nanochat gpt.py (MIT License, Copyright (c) 2025 Andrej Karpathy). The always-on pieces: - decoder-only transformer, causal attention - rotary position embeddings (base 100000, nanochat's half-split convention) - RMSNorm with no learnable parameters (after embedding, pre-attn, pre-MLP, final) - QK norm: queries/keys RMS-normalized AFTER rotary, no learnable weight - MLP with relu(x)^2 activation, no gating - no biases anywhere, untied input embedding / output head - logit softcap: logits = softcap * tanh(logits / softcap), in float32 The speedrun mechanisms, each enabled by its config field (see configuration_nanochat_gpt.py; all off = the clean d26-style architecture): - sliding-window attention (window_pattern, "S"/"L" tiled across layers) - value embeddings: per-layer token-embedding tables mixed into the attention values through a learned per-head sigmoid gate (ResFormer-style) - x0 re-injection and per-layer residual scaling (x0_lambdas, resid_lambdas) - smear: gated mix of the previous token's embedding into the current one - backout: subtract a scaled mid-layer residual before the final norm - QK sharpening: fixed scale on queries and keys after QK norm Numerical intent: weights are stored in bfloat16 and all matmuls run in bfloat16 (this matches training, where fp32 master weights were cast to bfloat16 for every forward). The logit softcap and the loss run in float32. Every mechanism follows the reference (ppriors/utils/gpt.py) operation by operation, in the same order and dtype flow, so logits reproduce the original model bit for bit on the same kernel (verified in verify.py). """ from typing import Optional import torch import torch.nn as nn import torch.nn.functional as F from transformers import PreTrainedModel from transformers.generation import GenerationMixin from transformers.cache_utils import Cache, DynamicCache from transformers.modeling_outputs import BaseModelOutputWithPast, CausalLMOutputWithPast from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS from .configuration_nanochat_gpt import NanochatGPTConfig # Attention implementations that take the REFERENCE path # (sliding_window_sdpa — bit-identical to nanochat's SDPA kernel, what # verify.py certifies against the training checkpoint). Anything else # (e.g. the "vllm" implementation vLLM's transformers backend patches into # config._attn_implementation) dispatches through HF's attention-interface # registry; whether that engine's numbers match is what the equivalence gate # (docs/archive/vllm-eval-acceptance.md) adjudicates. "eager" deliberately maps to # the reference path too: this export has always run one attention code # path, and silently switching kernels on an innocuous-looking config # default would invalidate the verify.py certificate. REFERENCE_ATTN_IMPLS = (None, "sdpa", "eager") class NanochatDynamicCache(DynamicCache): """DynamicCache plus the smear stash: the pre-smear embedding of the newest position, consumed by the next single-token decode. The stash must follow every batch-dimension shuffle of the k/v tensors. Beam search permutes the cache between steps via reorder_cache(beam_idx); a stash stored as a plain attribute on a stock DynamicCache stayed in the OLD beam order, so each beam smeared with another beam's embedding — silently wrong logits from the first reorder on (reviewer finding on 649aa3e: beam(3) diverged from the 3rd generated token). Cropping (assisted-decoding rollback) is refused: the stash holds only the newest position, so after a shrink the right embedding is gone. """ nanochat_prev_embedding = None # class default; instances stash their own def reorder_cache(self, beam_idx): super().reorder_cache(beam_idx) prev = self.nanochat_prev_embedding if prev is not None: self.nanochat_prev_embedding = prev.index_select(0, beam_idx.to(prev.device)) def batch_repeat_interleave(self, repeats): super().batch_repeat_interleave(repeats) prev = self.nanochat_prev_embedding if prev is not None: self.nanochat_prev_embedding = prev.repeat_interleave(repeats, dim=0) def batch_select_indices(self, indices): super().batch_select_indices(indices) prev = self.nanochat_prev_embedding if prev is not None: self.nanochat_prev_embedding = prev.index_select(0, indices.to(prev.device)) def crop(self, max_length): assert self.nanochat_prev_embedding is None or \ max_length >= self.get_seq_length(), ( "cropping a smear model's KV cache is unsupported: the cache " "stashes only the NEWEST position's pre-smear embedding, so a " "shrunk cache would smear with a stale embedding (silently wrong " "logits). Assisted decoding needs cropping; run without an " "assistant model." ) super().crop(max_length) def rms_norm(x): # RMSNorm without learnable parameters, computed by the framework kernel # (same call as nanochat) so results match the original bit-for-bit. return F.rms_norm(x, (x.size(-1),)) def apply_rotary_emb(x, cos, sin): # nanochat convention: rotates by -theta relative to the textbook # convention (only the relative q/k rotation matters, but q and k must # both use this exact form to reproduce the checkpoint). assert x.ndim == 4 # (B, T, H, D) d = x.shape[3] // 2 x1, x2 = x[..., :d], x[..., d:] y1 = x1 * cos + x2 * sin y2 = x1 * (-sin) + x2 * cos return torch.cat([y1, y2], 3) def compute_rotary_cos_sin(positions, head_dim, base, device, dtype): """cos/sin of shape (1, T, 1, head_dim/2), computed in fp32 then cast (nanochat computes its rotary cache the same way).""" channel_range = torch.arange(0, head_dim, 2, dtype=torch.float32, device=device) inv_freq = 1.0 / (base ** (channel_range / head_dim)) t = positions.to(device=device, dtype=torch.float32) freqs = torch.outer(t, inv_freq) cos, sin = freqs.cos(), freqs.sin() cos, sin = cos.to(dtype), sin.to(dtype) return cos[None, :, None, :], sin[None, :, None, :] def compute_window_sizes(config: NanochatGPTConfig): """Per-layer left attention window, ported from nanochat GPT._compute_window_sizes. The pattern string is tiled across layers; the final layer is always L. L = the full trained context (max_position_embeddings); S = quarter context, rounded up to a 128 multiple (nanochat rounds to the FA3 tile). """ pattern = config.window_pattern.upper() assert all(c in "SL" for c in pattern), f"Invalid window_pattern: {pattern}. Use only S and L." long_window = config.max_position_embeddings short_window = -(-long_window // 4 // 128) * 128 char_to_window = {"L": long_window, "S": short_window} window_sizes = [char_to_window[pattern[i % len(pattern)]] for i in range(config.num_hidden_layers)] window_sizes[-1] = long_window return window_sizes def sliding_window_sdpa(q, k, v, window, enable_gqa): """SDPA with nanochat's left-window semantics (a row attends to itself and the `window` previous positions). Ported from ppriors/utils/flash_attention._sdpa_attention, the kernel the reference model runs when Flash Attention 3 is unavailable (CPU verification). q, k, v are (B, H, T, D); k/v already include any cached positions. """ Tq = q.size(2) Tk = k.size(2) # Full context, same length if (window < 0 or window >= Tq) and Tq == Tk: return F.scaled_dot_product_attention(q, k, v, is_causal=True, enable_gqa=enable_gqa) # Single token generation if Tq == 1: if window >= 0 and window < Tk: # window is "left" tokens: include (window + 1) keys total start = max(0, Tk - (window + 1)) k = k[:, :, start:, :] v = v[:, :, start:, :] return F.scaled_dot_product_attention(q, k, v, is_causal=False, enable_gqa=enable_gqa) # Sliding window and/or chunked prefill on a cache: explicit bool mask device = q.device row_idx = (Tk - Tq) + torch.arange(Tq, device=device).unsqueeze(1) col_idx = torch.arange(Tk, device=device).unsqueeze(0) mask = col_idx <= row_idx if window >= 0 and window < Tk: mask = mask & ((row_idx - col_idx) <= window) return F.scaled_dot_product_attention(q, k, v, attn_mask=mask, enable_gqa=enable_gqa) class NanochatGPTAttention(nn.Module): def __init__(self, config: NanochatGPTConfig, layer_idx: int): super().__init__() self.layer_idx = layer_idx self.n_head = config.num_attention_heads self.n_kv_head = config.num_key_value_heads self.head_dim = config.head_dim self.q_proj = nn.Linear(config.hidden_size, self.n_head * self.head_dim, bias=False) self.k_proj = nn.Linear(config.hidden_size, self.n_kv_head * self.head_dim, bias=False) self.v_proj = nn.Linear(config.hidden_size, self.n_kv_head * self.head_dim, bias=False) self.o_proj = nn.Linear(config.hidden_size, config.hidden_size, bias=False) self.ve_gate_channels = config.ve_gate_channels self.ve_gate = ( nn.Linear(self.ve_gate_channels, self.n_kv_head, bias=False) if layer_idx in config.value_embedding_layers else None ) self.qk_sharpen_scale = config.qk_sharpen_scale self.window = None # left attention window (int), set by NanochatGPTModel # Attributes HF attention-interface implementations read off the module. self.config = config self.is_causal = True self.num_key_value_groups = self.n_head // self.n_kv_head self.scaling = self.head_dim**-0.5 # SDPA's default scale, made explicit def forward(self, x, ve, cos_sin, past_key_values: Optional[Cache], cache_position, **kwargs): B, T, C = x.size() q = self.q_proj(x).view(B, T, self.n_head, self.head_dim) k = self.k_proj(x).view(B, T, self.n_kv_head, self.head_dim) v = self.v_proj(x).view(B, T, self.n_kv_head, self.head_dim) # Value residual (ResFormer): mix in the value embedding with an # input-dependent gate per kv head, before rotary/QK norm (which do # not touch v anyway) — same point as the reference. assert (ve is None) == (self.ve_gate is None), ( f"layer {self.layer_idx}: value embedding and gate must appear together" ) if ve is not None: ve = ve.view(B, T, self.n_kv_head, self.head_dim) gate = 3 * torch.sigmoid(self.ve_gate(x[..., :self.ve_gate_channels])) # (B, T, n_kv_head), range (0, 3) v = v + gate.unsqueeze(-1) * ve cos, sin = cos_sin q, k = apply_rotary_emb(q, cos, sin), apply_rotary_emb(k, cos, sin) q, k = rms_norm(q), rms_norm(k) # QK norm, after rotary if self.qk_sharpen_scale is not None: q = q * self.qk_sharpen_scale # sharper attention, scale split between Q and K k = k * self.qk_sharpen_scale # SDPA layout (B, H, T, D) q = q.transpose(1, 2) k = k.transpose(1, 2) v = v.transpose(1, 2) if past_key_values is not None: k, v = past_key_values.update(k, v, self.layer_idx) assert self.window is not None, "window not set (NanochatGPTModel wires it)" impl = getattr(self.config, "_attn_implementation", None) if impl in REFERENCE_ATTN_IMPLS: # The reference path: exactly the kernel verify.py certifies. enable_gqa = self.n_kv_head != self.n_head y = sliding_window_sdpa(q, k, v, self.window, enable_gqa) y = y.transpose(1, 2).contiguous().view(B, T, -1) else: # Engine path (e.g. vLLM's "vllm" implementation): dispatch through # HF's attention-interface registry. The engine owns KV caching and # window/causality (vLLM: per-layer windows from config.layer_types # + config.sliding_window); q/k/v here carry everything upstream of # attention (rotary, QK norm, sharpening, value-embedding mix). # Interface convention: q/k/v in (B, H, T, D), output (B, T, H, D). attention_interface = ALL_ATTENTION_FUNCTIONS[impl] y, _ = attention_interface( self, q, k, v, None, scaling=self.scaling, **kwargs ) y = y.reshape(B, T, -1).contiguous() return self.o_proj(y) class NanochatGPTMLP(nn.Module): def __init__(self, config: NanochatGPTConfig): super().__init__() self.up_proj = nn.Linear(config.hidden_size, config.intermediate_size, bias=False) self.down_proj = nn.Linear(config.intermediate_size, config.hidden_size, bias=False) def forward(self, x): return self.down_proj(F.relu(self.up_proj(x)).square()) class NanochatGPTBlock(nn.Module): def __init__(self, config: NanochatGPTConfig, layer_idx: int): super().__init__() self.self_attn = NanochatGPTAttention(config, layer_idx) self.mlp = NanochatGPTMLP(config) def forward(self, x, ve, cos_sin, past_key_values, cache_position, **kwargs): x = x + self.self_attn(rms_norm(x), ve, cos_sin, past_key_values, cache_position, **kwargs) x = x + self.mlp(rms_norm(x)) return x class NanochatGPTPreTrainedModel(PreTrainedModel): config_class = NanochatGPTConfig base_model_prefix = "model" supports_gradient_checkpointing = False _no_split_modules = ["NanochatGPTBlock"] _supports_sdpa = True _supports_cache_class = True # Attention routes through HF's attention-interface registry when a # non-reference implementation is patched in (REFERENCE_ATTN_IMPLS above), # which is what vLLM's transformers backend requires # (is_backend_compatible reads this flag). _supports_attention_backend = True def _init_weights(self, module): # Export-only model: weights always come from a converted checkpoint. if isinstance(module, nn.Linear): module.weight.data.normal_(mean=0.0, std=0.02) elif isinstance(module, nn.Embedding): module.weight.data.normal_(mean=0.0, std=0.02) class NanochatGPTModel(NanochatGPTPreTrainedModel): def __init__(self, config: NanochatGPTConfig): super().__init__(config) self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size) self.layers = nn.ModuleList( [NanochatGPTBlock(config, layer_idx) for layer_idx in range(config.num_hidden_layers)] ) self.window_sizes = compute_window_sizes(config) for layer, window in zip(self.layers, self.window_sizes): layer.self_attn.window = window # The config's engine-facing layer_types/sliding_window (what vLLM # builds its attention from) must describe the SAME windows this # module enforces on the reference path — two derivations of one # pattern, pinned against drift here. layer_types = getattr(config, "layer_types", None) if layer_types is not None: long_window = config.max_position_embeddings expected = ["sliding_attention" if w < long_window else "full_attention" for w in self.window_sizes] short = [w for w in self.window_sizes if w < long_window] assert list(layer_types) == expected and \ all(w + 1 == config.sliding_window for w in short), ( "config.layer_types/sliding_window disagree with " "compute_window_sizes — the engine would attend differently " f"than the reference: {layer_types} vs {expected}, " f"sliding_window={getattr(config, 'sliding_window', None)}" ) # Mechanism parameters exist only when their mechanism is on, so the # clean-architecture state dict (older exports) still loads strictly. n_layer = config.num_hidden_layers if config.use_resid_lambdas: self.resid_lambdas = nn.Parameter(torch.ones(n_layer)) if config.use_x0_lambdas: self.x0_lambdas = nn.Parameter(torch.zeros(n_layer)) if config.use_smear: self.smear_gate = nn.Linear(config.smear_gate_channels, 1, bias=False) self.smear_lambda = nn.Parameter(torch.zeros(1)) if config.backout_layer is not None: self.backout_lambda = nn.Parameter(torch.zeros(1)) kv_dim = config.num_key_value_heads * config.head_dim self.value_embeds = nn.ModuleDict( {str(i): nn.Embedding(config.vocab_size, kv_dim) for i in config.value_embedding_layers} ) self.post_init() def _smear(self, x, past_key_values, cache_position): """Mix the previous token's (pre-smear) embedding into each position. Mirrors nanochat GPT.forward: full-sequence smear when every position is present; with a KV cache, the pre-smear embedding of the newest position is stashed on the cache object and consumed by the next single-token decode step. """ B, T, C = x.size() ch = self.config.smear_gate_channels if past_key_values is None: # Full sequence available (no cache): position 0 has no predecessor. assert T > 1, "smear on a full sequence needs T > 1" gate = self.smear_lambda.to(x.dtype) * torch.sigmoid(self.smear_gate(x[:, 1:, :ch])) return torch.cat([x[:, :1], x[:, 1:] + gate * x[:, :-1]], dim=1) prev = getattr(past_key_values, "nanochat_prev_embedding", None) past_key_values.nanochat_prev_embedding = x[:, -1:, :] # pre-smear, for the next step if T > 1: # Prefill: smear positions 1+, same as the full-sequence path. gate = self.smear_lambda.to(x.dtype) * torch.sigmoid(self.smear_gate(x[:, 1:, :ch])) return torch.cat([x[:, :1], x[:, 1:] + gate * x[:, :-1]], dim=1) if int(cache_position[0]) == 0: return x # single-token prefill at position 0: no predecessor exists # Single-token decode: the previous step must have stashed its embedding. # Refusing beats silently skipping the smear (wrong logits). assert prev is not None, ( "single-token decode past position 0 without a stashed previous " "embedding: the cache was not built by this model's forward" ) gate = self.smear_lambda.to(x.dtype) * torch.sigmoid(self.smear_gate(x[:, :, :ch])) return x + gate * prev def forward( self, input_ids: Optional[torch.LongTensor] = None, attention_mask: Optional[torch.Tensor] = None, past_key_values: Optional[Cache] = None, use_cache: Optional[bool] = None, cache_position: Optional[torch.LongTensor] = None, position_ids: Optional[torch.LongTensor] = None, inputs_embeds: Optional[torch.Tensor] = None, **kwargs, ): assert (input_ids is None) != (inputs_embeds is None), ( "pass exactly one of input_ids / inputs_embeds" ) B, T = input_ids.size() if input_ids is not None else inputs_embeds.shape[:2] if attention_mask is not None: assert bool(torch.all(attention_mask == 1)), ( "NanochatGPT does not support padded batches; use batch size 1 " "or unpadded sequences." ) # Under an engine (vLLM's transformers backend) requests are packed # into one flattened row: positions restart at each request boundary # inside dim 1, and the engine owns attention. Mechanisms that mix # information ACROSS positions in our own code (smear) or need token # ids we were not given (value embeddings under an inputs_embeds-only # call) would silently cross request boundaries or cannot run — refuse # loudly instead. engine_packed = "attention_instances" in kwargs if engine_packed: assert not self.config.use_smear, ( "smear models cannot run under an engine that packs requests " "into one row: the previous-token embedding mix would cross " "request boundaries (silently wrong logits). Run smear models " "on the HF path." ) if self.value_embeds: assert input_ids is not None, ( "value-embedding models need input_ids (per-layer token-id " "lookups); this call passed only inputs_embeds" ) if use_cache and past_key_values is None: past_key_values = NanochatDynamicCache() if use_cache and self.config.use_smear and \ not isinstance(past_key_values, NanochatDynamicCache): # generate() constructs a stock DynamicCache and passes it in. # The smear stash must follow beam-search reorder (and refuse # crop), so the EMPTY stock cache is grafted onto the stash-aware # subclass in place — keeping all internal state and the object # identity generate() holds. Any other cache cannot keep the # stash in sync; refusing beats silently smearing with another # batch row's embedding. assert type(past_key_values) is DynamicCache and \ past_key_values.get_seq_length() == 0, ( "smear models support only the default dynamic KV cache: pass " "past_key_values=None or a fresh DynamicCache, got " f"{type(past_key_values).__name__} with " f"{past_key_values.get_seq_length()} cached positions" ) past_key_values.__class__ = NanochatDynamicCache device = input_ids.device if input_ids is not None else inputs_embeds.device if cache_position is None: past_len = past_key_values.get_seq_length() if past_key_values is not None else 0 cache_position = torch.arange(past_len, past_len + T, device=device) cache = past_key_values if use_cache else None x = inputs_embeds if inputs_embeds is not None else self.embed_tokens(input_ids) x = rms_norm(x) if self.config.use_smear: x = self._smear(x, cache, cache_position) # Rotary positions: an engine passes explicit position_ids (packed # rows restart positions per request); the HF path derives them from # the cache. The rope table is built once per forward from 1-D # positions and broadcast over the batch, so distinct per-row # positions are refused rather than silently rotated wrong. if position_ids is not None: assert position_ids.ndim == 2, position_ids.shape assert bool(torch.all(position_ids == position_ids[0:1])), ( "per-row position_ids differ; this model broadcasts one " "rotary table over the batch" ) rope_positions = position_ids[0] else: rope_positions = cache_position cos_sin = compute_rotary_cos_sin( rope_positions, self.config.head_dim, self.config.rope_theta, x.device, x.dtype ) x0 = x # initial (post-smear) normalized embedding, for x0 re-injection use_resid = self.config.use_resid_lambdas use_x0 = self.config.use_x0_lambdas backout_layer = self.config.backout_layer x_backout = None for i, layer in enumerate(self.layers): # Same branch structure and expressions as the reference so the # bf16 rounding sequence is identical. if not use_resid and not use_x0: pass elif not use_x0: x = self.resid_lambdas[i] * x elif not use_resid: x = x + self.x0_lambdas[i] * x0 else: x = self.resid_lambdas[i] * x + self.x0_lambdas[i] * x0 ve = self.value_embeds[str(i)](input_ids).to(x.dtype) if str(i) in self.value_embeds else None x = layer(x, ve, cos_sin, cache, cache_position, **kwargs) if i == backout_layer: x_backout = x if backout_layer is not None: assert x_backout is not None x = x - self.backout_lambda.to(x.dtype) * x_backout x = rms_norm(x) return BaseModelOutputWithPast( last_hidden_state=x, past_key_values=past_key_values if use_cache else None, ) class NanochatGPTForCausalLM(NanochatGPTPreTrainedModel, GenerationMixin): _tied_weights_keys = [] def __init__(self, config: NanochatGPTConfig): super().__init__(config) self.model = NanochatGPTModel(config) self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) self.post_init() def get_input_embeddings(self): return self.model.embed_tokens def set_input_embeddings(self, value): self.model.embed_tokens = value def get_output_embeddings(self): return self.lm_head def forward( self, input_ids: Optional[torch.LongTensor] = None, attention_mask: Optional[torch.Tensor] = None, past_key_values: Optional[Cache] = None, labels: Optional[torch.LongTensor] = None, use_cache: Optional[bool] = None, cache_position: Optional[torch.LongTensor] = None, position_ids: Optional[torch.LongTensor] = None, inputs_embeds: Optional[torch.Tensor] = None, **kwargs, ): outputs = self.model( input_ids=input_ids, attention_mask=attention_mask, past_key_values=past_key_values, use_cache=use_cache, cache_position=cache_position, position_ids=position_ids, inputs_embeds=inputs_embeds, ) logits = self.lm_head(outputs.last_hidden_state) logits = logits.float() # fp32 for softcap and loss, as in training softcap = self.config.logit_softcap if softcap is not None and softcap > 0: logits = softcap * torch.tanh(logits / softcap) loss = None if labels is not None: loss = F.cross_entropy( logits[:, :-1].reshape(-1, logits.size(-1)), labels[:, 1:].reshape(-1), ignore_index=-100, ) return CausalLMOutputWithPast( loss=loss, logits=logits, past_key_values=outputs.past_key_values, )