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