| """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() |
|
|