"""验证 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")