File size: 4,700 Bytes
178f61f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""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()