Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
33 lines
1001 B
Python
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()
|