Files
dela 49aede9cb2 Fit 0.5b training on 32GB: SDPA MLA, block checkpoint, chunked CE
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.
2026-08-25 20:09:27 +08:00

92 lines
3.8 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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]