Download model.py from ahiok/looped-fineweb-10m: direct link, hf CLI and curl.
- Browser
- Download file 28.1 kB
-
https://huggingface.co/ahiok/looped-fineweb-10m/resolve/main/model.py
- Command line
-
hf download hf://ahiok/looped-fineweb-10m/model.py
-
curl -L -o model.py https://huggingface.co/ahiok/looped-fineweb-10m/resolve/main/model.py
28.1 kB
| """Looped decoder-only LM with a Qwen3-style block. | |
| Layout follows the prelude / recurrent / coda decomposition: | |
| x -> embed -> [prelude L_p layers] -> e | |
| s_0 = e | |
| s_r = Block(s_{r-1}, e, r, R) for r = 1..R (shared weights) | |
| logits = head(norm(coda(s_R))) | |
| Every research knob is a config flag so that one binary can produce the whole | |
| ablation ladder and every run is described by its config dict alone. | |
| """ | |
| from __future__ import annotations | |
| import math | |
| from dataclasses import asdict, dataclass, field | |
| from typing import Optional | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| import torch.utils.checkpoint | |
| # -------------------------------------------------------------------------------------- | |
| # config | |
| # -------------------------------------------------------------------------------------- | |
| class ModelConfig: | |
| # --- Qwen3-style backbone ----------------------------------------------------- | |
| vocab_size: int = 8192 | |
| d_model: int = 384 | |
| n_heads: int = 6 | |
| n_kv_heads: int = 2 | |
| head_dim: int = 64 | |
| d_ff: int = 1024 | |
| max_seq_len: int = 512 | |
| rope_theta: float = 10_000.0 | |
| rms_eps: float = 1e-6 | |
| tie_embeddings: bool = True | |
| # pre = Qwen3 default. sandwich = Huginn's block, which normalises after each | |
| # residual add as well; costs 2d per layer and bounds the residual stream. | |
| block_norm: str = "pre" # pre | sandwich | |
| # --- depth layout ------------------------------------------------------------- | |
| n_prelude: int = 1 | |
| n_recurrent: int = 2 | |
| n_coda: int = 1 | |
| # --- looping ------------------------------------------------------------------ | |
| n_loops: int = 8 # R used at train time (mean of the distribution if sampled) | |
| max_loops: int = 256 # size of the precomputed depth-embedding table | |
| state_init: str = "prelude" # prelude | randn | |
| state_init_std: float = 0.4 # only for state_init == "randn" | |
| input_injection: str = "add" # none | add | adapter | |
| state_norm: str = "none" # none | rms (normalise s at loop entry) | |
| # residual : s <- Block(s) (the usual looped transformer) | |
| # convex : s <- (1-a) s + a Block(s) (learned step size) | |
| # flow : s <- s + (gain/R) * Delta(s, r/R) (explicit Euler step of a learned flow) | |
| update_rule: str = "residual" | |
| # pre-sigmoid init of the convex step size. +3 starts at ~0.95, i.e. almost a | |
| # full replacement (the usual looped behaviour); -3 starts at ~0.05, so the | |
| # loop begins as a near-identity and has to earn its depth, which is what | |
| # makes very deep shared stacks trainable at all | |
| update_gate_init: float = 3.0 | |
| depth_cond: str = "none" # none | film | |
| depth_cond_input: str = "progress" # absolute | progress | both | |
| depth_cond_dim: int = 64 | |
| loop_noise: float = 0.0 # std of exploration noise injected at loop entry | |
| noise_schedule: str = "linear" # linear | const | cosine (annealed towards 0 at r=R) | |
| # learned halting, PonderNet style: a per-token probability of stopping after | |
| # each iteration, trained jointly with the language-model loss. Costs d + 1 | |
| # parameters and is independent of R, so the maximum depth stays a runtime knob. | |
| halting: str = "none" # none | ponder | |
| halt_prior: float = 0.1 # geometric prior on the halting step | |
| halt_kl_weight: float = 0.01 | |
| # --- init --------------------------------------------------------------------- | |
| init_std: float = 0.02 | |
| depth_scaled_init: bool = True | |
| def __post_init__(self) -> None: | |
| assert self.n_heads % self.n_kv_heads == 0 | |
| assert self.state_init in {"prelude", "randn"} | |
| assert self.input_injection in {"none", "add", "adapter"} | |
| assert self.block_norm in {"pre", "sandwich"} | |
| assert self.state_norm in {"none", "rms"} | |
| assert self.update_rule in {"residual", "convex", "flow"} | |
| assert self.depth_cond in {"none", "film"} | |
| assert self.depth_cond_input in {"absolute", "progress", "both"} | |
| assert self.noise_schedule in {"linear", "const", "cosine"} | |
| assert self.halting in {"none", "ponder"} | |
| def to_dict(self) -> dict: | |
| return asdict(self) | |
| # -------------------------------------------------------------------------------------- | |
| # primitives | |
| # -------------------------------------------------------------------------------------- | |
| class RMSNorm(nn.Module): | |
| def __init__(self, dim: int, eps: float = 1e-6, elementwise_affine: bool = True): | |
| super().__init__() | |
| self.eps = eps | |
| self.weight = nn.Parameter(torch.ones(dim)) if elementwise_affine else None | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| dtype = x.dtype | |
| x = x.float() | |
| x = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) | |
| x = x.to(dtype) | |
| return x * self.weight if self.weight is not None else x | |
| def build_rope_cache(seq_len: int, head_dim: int, theta: float, device, dtype=torch.float32): | |
| inv_freq = 1.0 / (theta ** (torch.arange(0, head_dim, 2, device=device, dtype=torch.float32) / head_dim)) | |
| t = torch.arange(seq_len, device=device, dtype=torch.float32) | |
| freqs = torch.outer(t, inv_freq) # (T, hd/2) | |
| emb = torch.cat((freqs, freqs), dim=-1) # (T, hd) | |
| return emb.cos().to(dtype), emb.sin().to(dtype) | |
| def rotate_half(x: torch.Tensor) -> torch.Tensor: | |
| x1, x2 = x.chunk(2, dim=-1) | |
| return torch.cat((-x2, x1), dim=-1) | |
| def apply_rope(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: | |
| # x: (B, H, T, hd); cos/sin: (T, hd) | |
| cos = cos[None, None, :, :] | |
| sin = sin[None, None, :, :] | |
| return x * cos + rotate_half(x) * sin | |
| class Attention(nn.Module): | |
| """Qwen3 attention: GQA, no qkv bias, RMSNorm on q and k heads.""" | |
| def __init__(self, cfg: ModelConfig): | |
| super().__init__() | |
| self.n_heads = cfg.n_heads | |
| self.n_kv_heads = cfg.n_kv_heads | |
| self.head_dim = cfg.head_dim | |
| self.q_proj = nn.Linear(cfg.d_model, cfg.n_heads * cfg.head_dim, bias=False) | |
| self.k_proj = nn.Linear(cfg.d_model, cfg.n_kv_heads * cfg.head_dim, bias=False) | |
| self.v_proj = nn.Linear(cfg.d_model, cfg.n_kv_heads * cfg.head_dim, bias=False) | |
| self.o_proj = nn.Linear(cfg.n_heads * cfg.head_dim, cfg.d_model, bias=False) | |
| self.q_norm = RMSNorm(cfg.head_dim, cfg.rms_eps) | |
| self.k_norm = RMSNorm(cfg.head_dim, cfg.rms_eps) | |
| def forward(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: | |
| B, T, _ = x.shape | |
| q = self.q_proj(x).view(B, T, self.n_heads, self.head_dim).transpose(1, 2) | |
| k = self.k_proj(x).view(B, T, self.n_kv_heads, self.head_dim).transpose(1, 2) | |
| v = self.v_proj(x).view(B, T, self.n_kv_heads, self.head_dim).transpose(1, 2) | |
| q = self.q_norm(q) | |
| k = self.k_norm(k) | |
| q = apply_rope(q, cos, sin) | |
| k = apply_rope(k, cos, sin) | |
| o = F.scaled_dot_product_attention(q, k, v, is_causal=True, enable_gqa=True) | |
| o = o.transpose(1, 2).contiguous().view(B, T, self.n_heads * self.head_dim) | |
| return self.o_proj(o) | |
| class MLP(nn.Module): | |
| def __init__(self, cfg: ModelConfig): | |
| super().__init__() | |
| self.gate_proj = nn.Linear(cfg.d_model, cfg.d_ff, bias=False) | |
| self.up_proj = nn.Linear(cfg.d_model, cfg.d_ff, bias=False) | |
| self.down_proj = nn.Linear(cfg.d_ff, cfg.d_model, bias=False) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x)) | |
| class DecoderLayer(nn.Module): | |
| """Qwen3 pre-norm layer, optionally with Huginn's sandwich norm. | |
| Pre-norm (`block_norm="pre"`) is the Qwen3 default: the stream is normalised | |
| on the way *into* each sublayer and the residual add is left alone, so the | |
| stream is free to grow. Section 4.2 measures that growth and identifies it as | |
| the reason late iterations stop mattering. | |
| Sandwich (`block_norm="sandwich"`) is what the Huginn recurrent block | |
| actually does: it normalises again *after* each residual add, which bounds | |
| the stream without removing the residual path itself. That distinction is the | |
| whole reason the loop-entry normalisation of 5.2 failed and this does not. | |
| """ | |
| def __init__(self, cfg: ModelConfig): | |
| super().__init__() | |
| self.input_layernorm = RMSNorm(cfg.d_model, cfg.rms_eps) | |
| self.self_attn = Attention(cfg) | |
| self.post_attention_layernorm = RMSNorm(cfg.d_model, cfg.rms_eps) | |
| self.mlp = MLP(cfg) | |
| self.sandwich = cfg.block_norm == "sandwich" | |
| if self.sandwich: | |
| self.post_attn_residual_norm = RMSNorm(cfg.d_model, cfg.rms_eps) | |
| self.post_mlp_residual_norm = RMSNorm(cfg.d_model, cfg.rms_eps) | |
| def forward(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: | |
| x = x + self.self_attn(self.input_layernorm(x), cos, sin) | |
| if self.sandwich: | |
| x = self.post_attn_residual_norm(x) | |
| x = x + self.mlp(self.post_attention_layernorm(x)) | |
| if self.sandwich: | |
| x = self.post_mlp_residual_norm(x) | |
| return x | |
| # -------------------------------------------------------------------------------------- | |
| # recurrent block | |
| # -------------------------------------------------------------------------------------- | |
| class RecurrentBlock(nn.Module): | |
| """The shared block applied R times. | |
| Everything that makes iteration r behave differently from iteration r+1 has | |
| to enter here, because the weights themselves are identical across r. | |
| """ | |
| def __init__(self, cfg: ModelConfig): | |
| super().__init__() | |
| self.cfg = cfg | |
| self.layers = nn.ModuleList([DecoderLayer(cfg) for _ in range(cfg.n_recurrent)]) | |
| if cfg.input_injection == "adapter": | |
| self.adapter = nn.Linear(2 * cfg.d_model, cfg.d_model, bias=False) | |
| if cfg.state_norm == "rms": | |
| self.entry_norm = RMSNorm(cfg.d_model, cfg.rms_eps) | |
| if cfg.depth_cond == "film": | |
| # sinusoidal features -> (scale, shift). Cost is O(d), independent of R, | |
| # which is what keeps this usable at larger scale. | |
| self.film = nn.Linear(cfg.depth_cond_dim, 2 * cfg.d_model, bias=True) | |
| nn.init.zeros_(self.film.weight) | |
| nn.init.zeros_(self.film.bias) | |
| if cfg.update_rule == "convex": | |
| # learned per-channel step size, sigmoid-gated, initialised near 1.0 so the | |
| # untouched model starts out identical to the plain residual update | |
| self.alpha = nn.Parameter(torch.full((cfg.d_model,), float(cfg.update_gate_init))) | |
| elif cfg.update_rule == "flow": | |
| # learned per-channel speed of the flow; the 1/R factor lives in forward() | |
| self.flow_gain = nn.Parameter(torch.ones(cfg.d_model)) | |
| def forward(self, s, e, cos, sin, depth_feat: Optional[torch.Tensor] = None, | |
| noise_std: Optional[torch.Tensor] = None, step_scale: Optional[torch.Tensor] = None): | |
| # noise_std and step_scale arrive as 0-dim tensors on purpose: as python | |
| # floats dynamo specialises the graph on their value and recompiles the | |
| # block for every distinct loop count and noise level. | |
| cfg = self.cfg | |
| h = s | |
| if cfg.input_injection == "add": | |
| h = h + e | |
| elif cfg.input_injection == "adapter": | |
| h = self.adapter(torch.cat([h, e], dim=-1)) | |
| if cfg.state_norm == "rms": | |
| h = self.entry_norm(h) | |
| if cfg.depth_cond == "film" and depth_feat is not None: | |
| mod = self.film(depth_feat) # (2d,) | |
| scale, shift = mod.chunk(2, dim=-1) | |
| h = h * (1.0 + scale) + shift | |
| if cfg.loop_noise > 0.0 and noise_std is not None: | |
| h = h + noise_std * torch.randn_like(h) | |
| inner = h | |
| for layer in self.layers: | |
| inner = layer(inner, cos, sin) | |
| if cfg.update_rule == "convex": | |
| a = torch.sigmoid(self.alpha) | |
| return (1.0 - a) * s + a * inner | |
| if cfg.update_rule == "flow": | |
| # explicit Euler step: s' = s + h_step * g(s, r/R). The loop count then | |
| # sets the integration resolution rather than the amount of drift, so | |
| # raising R at inference refines the same trajectory instead of | |
| # walking further along it. | |
| return s + (step_scale * self.flow_gain) * (inner - h) | |
| return inner | |
| # -------------------------------------------------------------------------------------- | |
| # full model | |
| # -------------------------------------------------------------------------------------- | |
| class LoopedLM(nn.Module): | |
| def __init__(self, cfg: ModelConfig): | |
| super().__init__() | |
| self.cfg = cfg | |
| self.embed_tokens = nn.Embedding(cfg.vocab_size, cfg.d_model) | |
| self.prelude = nn.ModuleList([DecoderLayer(cfg) for _ in range(cfg.n_prelude)]) | |
| self.block = RecurrentBlock(cfg) | |
| self.coda = nn.ModuleList([DecoderLayer(cfg) for _ in range(cfg.n_coda)]) | |
| self.norm = RMSNorm(cfg.d_model, cfg.rms_eps) | |
| if cfg.halting == "ponder": | |
| self.halt_head = nn.Linear(cfg.d_model, 1) | |
| self.lm_head = nn.Linear(cfg.d_model, cfg.vocab_size, bias=False) | |
| if cfg.tie_embeddings: | |
| self.lm_head.weight = self.embed_tokens.weight | |
| cos, sin = build_rope_cache(cfg.max_seq_len, cfg.head_dim, cfg.rope_theta, device="cpu") | |
| self.register_buffer("rope_cos", cos, persistent=False) | |
| self.register_buffer("rope_sin", sin, persistent=False) | |
| self._depth_cache: dict[tuple, torch.Tensor] = {} | |
| self.apply(self._init_weights) | |
| if cfg.depth_scaled_init: | |
| self._rescale_residual_projections() | |
| if cfg.depth_cond == "film": | |
| # zero-init the modulation so an untrained depth-conditioned model is | |
| # bit-identical to the unconditioned one at step 0 | |
| nn.init.zeros_(self.block.film.weight) | |
| nn.init.zeros_(self.block.film.bias) | |
| # -- init ------------------------------------------------------------------------ | |
| def _init_weights(self, module: nn.Module) -> None: | |
| std = self.cfg.init_std | |
| if isinstance(module, nn.Linear): | |
| nn.init.normal_(module.weight, mean=0.0, std=std) | |
| if module.bias is not None: | |
| nn.init.zeros_(module.bias) | |
| elif isinstance(module, nn.Embedding): | |
| nn.init.normal_(module.weight, mean=0.0, std=std) | |
| def _rescale_residual_projections(self) -> None: | |
| """GPT-2 style 1/sqrt(2 * depth) scaling, with depth counted through the loop. | |
| The flow update already divides every step by R, so counting the loop | |
| twice would leave the block effectively dead at initialisation. | |
| """ | |
| loops = 1 if self.cfg.update_rule == "flow" else self.cfg.n_loops | |
| depth = self.cfg.n_prelude + self.cfg.n_recurrent * loops + self.cfg.n_coda | |
| scale = 1.0 / math.sqrt(2.0 * max(depth, 1)) | |
| for mod in self.modules(): | |
| if isinstance(mod, DecoderLayer): | |
| mod.self_attn.o_proj.weight.data.mul_(scale) | |
| mod.mlp.down_proj.weight.data.mul_(scale) | |
| def _depth_table(self, R: int, device, dtype) -> torch.Tensor: | |
| """Sinusoidal encodings of every loop index, shape (R + 1, depth_cond_dim). | |
| Two things can be encoded: the absolute index r (tells the block how much | |
| work has been done) and the progress r/R (tells it how much is left). | |
| Which one matters is an experiment, not an assumption, hence the flag. | |
| Cached per (R, device, dtype) so the table is built once per run. | |
| """ | |
| cfg = self.cfg | |
| key = (R, str(device), str(dtype)) | |
| if key in self._depth_cache: | |
| return self._depth_cache[key] | |
| r = torch.arange(R + 1, device=device, dtype=torch.float32) | |
| vals = [] | |
| if cfg.depth_cond_input in {"absolute", "both"}: | |
| vals.append(r) | |
| if cfg.depth_cond_input in {"progress", "both"}: | |
| vals.append(r / max(R, 1) * 32.0) # rescale so low frequencies stay informative | |
| per = cfg.depth_cond_dim // (2 * len(vals)) | |
| idx = torch.arange(per, device=device, dtype=torch.float32) | |
| freq = torch.exp(-math.log(10_000.0) * idx / max(per - 1, 1)) | |
| feats = [] | |
| for v in vals: | |
| ang = v[:, None] * freq[None, :] | |
| feats.append(torch.cat([torch.sin(ang), torch.cos(ang)], dim=-1)) | |
| out = torch.cat(feats, dim=-1) | |
| if out.shape[-1] < cfg.depth_cond_dim: | |
| out = F.pad(out, (0, cfg.depth_cond_dim - out.shape[-1])) | |
| out = out.to(dtype) | |
| self._depth_cache[key] = out | |
| return out | |
| def _noise_std(self, r: int, R: int) -> float: | |
| cfg = self.cfg | |
| if cfg.loop_noise <= 0.0 or not self.training: | |
| return 0.0 | |
| if cfg.noise_schedule == "const": | |
| return cfg.loop_noise | |
| frac = (r - 1) / max(R - 1, 1) | |
| if cfg.noise_schedule == "linear": | |
| return cfg.loop_noise * (1.0 - frac) | |
| return cfg.loop_noise * 0.5 * (1.0 + math.cos(math.pi * frac)) | |
| # -- forward ---------------------------------------------------------------------- | |
| def _readout_hidden(self, s: torch.Tensor, cos, sin) -> torch.Tensor: | |
| """Coda output after the final norm; shared by the LM head and the halting head.""" | |
| h = s | |
| for layer in self.coda: | |
| h = layer(h, cos, sin) | |
| return self.norm(h) | |
| def _readout(self, s: torch.Tensor, cos, sin) -> torch.Tensor: | |
| return self.lm_head(self._readout_hidden(s, cos, sin)) | |
| def forward( | |
| self, | |
| idx: torch.Tensor, | |
| targets: Optional[torch.Tensor] = None, | |
| n_loops: Optional[int] = None, | |
| backprop_loops: int = 0, | |
| readout_loops: Optional[list[int]] = None, | |
| return_states: bool = False, | |
| grad_checkpoint: bool = False, | |
| readout_mode: str = "logits", | |
| ): | |
| """Run the model. | |
| Args: | |
| n_loops: R for this call (defaults to cfg.n_loops). | |
| backprop_loops: if > 0, only the last k iterations carry gradient. | |
| readout_loops: loop indices (1-based) whose intermediate logits are | |
| also returned, used for deep supervision and for the coda lens. | |
| return_states: also return the per-loop hidden states (diagnostics). | |
| grad_checkpoint: recompute each iteration's internals in the backward | |
| pass. Activation memory then stops growing with R, so a *full* | |
| backward through 32 or 64 loops fits, which truncation does not | |
| achieve without also changing what is being optimised. | |
| readout_mode: "logits" keeps every intermediate logit tensor, which is | |
| what deep supervision needs. "stats" reduces each one to per-token | |
| loss and confidence immediately and throws the logits away; a | |
| (B, T, 8192) tensor per loop is ~130 MB, so reading out all 32 | |
| loops for diagnostics costs gigabytes otherwise. "grad_stats" is | |
| the same reduction but keeps the graph, which is what the ponder | |
| objective needs: it weights every loop's loss by a learned halting | |
| probability and so requires all of them to be differentiable. | |
| """ | |
| cfg = self.cfg | |
| B, T = idx.shape | |
| R = n_loops if n_loops is not None else cfg.n_loops | |
| cos = self.rope_cos[:T].to(idx.device) | |
| sin = self.rope_sin[:T].to(idx.device) | |
| h = self.embed_tokens(idx) | |
| for layer in self.prelude: | |
| h = layer(h, cos, sin) | |
| e = h | |
| if cfg.state_init == "randn": | |
| s = torch.randn_like(e) * cfg.state_init_std | |
| else: | |
| s = e | |
| readout_set = set(readout_loops or []) | |
| aux_logits: dict[int, torch.Tensor] = {} | |
| aux_stats: dict[int, dict] = {} | |
| states = [s.detach()] if return_states else None | |
| halt_logits: dict[int, torch.Tensor] = {} | |
| def record(loop_idx: int, hidden: torch.Tensor) -> None: | |
| hn = self._readout_hidden(hidden, cos, sin) | |
| if cfg.halting == "ponder": | |
| halt_logits[loop_idx] = self.halt_head(hn).squeeze(-1).reshape(-1) | |
| lg = self.lm_head(hn) | |
| if readout_mode == "logits": | |
| aux_logits[loop_idx] = lg | |
| return | |
| flat = lg.float().view(-1, lg.size(-1)) | |
| tl = F.cross_entropy(flat, targets.reshape(-1), reduction="none") | |
| aux_stats[loop_idx] = { | |
| "token_loss": tl if readout_mode == "grad_stats" else tl.detach(), | |
| "confidence": F.softmax(flat, dim=-1).max(dim=-1).values.detach(), | |
| } | |
| no_grad_until = 0 | |
| if backprop_loops and backprop_loops < R: | |
| no_grad_until = R - backprop_loops | |
| depth_table = self._depth_table(R, s.device, s.dtype) if cfg.depth_cond == "film" else None | |
| step_scale = ( | |
| torch.tensor(1.0 / R, device=s.device, dtype=s.dtype) | |
| if cfg.update_rule == "flow" else None | |
| ) | |
| noise_table = ( | |
| torch.tensor([self._noise_std(r, R) for r in range(R + 1)], device=s.device, dtype=s.dtype) | |
| if cfg.loop_noise > 0.0 else None | |
| ) | |
| for r in range(1, R + 1): | |
| depth_feat = depth_table[r] if depth_table is not None else None | |
| noise = noise_table[r] if noise_table is not None else None | |
| if r <= no_grad_until: | |
| with torch.no_grad(): | |
| s = self.block(s, e, cos, sin, depth_feat, noise, step_scale) | |
| s = s.detach() | |
| elif grad_checkpoint and self.training: | |
| s = torch.utils.checkpoint.checkpoint( | |
| self.block, s, e, cos, sin, depth_feat, noise, step_scale, | |
| use_reentrant=False, | |
| ) | |
| else: | |
| s = self.block(s, e, cos, sin, depth_feat, noise, step_scale) | |
| if return_states: | |
| states.append(s.detach()) | |
| if r in readout_set and r != R: | |
| record(r, s) | |
| final_hidden = self._readout_hidden(s, cos, sin) | |
| logits = self.lm_head(final_hidden) | |
| if cfg.halting == "ponder": | |
| halt_logits[R] = self.halt_head(final_hidden).squeeze(-1).reshape(-1) | |
| if readout_mode in {"stats", "grad_stats"} and R in readout_set: | |
| flat = logits.float().view(-1, logits.size(-1)) | |
| tl = F.cross_entropy(flat, targets.reshape(-1), reduction="none") | |
| aux_stats[R] = { | |
| "token_loss": tl if readout_mode == "grad_stats" else tl.detach(), | |
| "confidence": F.softmax(flat, dim=-1).max(dim=-1).values.detach(), | |
| } | |
| loss = None | |
| if targets is not None: | |
| loss = F.cross_entropy( | |
| logits.float().view(-1, logits.size(-1)), targets.reshape(-1), ignore_index=-1 | |
| ) | |
| out = {"logits": logits, "loss": loss, "aux_logits": aux_logits, | |
| "aux_stats": aux_stats, "halt_logits": halt_logits, "n_loops": R} | |
| if return_states: | |
| out["states"] = states | |
| out["e"] = e.detach() | |
| return out | |
| # -- bookkeeping ------------------------------------------------------------------ | |
| def param_counts(self) -> dict: | |
| total = sum(p.numel() for p in self.parameters()) | |
| emb = self.embed_tokens.weight.numel() | |
| if not self.cfg.tie_embeddings: | |
| emb += self.lm_head.weight.numel() | |
| return {"total": total, "embedding": emb, "non_embedding": total - emb} | |
| def flops_per_token(self, n_loops: Optional[int] = None) -> float: | |
| """Forward FLOPs per token, counting matmuls only (attention scores included).""" | |
| cfg = self.cfg | |
| R = n_loops if n_loops is not None else cfg.n_loops | |
| d, hd = cfg.d_model, cfg.head_dim | |
| proj = 2 * d * (cfg.n_heads * hd) + 2 * d * (cfg.n_kv_heads * hd) # q,o and k,v | |
| mlp = 3 * d * cfg.d_ff | |
| attn_scores = 2 * cfg.n_heads * hd * cfg.max_seq_len / 2 # causal, averaged | |
| per_layer = 2 * (proj + mlp) + 2 * attn_scores | |
| n_layers = cfg.n_prelude + cfg.n_coda + cfg.n_recurrent * R | |
| return per_layer * n_layers + 2 * d * cfg.vocab_size | |
| def halting_distribution(halt_logits: torch.Tensor) -> torch.Tensor: | |
| """Per-token distribution over the halting step, PonderNet style. | |
| ``halt_logits`` is (R, N) pre-sigmoid. With ``lam_r`` the probability of | |
| stopping at r given that r was reached, | |
| p_r = lam_r * prod_{j<r} (1 - lam_j), | |
| and all remaining mass is forced onto r = R, since the loop cannot run | |
| further. Returns (R, N) summing to one along the loop axis. | |
| """ | |
| # float32 and a loose clamp on purpose: under bf16 autocast, 1 - 1e-6 rounds | |
| # to exactly 1.0, log1p(-1.0) is -inf, and the whole objective becomes NaN | |
| # within a few hundred steps. | |
| lam = torch.sigmoid(halt_logits.float()).clamp(1e-4, 1 - 1e-4) | |
| log_not = torch.log1p(-lam) | |
| # exclusive cumulative sum: log prod_{j<r} (1 - lam_j) | |
| cum = torch.cumsum(log_not, dim=0) - log_not | |
| p = lam * cum.exp() | |
| leftover = (cum[-1] + log_not[-1]).exp() | |
| return torch.cat([p[:-1], p[-1:] + leftover.unsqueeze(0)], dim=0) | |
| def ponder_loss(token_losses: torch.Tensor, halt_logits: torch.Tensor, | |
| prior: float = 0.1, kl_weight: float = 0.01): | |
| """PonderNet objective: expected loss under the halting distribution, plus a | |
| KL pull towards a geometric prior that sets the expected number of loops. | |
| ``token_losses`` and ``halt_logits`` are both (R, N). | |
| """ | |
| R = token_losses.shape[0] | |
| p = halting_distribution(halt_logits) | |
| token_losses = token_losses.float() | |
| expected = (p * token_losses).sum(0).mean() | |
| steps = torch.arange(1, R + 1, device=p.device, dtype=p.dtype).unsqueeze(1) | |
| prior_p = prior * (1.0 - prior) ** (steps - 1) | |
| prior_p = prior_p / prior_p.sum(0, keepdim=True) | |
| kl = (p * (p.clamp_min(1e-9).log() - prior_p.log())).sum(0).mean() | |
| expected_steps = (p * steps).sum(0).mean() | |
| return expected + kl_weight * kl, { | |
| "expected_loss": float(expected.detach()), | |
| "kl": float(kl.detach()), | |
| "expected_steps": float(expected_steps.detach()), | |
| } | |
| def q_exit(halt_logits: torch.Tensor, tau: float = 0.5) -> torch.Tensor: | |
| """Deterministic exit step per token: the first r whose cumulative halting | |
| probability reaches ``tau``. This is the PALBERT criterion, chosen over | |
| sampling from the halting distribution because sampling adds variance to the | |
| exit index for no benefit at inference. | |
| Returns a (N,) tensor of 0-based loop indices. | |
| """ | |
| p = halting_distribution(halt_logits) | |
| reached = p.cumsum(0) >= tau | |
| R = p.shape[0] | |
| return torch.where(reached.any(0), reached.float().argmax(0), | |
| torch.full((p.shape[1],), R - 1, device=p.device, dtype=torch.long)) | |
| def build_model(cfg: ModelConfig) -> LoopedLM: | |
| return LoopedLM(cfg) | |