Files
K3/tests/integration/test_train_overfit.py
T
dela 584f7e9e73 Initial K3 snapshot: 0.5B KDA/MLA/MoE train path
Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias
backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
2026-08-25 14:43:17 +08:00

33 lines
1001 B
Python

"""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()