Files
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

111 lines
3.6 KiB
Python

"""L7: toy training loop — overfit 起步.
target:
端到端验证模型 + 数据流 + optimizer + ckpt + generate.
toy data:
建一份 256-token vocab 的小数据集: e.g. 1000 个长度 32 随机 token 序列
起步只取 batch=4, 看能否在 ~320 steps 内把 loss 压到 < 0.1 (overfit 单 batch).
step:
optimizer = AdamW(lr=1e-3, wd=0.01)
loss.backward(); optimizer.step(); optimizer.zero_grad()
every N steps: 打印 loss
end: 保存 ckpt to ckpts/kda_toy.pt
ckpt:
save:
torch.save({"model_state": model.state_dict(), "config": asdict(config)}, path)
load:
torch.load -> model.load_state_dict
"""
from __future__ import annotations
import os
from dataclasses import asdict, fields
import torch
from ..models.causal_lm import CausalLM
from ..models.config import KDAConfig
from ..models.k3_config import K3Config
def make_toy_data(batch: int = 4, seq_len: int = 32, vocab: int = 256, seed: int = 42):
"""单 batch overfit 数据: 同一组序列循环."""
torch.manual_seed(seed)
seq = torch.randint(0, vocab, (batch, seq_len), dtype=torch.long)
return seq # 用作 input_ids 和 labels (shift one inside forward)
def train_one_batch(model, optimizer, input_ids, labels):
optimizer.zero_grad(set_to_none=True)
loss = model(input_ids, labels=labels)
loss.backward()
optimizer.step()
return loss.detach()
def save_ckpt(model, config, path: str):
os.makedirs(os.path.dirname(path) or ".", exist_ok=True)
torch.save({"model_state": model.state_dict(), "config": asdict(config)}, path)
#: The feed-forward submodule was named after its contents (``mlp`` in the
#: dense config, ``moe`` in K3) before both were unified under ``ffn``.
#: Checkpoints saved before that rename still carry the old prefixes.
_LEGACY_PREFIXES = {
".mlp.": ".ffn.",
".mlp_norm.": ".ffn_norm.",
".moe.": ".ffn.",
".moe_norm.": ".ffn_norm.",
}
def _rename_legacy_keys(state: dict) -> dict:
def fix(key: str) -> str:
for old, new in _LEGACY_PREFIXES.items():
if old in key:
return key.replace(old, new)
return key
return {fix(k): v for k, v in state.items()}
def _config_from(payload_config: dict) -> K3Config | KDAConfig:
"""Pick the config class the checkpoint was written with.
``moe_latent_size`` is a K3-only field, so its presence identifies the
hybrid K3 architecture; anything else is the dense KDA config.
"""
cls = K3Config if "moe_latent_size" in payload_config else KDAConfig
known = {item.name for item in fields(cls)}
return cls(**{k: v for k, v in payload_config.items() if k in known})
def load_ckpt(path: str, model: CausalLM | None = None) -> tuple[CausalLM, K3Config | KDAConfig]:
payload = torch.load(path, map_location="cpu", weights_only=False)
config = _config_from(payload["config"])
if model is None:
model = CausalLM(config)
model.load_state_dict(_rename_legacy_keys(payload["model_state"]))
return model, config
def main():
"""主入口: overfit 起步. 320 steps 期望 loss < 0.1."""
device = "cuda" if torch.cuda.is_available() else "cpu"
config = KDAConfig()
model = CausalLM(config).to(device)
tokens = make_toy_data(seq_len=32, vocab=config.vocab_size).to(device)
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=0.01)
for step in range(320):
loss = train_one_batch(model, optimizer, tokens, tokens)
if step % 64 == 0 or step == 319:
print(f"step {step:3d} loss {loss.item():.4f}")
save_ckpt(model, config, "ckpts/kda_toy.pt")
if __name__ == "__main__":
main()