Download scripts/assemble_pointkd_cache.py from thanhhuyvan/Gaze-LIPE: direct link, hf CLI and curl.
- Browser
- Download file 6.47 kB
-
https://huggingface.co/thanhhuyvan/Gaze-LIPE/resolve/main/scripts/assemble_pointkd_cache.py
- Command line
-
hf download hf://thanhhuyvan/Gaze-LIPE/scripts/assemble_pointkd_cache.py
-
curl -L -o assemble_pointkd_cache.py https://huggingface.co/thanhhuyvan/Gaze-LIPE/resolve/main/scripts/assemble_pointkd_cache.py
6.47 kB
| """Create an unambiguous continuous-vector KD cache from official 448 raw logits.""" | |
| from __future__ import annotations | |
| import argparse | |
| from datetime import datetime, timezone | |
| import hashlib | |
| 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 | |
| CACHE_ROOT = ROOT / "data" / "processed_kd_clean_v1" / "cache" | |
| def file_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 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(values: np.ndarray, roll_deg: np.ndarray) -> np.ndarray: | |
| angle = np.deg2rad(roll_deg) | |
| cosine, sine = np.cos(angle), np.sin(angle) | |
| result = values.copy() | |
| result[:, 0] = cosine * values[:, 0] - sine * values[:, 1] | |
| result[:, 1] = sine * values[:, 0] + cosine * values[:, 1] | |
| return result | |
| def angles(values: np.ndarray) -> tuple[np.ndarray, np.ndarray]: | |
| pitch = np.arcsin(np.clip(-values[:, 1], -1.0, 1.0)) | |
| yaw = np.arctan2(-values[:, 0], -values[:, 2]) | |
| return np.rad2deg(pitch), np.rad2deg(yaw) | |
| def angular_error(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: | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("--participant", required=True) | |
| parser.add_argument("--source-tag", default="official448_full") | |
| parser.add_argument("--output-tag", default="official448_pointkd") | |
| args = parser.parse_args() | |
| source = CACHE_ROOT / f"{args.participant}.{args.source_tag}.h5" | |
| source_processing = CACHE_ROOT / f"{args.participant}.{args.source_tag}.processing.jsonl" | |
| output = CACHE_ROOT / f"{args.participant}.{args.output_tag}.h5" | |
| output_processing = CACHE_ROOT / f"{args.participant}.{args.output_tag}.processing.jsonl" | |
| output_summary = CACHE_ROOT / f"{args.participant}.{args.output_tag}.summary.json" | |
| for path in (output, output_processing, output_summary): | |
| if path.exists(): | |
| raise FileExistsError(f"refusing to overwrite point-KD artifact: {path}") | |
| copied_fields = ( | |
| "sample_id", "relative_frame_path", "participant", "day", "frame_id", | |
| "raw_image_sha256", "source_index", "annotation_row", "left_patches", | |
| "right_patches", "landmarks", "left_gaze", "right_gaze", | |
| "left_affine_matrix", "right_affine_matrix", "left_roll_deg", "right_roll_deg", | |
| "teacher_pitch_logits_raw", "teacher_yaw_logits_raw", | |
| ) | |
| with h5py.File(source, "r") as source_h5, h5py.File(output, "x") as output_h5: | |
| for field in copied_fields: | |
| source_h5.copy(field, output_h5) | |
| raw_pitch = source_h5["teacher_pitch_logits_raw"][:] | |
| raw_yaw = source_h5["teacher_yaw_logits_raw"][:] | |
| pitch = softmax(raw_pitch) @ np.arange(90) * 4.0 - 180.0 | |
| yaw = softmax(raw_yaw) @ np.arange(90) * 4.0 - 180.0 | |
| raw_vector = vectors(pitch, yaw) | |
| target_vector = rotate_z(raw_vector, source_h5["left_roll_deg"][:]) | |
| target_pitch, target_yaw = angles(target_vector) | |
| gaze_deg = np.rad2deg(source_h5["left_gaze"][:].astype(np.float64)) | |
| label_vector = vectors(gaze_deg[:, 0], gaze_deg[:, 1]) | |
| error = angular_error(target_vector, label_vector) | |
| output_h5.create_dataset("teacher_target_vector", data=target_vector.astype(np.float32), compression="gzip") | |
| output_h5.create_dataset("teacher_target_pitch_deg", data=target_pitch.astype(np.float32), compression="gzip") | |
| output_h5.create_dataset("teacher_target_yaw_deg", data=target_yaw.astype(np.float32), compression="gzip") | |
| output_h5.create_dataset("teacher_target_error_deg", data=error.astype(np.float32), compression="gzip") | |
| for key, value in source_h5.attrs.items(): | |
| output_h5.attrs[key] = value | |
| output_h5.attrs["schema"] = "mpiigaze-point-kd-cache-v3-official448" | |
| output_h5.attrs["created_utc"] = datetime.now(timezone.utc).isoformat() | |
| output_h5.attrs["source_official448_cache_sha256"] = file_sha256(source) | |
| output_h5.attrs["pointkd_assembly_script_sha256"] = file_sha256(Path(__file__)) | |
| output_h5.attrs["teacher_target_definition"] = ( | |
| "expectation of named 4-degree fc_pitch/fc_yaw logits converted to 3D, then rotated " | |
| "by the same left-eye Z roll used for the training label" | |
| ) | |
| output_h5.attrs["approved_distillation"] = "continuous 3D vector loss only" | |
| output_h5.attrs["prohibited_distillation"] = "KL on roll-rebinned marginal logits" | |
| output_h5.attrs["status"] = "POINT_KD_READY_PENDING_VALIDATION" | |
| output_processing.write_bytes(source_processing.read_bytes()) | |
| decisions = [json.loads(line) for line in output_processing.read_text(encoding="utf-8").splitlines()] | |
| summary = { | |
| "schema": "mpiigaze-point-kd-cache-summary-v3", | |
| "participant": args.participant, | |
| "source_rows_examined": len(decisions), | |
| "accepted_rows": int(sum(bool(row["accepted"]) for row in decisions)), | |
| "rejected_rows": int(sum(not bool(row["accepted"]) for row in decisions)), | |
| "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_official448_cache_sha256": file_sha256(source), | |
| "teacher_target_error_mean_deg": float(error.mean()), | |
| "teacher_target_error_median_deg": float(np.median(error)), | |
| } | |
| 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() | |