DM-KDAGQA-1.7B-SD / modeling_bqalm.py
tturing's picture
Add Q35KDA_fp8base_ffn8_qkvo8_fp11_hd128_11785_fm128: KDA:GQA 1.7B, SD (stale) arm, seed 11785 (standalone, trust_remote_code)
4e9d3a8 verified
Raw History Blame Contribute Delete
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
@property
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
@property
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))
@staticmethod
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"}
@dataclass
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
@classmethod
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}