Gaze-LIPE / scripts /audit_p01_teacher_protocol_448.py
thanhhuyvan's picture
Publish KD reproducibility investigation
178f61f
Raw
History Blame Contribute Delete
8.3 kB
"""Evaluate the recovered checkpoint under the documented 448px/4-degree protocol.
This is a protocol audit, not a KD-ready cache generator. It evaluates all explicit
head-order/sign candidates and reports the documented candidate separately.
"""
from __future__ import annotations
import hashlib
import json
from pathlib import Path
import sys
ROOT = Path(r"E:\Gaze_estimation")
sys.path.insert(0, str(ROOT / ".codex_deps"))
sys.path.insert(0, str(ROOT))
import cv2
import h5py
import numpy as np
import torch
from src.models.teacher_strict import file_sha256, load_teacher_model_strict
from src.utils.preprocess import GazePreprocessor
CLEAN = ROOT / "data" / "processed_kd_clean_v1" / "cache" / "p01.full.h5"
ORIGINAL = ROOT / "data" / "MPIIGaze" / "MPIIGaze" / "MPIIGaze" / "Data" / "Original"
LANDMARK_MODEL = ROOT / "src" / "utils" / "face_landmarker.task"
CHECKPOINT = ROOT / "checkpoints" / "resnet50.pt"
OUT = ROOT / "artifacts" / "kd-teacher-trap-diagnostic"
def decode(values: np.ndarray) -> list[str]:
return [value.decode("utf-8") if isinstance(value, bytes) else str(value) for value in values]
def face_crop(frame: np.ndarray, landmarks, size: int = 448) -> np.ndarray:
height, width = frame.shape[:2]
coordinates = np.array([[point.x * width, point.y * height] for point in landmarks])
minimum, maximum = coordinates.min(axis=0), coordinates.max(axis=0)
center = (minimum + maximum) / 2.0
extent = float(np.max(maximum - minimum) * 1.5)
x1, y1 = np.maximum(0, (center - extent / 2.0).astype(int))
x2, y2 = min(width, int(center[0] + extent / 2.0)), min(height, int(center[1] + extent / 2.0))
crop = frame[y1:y2, x1:x2]
if crop.size == 0:
raise RuntimeError("empty face crop")
crop = cv2.resize(crop, (size, size))
crop = cv2.cvtColor(crop, cv2.COLOR_BGR2RGB).astype(np.float32) / 255.0
crop = (crop - np.array([0.485, 0.456, 0.406], dtype=np.float32)) / np.array(
[0.229, 0.224, 0.225], dtype=np.float32
)
return np.transpose(crop, (2, 0, 1))
def softmax(values: np.ndarray) -> np.ndarray:
shifted = values.astype(np.float64) - values.max(axis=1, keepdims=True)
exp = np.exp(shifted)
return exp / exp.sum(axis=1, keepdims=True)
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 rotate_z(vectors_: np.ndarray, roll_deg: np.ndarray) -> np.ndarray:
angle = np.deg2rad(roll_deg)
cosine, sine = np.cos(angle), np.sin(angle)
output = vectors_.copy()
output[:, 0] = cosine * vectors_[:, 0] - sine * vectors_[:, 1]
output[:, 1] = sine * vectors_[:, 0] + cosine * vectors_[:, 1]
return output
def angular_error(prediction: np.ndarray, target: np.ndarray) -> np.ndarray:
prediction /= np.linalg.norm(prediction, axis=1, keepdims=True)
target /= np.linalg.norm(target, axis=1, keepdims=True)
return np.rad2deg(np.arccos(np.clip(np.sum(prediction * target, axis=1), -1.0, 1.0)))
def main() -> None:
with h5py.File(CLEAN, "r") as handle:
paths = decode(handle["relative_frame_path"][:])
sample_ids = decode(handle["sample_id"][:])
left_roll = handle["left_roll_deg"][:].astype(np.float64)
target_deg = np.rad2deg(handle["left_gaze"][:].astype(np.float64))
target_vector = vectors(target_deg[:, 0], target_deg[:, 1])
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model, loader_audit = load_teacher_model_strict(CHECKPOINT, device=device)
preprocessor = GazePreprocessor(model_path=str(LANDMARK_MODEL))
fc_pitch_batches: list[np.ndarray] = []
fc_yaw_batches: list[np.ndarray] = []
batch: list[np.ndarray] = []
batch_size = 16
def infer_batch() -> None:
if not batch:
return
tensor = torch.from_numpy(np.stack(batch)).to(device)
with torch.inference_mode():
fc_pitch, fc_yaw = model(tensor)
fc_pitch_batches.append(fc_pitch.detach().cpu().numpy().astype(np.float32))
fc_yaw_batches.append(fc_yaw.detach().cpu().numpy().astype(np.float32))
batch.clear()
for index, relative in enumerate(paths):
frame = cv2.imread(str(ORIGINAL / Path(relative)))
if frame is None:
raise RuntimeError(f"failed to decode mapped frame {relative}")
landmarks = preprocessor.get_landmarks(frame)
if landmarks is None:
raise RuntimeError(f"mapped accepted frame no longer produces landmarks: {relative}")
batch.append(face_crop(frame, landmarks, size=448))
if len(batch) == batch_size:
infer_batch()
if (index + 1) % 200 == 0:
print(f"protocol448 {index + 1}/{len(paths)}", flush=True)
infer_batch()
fc_pitch = np.concatenate(fc_pitch_batches)
fc_yaw = np.concatenate(fc_yaw_batches)
npz_path = OUT / "p01_teacher_protocol_448_logits.npz"
np.savez_compressed(
npz_path,
sample_id=np.asarray(sample_ids, dtype="S64"),
relative_frame_path=np.asarray(paths, dtype="S128"),
fc_pitch_logits=fc_pitch,
fc_yaw_logits=fc_yaw,
left_roll_deg=left_roll.astype(np.float32),
target_pitch_yaw_rad=np.deg2rad(target_deg).astype(np.float32),
)
probabilities = {"fc_pitch": softmax(fc_pitch), "fc_yaw": softmax(fc_yaw)}
rows: list[dict] = []
for binwidth, offset in ((2.0, 90.0), (4.0, 180.0)):
decoded = {
name: probability @ np.arange(90, dtype=np.float64) * binwidth - offset
for name, probability in probabilities.items()
}
for swap in (False, True):
pitch_base = decoded["fc_yaw"] if swap else decoded["fc_pitch"]
yaw_base = decoded["fc_pitch"] if swap else decoded["fc_yaw"]
for pitch_sign in (-1, 1):
for yaw_sign in (-1, 1):
raw_vector = vectors(pitch_sign * pitch_base, yaw_sign * yaw_base)
for apply_roll in (False, True):
prediction = rotate_z(raw_vector, left_roll) if apply_roll else raw_vector.copy()
error = angular_error(prediction, target_vector.copy())
rows.append(
{
"input_size": 448,
"binwidth_deg": binwidth,
"offset_deg": offset,
"swap_fc_pitch_fc_yaw": swap,
"pitch_sign": pitch_sign,
"yaw_sign": yaw_sign,
"apply_left_roll": apply_roll,
"mean_3d_error_deg": float(error.mean()),
"median_3d_error_deg": float(np.median(error)),
}
)
rows.sort(key=lambda row: row["mean_3d_error_deg"])
documented = next(
row for row in rows
if row["binwidth_deg"] == 4.0
and not row["swap_fc_pitch_fc_yaw"]
and row["pitch_sign"] == 1
and row["yaw_sign"] == 1
and row["apply_left_roll"]
)
report = {
"status": "PROTOCOL_AUDIT_NOT_KD_CACHE",
"participant": "p01",
"samples": len(paths),
"checkpoint_sha256": loader_audit.checkpoint_sha256,
"strict_loader_inference_sha256": loader_audit.inference_sha256,
"input_size": 448,
"logits_npz": str(npz_path),
"logits_npz_sha256": file_sha256(npz_path),
"documented_named-head_4deg_roll_candidate": documented,
"best_posthoc_candidate": rows[0],
"top_ten_candidates": rows[:10],
"warning": (
"Head order/sign search is post-hoc and cannot define a confirmatory protocol. "
"The fc_yaw/fc_pitch checkpoint lineage and training-code semantics must be fixed before KD."
),
}
output = OUT / "p01_teacher_protocol_448_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()