File size: 4,965 Bytes
8a5ffa8 | 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 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 | #!/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()
|