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.
This commit is contained in:
@@ -0,0 +1,32 @@
|
||||
"""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()
|
||||
Reference in New Issue
Block a user