traj-mc / code /analysis /summary_diff.py
ttishere's picture
Publish code/analysis
3fffa60 verified
Raw History Blame Contribute Delete
4.48 kB
"""
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()