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