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,110 @@
|
||||
"""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()
|
||||
Reference in New Issue
Block a user