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:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user