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.
105 lines
2.9 KiB
Python
105 lines
2.9 KiB
Python
import torch
|
|
import torch.nn.functional as F
|
|
|
|
from kda.models.causal_lm import CausalLM
|
|
from kda.models.config import KDAConfig
|
|
from kda.models.k3_config import K3Config
|
|
|
|
|
|
def _tiny():
|
|
torch.manual_seed(4)
|
|
cfg = KDAConfig(
|
|
hidden_size=16,
|
|
num_hidden_layers=2,
|
|
num_heads=2,
|
|
num_value_heads=2,
|
|
head_dim=4,
|
|
chunk_size=4,
|
|
vocab_size=32,
|
|
intermediate_size=32,
|
|
kda_backend="reference",
|
|
)
|
|
return CausalLM(cfg), cfg
|
|
|
|
|
|
def test_ignore_index_skips_masked_positions():
|
|
model, _ = _tiny()
|
|
tokens = torch.tensor([[1, 2, 3, 4]])
|
|
labels = tokens.clone()
|
|
labels[:, 1:3] = -100
|
|
with torch.no_grad():
|
|
logits = model(tokens)
|
|
actual = model(tokens, labels=labels)
|
|
expected = F.cross_entropy(
|
|
logits[:, :-1].reshape(-1, logits.size(-1)),
|
|
labels[:, 1:].reshape(-1),
|
|
ignore_index=-100,
|
|
)
|
|
torch.testing.assert_close(actual, expected)
|
|
|
|
|
|
def test_gradient_checkpointing_matches_eager_grad():
|
|
torch.manual_seed(8)
|
|
tokens = torch.randint(0, 32, (2, 8))
|
|
m1, cfg = _tiny()
|
|
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_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
|