"""Assemble the p01 official-448 KD cache from validated rows and sealed raw logits.""" from __future__ import annotations from datetime import datetime, timezone 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 h5py import numpy as np from scripts.generate_kd_clean_cache import ( angles_from_vectors, angular_error_deg, create_h5, file_sha256, gaze_vector, roll_align_teacher_distribution, ) SOURCE_CACHE = ROOT / "data" / "processed_kd_clean_v1" / "cache" / "p01.full.h5" SOURCE_PROCESSING = ROOT / "data" / "processed_kd_clean_v1" / "cache" / "p01.full.processing.jsonl" LOGITS = ROOT / "artifacts" / "kd-teacher-trap-diagnostic" / "p01_teacher_protocol_448_logits.npz" STRICT_AUDIT = ROOT / "artifacts" / "kd-teacher-trap-diagnostic" / "strict_teacher_loader_audit.json" OUTPUT = ROOT / "data" / "processed_kd_clean_v1" / "cache" / "p01.official448_full.h5" OUTPUT_PROCESSING = ( ROOT / "data" / "processed_kd_clean_v1" / "cache" / "p01.official448_full.processing.jsonl" ) OUTPUT_SUMMARY = ROOT / "data" / "processed_kd_clean_v1" / "cache" / "p01.official448_full.summary.json" def decode(values: np.ndarray) -> list[str]: return [value.decode("utf-8") if isinstance(value, bytes) else str(value) for value in values] def main() -> None: for path in (OUTPUT, OUTPUT_PROCESSING, OUTPUT_SUMMARY): if path.exists(): raise FileExistsError(f"refusing to overwrite clean artifact: {path}") logits = np.load(LOGITS) strict_audit = json.loads(STRICT_AUDIT.read_text(encoding="utf-8")) if not strict_audit.get("gate_a_loader_pass"): raise RuntimeError("current strict-loader audit does not pass") with h5py.File(SOURCE_CACHE, "r") as source: sample_ids = decode(source["sample_id"][:]) paths = decode(source["relative_frame_path"][:]) if sample_ids != decode(logits["sample_id"]): raise RuntimeError("448 logits sample IDs do not exactly match validated cache rows") if paths != decode(logits["relative_frame_path"]): raise RuntimeError("448 logits paths do not exactly match validated cache rows") row_count = len(sample_ids) rows: list[dict] = [] text_fields = ("sample_id", "relative_frame_path", "participant", "day", "frame_id", "raw_image_sha256") scalar_fields = ("source_index", "annotation_row") copied_numeric = ( "left_patches", "right_patches", "landmarks", "left_gaze", "right_gaze", "left_affine_matrix", "right_affine_matrix", "left_roll_deg", "right_roll_deg", ) text_values = {field: decode(source[field][:]) for field in text_fields} scalar_values = {field: source[field][:] for field in scalar_fields} numeric_values = {field: source[field][:] for field in copied_numeric} raw_pitch = logits["fc_pitch_logits"].astype(np.float32) raw_yaw = logits["fc_yaw_logits"].astype(np.float32) for index in range(row_count): aligned_pitch, aligned_yaw, teacher_vector = roll_align_teacher_distribution( raw_pitch[index], raw_yaw[index], float(numeric_values["left_roll_deg"][index]) ) pitch_deg, yaw_deg = angles_from_vectors(teacher_vector[None, :]) left_gaze = numeric_values["left_gaze"][index] target_vector = gaze_vector(float(left_gaze[0]), float(left_gaze[1])) row = {field: text_values[field][index] for field in text_fields} row.update({field: scalar_values[field][index] for field in scalar_fields}) row.update({field: numeric_values[field][index] for field in copied_numeric}) row.update( { "teacher_pitch_logits_raw": raw_pitch[index], "teacher_yaw_logits_raw": raw_yaw[index], "teacher_pitch_logits": aligned_pitch, "teacher_yaw_logits": aligned_yaw, "teacher_aligned_vector": teacher_vector.astype(np.float32), "teacher_pitch_deg": np.float32(pitch_deg[0]), "teacher_yaw_deg": np.float32(yaw_deg[0]), "teacher_error_deg": np.float32(angular_error_deg(teacher_vector, target_vector)), } ) rows.append(row) attributes = { "schema": "mpiigaze-kd-clean-cache-v2-official448", "created_utc": datetime.now(timezone.utc).isoformat(), "participant": "p01", "partial": False, "source_manifest_sha256": source.attrs["source_manifest_sha256"], "teacher_loader_audit": strict_audit, "teacher_checkpoint_sha256": strict_audit["checkpoint_sha256"], "assembly_script_sha256": file_sha256(Path(__file__)), "source_preprocessing_cache_sha256": file_sha256(SOURCE_CACHE), "source_448_logits_sha256": file_sha256(LOGITS), "patch_size": 16, "teacher_input_size": 448, "teacher_bin_centers_deg": "index * 4 - 180", "training_target": "left-eye gaze rotated by left eye affine roll angle", "teacher_alignment": "independent pitch/yaw joint distribution rotated by same left-eye Z roll; marginals rebinned to 90 bins", "status": "KD_READY_PENDING_VALIDATION", } create_h5(OUTPUT, rows, attributes) OUTPUT_PROCESSING.write_bytes(SOURCE_PROCESSING.read_bytes()) summary = { "schema": "mpiigaze-kd-clean-cache-summary-v2-official448", "participant": "p01", "source_rows_examined": sum(1 for _ in SOURCE_PROCESSING.open("r", encoding="utf-8")), "accepted_rows": len(rows), "rejected_rows": sum(1 for line in SOURCE_PROCESSING.open("r", encoding="utf-8") if not json.loads(line)["accepted"]), "cache_path": str(OUTPUT.resolve()), "cache_sha256": file_sha256(OUTPUT), "processing_manifest_path": str(OUTPUT_PROCESSING.resolve()), "processing_manifest_sha256": file_sha256(OUTPUT_PROCESSING), "source_manifest_sha256": attributes["source_manifest_sha256"], "teacher_checkpoint_sha256": attributes["teacher_checkpoint_sha256"], "strict_loader_inference_sha256": strict_audit["inference_sha256"], "teacher_error_mean_deg": float(np.mean([row["teacher_error_deg"] for row in rows])), "teacher_error_median_deg": float(np.median([row["teacher_error_deg"] for row in rows])), } OUTPUT_SUMMARY.write_text(json.dumps(summary, indent=2) + "\n", encoding="utf-8", newline="\n") print(json.dumps(summary, indent=2)) if __name__ == "__main__": main()