File size: 6,671 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
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
"""Read-only go/no-go audit for legacy caches before a clean KD experiment.

This script never edits H5 data or checkpoints. It writes one diagnostic JSON report.
"""
from __future__ import annotations

import hashlib
import json
import math
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


SUBJECTS = ("p01", "p08", "p11")
OUT_DIR = ROOT / "artifacts" / "kd-teacher-trap-diagnostic"
LOADER_AUDIT = OUT_DIR / "teacher_loader_key_audit.json"
OUTPUT = OUT_DIR / "clean_kd_gate_audit.json"

IDENTITY_FIELDS = (
    "sample_id",
    "relative_frame_path",
    "participant",
    "day",
    "frame_id",
    "annotation_row",
    "raw_image_sha256",
)
MODEL_FIELDS = (
    "left_patches",
    "right_patches",
    "landmarks",
    "left_gaze",
    "right_gaze",
    "teacher_pitch_logits",
    "teacher_yaw_logits",
)


def sha256(path: Path) -> str:
    digest = hashlib.sha256()
    with path.open("rb") as stream:
        for block in iter(lambda: stream.read(1024 * 1024), b""):
            digest.update(block)
    return digest.hexdigest().upper()


def expectation_deg(logits: np.ndarray) -> np.ndarray:
    shifted = logits.astype(np.float64) - logits.max(axis=1, keepdims=True)
    probability = np.exp(shifted)
    probability /= probability.sum(axis=1, keepdims=True)
    return (probability * np.arange(90, dtype=np.float64)).sum(axis=1) * 2.0 - 90.0


def gaze_vectors(pitch_deg: np.ndarray, yaw_deg: np.ndarray) -> np.ndarray:
    pitch = np.deg2rad(pitch_deg)
    yaw = 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_error_deg(prediction: np.ndarray, target: np.ndarray) -> np.ndarray:
    dots = np.sum(prediction * target, axis=1)
    return np.rad2deg(np.arccos(np.clip(dots, -1.0, 1.0)))


def audit_h5(subject: str) -> dict:
    path = ROOT / "data" / "processed" / f"{subject}_v16.h5"
    result = {
        "subject": subject,
        "path": str(path),
        "sha256": sha256(path),
    }
    with h5py.File(path, "r") as handle:
        keys = sorted(handle.keys())
        result["datasets"] = keys
        result["missing_identity_fields"] = [name for name in IDENTITY_FIELDS if name not in handle]
        result["missing_model_fields"] = [name for name in MODEL_FIELDS if name not in handle]
        counts = {name: int(handle[name].shape[0]) for name in MODEL_FIELDS if name in handle}
        result["row_counts"] = counts
        result["row_counts_equal"] = len(set(counts.values())) == 1
        result["teacher_logit_shapes_valid"] = all(
            name in handle and handle[name].ndim == 2 and handle[name].shape[1] == 90
            for name in ("teacher_pitch_logits", "teacher_yaw_logits")
        )
        finite = {}
        for name in MODEL_FIELDS:
            if name in handle:
                finite[name] = bool(np.isfinite(handle[name][:]).all())
        result["all_numeric_values_finite"] = bool(finite) and all(finite.values())
        result["finite_by_dataset"] = finite

        if not result["missing_model_fields"] and result["teacher_logit_shapes_valid"]:
            pitch = expectation_deg(handle["teacher_pitch_logits"][:])
            yaw = expectation_deg(handle["teacher_yaw_logits"][:])
            target_deg = np.rad2deg(handle["left_gaze"][:].astype(np.float64))
            teacher_vec = gaze_vectors(pitch, yaw)
            target_vec = gaze_vectors(target_deg[:, 0], target_deg[:, 1])
            error_3d = angular_error_deg(teacher_vec, target_vec)
            axis_error = np.abs(np.column_stack((pitch, yaw)) - target_deg).mean(axis=1)
            result["samples"] = int(len(error_3d))
            result["teacher_axis_mae_deg"] = float(axis_error.mean())
            result["teacher_3d_angular_error_deg"] = float(error_3d.mean())
            result["teacher_3d_error_median_deg"] = float(np.median(error_3d))
            result["teacher_3d_error_p90_deg"] = float(np.quantile(error_3d, 0.90))

    result["cache_integrity_pass"] = bool(
        not result["missing_identity_fields"]
        and not result["missing_model_fields"]
        and result["row_counts_equal"]
        and result["teacher_logit_shapes_valid"]
        and result["all_numeric_values_finite"]
    )
    return result


def main() -> None:
    loader = json.loads(LOADER_AUDIT.read_text(encoding="utf-8"))
    loader_pass = bool(
        loader["historical_loader_loaded_key_count"] == loader["checkpoint_tensor_keys"]
        and loader["historical_loader_unexpected_key_count"] == 0
        and set(loader["historical_loader_missing_keys"]) <= {"idx_tensor"}
    )
    caches = [audit_h5(subject) for subject in SUBJECTS]
    cache_pass = all(item["cache_integrity_pass"] for item in caches)
    teacher_errors = [item.get("teacher_3d_angular_error_deg", math.nan) for item in caches]
    report = {
        "audit_type": "read-only clean-KD execution gate",
        "legacy_files_modified": False,
        "gate_a_loader": {
            "pass": loader_pass,
            "loaded_checkpoint_tensors": loader["historical_loader_loaded_key_count"],
            "checkpoint_tensors": loader["checkpoint_tensor_keys"],
            "unexpected_checkpoint_tensors": loader["historical_loader_unexpected_key_count"],
            "reason": "Recovered historical loader does not strictly load the trained backbone."
            if not loader_pass
            else "Strict loading requirements satisfied.",
        },
        "gate_b_cache_integrity": {
            "pass": cache_pass,
            "required_identity_fields": list(IDENTITY_FIELDS),
            "subjects": caches,
            "reason": "Legacy caches have no row-level source identity/provenance fields."
            if not cache_pass
            else "Row-level provenance and numeric integrity requirements satisfied.",
        },
        "gate_c_teacher_quality": {
            "pass": None,
            "status": "DESCRIPTIVE_ONLY",
            "mean_of_subjects_teacher_3d_angular_error_deg": float(np.nanmean(teacher_errors)),
            "reason": "A matched clean control is required to apply the preregistered relative-quality gate.",
        },
        "decision": "STOP_BEFORE_TRAINING" if not (loader_pass and cache_pass) else "READY_FOR_TEACHER_QUALITY_GATE",
        "next_action": "Create a new strict-loader, manifest-backed cache; do not repair or overwrite legacy H5 files.",
    }
    OUTPUT.write_text(json.dumps(report, indent=2) + "\n", encoding="utf-8")
    print(json.dumps(report, indent=2))


if __name__ == "__main__":
    main()