Files
K3/kda/layers/mla.py
T
dela 584f7e9e73 Initial K3 snapshot: 0.5B KDA/MLA/MoE train path
Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias
backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
2026-08-25 14:43:17 +08:00

98 lines
4.0 KiB
Python
Raw 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
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]