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:
+14
-1
@@ -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
@@ -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
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -59,6 +59,46 @@ def test_gradient_checkpointing_matches_eager_grad():
|
||||
torch.testing.assert_close(p1.grad, p2.grad, atol=1e-4, rtol=1e-4)
|
||||
|
||||
|
||||
def test_attnres_block_checkpoint_matches_eager_grad():
|
||||
torch.manual_seed(8)
|
||||
cfg = K3Config(
|
||||
hidden_size=32,
|
||||
num_hidden_layers=4,
|
||||
num_heads=4,
|
||||
head_dim=8,
|
||||
chunk_size=4,
|
||||
vocab_size=32,
|
||||
moe_latent_size=16,
|
||||
moe_d_ff=16,
|
||||
n_routed=4,
|
||||
kv_lora_rank=16,
|
||||
q_lora_rank=32,
|
||||
qk_nope_head_dim=8,
|
||||
v_head_dim=8,
|
||||
attnres="block",
|
||||
attnres_block_size=1,
|
||||
moe_aux_loss_coef=0.0,
|
||||
moe_z_loss_coef=0.0,
|
||||
)
|
||||
tokens = torch.randint(0, cfg.vocab_size, (2, 8))
|
||||
m1 = CausalLM(cfg)
|
||||
m2 = CausalLM(cfg)
|
||||
m2.load_state_dict(m1.state_dict())
|
||||
m2.gradient_checkpointing = True
|
||||
m1.train()
|
||||
m2.train()
|
||||
l1 = m1(tokens, labels=tokens)
|
||||
l2 = m2(tokens, labels=tokens)
|
||||
torch.testing.assert_close(l1, l2, atol=1e-5, rtol=1e-5)
|
||||
l1.backward()
|
||||
l2.backward()
|
||||
for p1, p2 in zip(m1.parameters(), m2.parameters()):
|
||||
if p1.grad is None:
|
||||
assert p2.grad is None
|
||||
continue
|
||||
torch.testing.assert_close(p1.grad, p2.grad, atol=1e-4, rtol=1e-4)
|
||||
|
||||
|
||||
def test_0_5b_preset_enables_checkpointing():
|
||||
assert K3Config.preset("0.5b").gradient_checkpointing is True
|
||||
assert K3Config.preset("toy").gradient_checkpointing is False
|
||||
|
||||
@@ -9,13 +9,21 @@ def test_warmup_then_cosine_floor():
|
||||
assert abs(end - 0.1) < 1e-6
|
||||
|
||||
|
||||
def test_horizon_prefers_the_earlier_stop():
|
||||
# 8.2M tokens @ batch 2 seq 2048 acc 8 -> 250 opt
|
||||
def test_horizon_max_tokens_overrides_micro_cap():
|
||||
# 8.2M tokens @ batch 2 seq 2048 acc 8 -> 250 opt even if --steps is larger
|
||||
opt_from_tokens = total_opt_steps(
|
||||
max_tokens=8_192_000, max_micro=10_000, batch=2, seq_len=2048, grad_acc=8
|
||||
)
|
||||
assert opt_from_tokens == 250
|
||||
# 1B-token run must not inherit the default --steps 2000 cap (250 opt)
|
||||
opt_1b = total_opt_steps(
|
||||
max_tokens=10**9, max_micro=2000, batch=2, seq_len=2048, grad_acc=8
|
||||
)
|
||||
assert opt_1b == 30518
|
||||
|
||||
|
||||
def test_horizon_micro_when_tokens_unset():
|
||||
opt_from_micro = total_opt_steps(
|
||||
max_tokens=10**12, max_micro=2000, batch=2, seq_len=2048, grad_acc=8
|
||||
max_tokens=None, max_micro=2000, batch=2, seq_len=2048, grad_acc=8
|
||||
)
|
||||
assert opt_from_micro == 250
|
||||
|
||||
+9
-4
@@ -247,6 +247,7 @@ def main() -> None:
|
||||
help="router z-loss weight (default 0.001; 0 disables)",
|
||||
)
|
||||
args = p.parse_args()
|
||||
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
|
||||
if args.gen_prefix is None:
|
||||
args.gen_prefix = ["人工智能的发展", "The history of computing"]
|
||||
|
||||
@@ -387,9 +388,10 @@ def main() -> None:
|
||||
t0 = time.perf_counter()
|
||||
tokens_at_t0 = tokens
|
||||
for chunk_index, x, y in iter_indexed(train_chunks, start=chunk_index):
|
||||
if args.max_tokens is not None and tokens >= args.max_tokens:
|
||||
break
|
||||
if micro_step >= args.steps:
|
||||
if args.max_tokens is not None:
|
||||
if tokens >= args.max_tokens:
|
||||
break
|
||||
elif micro_step >= args.steps:
|
||||
break
|
||||
x, y = x.to(device), y.to(device)
|
||||
scale = lr_scale(opt_step, args.warmup, horizon)
|
||||
@@ -431,7 +433,7 @@ def main() -> None:
|
||||
micro_step % args.eval_every == 0
|
||||
or micro_step == 1
|
||||
or (args.max_tokens is not None and tokens >= args.max_tokens)
|
||||
or micro_step >= args.steps
|
||||
or (args.max_tokens is None and micro_step >= args.steps)
|
||||
)
|
||||
if log_now:
|
||||
held = _heldout_loss(model, held_chunks, device, use_bf16)
|
||||
@@ -474,6 +476,9 @@ def main() -> None:
|
||||
print(
|
||||
f" best held-out {best_heldout:.4f} -> {_sibling(args.out, '_best')}"
|
||||
)
|
||||
del payload
|
||||
if device == "cuda":
|
||||
torch.cuda.empty_cache()
|
||||
if tracker is not None:
|
||||
tracker.log(metrics, step=micro_step)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user