File size: 6,278 Bytes
a9b5915
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
#!/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)