Whole-mixer checkpoint plus T×T MLA scores OOM'd a 31GB GPU on backward. Checkpoint each AttnRes block, run absorbed MLA through SDPA, and compute CE in vocab chunks so [B,T,V] logits are never materialized. --max-tokens is now the training budget; default --steps 2000 no longer caps a 1B-token run at 250 optimizer steps.
92 lines
3.8 KiB
Python
92 lines
3.8 KiB
Python
"""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, scale=1 matches the unscaled einsum.
|
||
q_absorb = torch.einsum("bthd,hdj->bthj", q, w_uk) # [B, T, H, r]
|
||
q_h = q_absorb.transpose(1, 2) # [B, H, T, r]
|
||
kv = c.unsqueeze(1).expand(B, H, T, r)
|
||
latent_out = F.scaled_dot_product_attention(
|
||
q_h, kv, kv, is_causal=True, scale=1.0
|
||
) # [B, H, T, r]
|
||
o_heads = torch.einsum("bhtr,hvr->bhtv", latent_out, w_uv)
|
||
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]
|