"""L7: toy overfit smoke test. 320 steps loss < 0.1.""" import torch from kda.models.causal_lm import CausalLM from kda.models.config import KDAConfig def test_overfit(): cfg = KDAConfig() # 起步默认 toy 配置 torch.manual_seed(30) model = CausalLM(cfg).cuda() x = torch.randint(0, cfg.vocab_size, (4, 16), device="cuda") labels = x.clone() # A single repeated batch is an optimizer/dataflow smoke test, so converge it quickly. optim = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=0.01) for step in range(320): optim.zero_grad() loss = model(x, labels=labels) loss.backward() optim.step() if step % 64 == 0 or step == 319: print(f" step {step:3d} loss {loss.item():.4f}") final = loss.item() assert final < 0.1, f"final loss {final:.4f} > 0.1" print(f"L7 overfit: PASSED (final loss {final:.4f})") if __name__ == "__main__": test_overfit()