Download modeling_bqalm.py from tturing/DM-KDAGQA-1.7B-SD: direct link, hf CLI and curl.
- Browser
- Download file 31.4 kB
-
https://huggingface.co/tturing/DM-KDAGQA-1.7B-SD/resolve/main/modeling_bqalm.py
- Command line
-
hf download hf://tturing/DM-KDAGQA-1.7B-SD/modeling_bqalm.py
-
curl -L -o modeling_bqalm.py https://huggingface.co/tturing/DM-KDAGQA-1.7B-SD/resolve/main/modeling_bqalm.py
31.4 kB
| """BqaLM for `transformers` (remote code): the hybrid decoder of the DeltaMatching study. | |
| Self-contained copy of the bqa codebase's inference path (`src/pretrain/modeling` + `src/pretrain/hf`) for the mixers the | |
| published checkpoints use: GQA and MLA attention, Mamba2, GatedDeltaNet, KDA. Training-only paths (FP8 FlashMatch | |
| attention, fused FP8 producers, TransformerEngine, the Liger fused kernels) are left out; what remains is the eager / | |
| SDPA code the study's evaluation ran, so the logits match the bqa loader bit for bit. Parameter names are the | |
| checkpoint's: `model.embed_tokens`, `model.layers.N.{input_layernorm,mixer,post_attention_layernorm,mlp}`, | |
| `model.norm`, `lm_head` (tied to the embedding). | |
| Requirements: torch and transformers; Mamba2 layers also need `mamba_ssm` + `causal_conv1d`, GatedDeltaNet and KDA | |
| layers `flash-linear-attention` (`fla`). Each is imported only by the layers that use it, with an install hint if missing. | |
| `generate()` keeps a cache (`BqaLMCache`: K/V of the attention layers, mamba_ssm's conv / SSM states, fla's | |
| recurrent and conv states), so each new token costs one step. It supports greedy decoding and sampling (num_beams=1); | |
| a batch of prompts must share a length, since the recurrent layers have no padding mask. An all-ones attention mask | |
| is the same as none. A plain forward (no `use_cache`) is the uncached evaluation path. | |
| """ | |
| from __future__ import annotations | |
| import math | |
| from dataclasses import dataclass | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from transformers import GenerationMixin, PreTrainedModel | |
| from transformers.utils import ModelOutput | |
| from .configuration_bqalm import BqaLMConfig, ModelConfig | |
| # ----------------------------------------------------------------------------------------------- primitive layers | |
| class RMSNorm(nn.Module): | |
| """x / sqrt(mean(x^2) + eps) * w, reduced in fp32 when `in_fp32` (then cast back before the scale).""" | |
| def __init__(self, dim: int, eps: float = 1e-5, in_fp32: bool = True): | |
| super().__init__() | |
| self.eps = eps | |
| self.in_fp32 = in_fp32 | |
| self.weight = nn.Parameter(torch.ones(dim)) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| dtype = x.dtype | |
| if self.in_fp32: | |
| x = x.float() | |
| var = x.pow(2).mean(dim=-1, keepdim=True) | |
| x = x * torch.rsqrt(var + self.eps) | |
| return (self.weight * x.to(dtype)) if self.in_fp32 else (self.weight * x) | |
| def _is_yarn(rope_scaling: dict | None) -> bool: | |
| return bool(rope_scaling) and str(rope_scaling.get("type", rope_scaling.get("rope_type", ""))).lower() == "yarn" | |
| def yarn_mscale(rope_scaling: dict | None) -> float: | |
| """YaRN attention temperature applied to post-RoPE q (1.0 == no scaling).""" | |
| if not _is_yarn(rope_scaling): | |
| return 1.0 | |
| if rope_scaling.get("mscale") is not None: | |
| return float(rope_scaling["mscale"]) | |
| factor = float(rope_scaling["factor"]) | |
| if factor <= 1.0: | |
| return 1.0 | |
| return 0.1 * math.log(factor) + 1.0 | |
| def _yarn_find_dim(num_rotations: float, head_dim: int, theta: float, max_pos: int) -> float: | |
| return (head_dim * math.log(max_pos / (num_rotations * 2 * math.pi))) / (2 * math.log(theta)) | |
| def yarn_inv_freq(head_dim: int, theta: float, rope_scaling: dict, device, dtype=torch.float32) -> torch.Tensor: | |
| """NTK-by-parts interpolated inv_freq, shape [head_dim / 2] (fp32).""" | |
| factor = float(rope_scaling["factor"]) | |
| orig_max = int(rope_scaling.get("original_max_position_embeddings", 8192)) | |
| beta_fast = float(rope_scaling.get("beta_fast", 32)) | |
| beta_slow = float(rope_scaling.get("beta_slow", 1)) | |
| pos_freqs = theta ** (torch.arange(0, head_dim, 2, device=device, dtype=torch.float32) / head_dim) | |
| inv_freq_extrap = 1.0 / pos_freqs | |
| inv_freq_interp = 1.0 / (factor * pos_freqs) | |
| low = math.floor(_yarn_find_dim(beta_fast, head_dim, theta, orig_max)) | |
| high = math.ceil(_yarn_find_dim(beta_slow, head_dim, theta, orig_max)) | |
| low = max(low, 0) | |
| high = min(high, head_dim // 2 - 1) | |
| if low == high: | |
| high += 0.001 | |
| ramp = (torch.arange(head_dim // 2, device=device, dtype=torch.float32) - low) / (high - low) | |
| ramp = torch.clamp(ramp, 0.0, 1.0) | |
| extrap_factor = 1.0 - ramp | |
| inv = inv_freq_interp * (1.0 - extrap_factor) + inv_freq_extrap * extrap_factor | |
| return inv.to(dtype) | |
| class RotaryEmbedding(nn.Module): | |
| """cos / sin computed per forward from inv_freq (NeoX layout, fp32), over the rotated width `rotary_dim`. | |
| No inv_freq buffer on purpose: `from_pretrained` would materialize a non-persistent buffer uninitialized.""" | |
| def __init__(self, head_dim: int, max_seq_len: int, theta: float = 10000.0, | |
| rope_scaling: dict | None = None, rotary_dim: int | None = None): | |
| super().__init__() | |
| head_dim = int(rotary_dim) if rotary_dim else head_dim | |
| assert head_dim % 2 == 0, "RoPE needs an even head_dim" | |
| self.head_dim = head_dim | |
| self.theta = theta | |
| self.max_seq_len = max_seq_len | |
| self.rope_scaling = rope_scaling | |
| self._use_yarn = _is_yarn(rope_scaling) | |
| def forward(self, position_ids: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: | |
| if self._use_yarn: | |
| inv = yarn_inv_freq(self.head_dim, self.theta, self.rope_scaling, position_ids.device) | |
| else: | |
| inv = 1.0 / (self.theta ** (torch.arange(0, self.head_dim, 2, device=position_ids.device, | |
| dtype=torch.float32) / self.head_dim)) | |
| freqs = torch.einsum("bt,d->btd", position_ids.float(), inv) | |
| emb = torch.cat([freqs, freqs], dim=-1) | |
| return emb.cos(), emb.sin() | |
| def rotate_half(x: torch.Tensor) -> torch.Tensor: | |
| x1, x2 = x.chunk(2, dim=-1) | |
| return torch.cat([-x2, x1], dim=-1) | |
| def apply_rotary(q: torch.Tensor, k: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor): | |
| """q, k: [B, H, T, D]; cos, sin: [B, T, R] with R <= D (partial RoPE rotates the leading R channels).""" | |
| cos = cos.unsqueeze(1) | |
| sin = sin.unsqueeze(1) | |
| rd = cos.shape[-1] | |
| if rd == q.shape[-1]: | |
| qf, kf = q.float(), k.float() | |
| q_out = qf * cos + rotate_half(qf) * sin | |
| k_out = kf * cos + rotate_half(kf) * sin | |
| return q_out.to(q.dtype), k_out.to(k.dtype) | |
| def _split_rot(x): | |
| xr, xp = x[..., :rd].float(), x[..., rd:] | |
| out = xr * cos + rotate_half(xr) * sin | |
| return torch.cat([out.to(x.dtype), xp], dim=-1) | |
| return _split_rot(q), _split_rot(k) | |
| def q_proj_out_features(cfg: ModelConfig) -> int: | |
| """q_proj width, doubled when the attention output gate is on (Qwen3.5 / Qwen3-Next idiom).""" | |
| return cfg.q_dim * (2 if cfg.attn_output_gate else 1) | |
| def split_q_gate(qg: torch.Tensor, n_heads: int, head_dim: int, gated: bool): | |
| """[B, T, q_dim * (2 if gated)] -> (q, gate), each [B, T, n_heads, head_dim]; the split is per head.""" | |
| if not gated: | |
| return qg.view(*qg.shape[:2], n_heads, head_dim), None | |
| q, g = qg.view(*qg.shape[:2], n_heads, 2 * head_dim).chunk(2, dim=-1) | |
| return q, g | |
| def apply_output_gate(o: torch.Tensor, gate: torch.Tensor | None) -> torch.Tensor: | |
| if gate is None: | |
| return o | |
| return o * torch.sigmoid(gate.reshape(o.shape).to(o.dtype)) | |
| class SwiGLUMLP(nn.Module): | |
| """down(silu(gate(x)) * up(x)).""" | |
| def __init__(self, d_model: int, intermediate_size: int, dropout: float = 0.0): | |
| super().__init__() | |
| self.gate_proj = nn.Linear(d_model, intermediate_size, bias=False) | |
| self.up_proj = nn.Linear(d_model, intermediate_size, bias=False) | |
| self.down_proj = nn.Linear(intermediate_size, d_model, bias=False) | |
| self.drop = nn.Dropout(dropout) if dropout > 0 else nn.Identity() | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| return self.drop(self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x))) | |
| # ------------------------------------------------------------------------------------------------ generation cache | |
| class BqaLMCache: | |
| """Decoding state. `kv`: K/V of the attention layers (post norm and RoPE). The recurrent layers' states live in | |
| their libraries' own containers, created on first use: mamba_ssm's InferenceParams (Mamba2 conv / SSM states, | |
| updated in place) and fla's FLACache (GatedDeltaNet / KDA recurrent and conv states, indexed by layer).""" | |
| is_compileable = False | |
| def __init__(self, batch_size: int, max_seqlen: int): | |
| self.kv = {} | |
| self.batch_size = batch_size | |
| self.max_seqlen = max_seqlen | |
| self.seen = 0 | |
| self._mamba = None | |
| self._fla = None | |
| def mamba_params(self): | |
| if self._mamba is None: | |
| try: | |
| from mamba_ssm.utils.generation import InferenceParams | |
| except ImportError as e: | |
| raise ImportError("this checkpoint's Mamba2 layers need `pip install mamba-ssm causal-conv1d`") from e | |
| self._mamba = InferenceParams(max_seqlen=self.max_seqlen, max_batch_size=self.batch_size) | |
| self._mamba.seqlen_offset = self.seen | |
| return self._mamba | |
| def fla(self): | |
| if self._fla is None: | |
| try: | |
| from fla.models.utils import FLACache | |
| except ImportError as e: | |
| raise ImportError("this checkpoint's GatedDeltaNet / KDA layers need `pip install flash-linear-attention`") from e | |
| self._fla = FLACache() | |
| return self._fla | |
| def get_seq_length(self, layer_idx: int = 0) -> int: | |
| return self.seen | |
| def advance(self, n: int) -> None: | |
| self.seen += n | |
| if self._mamba is not None: | |
| self._mamba.seqlen_offset += n | |
| # ---------------------------------------------------------------------------------------------------- attention | |
| def _pick_flash(): | |
| """flash-attn's flash_attn_func if installed (accepted only if it takes `dropout_p`), else None.""" | |
| import inspect | |
| for mod in ("flash_attn_interface", "flash_attn"): | |
| try: | |
| fn = getattr(__import__(mod, fromlist=["flash_attn_func"]), "flash_attn_func") | |
| if "dropout_p" in inspect.signature(fn).parameters: | |
| return fn | |
| except Exception: | |
| continue | |
| return None | |
| _FLASH_FN = _pick_flash() | |
| class GQAAttention(nn.Module): | |
| """Grouped-query attention: qk-norm before (partial) RoPE, optional sigmoid output gate, SDPA or flash-attn.""" | |
| def __init__(self, cfg: ModelConfig, layer_idx: int): | |
| super().__init__() | |
| self.layer_idx = layer_idx | |
| self.n_heads = cfg.n_heads | |
| self.n_kv_heads = cfg.n_kv_heads | |
| self.n_rep = cfg.n_rep | |
| self.head_dim = cfg.head_dim | |
| self.attn_dropout = cfg.attn_dropout | |
| impl = "flash" if (cfg.attn_impl == "auto" and _FLASH_FN is not None) else cfg.attn_impl | |
| if impl not in ("flash", "sdpa", "auto"): | |
| raise ValueError(f"this checkpoint's remote code supports attn_impl sdpa | flash, got {impl!r}") | |
| if impl == "flash" and _FLASH_FN is None: | |
| raise RuntimeError("attn_impl='flash' but flash-attn is not importable; use attn_implementation='sdpa'") | |
| self.impl = "sdpa" if impl == "auto" else impl | |
| self.attn_mscale = yarn_mscale(cfg.rope_scaling) | |
| self.attn_output_gate = cfg.attn_output_gate | |
| self.q_proj = nn.Linear(cfg.d_model, q_proj_out_features(cfg), bias=False) | |
| self.k_proj = nn.Linear(cfg.d_model, cfg.kv_dim, bias=False) | |
| self.v_proj = nn.Linear(cfg.d_model, cfg.kv_dim, bias=False) | |
| self.o_proj = nn.Linear(cfg.q_dim, cfg.d_model, bias=False) | |
| self.qk_norm = cfg.qk_norm | |
| if self.qk_norm: | |
| self.q_norm = RMSNorm(self.head_dim, cfg.norm_eps, cfg.rms_norm_in_fp32) | |
| self.k_norm = RMSNorm(self.head_dim, cfg.norm_eps, cfg.rms_norm_in_fp32) | |
| def forward(self, x, cos, sin, attention_mask=None, cache=None): | |
| B, T, _ = x.shape | |
| q, gate = split_q_gate(self.q_proj(x), self.n_heads, self.head_dim, self.attn_output_gate) | |
| q = q.transpose(1, 2) # [B, H, T, D] | |
| k = self.k_proj(x).view(B, T, self.n_kv_heads, self.head_dim).transpose(1, 2) # [B, Hkv, T, D] | |
| v = self.v_proj(x).view(B, T, self.n_kv_heads, self.head_dim).transpose(1, 2) | |
| if self.qk_norm: | |
| q, k = self.q_norm(q), self.k_norm(k) | |
| q, k = apply_rotary(q, k, cos, sin) | |
| if self.attn_mscale != 1.0: | |
| q = q * (self.attn_mscale * self.attn_mscale) | |
| drop = self.attn_dropout if self.training else 0.0 | |
| past = None | |
| if cache is not None: | |
| past = cache.kv.get(self.layer_idx) | |
| if past is not None: | |
| k = torch.cat([past[0], k], dim=2) | |
| v = torch.cat([past[1], v], dim=2) | |
| cache.kv[self.layer_idx] = (k, v) | |
| if past is not None: # decoding: the T new queries sit at the end of the S cached positions | |
| S = k.shape[2] | |
| gqa = self.n_rep > 1 | |
| if attention_mask is None and T == 1: | |
| out = F.scaled_dot_product_attention(q, k, v, enable_gqa=gqa) | |
| else: | |
| pos_q = torch.arange(S - T, S, device=q.device) | |
| keep = (torch.arange(S, device=q.device)[None, :] <= pos_q[:, None])[None, None] | |
| if attention_mask is not None: | |
| keep = keep & attention_mask.bool()[:, None, None, -S:] | |
| bias = torch.zeros(keep.shape, dtype=q.dtype, device=q.device) | |
| bias.masked_fill_(~keep, torch.finfo(q.dtype).min) | |
| out = F.scaled_dot_product_attention(q, k, v, attn_mask=bias, enable_gqa=gqa) | |
| out = out.transpose(1, 2).reshape(B, T, self.n_heads * self.head_dim) | |
| elif self.impl == "flash" and attention_mask is None: | |
| out = _FLASH_FN(q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2), dropout_p=drop, causal=True) | |
| if isinstance(out, tuple): | |
| out = out[0] | |
| out = out.reshape(B, T, self.n_heads * self.head_dim) | |
| else: | |
| gqa = self.n_rep > 1 | |
| if attention_mask is None: | |
| out = F.scaled_dot_product_attention(q, k, v, is_causal=True, dropout_p=drop, enable_gqa=gqa) | |
| else: | |
| bias = self._build_additive_mask(attention_mask, T, q.dtype, q.device) | |
| out = F.scaled_dot_product_attention(q, k, v, attn_mask=bias, dropout_p=drop, enable_gqa=gqa) | |
| out = out.transpose(1, 2).reshape(B, T, self.n_heads * self.head_dim) | |
| return self.o_proj(apply_output_gate(out, gate)) | |
| def _build_additive_mask(attention_mask, T, dtype, device): | |
| causal = torch.tril(torch.ones(T, T, dtype=torch.bool, device=device)) | |
| keep = attention_mask.bool()[:, None, None, :] & causal[None, None] | |
| bias = torch.zeros(keep.shape, dtype=dtype, device=device) | |
| bias.masked_fill_(~keep, torch.finfo(dtype).min) | |
| return bias | |
| # ------------------------------------------------------------------------------------------------------- Mamba2 | |
| class Mamba2Mixer(nn.Module): | |
| """mamba_ssm.Mamba2 on its fused kernel path (conv1d + SSD scan); RoPE inputs are accepted and ignored.""" | |
| def __init__(self, cfg: ModelConfig, layer_idx: int): | |
| super().__init__() | |
| try: | |
| from mamba_ssm import Mamba2 | |
| import causal_conv1d # noqa: F401 (the fused kernel path needs it) | |
| except ImportError as e: | |
| raise ImportError("this checkpoint's Mamba2 layers need `pip install mamba-ssm causal-conv1d`") from e | |
| headdim = int(cfg.mamba2_headdim) if getattr(cfg, "mamba2_headdim", None) else cfg.head_dim | |
| self.mamba = Mamba2( | |
| d_model=cfg.d_model, | |
| headdim=headdim, | |
| d_state=int(getattr(cfg, "mamba2_d_state", 128)), | |
| expand=int(getattr(cfg, "mamba2_expand", 2)), | |
| ngroups=int(getattr(cfg, "mamba2_ngroups", 1)), | |
| chunk_size=int(getattr(cfg, "mamba2_chunk_size", 256)), | |
| layer_idx=layer_idx, | |
| ) | |
| def forward(self, x, cos=None, sin=None, attention_mask=None, cache=None): | |
| return self.mamba(x, inference_params=cache.mamba_params if cache is not None else None) | |
| # ------------------------------------------------------------------------------------------------ GatedDeltaNet | |
| class GatedDeltaNetMixer(nn.Module): | |
| """fla's GatedDeltaNet (chunk mode, gate, short convolution); RoPE inputs are accepted and ignored.""" | |
| def __init__(self, cfg: ModelConfig, layer_idx: int): | |
| super().__init__() | |
| try: | |
| from fla.layers import GatedDeltaNet | |
| except ImportError as e: | |
| raise ImportError("this checkpoint's GatedDeltaNet layers need `pip install flash-linear-attention`") from e | |
| head_dim = cfg.gdn_head_dim or cfg.head_dim | |
| num_heads = cfg.gdn_num_heads or max(1, cfg.d_model // head_dim) | |
| extra = {} | |
| if getattr(cfg, "gdn_num_v_heads", None): # omitted when unset: passing None is not equivalent in older fla | |
| extra["num_v_heads"] = cfg.gdn_num_v_heads | |
| self.gdn = GatedDeltaNet(hidden_size=cfg.d_model, head_dim=head_dim, num_heads=num_heads, **extra, | |
| mode="chunk", use_gate=True, use_short_conv=True, layer_idx=layer_idx) | |
| def forward(self, x, cos=None, sin=None, attention_mask=None, cache=None): | |
| out = self.gdn(x) if cache is None else self.gdn(x, past_key_values=cache.fla, use_cache=True) | |
| return out[0] if isinstance(out, tuple) else out | |
| # ---------------------------------------------------------------------------------------------------------- KDA | |
| class KDAMixer(nn.Module): | |
| """fla's Kimi Delta Attention (chunk mode, short convolution, per-channel decay); RoPE inputs are accepted and | |
| ignored. head_dim / num_heads fall back to the gdn_* fields, as in training.""" | |
| def __init__(self, cfg: ModelConfig, layer_idx: int): | |
| super().__init__() | |
| try: | |
| from fla.layers.kda import KimiDeltaAttention | |
| except ImportError as e: | |
| raise ImportError("this checkpoint's KDA layers need `pip install flash-linear-attention`") from e | |
| head_dim = cfg.kda_head_dim or cfg.gdn_head_dim or cfg.head_dim | |
| num_heads = cfg.kda_num_heads or cfg.gdn_num_heads or max(1, cfg.d_model // head_dim) | |
| extra = {} | |
| if getattr(cfg, "kda_num_v_heads", None): | |
| extra["num_v_heads"] = int(cfg.kda_num_v_heads) | |
| if getattr(cfg, "kda_lower_bound", None) is not None: | |
| extra["lower_bound"] = float(cfg.kda_lower_bound) | |
| self.kda = KimiDeltaAttention(hidden_size=cfg.d_model, head_dim=head_dim, num_heads=num_heads, | |
| expand_v=float(getattr(cfg, "kda_expand_v", 1.0)), mode="chunk", | |
| use_short_conv=True, conv_size=int(getattr(cfg, "kda_conv_size", 4)), | |
| allow_neg_eigval=bool(getattr(cfg, "kda_allow_neg_eigval", False)), | |
| safe_gate=bool(getattr(cfg, "kda_safe_gate", False)), layer_idx=layer_idx, | |
| **extra) | |
| def forward(self, x, cos=None, sin=None, attention_mask=None, cache=None): | |
| out = self.kda(x) if cache is None else self.kda(x, past_key_values=cache.fla, use_cache=True) | |
| return out[0] if isinstance(out, tuple) else out | |
| # ---------------------------------------------------------------------------------------------------------- MLA | |
| def _rope_leading(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, rd: int) -> torch.Tensor: | |
| """Rotate the leading `rd` channels of x [B, h, T, D] (h = 1 or H) with cos / sin [B, T, rd]; the tail passes.""" | |
| c, s = cos.unsqueeze(1), sin.unsqueeze(1) | |
| xr = x[..., :rd].float() | |
| out = (xr * c + rotate_half(xr) * s).to(x.dtype) | |
| return out if rd == x.shape[-1] else torch.cat([out, x[..., rd:]], dim=-1) | |
| class MLAAttention(nn.Module): | |
| """Multi-head latent attention at one head dim for q, k and v (the bf16 core the study's evaluation ran). | |
| q = q_proj(x) per head [rope | nope]; [c ; k_r] = kv_a_proj(x); c = kv_a_norm(c); k_nope = k_proj(c); | |
| v = v_proj(c); four sub-vector norms (qk_norm); k_r rotated once and shared by every head; MHA on the wire.""" | |
| def __init__(self, cfg: ModelConfig, layer_idx: int): | |
| super().__init__() | |
| self.layer_idx = layer_idx | |
| self.n_heads = cfg.n_heads | |
| self.head_dim = cfg.head_dim | |
| self.rotary_dim = cfg.rotary_dim | |
| self.nope_dim = cfg.head_dim - cfg.rotary_dim | |
| self.kv_lora_rank = int(cfg.kv_lora_rank) | |
| self.attn_output_gate = cfg.attn_output_gate | |
| self.qk_norm = cfg.qk_norm | |
| assert cfg.n_kv_heads == cfg.n_heads, "MLA is MHA on the wire: n_kv_heads must equal n_heads" | |
| assert 0 < self.rotary_dim < self.head_dim, "MLA needs 0 < rotary_dim < head_dim" | |
| self.attn_mscale = yarn_mscale(cfg.rope_scaling) | |
| d, H, D, rd, dn, dc = cfg.d_model, self.n_heads, self.head_dim, self.rotary_dim, self.nope_dim, self.kv_lora_rank | |
| self.q_proj = nn.Linear(d, q_proj_out_features(cfg), bias=False) | |
| self.kv_a_proj = nn.Linear(d, dc + rd, bias=False) | |
| self.kv_a_norm = RMSNorm(dc, cfg.norm_eps, cfg.rms_norm_in_fp32) | |
| self.k_proj = nn.Linear(dc, H * dn, bias=False) | |
| self.v_proj = nn.Linear(dc, H * D, bias=False) | |
| self.o_proj = nn.Linear(H * D, d, bias=False) | |
| if self.qk_norm: | |
| self.q_rope_norm = RMSNorm(rd, cfg.norm_eps, cfg.rms_norm_in_fp32) | |
| self.q_nope_norm = RMSNorm(dn, cfg.norm_eps, cfg.rms_norm_in_fp32) | |
| self.k_rope_norm = RMSNorm(rd, cfg.norm_eps, cfg.rms_norm_in_fp32) | |
| self.k_nope_norm = RMSNorm(dn, cfg.norm_eps, cfg.rms_norm_in_fp32) | |
| impl = "flash" if (cfg.attn_impl == "auto" and _FLASH_FN is not None) else cfg.attn_impl | |
| if impl == "flash" and _FLASH_FN is None: | |
| raise RuntimeError("attn_impl='flash' but flash-attn is not importable; use attn_implementation='sdpa'") | |
| self.impl = impl if impl == "flash" else "sdpa" | |
| def _qkv(self, x, cos, sin): | |
| B, T, _ = x.shape | |
| H, D, rd, dn, dc = self.n_heads, self.head_dim, self.rotary_dim, self.nope_dim, self.kv_lora_rank | |
| q, gate = split_q_gate(self.q_proj(x), H, D, self.attn_output_gate) | |
| q = q.transpose(1, 2) # [B, H, T, D] | |
| ckr = self.kv_a_proj(x) | |
| c, k_r = ckr[..., :dc], ckr[..., dc:] # [B, T, dc], [B, T, rd] | |
| c = self.kv_a_norm(c) | |
| k_n = self.k_proj(c).view(B, T, H, dn).transpose(1, 2) # [B, H, T, dn] | |
| v = self.v_proj(c).view(B, T, H, D).transpose(1, 2) # [B, H, T, D] | |
| k_r = k_r.unsqueeze(1) # [B, 1, T, rd] | |
| if self.qk_norm: | |
| q = torch.cat([self.q_rope_norm(q[..., :rd]), self.q_nope_norm(q[..., rd:])], dim=-1) | |
| k_r = self.k_rope_norm(k_r) | |
| k_n = self.k_nope_norm(k_n) | |
| q = _rope_leading(q, cos, sin, rd) | |
| if self.attn_mscale != 1.0: | |
| q = q * (self.attn_mscale * self.attn_mscale) | |
| k_r = _rope_leading(k_r, cos, sin, rd) | |
| k = torch.cat([k_r.expand(B, H, T, rd), k_n], dim=-1) | |
| return q, k, v, gate | |
| def forward(self, x, cos, sin, attention_mask=None, cache=None): | |
| B, T, _ = x.shape | |
| q, k, v, gate = self._qkv(x, cos, sin) | |
| past = None | |
| if cache is not None: | |
| past = cache.kv.get(self.layer_idx) | |
| if past is not None: | |
| k = torch.cat([past[0], k], dim=2) | |
| v = torch.cat([past[1], v], dim=2) | |
| cache.kv[self.layer_idx] = (k, v) | |
| S = k.shape[2] | |
| if attention_mask is None and past is None and self.impl == "flash": | |
| out = _FLASH_FN(q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2), causal=True) | |
| if isinstance(out, tuple): | |
| out = out[0] | |
| out = out.reshape(B, T, self.n_heads * self.head_dim) | |
| else: | |
| if attention_mask is None and past is None: | |
| out = F.scaled_dot_product_attention(q, k, v, is_causal=True) | |
| elif attention_mask is None and T == 1: | |
| out = F.scaled_dot_product_attention(q, k, v) | |
| else: | |
| pos_q = torch.arange(S - T, S, device=q.device) | |
| keep = (torch.arange(S, device=q.device)[None, :] <= pos_q[:, None])[None, None] | |
| if attention_mask is not None: | |
| keep = keep & attention_mask.bool()[:, None, None, -S:] | |
| bias = torch.zeros(keep.shape, dtype=q.dtype, device=q.device) | |
| bias.masked_fill_(~keep, torch.finfo(q.dtype).min) | |
| out = F.scaled_dot_product_attention(q, k, v, attn_mask=bias) | |
| out = out.transpose(1, 2).reshape(B, T, self.n_heads * self.head_dim) | |
| return self.o_proj(apply_output_gate(out, gate)) | |
| MIXER_REGISTRY = {"gqa": GQAAttention, "mla": MLAAttention, "mamba2": Mamba2Mixer, "gated_deltanet": GatedDeltaNetMixer, | |
| "kda": KDAMixer} | |
| def build_mixer(name: str, cfg: ModelConfig, layer_idx: int) -> nn.Module: | |
| if name not in MIXER_REGISTRY: | |
| raise ValueError(f"mixer {name!r} is not shipped with this checkpoint's code (have {sorted(MIXER_REGISTRY)})") | |
| return MIXER_REGISTRY[name](cfg, layer_idx) | |
| # ------------------------------------------------------------------------------------------------------ backbone | |
| class DecoderBlock(nn.Module): | |
| """Pre-norm block: x += mixer(norm(x)); x += mlp(norm(x)).""" | |
| def __init__(self, cfg: ModelConfig, layer_idx: int): | |
| super().__init__() | |
| self.layer_idx = layer_idx | |
| self.input_layernorm = RMSNorm(cfg.d_model, cfg.norm_eps, cfg.rms_norm_in_fp32) | |
| self.mixer = build_mixer(cfg.mixer_for_layer(layer_idx), cfg, layer_idx) | |
| self.post_attention_layernorm = RMSNorm(cfg.d_model, cfg.norm_eps, cfg.rms_norm_in_fp32) | |
| self.mlp = SwiGLUMLP(cfg.d_model, cfg.intermediate_size, cfg.resid_dropout) | |
| self.resid_drop = nn.Dropout(cfg.resid_dropout) if cfg.resid_dropout > 0 else nn.Identity() | |
| def forward(self, x, cos, sin, attention_mask=None, cache=None): | |
| x = x + self.resid_drop(self.mixer(self.input_layernorm(x), cos, sin, attention_mask, cache)) | |
| x = x + self.resid_drop(self.mlp(self.post_attention_layernorm(x))) | |
| return x | |
| class Transformer(nn.Module): | |
| def __init__(self, cfg: ModelConfig): | |
| super().__init__() | |
| self.cfg = cfg | |
| self.embed_tokens = nn.Embedding(cfg.vocab_size, cfg.d_model) | |
| self.layers = nn.ModuleList([DecoderBlock(cfg, i) for i in range(cfg.n_layers)]) | |
| self.norm = RMSNorm(cfg.d_model, cfg.norm_eps, cfg.rms_norm_in_fp32) | |
| self.rotary = RotaryEmbedding(cfg.head_dim, cfg.max_seq_len, cfg.rope_theta, | |
| rope_scaling=cfg.rope_scaling, rotary_dim=cfg.rotary_dim) | |
| def forward(self, input_ids, position_ids, attention_mask=None, cache=None): | |
| if attention_mask is not None and bool(attention_mask.all()): | |
| attention_mask = None # no padding: take the maskless kernels, the same path as a plain forward | |
| h = self.embed_tokens(input_ids) | |
| cos, sin = self.rotary(position_ids) | |
| if getattr(self.cfg, "nope", False): | |
| cos, sin = torch.ones_like(cos), torch.zeros_like(sin) | |
| cos, sin = cos.to(h.dtype), sin.to(h.dtype) | |
| for layer in self.layers: | |
| h = layer(h, cos, sin, attention_mask, cache) | |
| if cache is not None: | |
| cache.advance(input_ids.shape[1]) | |
| return self.norm(h) | |
| # ------------------------------------------------------------------------------------------------ HF causal LM | |
| _ATTN_MAP = {"flash_attention_2": "flash", "sdpa": "sdpa", "eager": "sdpa"} | |
| class BqaLMCausalLMOutput(ModelOutput): | |
| loss: torch.FloatTensor | None = None | |
| logits: torch.FloatTensor | None = None | |
| cache_params: BqaLMCache | None = None | |
| class BqaLMForCausalLM(PreTrainedModel, GenerationMixin): | |
| config_class = BqaLMConfig | |
| base_model_prefix = "model" | |
| supports_gradient_checkpointing = True | |
| _supports_flash_attn = True | |
| _supports_flash_attn_2 = True | |
| _supports_sdpa = True | |
| _supports_attention_backend = True | |
| _tied_weights_keys = {"lm_head.weight": "model.embed_tokens.weight"} | |
| _is_stateful = True # the Mamba2 state cannot be rewound, so no assisted generation | |
| def _supports_default_dynamic_cache(cls) -> bool: | |
| return False # the model keeps its own state in `cache_params` (BqaLMCache) | |
| def __init__(self, config: BqaLMConfig): | |
| super().__init__(config) | |
| mc = config.to_model_config() | |
| hf_impl = getattr(config, "_attn_implementation", None) | |
| if hf_impl in _ATTN_MAP: | |
| mc.attn_impl = _ATTN_MAP[hf_impl] | |
| self._mc = mc | |
| self.model = Transformer(mc) | |
| self.lm_head = nn.Linear(mc.d_model, mc.vocab_size, bias=False) | |
| if mc.tie_embeddings: | |
| self.lm_head.weight = self.model.embed_tokens.weight | |
| 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 set_output_embeddings(self, value): | |
| self.lm_head = value | |
| def forward(self, input_ids=None, attention_mask=None, position_ids=None, labels=None, | |
| cache_params=None, use_cache=None, past_key_values=None, output_attentions=None, | |
| output_hidden_states=None, return_dict=None, **kwargs): | |
| B, T = input_ids.shape | |
| if use_cache and cache_params is None: | |
| cache_params = BqaLMCache(B, self._mc.max_seq_len) | |
| past = cache_params.seen if cache_params is not None else 0 | |
| if position_ids is None: | |
| position_ids = torch.arange(past, past + T, device=input_ids.device).unsqueeze(0).expand(B, -1) | |
| hidden = self.model(input_ids, position_ids, attention_mask, cache=cache_params) | |
| logits = self.lm_head(hidden) | |
| loss = None | |
| if labels is not None: | |
| sl = logits[:, :-1, :].contiguous().float() | |
| lb = labels[:, 1:].contiguous() | |
| loss = F.cross_entropy(sl.view(-1, sl.size(-1)), lb.view(-1), ignore_index=-100) | |
| return BqaLMCausalLMOutput(loss=loss, logits=logits, cache_params=cache_params) | |
| def prepare_inputs_for_generation(self, input_ids, attention_mask=None, cache_params=None, use_cache=None, | |
| **kwargs): | |
| # with a cache only the newest token is fed; use_cache=False recomputes the whole prefix every step | |
| if cache_params is not None and cache_params.seen > 0: | |
| input_ids = input_ids[:, -1:] | |
| return {"input_ids": input_ids, "attention_mask": attention_mask, "cache_params": cache_params, | |
| "use_cache": use_cache} | |