#!/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 """ 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)