Gaze-LIPE / scripts /analyze_pointkd_gating.py
thanhhuyvan's picture
Publish KD reproducibility investigation
178f61f
Raw
History Blame Contribute Delete
4.7 kB
"""Compare point-KD teacher quality with the reproduced no-KD student by participant."""
from __future__ import annotations
import csv
import json
from pathlib import Path
import sys
ROOT = Path(r"E:\Gaze_estimation")
sys.path.insert(0, str(ROOT / ".codex_deps"))
import h5py
import numpy as np
import pandas as pd
CACHE_ROOT = ROOT / "data" / "processed_kd_clean_v1" / "cache"
SAMPLES = ROOT / "artifacts" / "kd-teacher-trap-diagnostic" / "per_sample_diagnostics.csv"
OUT = ROOT / "artifacts" / "kd-teacher-trap-diagnostic"
def vectors(pitch_deg: np.ndarray, yaw_deg: np.ndarray) -> np.ndarray:
pitch, yaw = np.deg2rad(pitch_deg), np.deg2rad(yaw_deg)
return np.column_stack(
(-np.cos(pitch) * np.sin(yaw), -np.sin(pitch), -np.cos(pitch) * np.cos(yaw))
)
def angular(first: np.ndarray, second: np.ndarray) -> np.ndarray:
first = first / np.linalg.norm(first, axis=1, keepdims=True)
second = second / np.linalg.norm(second, axis=1, keepdims=True)
return np.rad2deg(np.arccos(np.clip(np.sum(first * second, axis=1), -1.0, 1.0)))
def main() -> None:
samples = pd.read_csv(SAMPLES)
subject_reports = []
quartile_rows = []
available = sorted(path.name.split(".")[0] for path in CACHE_ROOT.glob("p??.official448_pointkd.h5"))
for subject in available:
cache_path = CACHE_ROOT / f"{subject}.official448_pointkd.h5"
with h5py.File(cache_path, "r") as cache:
teacher_vector = cache["teacher_target_vector"][:].astype(np.float64)
target_deg = np.rad2deg(cache["left_gaze"][:].astype(np.float64))
target_vector = vectors(target_deg[:, 0], target_deg[:, 1])
teacher_error = angular(teacher_vector, target_vector)
student = samples[(samples.subject == subject) & (samples.id == 5)].sort_values("row_index")
student_vector = vectors(student.student_pitch_deg.to_numpy(), student.student_yaw_deg.to_numpy())
student_error = angular(student_vector, target_vector)
if len(student_error) != len(teacher_error):
raise RuntimeError(f"row mismatch for {subject}")
quartiles = np.asarray(pd.qcut(teacher_error, 4, labels=False))
for quartile in range(4):
mask = quartiles == quartile
quartile_rows.append(
{
"subject": subject,
"teacher_error_quartile": quartile + 1,
"samples": int(mask.sum()),
"teacher_3d_error_mean_deg": float(teacher_error[mask].mean()),
"student_id5_3d_error_mean_deg": float(student_error[mask].mean()),
"teacher_minus_student_deg": float(
teacher_error[mask].mean() - student_error[mask].mean()
),
"teacher_better_fraction": float(
(teacher_error[mask] < student_error[mask]).mean()
),
}
)
subject_reports.append(
{
"subject": subject,
"samples": len(teacher_error),
"teacher_3d_error_mean_deg": float(teacher_error.mean()),
"student_id5_3d_error_mean_deg": float(student_error.mean()),
"teacher_minus_student_deg": float(teacher_error.mean() - student_error.mean()),
"teacher_better_fraction": float((teacher_error < student_error).mean()),
"teacher_student_error_pearson": float(np.corrcoef(teacher_error, student_error)[0, 1]),
}
)
csv_path = OUT / "pointkd_gating_quartiles.csv"
with csv_path.open("w", encoding="utf-8", newline="") as stream:
writer = csv.DictWriter(stream, fieldnames=quartile_rows[0].keys())
writer.writeheader()
writer.writerows(quartile_rows)
report = {
"teacher_protocol": (
"strict 322/322 load; 448px; named heads; 4-degree bins; expected 3D vector "
"rotated into left-eye target coordinates"
),
"subjects": subject_reports,
"quartiles": quartile_rows,
"standard_kd_decision": "FAIL: teacher is globally worse than the reproduced no-KD student on every audited subject",
"quality_gated_kd_decision": (
"PREFLIGHT_JUSTIFIED_ONLY: low-error teacher strata may help, but gate thresholds "
"must be selected using training participants and evaluated on held-out participants."
),
}
output = OUT / "pointkd_gating_audit.json"
output.write_text(json.dumps(report, indent=2) + "\n", encoding="utf-8", newline="\n")
print(json.dumps(report, indent=2))
if __name__ == "__main__":
main()