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