Gaze-LIPE / scripts /teacher_shift_audit.py
thanhhuyvan's picture
Publish KD reproducibility investigation
178f61f
Raw
History Blame Contribute Delete
1.54 kB
"""Test whether a constant row offset materially improves cached teacher accuracy."""
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
OUT = ROOT / "artifacts" / "kd-teacher-trap-diagnostic"
def expectation(logits):
x = logits - logits.max(1, keepdims=True)
p = np.exp(x); p /= p.sum(1, keepdims=True)
return (p * np.arange(90)).sum(1) * 2 - 90
rows = []
for subject in ("p01", "p08", "p11"):
with h5py.File(ROOT / "data" / "processed" / f"{subject}_v16.h5", "r") as f:
pred = np.column_stack([expectation(f["teacher_pitch_logits"][:]), expectation(f["teacher_yaw_logits"][:])])
gt = np.rad2deg(f["left_gaze"][:])
for shift in range(-250, 251):
if shift >= 0:
p, g = pred[:len(pred)-shift or None], gt[shift:]
else:
p, g = pred[-shift:], gt[:len(gt)+shift]
mae = np.abs(p - g).mean()
rows.append({"subject": subject, "shift_gt_minus_teacher": shift, "overlap": len(p), "axis_mae_deg": mae})
df = pd.DataFrame(rows)
df.to_csv(OUT / "teacher_constant_shift_audit.csv", index=False)
for subject in ("p01", "p08", "p11"):
cur = df[df.subject == subject].sort_values("axis_mae_deg")
zero = float(cur[cur.shift_gt_minus_teacher == 0].axis_mae_deg.iloc[0])
best = cur.iloc[0]
print(subject, "zero", zero, "best", int(best.shift_gt_minus_teacher), float(best.axis_mae_deg), "gain", zero-float(best.axis_mae_deg))