| """~1B dense reasoning LLM. |
| |
| Components: RMSNorm (pre-norm), partial RoPE, QK-norm, GQA, SwiGLU, tied embeddings, |
| attention + final logit soft-cap, depth-scaled residual init. |
| |
| Our twists live here: |
| - hybrid attention: local (sliding-window) layers + global layers in the top third |
| - NoPE on global layers |
| - Tier A: z-loss, value-residual (ResFormer), partial RoPE |
| - optional MTP-2 auxiliary head (mid-training only) |
| |
| Attention uses flash-attn: flash_attn_func supports GQA, causal masking, |
| sliding windows, and attention-logit softcapping natively. |
| """ |
| from __future__ import annotations |
| import math |
| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
|
|
| from config import ModelConfig |
|
|
| try: |
| from flash_attn import flash_attn_func |
| except Exception: |
| flash_attn_func = None |
| _HAS_FLASH = flash_attn_func is not None |
|
|
|
|
| |
| class RMSNorm(nn.Module): |
| def __init__(self, dim: int, eps: float = 1e-6): |
| super().__init__() |
| self.eps = eps |
| self.weight = nn.Parameter(torch.ones(dim)) |
|
|
| def forward(self, x: torch.Tensor) -> torch.Tensor: |
| dt = x.dtype |
| x = x.float() |
| x = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) |
| return (x * self.weight.float()).to(dt) |
|
|
|
|
| |
| def build_rope_cache(seq_len: int, rope_dim: int, theta: float, device, dtype): |
| inv_freq = 1.0 / (theta ** (torch.arange(0, rope_dim, 2, device=device).float() / rope_dim)) |
| t = torch.arange(seq_len, device=device).float() |
| freqs = torch.outer(t, inv_freq) |
| return torch.cos(freqs).to(dtype), torch.sin(freqs).to(dtype) |
|
|
|
|
| def apply_rope(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, rope_dim: int) -> torch.Tensor: |
| |
| xr, xp = x[..., :rope_dim], x[..., rope_dim:] |
| x1, x2 = xr[..., 0::2], xr[..., 1::2] |
| cos = cos[None, :, None, :] |
| sin = sin[None, :, None, :] |
| o1 = x1 * cos - x2 * sin |
| o2 = x1 * sin + x2 * cos |
| xr = torch.stack((o1, o2), dim=-1).flatten(-2) |
| return torch.cat([xr, xp], dim=-1).to(x.dtype) |
|
|
|
|
| |
| def compute_think_mask(idx: torch.Tensor, start_id: int, end_id: int) -> torch.Tensor: |
| """Bool (B,S): True for tokens strictly inside a <think> ... </think> span. |
| depth = cumsum(open - close); a position is 'inside' where depth>0. The <think> |
| opener itself gets depth 1 (inside), the </think> closer gets depth 0 (outside).""" |
| opens = (idx == start_id).int() |
| closes = (idx == end_id).int() |
| depth = torch.cumsum(opens - closes, dim=1) |
| return depth > 0 |
|
|
|
|
| |
| class Attention(nn.Module): |
| def __init__(self, cfg: ModelConfig, layer_idx: int, is_global: bool, rope_dim: int): |
| super().__init__() |
| self.cfg = cfg |
| self.layer_idx = layer_idx |
| self.is_global = is_global |
| self.rope_dim = rope_dim |
| self.use_rope = not (is_global and cfg.nope_on_global) |
| self.nh, self.nkv, self.hd = cfg.n_heads, cfg.n_kv_heads, cfg.head_dim |
|
|
| self.wq = nn.Linear(cfg.d_model, self.nh * self.hd, bias=False) |
| self.wk = nn.Linear(cfg.d_model, self.nkv * self.hd, bias=False) |
| self.wv = nn.Linear(cfg.d_model, self.nkv * self.hd, bias=False) |
| self.wo = nn.Linear(self.nh * self.hd, cfg.d_model, bias=False) |
|
|
| if cfg.qk_norm: |
| self.q_norm = RMSNorm(self.hd, cfg.rms_eps) |
| self.k_norm = RMSNorm(self.hd, cfg.rms_eps) |
| else: |
| self.q_norm = self.k_norm = None |
|
|
| |
| self.value_residual = cfg.value_residual and layer_idx > 0 |
| if self.value_residual: |
| self.v_lambda = nn.Parameter(torch.tensor(float(cfg.value_residual_init))) |
|
|
| self.window = (-1, -1) if is_global else (cfg.sliding_window - 1, 0) |
| self.softcap = cfg.attn_logit_softcap |
|
|
| def forward(self, x, cos, sin, v_first): |
| B, S, _ = x.shape |
| q = self.wq(x).view(B, S, self.nh, self.hd) |
| k = self.wk(x).view(B, S, self.nkv, self.hd) |
| v = self.wv(x).view(B, S, self.nkv, self.hd) |
| raw_v = v |
|
|
| if self.q_norm is not None: |
| q = self.q_norm(q) |
| k = self.k_norm(k) |
| if self.use_rope: |
| q = apply_rope(q, cos, sin, self.rope_dim) |
| k = apply_rope(k, cos, sin, self.rope_dim) |
| if self.value_residual and v_first is not None: |
| lam = torch.sigmoid(self.v_lambda) |
| v = lam * v + (1.0 - lam) * v_first |
|
|
| if _HAS_FLASH: |
| out = flash_attn_func( |
| q, k, v, |
| causal=True, |
| window_size=self.window, |
| softcap=self.softcap if self.softcap else 0.0, |
| ) |
| out = out.reshape(B, S, self.nh * self.hd) |
| else: |
| |
| |
| |
| |
| rep = self.nh // self.nkv |
| qs = q.transpose(1, 2) |
| ks = k.transpose(1, 2).repeat_interleave(rep, dim=1) |
| vs = v.transpose(1, 2).repeat_interleave(rep, dim=1) |
| out = F.scaled_dot_product_attention(qs, ks, vs, is_causal=True) |
| out = out.transpose(1, 2).reshape(B, S, self.nh * self.hd) |
| return self.wo(out), raw_v |
|
|
|
|
| |
| class SwiGLU(nn.Module): |
| def __init__(self, cfg: ModelConfig): |
| super().__init__() |
| self.w_gate = nn.Linear(cfg.d_model, cfg.d_ff, bias=False) |
| self.w_up = nn.Linear(cfg.d_model, cfg.d_ff, bias=False) |
| self.w_down = nn.Linear(cfg.d_ff, cfg.d_model, bias=False) |
|
|
| def forward(self, x): |
| return self.w_down(F.silu(self.w_gate(x)) * self.w_up(x)) |
|
|
|
|
| |
| class Block(nn.Module): |
| def __init__(self, cfg: ModelConfig, layer_idx: int, rope_dim: int): |
| super().__init__() |
| is_global = layer_idx in cfg.global_layers() |
| self.attn_norm = RMSNorm(cfg.d_model, cfg.rms_eps) |
| self.attn = Attention(cfg, layer_idx, is_global, rope_dim) |
| self.mlp_norm = RMSNorm(cfg.d_model, cfg.rms_eps) |
| self.mlp = SwiGLU(cfg) |
|
|
| def forward(self, x, cos, sin, v_first): |
| a, raw_v = self.attn(self.attn_norm(x), cos, sin, v_first) |
| x = x + a |
| x = x + self.mlp(self.mlp_norm(x)) |
| return x, raw_v |
|
|
|
|
| |
| class LLM(nn.Module): |
| def __init__(self, cfg: ModelConfig): |
| super().__init__() |
| self.cfg = cfg |
| |
| self.rope_dim = int(cfg.head_dim * cfg.rope_fraction) |
| self.rope_dim -= self.rope_dim % 2 |
|
|
| self.tok_emb = nn.Embedding(cfg.vocab_size, cfg.d_model) |
| self.blocks = nn.ModuleList([Block(cfg, i, self.rope_dim) for i in range(cfg.n_layers)]) |
| self.final_norm = RMSNorm(cfg.d_model, cfg.rms_eps) |
| self.lm_head = nn.Linear(cfg.d_model, cfg.vocab_size, bias=False) |
| if cfg.tie_embeddings: |
| self.lm_head.weight = self.tok_emb.weight |
|
|
| |
| self.mtp_norm = RMSNorm(cfg.d_model, cfg.rms_eps) if cfg.use_mtp else None |
|
|
| |
| self.ponder_halt = nn.Linear(cfg.d_model, 1, bias=True) if cfg.adaptive_compute else None |
|
|
| |
| |
| |
| self.fused_ce = None |
| try: |
| from liger_kernel.transformers.fused_linear_cross_entropy import ( |
| LigerFusedLinearCrossEntropyLoss) |
| self.fused_ce = LigerFusedLinearCrossEntropyLoss( |
| ignore_index=-100, |
| lse_square_scale=cfg.z_loss_weight, |
| softcap=cfg.final_logit_softcap if cfg.final_logit_softcap else None) |
| except Exception as e: |
| print(f"[model] Liger fused CE unavailable ({e}); using full-logits loss", flush=True) |
|
|
| self._rope_cache = None |
| self.apply(self._init_weights) |
| if cfg.depth_scaled_residual: |
| self._scale_residual_branches() |
| print(f"[model] attention backend: {'flash_attn' if _HAS_FLASH else 'torch SDPA (no flash pkg)'}", |
| flush=True) |
|
|
| |
| def _init_weights(self, m): |
| if isinstance(m, nn.Linear): |
| nn.init.normal_(m.weight, mean=0.0, std=self.cfg.init_std) |
| if m.bias is not None: |
| nn.init.zeros_(m.bias) |
| elif isinstance(m, nn.Embedding): |
| nn.init.normal_(m.weight, mean=0.0, std=self.cfg.init_std) |
|
|
| def _scale_residual_branches(self): |
| |
| scale = 1.0 / math.sqrt(2 * self.cfg.n_layers) |
| for blk in self.blocks: |
| with torch.no_grad(): |
| blk.attn.wo.weight.mul_(scale) |
| blk.mlp.w_down.weight.mul_(scale) |
|
|
| def _rope(self, S, device, dtype): |
| if self._rope_cache is None or self._rope_cache[0].shape[0] < S: |
| self._rope_cache = build_rope_cache( |
| max(S, self.cfg.max_seq_len), self.rope_dim, self.cfg.rope_theta, device, dtype) |
| cos, sin = self._rope_cache |
| return cos[:S], sin[:S] |
|
|
| def _softcap(self, logits): |
| c = self.cfg.final_logit_softcap |
| return c * torch.tanh(logits / c) if c else logits |
|
|
| def _ponder(self, x, cos, sin, v_first, think_mask): |
| """Loop the shared reasoning core [core_start, core_end) up to ponder_max_steps |
| times. PonderNet halting decides per-token depth; outside <think> every token |
| halts at step 1 (== plain single pass). Returns (weighted_hidden, ponder_loss).""" |
| cfg = self.cfg |
| core = self.blocks[cfg.core_start:cfg.core_end] |
| N = cfg.ponder_max_steps |
|
|
| hs, lambdas = [], [] |
| cur = x |
| for step in range(N): |
| for blk in core: |
| cur, _ = blk(cur, cos, sin, v_first) |
| hs.append(cur) |
| lam = torch.sigmoid(self.ponder_halt(cur).squeeze(-1)) |
| lambdas.append(lam) |
|
|
| |
| ps, remain = [], torch.ones_like(lambdas[0]) |
| for n in range(N): |
| lam = torch.ones_like(lambdas[n]) if n == N - 1 else lambdas[n] |
| ps.append(remain * lam) |
| remain = remain * (1.0 - lam) |
|
|
| |
| if think_mask is not None: |
| tm = think_mask.to(ps[0].dtype) |
| ps[0] = ps[0] * tm + (1.0 - tm) |
| for n in range(1, N): |
| ps[n] = ps[n] * tm |
|
|
| out = sum(ps[n].unsqueeze(-1) * hs[n] for n in range(N)) |
|
|
| |
| ponder_loss = x.new_zeros(()) |
| if think_mask is not None and think_mask.any(): |
| p = torch.stack(ps, dim=-1).clamp_min(1e-8) |
| n_idx = torch.arange(N, device=x.device) |
| prior = (cfg.ponder_prior * (1 - cfg.ponder_prior) ** n_idx) |
| prior = (prior / prior.sum()).clamp_min(1e-8) |
| kl = (p * (p.log() - prior.log())).sum(-1) |
| ponder_loss = kl[think_mask].mean() |
| return out, ponder_loss |
|
|
| def forward(self, idx: torch.Tensor, targets: torch.Tensor | None = None, |
| think_mask: torch.Tensor | None = None): |
| cfg = self.cfg |
| B, S = idx.shape |
| x = self.tok_emb(idx) |
| cos, sin = self._rope(S, idx.device, x.dtype) |
|
|
| ponder_loss = None |
| if cfg.adaptive_compute: |
| if think_mask is None: |
| think_mask = compute_think_mask(idx, cfg.think_start_id, cfg.think_end_id) |
| v_first = None |
| for i in range(cfg.core_start): |
| x, raw_v = self.blocks[i](x, cos, sin, v_first) |
| if i == 0: |
| v_first = raw_v |
| x, ponder_loss = self._ponder(x, cos, sin, v_first, think_mask) |
| for i in range(cfg.core_end, cfg.n_layers): |
| x, _ = self.blocks[i](x, cos, sin, v_first) |
| else: |
| v_first = None |
| for i, blk in enumerate(self.blocks): |
| x, raw_v = blk(x, cos, sin, v_first) |
| if i == 0: |
| v_first = raw_v |
| h = self.final_norm(x) |
|
|
| if targets is None: |
| return self._softcap(self.lm_head(h[:, -1:])) |
|
|
| |
| |
| need_full = (self.mtp_norm is not None) or cfg.lookahead_consistency or (self.fused_ce is None) |
| if not need_full: |
| D = h.size(-1) |
| loss = self.fused_ce(self.lm_head.weight, h.reshape(-1, D), targets.reshape(-1)) |
| if ponder_loss is not None: |
| loss = loss + cfg.ponder_loss_weight * ponder_loss |
| return None, loss |
|
|
| logits = self._softcap(self.lm_head(h)) |
| loss = F.cross_entropy( |
| logits.view(-1, logits.size(-1)).float(), targets.view(-1), ignore_index=-100) |
|
|
| |
| if cfg.z_loss_weight: |
| valid = (targets.view(-1) != -100) |
| lse = torch.logsumexp(logits.view(-1, logits.size(-1)).float(), dim=-1) |
| if valid.any(): |
| loss = loss + cfg.z_loss_weight * (lse[valid] ** 2).mean() |
|
|
| mtp_logits = None |
| if self.mtp_norm is not None: |
| off = cfg.mtp_offset |
| mtp_logits = self._softcap(self.lm_head(self.mtp_norm(h))) |
| mtp_loss = F.cross_entropy( |
| mtp_logits[:, :-off].reshape(-1, mtp_logits.size(-1)).float(), |
| idx[:, off:].reshape(-1)) |
| loss = loss + cfg.mtp_loss_weight * mtp_loss |
|
|
| |
| if cfg.lookahead_consistency and mtp_logits is not None: |
| sh = cfg.mtp_offset - 1 |
| if sh >= 1 and S > sh: |
| teacher = F.softmax(logits[:, sh:].detach().float(), dim=-1) |
| student = F.log_softmax(mtp_logits[:, :-sh].float(), dim=-1) |
| kl = F.kl_div(student, teacher, reduction="batchmean") |
| loss = loss + cfg.lookahead_weight * kl |
|
|
| if ponder_loss is not None: |
| loss = loss + cfg.ponder_loss_weight * ponder_loss |
|
|
| return logits, loss |
|
|
| def num_params(self, non_embedding=False): |
| n = sum(p.numel() for p in self.parameters()) |
| if non_embedding and self.cfg.tie_embeddings: |
| n -= self.tok_emb.weight.numel() |
| return n |
|
|
|
|
| if __name__ == "__main__": |
| print("=== budgie-1b (full) ===") |
| cfg = ModelConfig() |
| print("global layers:", sorted(cfg.global_layers())) |
| print(f"reasoning core span: [{cfg.core_start}, {cfg.core_end})") |
| m = LLM(cfg) |
| print(f"rope_dim (partial): {m.rope_dim}/{cfg.head_dim}") |
| print(f"total params: {m.num_params()/1e9:.3f}B") |
| print(f"non-embedding params: {m.num_params(non_embedding=True)/1e9:.3f}B") |
|
|
| print("\n=== budgie-mini (ablation, signature ON) ===") |
| mc = ModelConfig.mini(use_mtp=True, adaptive_compute=True, lookahead_consistency=True) |
| print("global layers:", sorted(mc.global_layers()), "core:", (mc.core_start, mc.core_end)) |
| mm = LLM(mc) |
| print(f"total params: {mm.num_params()/1e6:.1f}M") |
|
|
| |
| if torch.cuda.is_available() and _HAS_FLASH: |
| dev = "cuda" |
| mm = mm.to(dev).to(torch.bfloat16) |
| idx = torch.randint(0, mc.vocab_size, (2, 128), device=dev) |
| |
| idx[:, 10] = mc.think_start_id |
| idx[:, 40] = mc.think_end_id |
| tgt = torch.randint(0, mc.vocab_size, (2, 128), device=dev) |
| logits, loss = mm(idx, tgt) |
| print(f"mini forward OK: logits {tuple(logits.shape)}, loss {loss.item():.3f}") |