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.
This commit is contained in:
+7
-13
@@ -78,20 +78,14 @@ class GatedMLA(nn.Module):
|
||||
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
|
||||
# 吸收 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]
|
||||
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]
|
||||
|
||||
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]
|
||||
|
||||
Reference in New Issue
Block a user