File size: 4,480 Bytes
3fffa60
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
summary_diff.py -- HARD GATE: BASE vs OURS compression_summary.json must be
identical except the noise-related / runtime fields. If anything else differs
(replaced-layer count, per-layer rank, param counts, rank formula, layer range,
lm_head status, kept_fraction, decomp, model_id), the two arms are NOT a clean
single-variable comparison and MUST NOT proceed to eval.

Exit 0 = identical (gate PASS). Exit 1 = mismatch (gate FAIL) -> used as a Slurm
dependency (afterok): a non-zero exit auto-cancels the downstream eval jobs.

Fields ALLOWED to differ (arm identity + runtime, not part of the comparison):
  arm, calib_file, save_path, peak_rss_gb, wall_clock_sec
Everything else (incl. the full per_layer rank map) must match byte-for-byte.
"""
import os
import sys
import json
import argparse

sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import common as C

ALLOWED_DIFFER = {"arm", "calib_file", "save_path", "peak_rss_gb", "wall_clock_sec"}


def _load(path):
    if os.path.isdir(path):
        path = os.path.join(path, "compression_summary.json")
    return C.load_json(path), path


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--base", required=True, help="BASE weights dir or summary.json")
    ap.add_argument("--ours", required=True, help="OURS weights dir or summary.json")
    ap.add_argument("--out", default=None)
    args = ap.parse_args()

    b, bp = _load(args.base)
    o, op = _load(args.ours)
    keys = (set(b) | set(o)) - ALLOWED_DIFFER
    mism = []
    for k in sorted(keys):
        if b.get(k) != o.get(k):
            # summarize per_layer mismatch compactly
            if k == "per_layer":
                bl, ol = b.get(k, {}), o.get(k, {})
                diff_layers = [n for n in (set(bl) | set(ol)) if bl.get(n) != ol.get(n)]
                mism.append(("per_layer", f"{len(diff_layers)} layers differ: "
                                          f"{diff_layers[:5]}{'...' if len(diff_layers)>5 else ''}"))
            else:
                mism.append((k, f"base={b.get(k)!r} ours={o.get(k)!r}"))

    result = {
        "base_summary": bp, "ours_summary": op,
        "allowed_differ": sorted(ALLOWED_DIFFER),
        "n_mismatch": len(mism),
        "mismatches": {k: v for k, v in mism},
        "shared": {"ratio_mode": b.get("ratio_mode"),
                   "model_ratio": b.get("model_ratio"),
                   "layer_ratio": b.get("layer_ratio"),
                   "ratio": b.get("ratio"), "layer_type": b.get("layer_type"),
                   "decomp": b.get("decomp"), "model_id": b.get("model_id"),
                   "n_compressed": b.get("n_compressed"),
                   "rank_plan_hash": b.get("rank_plan_hash"),
                   "effective_model_ratio": b.get("effective_model_ratio"),
                   "kept_fraction_base": b.get("kept_fraction"),
                   "kept_fraction_ours": o.get("kept_fraction"),
                   "git_hash": b.get("git_hash")},
        "gate": "PASS" if not mism else "FAIL",
    }
    if args.out:
        C.dump_json(result, args.out)
    else:
        root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
        C.dump_json(result, os.path.join(root, "results", "compress",
                                         f"{C.git_hash()}_summary_diff_{b.get('layer_type')}_r{b.get('ratio')}.json"))
    if mism:
        print(f"[summary_diff] GATE FAIL: {len(mism)} field(s) differ beyond noise:")
        for k, v in mism:
            print(f"    {k}: {v}")
        sys.exit(1)
    print(f"[summary_diff] GATE PASS: BASE/OURS identical except {sorted(ALLOWED_DIFFER)}")
    if b.get("ratio_mode") == "model":
        # the full-model number is the one the paper reports; show it explicitly so
        # a PASS also confirms WHICH working point the two arms agree on
        print(f"    ratio_mode=model  model_ratio={b.get('model_ratio')}  "
              f"layer_ratio={b.get('layer_ratio'):.6f}  "
              f"effective_model_ratio={b.get('effective_model_ratio'):.6f}")
        print(f"    rank_plan_hash={b.get('rank_plan_hash')} (identical in both arms)")
    else:
        print(f"    ratio_mode=layer (LEGACY) ratio={b.get('ratio')}")
    print(f"    layer_type={b.get('layer_type')} "
          f"n_compressed={b.get('n_compressed')} "
          f"kept_fraction base={b.get('kept_fraction'):.4f} ours={o.get('kept_fraction'):.4f}")
    sys.exit(0)


if __name__ == "__main__":
    main()