From 49aede9cb286608195a823cd19e431f54a21d611 Mon Sep 17 00:00:00 2001 From: dela Date: Tue, 25 Aug 2026 20:09:27 +0800 Subject: [PATCH] Fit 0.5b training on 32GB: SDPA MLA, block checkpoint, chunked CE MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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. --- kda/layers/attn_res.py | 15 +++++++- kda/layers/mla.py | 20 ++++------- kda/models/causal_lm.py | 35 +++++++++++++----- kda/training/schedule.py | 16 +++++---- tests/integration/test_causal_checkpoint.py | 40 +++++++++++++++++++++ tests/integration/test_schedule.py | 14 ++++++-- train_k3.py | 13 ++++--- 7 files changed, 117 insertions(+), 36 deletions(-) diff --git a/kda/layers/attn_res.py b/kda/layers/attn_res.py index f70d11b..bfb8a73 100644 --- a/kda/layers/attn_res.py +++ b/kda/layers/attn_res.py @@ -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 ( diff --git a/kda/layers/mla.py b/kda/layers/mla.py index 5694d0f..b718ae1 100644 --- a/kda/layers/mla.py +++ b/kda/layers/mla.py @@ -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] diff --git a/kda/models/causal_lm.py b/kda/models/causal_lm.py index 9d1fd45..0b6f719 100644 --- a/kda/models/causal_lm.py +++ b/kda/models/causal_lm.py @@ -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() diff --git a/kda/training/schedule.py b/kda/training/schedule.py index f56e3c5..a23bf17 100644 --- a/kda/training/schedule.py +++ b/kda/training/schedule.py @@ -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 diff --git a/tests/integration/test_causal_checkpoint.py b/tests/integration/test_causal_checkpoint.py index 78c8527..89c55e3 100644 --- a/tests/integration/test_causal_checkpoint.py +++ b/tests/integration/test_causal_checkpoint.py @@ -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 diff --git a/tests/integration/test_schedule.py b/tests/integration/test_schedule.py index fe52ba9..66212a7 100644 --- a/tests/integration/test_schedule.py +++ b/tests/integration/test_schedule.py @@ -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 diff --git a/train_k3.py b/train_k3.py index aea1982..3eb3143 100644 --- a/train_k3.py +++ b/train_k3.py @@ -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)