Files
K3/notes/verify_gref_overflow.py
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

107 lines
4.5 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""验证 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")