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:
dela
2026-08-25 20:09:27 +08:00
parent 7a12f61de1
commit 49aede9cb2
7 changed files with 117 additions and 36 deletions
+14 -1
View File
@@ -20,6 +20,7 @@ import torch
import torch.nn.functional as F
from einops import rearrange
from torch import Tensor, nn
from torch.utils.checkpoint import checkpoint as activation_checkpoint
ATTNRES_MODES = ("off", "full", "block")
@@ -247,6 +248,7 @@ class BlockAttnResStack(nn.Module):
if is_final_aggregate
else None
)
self.gradient_checkpointing = False
def forward_naive(self, x: Tensor) -> Tensor:
blocks = [x] # b_0=embedding/input representation
@@ -294,9 +296,20 @@ class BlockAttnResStack(nn.Module):
blocks = [x]
depth = len(self.layers)
start = 0
use_ckpt = (
self.gradient_checkpointing and self.training and torch.is_grad_enabled()
)
while start < depth:
end = min(start + self.block_size, depth)
blocks.append(self._run_block_two_phase(blocks, start, end))
if use_ckpt:
def _run(*srcs, _start=start, _end=end):
return self._run_block_two_phase(list(srcs), _start, _end)
blocks.append(
activation_checkpoint(_run, *blocks, use_reentrant=False)
)
else:
blocks.append(self._run_block_two_phase(blocks, start, end))
start = end
return (
+7 -13
View File
@@ -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]
+27 -8
View File
@@ -26,6 +26,28 @@ from ..layers.block import DecoderBlock
from ..layers.rmsnorm import RMSNorm
def _chunked_linear_cross_entropy(
hidden: torch.Tensor,
weight: torch.Tensor,
labels: torch.Tensor,
ignore_index: int = -100,
chunk_size: int = 256,
) -> torch.Tensor:
"""CE without materializing [B, T, vocab]. Match mean reduction over valid labels."""
features = hidden[:, :-1].reshape(-1, hidden.size(-1))
targets = labels[:, 1:].reshape(-1)
total = hidden.new_zeros(())
n_valid = hidden.new_zeros((), dtype=torch.long)
for start in range(0, features.size(0), chunk_size):
sl = slice(start, start + chunk_size)
logits = F.linear(features[sl], weight)
total = total + F.cross_entropy(
logits, targets[sl], ignore_index=ignore_index, reduction="sum"
)
n_valid = n_valid + (targets[sl] != ignore_index).sum()
return total / n_valid.clamp_min(1).to(dtype=total.dtype)
def _build_mixer(config, blocks: nn.ModuleList):
mode = getattr(config, "attnres", "off")
if mode == "off":
@@ -89,17 +111,14 @@ class CausalLM(nn.Module):
x = activation_checkpoint(block, x, use_reentrant=False)
else:
x = block(x)
elif self.gradient_checkpointing and self.training:
x = activation_checkpoint(self.mixer, x, use_reentrant=False)
else:
self.mixer.gradient_checkpointing = self.gradient_checkpointing
x = self.mixer(x)
logits = self.lm_head(self.norm(x))
hidden = self.norm(x)
if labels is None:
return logits
return F.cross_entropy(
logits[:, :-1].reshape(-1, logits.size(-1)),
labels[:, 1:].reshape(-1),
ignore_index=ignore_index,
return self.lm_head(hidden)
return _chunked_linear_cross_entropy(
hidden, self.lm_head.weight, labels, ignore_index=ignore_index
)
@torch.inference_mode()
+9 -7
View File
@@ -34,14 +34,16 @@ def total_opt_steps(
seq_len: int,
grad_acc: int,
) -> int:
"""Optimizer-step horizon used by cosine. At least 1."""
"""Optimizer-step horizon used by cosine. At least 1.
``max_tokens`` is the training budget when set; ``max_micro`` is only used
when ``max_tokens`` is None. Otherwise a default ``--steps 2000`` would
shrink a 1B-token cosine to 250 opt steps.
"""
acc = max(grad_acc, 1)
candidates: list[int] = []
if max_tokens is not None and max_tokens > 0:
tpm = max(tokens_per_micro(batch, seq_len), 1)
candidates.append(math.ceil(max_tokens / (tpm * acc)))
return max(math.ceil(max_tokens / (tpm * acc)), 1)
if max_micro is not None and max_micro > 0:
candidates.append(math.ceil(max_micro / acc))
if not candidates:
return 1
return max(min(candidates), 1)
return max(math.ceil(max_micro / acc), 1)
return 1