""" 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()