kingjones777's picture
Add files using upload-large-folder tool
a9b5915 verified
Raw History Blame
6.28 kB
#!/usr/bin/env python3
"""Independent gate for fold_agnes.py. Re-derives every expected tensor from the SOURCE.
G1 name/count: every source tensor maps to exactly one output tensor, no extras
G2 bit-exact : passthrough tensors torch.equal; folded tensors equal main/parallel slices;
MTP pad region exactly zero
G3 forward : reference Agnes MLP (main + parallel, ACT2FN[hidden_act]) == folded MLP, fp32
G4 config : ids, layer plan, interval, widths
exit 0 only if every gate passes.
usage: verify_fold.py <agnes_hf_dir> <folded_dir>"""
import json, os, re, sys
import torch
from safetensors import safe_open
from transformers.activations import ACT2FN
SRC, DST = sys.argv[1], sys.argv[2]
fails = []
def check(ok, msg):
if not ok:
fails.append(msg); print(" FAIL", msg, flush=True)
scfg = json.load(open(f"{SRC}/config.json")); dcfg = json.load(open(f"{DST}/config.json"))
st, dt = scfg["text_config"], dcfg["text_config"]
MAIN, PAR = st["intermediate_size"], st["parallel_ffn_intermediate_size"]; FOLD = MAIN + PAR
act = ACT2FN[st["hidden_act"]]
# ---- G4 config ----
check(dcfg["model_type"] == "qwen3_5", "G4 model_type")
check(dcfg["architectures"] == ["Qwen3_5ForConditionalGeneration"], "G4 architectures")
check("auto_map" not in dcfg, "G4 auto_map removed")
check(dt["model_type"] == "qwen3_5_text", "G4 text model_type")
check(dt["intermediate_size"] == FOLD, f"G4 intermediate_size {dt['intermediate_size']} != {FOLD}")
check("parallel_ffn_intermediate_size" not in dt, "G4 parallel key removed")
check(dt["full_attention_interval"] == st["global_attention_interval"], "G4 interval")
m = {"agnes_delta_attention": "linear_attention", "agnes_global_attention": "full_attention"}
check(dt["layer_types"] == [m[x] for x in st["layer_types"]], "G4 layer_types")
for k in st:
if k in ("model_type", "layer_types", "global_attention_interval", "intermediate_size", "parallel_ffn_intermediate_size"):
continue
check(dt.get(k) == st[k], f"G4 text key changed unexpectedly: {k}")
check(dcfg["vision_config"]["model_type"] == "qwen3_5", "G4 vision model_type")
for k in scfg["vision_config"]:
if k != "model_type":
check(dcfg["vision_config"][k] == scfg["vision_config"][k], f"G4 vision key changed: {k}")
print(f"G4 config: {'ok' if not fails else 'FAILED'}", flush=True)
# ---- G1/G2 ----
swm = json.load(open(f"{SRC}/model.safetensors.index.json"))["weight_map"]
dwm = json.load(open(f"{DST}/model.safetensors.index.json"))["weight_map"]
sh, dh = {}, {}
def S(n):
f = swm[n]; sh.setdefault(f, safe_open(f"{SRC}/{f}", "pt")); return sh[f].get_tensor(n)
def D(n):
f = dwm[n]; dh.setdefault(f, safe_open(f"{DST}/{f}", "pt")); return dh[f].get_tensor(n)
def ren(n): return n.replace(".delta_attn.", ".linear_attn.").replace(".global_attn.", ".self_attn.")
RM = re.compile(r"^(model\.language_model\.layers\.\d+\.mlp)\.(gate_proj|up_proj|down_proj)\.weight$")
RT = re.compile(r"^mtp\.layers\.\d+\.mlp\.(gate_proj|up_proj|down_proj)\.weight$")
expected = {ren(n) for n in swm if ".parallel_ffn." not in n}
check(set(dwm) == expected, f"G1 name set mismatch: missing={sorted(expected-set(dwm))[:5]} extra={sorted(set(dwm)-expected)[:5]}")
check(not any(s in n for n in dwm for s in (".delta_attn.", ".global_attn.", ".parallel_ffn.")), "G1 forbidden substring survived")
print(f"G1 names: {len(dwm)} out vs {len(expected)} expected", flush=True)
nf = nt = npz = 0
for i, n in enumerate(sorted(swm)):
if ".parallel_ffn." in n:
continue
s, d = S(n), D(ren(n))
check(s.dtype == d.dtype, f"G2 dtype {n}")
mm, mt = RM.match(n), RT.match(n)
if mm:
pre, kind = mm.groups(); p = S(f"{pre}.parallel_ffn.{kind}.weight")
if kind == "down_proj":
check(tuple(d.shape) == (s.shape[0], FOLD), f"G2 shape {n} {tuple(d.shape)}")
check(torch.equal(d[:, :MAIN], s) and torch.equal(d[:, MAIN:], p), f"G2 fold content {n}")
else:
check(tuple(d.shape) == (FOLD, s.shape[1]), f"G2 shape {n} {tuple(d.shape)}")
check(torch.equal(d[:MAIN], s) and torch.equal(d[MAIN:], p), f"G2 fold content {n}")
nf += 1
elif mt:
kind = mt.group(1)
if kind == "down_proj":
check(tuple(d.shape) == (s.shape[0], FOLD) and torch.equal(d[:, :MAIN], s) and not d[:, MAIN:].any(), f"G2 mtp pad {n}")
else:
check(tuple(d.shape) == (FOLD, s.shape[1]) and torch.equal(d[:MAIN], s) and not d[MAIN:].any(), f"G2 mtp pad {n}")
npz += 1
else:
check(tuple(s.shape) == tuple(d.shape) and torch.equal(s, d), f"G2 passthrough {n}")
nt += 1
if i % 150 == 0:
print(f" G2 progress {i}/{len(swm)}", flush=True)
print(f"G2 bit-exact: folded={nf} mtp_padded={npz} passthrough={nt}", flush=True)
# ---- G3 forward equivalence ----
torch.manual_seed(0)
H = st["hidden_size"]
for L in (0, st["num_hidden_layers"] // 2, st["num_hidden_layers"] - 1):
b = f"model.language_model.layers.{L}.mlp"
g, u, dn = (S(f"{b}.{k}.weight").float() for k in ("gate_proj", "up_proj", "down_proj"))
pg, pu, pd = (S(f"{b}.parallel_ffn.{k}.weight").float() for k in ("gate_proj", "up_proj", "down_proj"))
fg, fu, fd = (D(f"{b}.{k}.weight").float() for k in ("gate_proj", "up_proj", "down_proj"))
x = torch.randn(16, H)
ref = (act(x @ g.T) * (x @ u.T)) @ dn.T + (act(x @ pg.T) * (x @ pu.T)) @ pd.T
fol = (act(x @ fg.T) * (x @ fu.T)) @ fd.T
rel = ((ref - fol).abs().max() / ref.abs().max()).item()
check(rel < 1e-5, f"G3 layer {L} rel err {rel:.3e}")
print(f"G3 layer {L:2d}: max rel err {rel:.3e} (|ref|max {ref.abs().max().item():.3f})", flush=True)
b = "mtp.layers.0.mlp"
g, u, dn = (S(f"{b}.{k}.weight").float() for k in ("gate_proj", "up_proj", "down_proj"))
fg, fu, fd = (D(f"{b}.{k}.weight").float() for k in ("gate_proj", "up_proj", "down_proj"))
x = torch.randn(16, H)
ref = (act(x @ g.T) * (x @ u.T)) @ dn.T
fol = (act(x @ fg.T) * (x @ fu.T)) @ fd.T
rel = ((ref - fol).abs().max() / ref.abs().max()).item()
check(rel < 1e-5, f"G3 mtp rel err {rel:.3e}")
print(f"G3 mtp : max rel err {rel:.3e}", flush=True)
print(f"\nRESULT: {'PASS' if not fails else 'FAIL'} ({len(fails)} failures)", flush=True)
sys.exit(0 if not fails else 1)