"""Gated MLA (K3): NoPE, latent KV compression, matrix absorption, full-rank output gate. K3 相对 DeepSeek MLA 的三个改动 (对照 learning/kimi-k3-notes): 1. NoPE — 不显式 RoPE; 位置感交给夹层 KDA 的 decay/gate。 2. 矩阵吸收 — 训练/推理都不解压 K/V: q 吸收 W_UK 后直接与 latent c 内积, 输出先在 latent 加权再乘 W_UV 还原 (v2 吸收版)。 3. Full-rank 输出门 — y = W_o[ σ(W_g x) ⊙ õ ]。 形状 (小规模 toy, d 为 hidden): c = RMSNorm(kv_down(x)) [B, T, r] latent q = q_up(RMSNorm(q_down(x))) [B, T, H, d_q] d_q = d_nope (NoPE) W_UK = kv_up[.., :H*d_q].view(H,d_q,r) W_UV = kv_up[.., H*d_q:].view(H,d_v,r) score = (q @ W_UK^T) @ c^T [B, H, T, T] causal õ = (softmax(score) @ c) @ W_UV^T [B, T, H, d_v] y = o_proj( σ(W_g x) ⊙ õ_head ) [B, T, d] """ from __future__ import annotations import torch import torch.nn.functional as F from torch import nn from .rmsnorm import RMSNorm class GatedMLA(nn.Module): def __init__( self, hidden_size: int, num_heads: int, kv_lora_rank: int, q_lora_rank: int, qk_nope_head_dim: int, v_head_dim: int, ): super().__init__() self.hidden_size = hidden_size self.num_heads = num_heads self.qk_nope_head_dim = qk_nope_head_dim self.v_head_dim = v_head_dim # Q 低秩路径 (NoPE, 只有 nope 段) self.q_down = nn.Linear(hidden_size, q_lora_rank, bias=False) self.q_norm = RMSNorm(q_lora_rank) self.q_up = nn.Linear(q_lora_rank, num_heads * qk_nope_head_dim, bias=False) # KV latent 压缩 + 解压 (W_UK | W_UV 拼接在同一矩阵里) self.kv_down = nn.Linear(hidden_size, kv_lora_rank, bias=False) self.kv_norm = RMSNorm(kv_lora_rank) self.kv_up = nn.Linear( kv_lora_rank, num_heads * (qk_nope_head_dim + v_head_dim), bias=False ) # Full-rank 输出门: σ(W_g x) 与 õ (H*d_v 维) 逐元素相乘 self.gate = nn.Linear(hidden_size, num_heads * v_head_dim, bias=False) self.o_proj = nn.Linear(num_heads * v_head_dim, hidden_size, bias=False) @classmethod def from_config(cls, config) -> GatedMLA: return cls( config.hidden_size, config.num_heads, config.kv_lora_rank, config.q_lora_rank, config.qk_nope_head_dim, config.v_head_dim, ) def forward(self, x: torch.Tensor): B, T, _ = x.shape H, r = self.num_heads, self.kv_up.in_features c = self.kv_norm(self.kv_down(x)) # [B, T, r] q = self.q_up(self.q_norm(self.q_down(x))) # [B, T, H*d_q] q = q.view(B, T, H, self.qk_nope_head_dim) # [B, T, H, d_q] w = self.kv_up.weight # [H*(d_q+d_v), r] w_uk = w[: H * self.qk_nope_head_dim].view(H, self.qk_nope_head_dim, r) w_uv = w[H * self.qk_nope_head_dim :].view(H, self.v_head_dim, r) # 吸收 W_UK 进 query: score = (q @ W_UK^T) @ c^T q_absorb = torch.einsum("bthd,hdj->bthj", q, w_uk) # [B, T, H, r] scores = torch.einsum("bthj,bsj->bhts", q_absorb, c) # [B, H, T, T] mask = torch.triu( torch.ones(T, T, dtype=torch.bool, device=x.device), diagonal=1 ) scores = scores.masked_fill(mask, float("-inf")) attn = F.softmax(scores, dim=-1) # [B, H, T, T] # 先在 latent 加权, 再乘 W_UV^T 还原 v —— 永不解压 latent_out = torch.einsum("bhts,bsj->bhtj", attn, c) # [B, H, T, r] o_heads = torch.einsum("bhtj,hvj->bhtv", latent_out, w_uv) # [B, H, T, d_v] o_heads = o_heads.transpose(1, 2).reshape(B, T, H * self.v_head_dim) gate = torch.sigmoid(self.gate(x)) # [B, T, H*d_v] return self.o_proj(gate * o_heads) # [B, T, d]