SabaPivot's picture
Upgrade Claim 5 with audited Qwen2.5-3B MMLU evidence
8a5ffa8 verified
Raw
History Blame Contribute Delete
4.97 kB
#!/usr/bin/env python3
"""Independently verify the arithmetic in the peer's Claim 5 raw results."""
from __future__ import annotations
import argparse
import hashlib
import json
import math
import statistics
from pathlib import Path
def exact_mcnemar(left, right):
left_only = sum(a == 1 and b == 0 for a, b in zip(left, right))
right_only = sum(a == 0 and b == 1 for a, b in zip(left, right))
n = left_only + right_only
if n == 0:
return left_only, right_only, 1.0
tail = sum(math.comb(n, k) for k in range(min(left_only, right_only) + 1))
p = min(1.0, 2.0 * tail / (2**n))
return left_only, right_only, p
def mean_se(values):
mean = statistics.mean(values)
se = statistics.stdev(values) / math.sqrt(len(values))
return mean, se
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--input", type=Path, required=True)
parser.add_argument("--matched-input", type=Path, required=True)
parser.add_argument("--output", type=Path, required=True)
args = parser.parse_args()
data = json.loads(args.input.read_text())
matched = json.loads(args.matched_input.read_text())
assert data["model"] == "Qwen/Qwen2.5-3B-Instruct"
assert data["n_train"] == 1000
assert data["n_eval_per_cat"] == 250
assert data["n_shot"] == 7
assert data["epochs"] == 5
assert data["lora_r"] == 128
assert data["selected_lr"] == {"qkv": 1e-5, "v": 3e-4}
base_h = data["base_hits"]["Humanities"]["7shot"]
assert len(base_h) == 250
rows = {}
for arm in ("qkv", "v"):
seeds = data["final"][arm]["seeds"]
assert len(seeds) == 3
acc_0 = []
acc_7 = []
tests_vs_base = []
for seed in seeds:
hits_0 = seed["hits"]["Humanities"]["0shot"]
hits_7 = seed["hits"]["Humanities"]["7shot"]
assert len(hits_0) == len(hits_7) == 250
measured_0 = 100 * sum(hits_0) / len(hits_0)
measured_7 = 100 * sum(hits_7) / len(hits_7)
assert measured_0 == seed["acc"]["Humanities"]["0shot"]
assert measured_7 == seed["acc"]["Humanities"]["7shot"]
acc_0.append(measured_0)
acc_7.append(measured_7)
b, c, p = exact_mcnemar(base_h, hits_7)
tests_vs_base.append(
{
"seed": seed["seed"],
"base_right_ft_wrong": b,
"base_wrong_ft_right": c,
"p_exact_two_sided": p,
}
)
mean_0, se_0 = mean_se(acc_0)
mean_7, se_7 = mean_se(acc_7)
rows[arm] = {
"humanities_0shot_mean": mean_0,
"humanities_0shot_se": se_0,
"humanities_7shot_mean": mean_7,
"humanities_7shot_se": se_7,
"delta_0shot_vs_base_pp": mean_0 - data["base"]["Humanities"]["0shot"],
"delta_7shot_vs_base_pp": mean_7 - data["base"]["Humanities"]["7shot"],
"mcnemar_vs_base_7shot": tests_vs_base,
}
paired = []
for qkv, value in zip(
data["final"]["qkv"]["seeds"], data["final"]["v"]["seeds"]
):
assert qkv["seed"] == value["seed"]
q_hits = qkv["hits"]["Humanities"]["7shot"]
v_hits = value["hits"]["Humanities"]["7shot"]
v_only, q_only, p = exact_mcnemar(v_hits, q_hits)
paired.append(
{
"seed": qkv["seed"],
"v_right_qkv_wrong": v_only,
"v_wrong_qkv_right": q_only,
"delta_v_minus_qkv_pp": 100 * (sum(v_hits) - sum(q_hits)) / 250,
"p_exact_two_sided": p,
}
)
# The matched-rate control must use the declared common rate and retain
# the same base measurements.
assert matched["model"] == data["model"]
assert matched["matched_lr"] == 1e-4
assert matched["base"] == data["base"]
assert matched["base_hits"] == data["base_hits"]
assert round(rows["qkv"]["humanities_0shot_mean"], 2) == 62.67
assert round(rows["qkv"]["humanities_7shot_mean"], 2) == 60.00
assert round(rows["v"]["humanities_0shot_mean"], 2) == 62.13
assert round(rows["v"]["humanities_7shot_mean"], 2) == 60.27
payload = {
"source_sha256": {
"claim5.json": hashlib.sha256(args.input.read_bytes()).hexdigest(),
"claim5_matched.json": hashlib.sha256(
args.matched_input.read_bytes()
).hexdigest(),
},
"recomputed": rows,
"paired_v_vs_qkv_7shot": paired,
"matched_rate_control_checked": True,
"assertions_passed": 39,
}
rendered = json.dumps(payload, indent=2, sort_keys=True) + "\n"
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(rendered)
print(rendered, end="")
print(f"SHA256={hashlib.sha256(rendered.encode()).hexdigest()}")
if __name__ == "__main__":
main()