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:
dela
2026-08-25 14:43:17 +08:00
commit 584f7e9e73
140 changed files with 21592 additions and 0 deletions
+106
View File
@@ -0,0 +1,106 @@
"""验证 reference `_decayed_dot` 的 g_ref 因子在真实训练门控量级下是否溢出。
_decayed_dot 把 exp(g_i-g_j) 拆成 (x*exp(g_i-g_ref)) @ (k*exp(g_ref-g_j)),
g_ref = g_cumsum[..., 0, :]。因 g<0,最大中间因子为
exp(g_ref - g_{C-1}) = exp(sum_{t=1}^{C-1} |g_t|)
溢出阈值: sum|g| > ln(3.39e38)=88.7 (fp32/bf16) ; > ln(65504)=11.09 (fp16)
门控: safe_gate, g = lower_bound * sigmoid(exp(A_log) * (g_raw + dt_bias))
lower_bound=-5 => |g| ∈ (0, 5) 逐步硬上界
"""
import math
import sys
import torch
import kda.ops.api as kapi
from kda.models.causal_lm import CausalLM
from kda.models.config import KDAConfig
from kda.models.k3_config import K3Config
dev = "cuda"
LIM = {"fp32/bf16": math.log(3.39e38), "fp16": math.log(65504.0)}
# ---------- 1) 解析上界: 每个 chunk_size 需要的平均 |g| ----------
print("=== 解析: 溢出所需的 chunk 内平均 |g|/step (硬上界 |g|<5) ===")
print(f"{'C':>4} | {'fp32/bf16 阈值':>15} | {'可达?':>6} | {'fp16 阈值':>10} | {'可达?':>6}")
for C in (16, 32, 64, 128):
a, b = LIM["fp32/bf16"] / (C - 1), LIM["fp16"] / (C - 1)
print(f"{C:4d} | {a:15.2f} | {'YES' if a < 5 else 'no':>6} | "
f"{b:10.2f} | {'YES' if b < 5 else 'no':>6}")
# ---------- 2) 实测: 真实 checkpoint 上的 chunk 内 sum|g| ----------
_orig = kapi.naive_chunk_kda # api.py 在 import 时绑定, 必须 patch 这里
stats = []
def max_sum_abs_g(g, C):
"""chunk 内最大 sum_{t=1..C-1}|g_t| (= max exp(g_ref-g_j) 的指数)。"""
T = g.shape[1]
gg = g[:, : T - T % C].reshape(g.shape[0], -1, C, *g.shape[2:])
return gg[:, :, 1:].abs().sum(dim=2).max().item()
def patched(q, k, v, g, beta, **kw):
stats.append(g.detach().float())
return _orig(q, k, v, g, beta, **kw)
kapi.naive_chunk_kda = patched
def build(cfg_dict):
"""checkpoint 的 config 是 dict, 按字段集合判断是 K3 还是纯 KDA。"""
cls = K3Config if "moe_latent_size" in cfg_dict else KDAConfig
return cls(**{k: v for k, v in cfg_dict.items()
if k in cls.__dataclass_fields__})
for path in ("ckpts/k3_wiki.pt", "ckpts/kda_toy.pt"):
ck = torch.load(path, map_location="cpu", weights_only=False)
cfg = build(ck["config"])
model = CausalLM(cfg).to(dev).eval()
# 旧 checkpoint 用 moe/moe_norm 命名, 现已重命名为 ffn/ffn_norm
sd = {k.replace(".moe_norm.", ".ffn_norm.").replace(".moe.", ".ffn.")
.replace(".mlp_norm.", ".ffn_norm.").replace(".mlp.", ".ffn."): v
for k, v in ck["model_state"].items()}
missing, unexpected = model.load_state_dict(sd, strict=False)
assert not missing and not unexpected, (path, missing[:5], unexpected[:5])
stats.clear()
ids = torch.randint(0, cfg.vocab_size, (2, 512), device=dev)
with torch.no_grad():
model(ids)
if not stats:
print(f"\n{path}: 未走 reference 路径 (backend={cfg.kda_backend})")
continue
gs = list(stats)
print(f"\n=== {path} (训练用 C={cfg.chunk_size}, KDA 层数 {len(gs)}) ===")
print(f" |g| mean/step: {sum(g.abs().mean().item() for g in gs)/len(gs):.4f} "
f"|g| max/step: {max(g.abs().max().item() for g in gs):.4f} "
f"(硬上界 {abs(cfg.lower_bound)})")
print(f" {'C':>4} | {'max chunk sum|g|':>16} | {'max exp 因子':>12} | "
f"{'fp32/bf16':>18} | {'fp16':>12}")
for C in (16, 32, 64, 128):
mg = max(max_sum_abs_g(g, C) for g in gs)
fac = math.exp(mg) if mg < 709 else float("inf")
f32 = f"{LIM['fp32/bf16']/mg:.2f}x OK" if mg < LIM["fp32/bf16"] else "OVERFLOW"
f16 = f"{LIM['fp16']/mg:.2f}x OK" if mg < LIM["fp16"] else "OVERFLOW"
mark = " <- 训练配置" if C == cfg.chunk_size else ""
print(f" {C:4d} | {mg:16.2f} | {fac:12.3e} | {f32:>18} | {f16:>12}{mark}")
# ---------- 3) 随机初始化模型 (未训练) 同样测一遍 ----------
for preset in ("toy", "0.5b"):
cfg = K3Config.preset(preset)
cfg.vocab_size = 2048
cfg.kda_backend = "reference"
model = CausalLM(cfg).to(dev).eval()
stats.clear()
with torch.no_grad():
model(torch.randint(0, cfg.vocab_size, (2, 512), device=dev))
if stats:
gs = list(stats)
mg = max(max_sum_abs_g(g, cfg.chunk_size) for g in gs)
print(f"\n=== init preset={preset} (C={cfg.chunk_size}) ===")
print(f" |g| mean/step {sum(g.abs().mean().item() for g in gs)/len(gs):.4f} "
f"max chunk sum|g| {mg:.3f} fp32 余量 {LIM['fp32/bf16']/max(mg,1e-9):.0f}x")