Gaze-LIPE / scripts /analyze_matched_pointkd_runs.py
thanhhuyvan's picture
Publish KD reproducibility investigation
178f61f
Raw
History Blame Contribute Delete
7.03 kB
"""Aggregate complete matched point-KD triplets without selecting best epochs."""
from __future__ import annotations
import csv
import hashlib
import itertools
import json
from pathlib import Path
import numpy as np
ROOT = Path(r"E:\Gaze_estimation")
RUN_ROOT = ROOT / "artifacts" / "kd-teacher-trap-diagnostic" / "matched-runs-v1"
OUT = ROOT / "artifacts" / "kd-teacher-trap-diagnostic"
ARMS = ("control", "point_kd", "quality_gated_point_kd", "shuffled_teacher_kd")
def sha256(path: Path):
value = hashlib.sha256()
with path.open("rb") as stream:
for block in iter(lambda: stream.read(1024 * 1024), b""):
value.update(block)
return value.hexdigest().upper()
def read_run(path: Path):
summary_path, config_path, epochs_path = path / "summary.json", path / "config.json", path / "epochs.jsonl"
summary = json.loads(summary_path.read_text(encoding="utf-8"))
config = json.loads(config_path.read_text(encoding="utf-8"))
if summary["config_sha256"] != sha256(config_path) or summary["epoch_log_sha256"] != sha256(epochs_path):
raise RuntimeError(f"run hashes changed: {path}")
epochs = [json.loads(line) for line in epochs_path.read_text(encoding="utf-8").splitlines()]
if epochs[-1] != summary["primary_result"]:
raise RuntimeError(f"primary result is not final epoch: {path}")
return summary, config, epochs
def sign_flip(values, seed=20260809):
values = np.asarray(values, dtype=float)
observed = abs(values.mean())
if len(values) <= 20:
outcomes = [abs((values * np.asarray(signs)).mean()) for signs in itertools.product((-1, 1), repeat=len(values))]
return float(np.mean(np.asarray(outcomes) >= observed)), "exact"
rng = np.random.default_rng(seed)
outcomes = np.abs((rng.choice((-1, 1), size=(200000, len(values))) * values).mean(axis=1))
return float((np.sum(outcomes >= observed) + 1) / (len(outcomes) + 1)), "monte_carlo_200000"
def contrast(rows, arm, field=None, interpretation="negative favors arm; positive favors control"):
field = field or f"{arm}_minus_control_deg"
differences = np.asarray([row[field] for row in rows])
by_participant = {}
for row in rows: by_participant.setdefault(row["held_out"], []).append(row[field])
participant_values = np.asarray([np.mean(value) for value in by_participant.values()])
rng = np.random.default_rng(20260809)
keys = list(by_participant)
boot = []
for _ in range(20000):
selected = rng.choice(keys, size=len(keys), replace=True)
boot.append(np.mean([np.mean(by_participant[key]) for key in selected]))
p_value, method = sign_flip(participant_values)
return {
"arm": arm,
"fold_seed_pairs": len(rows),
"participants": len(keys),
"mean_paired_difference_deg": float(differences.mean()),
"median_paired_difference_deg": float(np.median(differences)),
"participant_cluster_bootstrap_95ci_deg": [float(x) for x in np.quantile(boot, (0.025, 0.975))],
"participant_sign_flip_p_two_sided": p_value,
"participant_sign_flip_method": method,
"interpretation": interpretation,
}
def main():
grouped = {}
for summary_path in RUN_ROOT.glob("p??/seed-*/?*/summary.json"):
run_dir = summary_path.parent
summary, config, epochs = read_run(run_dir)
key = (summary["held_out"], int(summary["seed"]))
grouped.setdefault(key, {})[summary["arm"]] = (summary, config, epochs)
complete = {key: value for key, value in grouped.items() if set(value) == set(ARMS)}
if not complete:
raise RuntimeError("no complete control/point-KD/gated-point-KD triplets")
rows = []
for (held_out, seed), runs in sorted(complete.items()):
configs = [runs[arm][1] for arm in ARMS]
for field in ("initial_state_sha256", "sampler_orders_sha256", "shuffled_teacher_permutation_sha256", "effective_lambda_kd", "gate_tau_good_deg", "gate_tau_bad_deg", "epochs", "lr", "batch_size"):
if len({json.dumps(config[field], sort_keys=True) for config in configs}) != 1:
raise RuntimeError(f"unmatched {field} for {held_out}/seed-{seed}")
control = runs["control"][0]["primary_result"]["student_3d_error_mean_deg"]
point = runs["point_kd"][0]["primary_result"]["student_3d_error_mean_deg"]
gated = runs["quality_gated_point_kd"][0]["primary_result"]["student_3d_error_mean_deg"]
shuffled = runs["shuffled_teacher_kd"][0]["primary_result"]["student_3d_error_mean_deg"]
rows.append({
"held_out": held_out, "seed": seed,
"control_final_3d_error_deg": control,
"point_kd_final_3d_error_deg": point,
"quality_gated_point_kd_final_3d_error_deg": gated,
"shuffled_teacher_kd_final_3d_error_deg": shuffled,
"point_kd_minus_control_deg": point - control,
"quality_gated_point_kd_minus_control_deg": gated - control,
"shuffled_teacher_kd_minus_control_deg": shuffled - control,
"point_kd_minus_shuffled_teacher_kd_deg": point - shuffled,
"heldout_teacher_3d_error_deg": runs["control"][0]["primary_result"]["teacher_3d_error_mean_deg"],
"initial_state_sha256": configs[0]["initial_state_sha256"],
"sampler_orders_sha256": configs[0]["sampler_orders_sha256"],
"effective_lambda_kd": configs[0]["effective_lambda_kd"],
})
csv_path = OUT / "matched_pointkd_primary_results.csv"
with csv_path.open("w", encoding="utf-8", newline="") as stream:
writer = csv.DictWriter(stream, fieldnames=rows[0].keys()); writer.writeheader(); writer.writerows(rows)
participants = sorted({row["held_out"] for row in rows})
report = {
"schema": "matched-pointkd-analysis-v1",
"primary_epoch_policy": "fixed final epoch; no held-out early stopping",
"participants": participants,
"seeds": sorted({row["seed"] for row in rows}),
"complete_triplets": len(rows),
"study_status": "CONFIRMATORY_COMPLETE" if len(participants) == 15 and min(sum(row["held_out"] == p for row in rows) for p in participants) >= 3 else "PREFLIGHT_ONLY",
"contrasts": [
contrast(rows, "point_kd"),
contrast(rows, "quality_gated_point_kd"),
contrast(rows, "shuffled_teacher_kd"),
contrast(
rows,
"point_kd_vs_shuffled_teacher_kd",
field="point_kd_minus_shuffled_teacher_kd_deg",
interpretation="negative favors real teacher correspondence; positive favors shuffled teacher",
),
],
"primary_results_csv": str(csv_path.resolve()),
"primary_results_csv_sha256": sha256(csv_path),
}
output = OUT / "matched_pointkd_analysis.json"
output.write_text(json.dumps(report, indent=2) + "\n", encoding="utf-8")
print(json.dumps(report, indent=2))
if __name__ == "__main__":
main()