| """Matched, leakage-resistant control versus point-KD experiment. |
| |
| This runner only accepts validated ``official448_pointkd`` caches. It never opens |
| or modifies legacy V16 files. A fold/seed shares one serialized initialization, |
| one deterministic sampler schedule, and deterministic per-sample augmentations |
| across all arms. The fixed final epoch is the primary result; the held-out |
| participant is never used for early stopping or gate selection. |
| """ |
| from __future__ import annotations |
|
|
| import argparse |
| import hashlib |
| import json |
| import math |
| import os |
| import random |
| import sys |
| from datetime import datetime, timezone |
| from pathlib import Path |
|
|
| ROOT = Path(r"E:\Gaze_estimation") |
| sys.path.insert(0, str(ROOT / ".codex_deps")) |
| sys.path.insert(0, str(ROOT)) |
|
|
| import cv2 |
| import h5py |
| import numpy as np |
| import torch |
| import torch.nn as nn |
| from torch.utils.data import DataLoader, Dataset, Sampler |
|
|
|
|
| CACHE_ROOT = ROOT / "data" / "processed_kd_clean_v1" / "cache" |
| RUN_ROOT = ROOT / "artifacts" / "kd-teacher-trap-diagnostic" / "matched-runs-v1" |
| PARTICIPANTS = tuple(f"p{i:02d}" for i in range(15)) |
| ARMS = ("control", "point_kd", "quality_gated_point_kd", "shuffled_teacher_kd") |
|
|
|
|
| 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 json_sha256(value) -> str: |
| return hashlib.sha256(json.dumps(value, sort_keys=True, separators=(",", ":")).encode()).hexdigest().upper() |
|
|
|
|
| def seed_everything(seed: int) -> None: |
| random.seed(seed) |
| np.random.seed(seed) |
| torch.manual_seed(seed) |
| torch.use_deterministic_algorithms(True) |
|
|
|
|
| class SmoothAWLoss(nn.Module): |
| def __init__(self, omega=8.0, alpha=1.5, theta=0.5, epsilon=1.0): |
| super().__init__() |
| self.omega, self.alpha, self.theta, self.epsilon = omega, alpha, theta, epsilon |
|
|
| def forward(self, prediction, target): |
| delta = (target - prediction).abs() |
| theta_eps = torch.as_tensor(self.theta / self.epsilon, device=prediction.device) |
| a = self.omega / (1.0 + theta_eps.pow(self.alpha)) |
| a = a * self.alpha * theta_eps.pow(self.alpha - 1.0) / self.epsilon |
| b = a * self.theta - self.omega * torch.log1p(theta_eps.pow(self.alpha)) |
| return torch.where( |
| delta < self.theta, |
| self.omega * torch.log1p((delta / self.epsilon).pow(self.alpha)), |
| a * delta - b, |
| ).mean() |
|
|
|
|
| class FlexibleMiniConv(nn.Module): |
| def __init__(self): |
| super().__init__() |
| self.conv = nn.Sequential( |
| nn.Conv2d(1, 16, 3), nn.ReLU(inplace=True), |
| nn.Conv2d(16, 32, 3), nn.ReLU(inplace=True), |
| nn.Conv2d(32, 64, 3), nn.ReLU(inplace=True), |
| ) |
| self.avg_pool = nn.AdaptiveAvgPool2d(1) |
| self.max_pool = nn.AdaptiveMaxPool2d(1) |
|
|
| def forward(self, value): |
| value = self.conv(value) |
| return torch.cat((self.avg_pool(value), self.max_pool(value)), dim=1).flatten(1) |
|
|
|
|
| class MatchedStudent(nn.Module): |
| """The legacy ID5/ID8 dual-pool ablation architecture, frozen locally.""" |
| def __init__(self): |
| super().__init__() |
| self.app_net = FlexibleMiniConv() |
| self.geo_net = nn.Sequential( |
| nn.Linear(956, 256), nn.LayerNorm(256), nn.ReLU(inplace=True), |
| nn.Linear(256, 256), nn.ReLU(inplace=True), |
| ) |
| self.post_concat_bn = nn.BatchNorm1d(768) |
| self.fusion = nn.Sequential( |
| nn.Linear(768, 256), nn.ReLU(inplace=True), nn.Dropout(0.1), |
| nn.Linear(256, 128), nn.ReLU(inplace=True), |
| ) |
| self.pitch_head = nn.Linear(128, 90) |
| self.yaw_head = nn.Linear(128, 90) |
|
|
| def forward(self, patches, landmarks): |
| batch = patches.shape[0] |
| appearance = self.app_net(patches.reshape(-1, 1, patches.shape[2], patches.shape[3])).reshape(batch, -1) |
| geometry = self.geo_net(landmarks) |
| fused = self.fusion(self.post_concat_bn(torch.cat((appearance, geometry), dim=1))) |
| return self.pitch_head(fused), self.yaw_head(fused) |
|
|
|
|
| def angles_to_vectors(angles_deg: torch.Tensor) -> torch.Tensor: |
| pitch, yaw = torch.deg2rad(angles_deg[:, 0]), torch.deg2rad(angles_deg[:, 1]) |
| return torch.stack((-torch.cos(pitch) * torch.sin(yaw), -torch.sin(pitch), -torch.cos(pitch) * torch.cos(yaw)), dim=1) |
|
|
|
|
| def logits_to_angles(pitch_logits, yaw_logits): |
| bins = torch.arange(90, dtype=pitch_logits.dtype, device=pitch_logits.device) |
| pitch = (pitch_logits.softmax(1) * bins).sum(1) * 2.0 - 90.0 |
| yaw = (yaw_logits.softmax(1) * bins).sum(1) * 2.0 - 90.0 |
| return torch.stack((pitch, yaw), dim=1) |
|
|
|
|
| def angular_error(first, second): |
| first = nn.functional.normalize(first, dim=1) |
| second = nn.functional.normalize(second, dim=1) |
| return torch.rad2deg(torch.acos((first * second).sum(1).clamp(-1.0 + 1e-7, 1.0 - 1e-7))) |
|
|
|
|
| def correlation(first, second): |
| return float(np.corrcoef(first, second)[0, 1]) |
|
|
|
|
| def spearman(first, second): |
| |
| first_rank = np.argsort(np.argsort(first, kind="mergesort"), kind="mergesort") |
| second_rank = np.argsort(np.argsort(second, kind="mergesort"), kind="mergesort") |
| return correlation(first_rank, second_rank) |
|
|
|
|
| def deterministic_hardening(patch, landmarks, key): |
| rng = np.random.RandomState(key & 0xFFFFFFFF) |
| patch, landmarks = patch.copy(), landmarks.copy() |
| if rng.rand() > 0.5: |
| for index in range(4): |
| small = cv2.resize(patch[index], (8, 8), interpolation=cv2.INTER_CUBIC) |
| patch[index] = cv2.resize(small, (patch.shape[2], patch.shape[1]), interpolation=cv2.INTER_CUBIC) |
| for index in range(4): |
| patch[index] = cv2.bilateralFilter(patch[index], 5, 20, 20) |
| clahe = cv2.createCLAHE(clipLimit=1.1, tileGridSize=(4, 4)) |
| for index in range(4): |
| patch[index] = clahe.apply(patch[index]) |
| if rng.rand() > 0.5: |
| patch = np.clip(patch.astype(np.float32) + rng.normal(0, 3, patch.shape), 0, 255).astype(np.uint8) |
| if rng.rand() > 0.5: |
| landmarks += rng.normal(0, 0.003, landmarks.shape).astype(np.float32) |
| return patch, landmarks |
|
|
|
|
| class PointKDDataset(Dataset): |
| def __init__(self, paths, augment, seed): |
| self.paths = tuple(map(str, paths)) |
| self.augment, self.seed, self.epoch = augment, seed, 0 |
| self.teacher_index_map = None |
| self.handles = {} |
| self.index = [] |
| for file_index, path in enumerate(self.paths): |
| with h5py.File(path, "r") as handle: |
| self.index.extend((file_index, row) for row in range(len(handle["left_gaze"]))) |
|
|
| def __len__(self): |
| return len(self.index) |
|
|
| def close(self): |
| for handle in self.handles.values(): |
| handle.close() |
| self.handles.clear() |
|
|
| def __getitem__(self, global_index): |
| file_index, row = self.index[global_index] |
| if file_index not in self.handles: |
| self.handles[file_index] = h5py.File(self.paths[file_index], "r") |
| handle = self.handles[file_index] |
| teacher_file_index, teacher_row = file_index, row |
| if self.teacher_index_map is not None: |
| teacher_file_index, teacher_row = self.index[self.teacher_index_map[global_index]] |
| if teacher_file_index not in self.handles: |
| self.handles[teacher_file_index] = h5py.File(self.paths[teacher_file_index], "r") |
| teacher_handle = self.handles[teacher_file_index] |
| patch, landmarks = handle["left_patches"][row], handle["landmarks"][row] |
| if self.augment: |
| key = self.seed * 1_000_003 + self.epoch * 100_003 + global_index |
| patch, landmarks = deterministic_hardening(patch, landmarks, key) |
| return ( |
| torch.from_numpy(patch).float() / 255.0, |
| torch.from_numpy(landmarks).float().reshape(-1), |
| torch.from_numpy(handle["left_gaze"][row]).float() * (180.0 / math.pi), |
| torch.from_numpy(teacher_handle["teacher_target_vector"][teacher_row]).float(), |
| torch.as_tensor(handle["teacher_target_error_deg"][row], dtype=torch.float32), |
| ) |
|
|
|
|
| def within_participant_derangement(dataset, seed): |
| """Map each row to a different teacher row from the same participant/cache.""" |
| mapping = np.arange(len(dataset), dtype=np.int64) |
| rng = np.random.RandomState(seed & 0xFFFFFFFF) |
| for file_index in range(len(dataset.paths)): |
| indices = np.asarray([index for index, pair in enumerate(dataset.index) if pair[0] == file_index]) |
| if len(indices) < 2: |
| raise RuntimeError(f"cannot derange participant cache with {len(indices)} row(s)") |
| candidate = indices.copy() |
| while True: |
| rng.shuffle(candidate) |
| if np.all(candidate != indices): |
| break |
| mapping[indices] = candidate |
| if np.any(mapping == np.arange(len(dataset))): |
| raise RuntimeError("shuffled-teacher mapping contains a fixed point") |
| return mapping.tolist() |
|
|
|
|
| class FixedOrderSampler(Sampler): |
| def __init__(self, order): self.order = order |
| def __iter__(self): return iter(self.order) |
| def __len__(self): return len(self.order) |
|
|
|
|
| def validated_cache(participant): |
| cache = CACHE_ROOT / f"{participant}.official448_pointkd.h5" |
| validation = CACHE_ROOT / f"{participant}.official448_pointkd.validation.json" |
| if not cache.is_file() or not validation.is_file(): |
| raise FileNotFoundError(f"missing cache or validation for {participant}") |
| report = json.loads(validation.read_text(encoding="utf-8")) |
| if not report.get("pass") or report.get("cache_sha256") != sha256(cache): |
| raise RuntimeError(f"invalid or changed point-KD cache for {participant}") |
| return cache |
|
|
|
|
| def gradient_cosine(model, hard, kd): |
| hard_grad = torch.autograd.grad(hard, model.parameters(), retain_graph=True, allow_unused=True) |
| kd_grad = torch.autograd.grad(kd, model.parameters(), retain_graph=True, allow_unused=True) |
| pairs = [(a.reshape(-1), b.reshape(-1)) for a, b in zip(hard_grad, kd_grad) if a is not None and b is not None] |
| if not pairs: return float("nan") |
| a, b = torch.cat([x for x, _ in pairs]), torch.cat([y for _, y in pairs]) |
| return float((torch.dot(a, b) / (a.norm() * b.norm()).clamp_min(1e-12)).detach().cpu()) |
|
|
|
|
| def gradient_norms(model, hard, kd): |
| hard_grad = torch.autograd.grad(hard, model.parameters(), retain_graph=True, allow_unused=True) |
| kd_grad = torch.autograd.grad(kd, model.parameters(), allow_unused=True) |
| hard_norm = torch.sqrt(sum((value * value).sum() for value in hard_grad if value is not None)) |
| kd_norm = torch.sqrt(sum((value * value).sum() for value in kd_grad if value is not None)) |
| return float(hard_norm.detach().cpu()), float(kd_norm.detach().cpu()) |
|
|
|
|
| def evaluate(model, loader, device): |
| model.eval(); values = {key: [] for key in ("error", "axis", "disagreement", "teacher_error", "pitch_s", "pitch_t", "yaw_s", "yaw_t")} |
| with torch.inference_mode(): |
| for patch, landmarks, target_angles, teacher_vector, teacher_error in loader: |
| patch, landmarks = patch.to(device), landmarks.to(device) |
| target_angles, teacher_vector = target_angles.to(device), teacher_vector.to(device) |
| prediction = logits_to_angles(*model(patch, landmarks)) |
| student_vector, target_vector = angles_to_vectors(prediction), angles_to_vectors(target_angles) |
| teacher_angles = torch.stack((torch.rad2deg(torch.asin((-teacher_vector[:, 1]).clamp(-1, 1))), torch.rad2deg(torch.atan2(-teacher_vector[:, 0], -teacher_vector[:, 2]))), 1) |
| batch_values = { |
| "error": angular_error(student_vector, target_vector), |
| "axis": (prediction - target_angles).abs().mean(1), |
| "disagreement": angular_error(student_vector, teacher_vector), |
| "teacher_error": teacher_error, |
| "pitch_s": prediction[:, 0], "pitch_t": teacher_angles[:, 0], |
| "yaw_s": prediction[:, 1], "yaw_t": teacher_angles[:, 1], |
| } |
| for key, value in batch_values.items(): values[key].append(value.cpu().numpy()) |
| values = {key: np.concatenate(value) for key, value in values.items()} |
| return { |
| "student_3d_error_mean_deg": float(values["error"].mean()), |
| "student_axis_mae_deg": float(values["axis"].mean()), |
| "teacher_3d_error_mean_deg": float(values["teacher_error"].mean()), |
| "student_teacher_disagreement_mean_deg": float(values["disagreement"].mean()), |
| "pitch_pearson": correlation(values["pitch_s"], values["pitch_t"]), |
| "yaw_pearson": correlation(values["yaw_s"], values["yaw_t"]), |
| "pitch_spearman": spearman(values["pitch_s"], values["pitch_t"]), |
| "yaw_spearman": spearman(values["yaw_s"], values["yaw_t"]), |
| } |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser() |
| parser.add_argument("--held-out", required=True, choices=PARTICIPANTS) |
| parser.add_argument("--arm", required=True, choices=ARMS) |
| parser.add_argument("--seed", required=True, type=int) |
| parser.add_argument("--epochs", type=int, default=50) |
| parser.add_argument("--batch-size", type=int, default=32) |
| parser.add_argument("--lr", type=float, default=1e-4) |
| parser.add_argument("--lambda-kd", type=float, help="fixed override; default calibrates on training data") |
| parser.add_argument("--kd-gradient-ratio", type=float, default=0.5, |
| help="target ||lambda*grad(K)||/||grad(H)|| at initialization") |
| parser.add_argument("--workers", type=int, default=4, choices=range(9)) |
| parser.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu") |
| args = parser.parse_args() |
| seed_everything(args.seed) |
| train_participants = [p for p in PARTICIPANTS if p != args.held_out] |
| train_paths = [validated_cache(p) for p in train_participants] |
| heldout_path = validated_cache(args.held_out) |
| config = vars(args) | {"train_participants": train_participants, "primary_checkpoint": "fixed_final_epoch"} |
| run_dir = RUN_ROOT / args.held_out / f"seed-{args.seed}" / args.arm |
| run_dir.mkdir(parents=True, exist_ok=False) |
| common_dir = run_dir.parent / "common" |
| common_dir.mkdir(exist_ok=True) |
|
|
| train_data = PointKDDataset(train_paths, augment=True, seed=args.seed) |
| heldout_data = PointKDDataset([heldout_path], augment=False, seed=args.seed) |
| teacher_errors = [] |
| for path in train_paths: |
| with h5py.File(path, "r") as handle: teacher_errors.append(handle["teacher_target_error_deg"][:]) |
| teacher_errors = np.concatenate(teacher_errors) |
| tau_good, tau_bad = map(float, np.quantile(teacher_errors, (0.25, 0.75))) |
| config["gate_tau_good_deg"], config["gate_tau_bad_deg"] = tau_good, tau_bad |
| config["cache_sha256"] = {p: sha256(path) for p, path in zip(train_participants, train_paths)} | {args.held_out: sha256(heldout_path)} |
| config["runner_sha256"] = sha256(Path(__file__)) |
| (run_dir / "config.json").write_text(json.dumps(config, indent=2) + "\n", encoding="utf-8") |
|
|
| initial_path = common_dir / "initial_state.pt" |
| if not initial_path.exists(): |
| model = MatchedStudent(); torch.save(model.state_dict(), initial_path) |
| model = MatchedStudent().to(args.device) |
| model.load_state_dict(torch.load(initial_path, map_location=args.device, weights_only=True), strict=True) |
| config["initial_state_sha256"] = sha256(initial_path) |
|
|
| orders_path = common_dir / "sampler_orders.json" |
| orders = [torch.randperm(len(train_data), generator=torch.Generator().manual_seed(args.seed * 1009 + epoch)).tolist() for epoch in range(args.epochs)] |
| order_hash = json_sha256(orders) |
| if orders_path.exists() and json.loads(orders_path.read_text())["sha256"] != order_hash: |
| raise RuntimeError("shared sampler schedule mismatch") |
| if not orders_path.exists(): orders_path.write_text(json.dumps({"sha256": order_hash, "orders": orders}) + "\n", encoding="utf-8") |
| config["sampler_orders_sha256"] = order_hash |
|
|
| shuffled_mapping = within_participant_derangement(train_data, args.seed * 7919 + 17) |
| shuffled_hash = json_sha256(shuffled_mapping) |
| shuffled_path = common_dir / "shuffled_teacher_mapping.json" |
| if shuffled_path.exists() and json.loads(shuffled_path.read_text())["sha256"] != shuffled_hash: |
| raise RuntimeError("shared shuffled-teacher mapping mismatch") |
| if not shuffled_path.exists(): |
| shuffled_path.write_text(json.dumps({ |
| "schema": "within-participant-teacher-derangement-v1", |
| "sha256": shuffled_hash, |
| "fixed_points": 0, |
| "mapping": shuffled_mapping, |
| }) + "\n", encoding="utf-8") |
| config["shuffled_teacher_permutation_sha256"] = shuffled_hash |
| config["shuffled_teacher_policy"] = "fixed seeded derangement within each training participant" |
|
|
| |
| |
| train_data.epoch = 1 |
| diagnostic_indices = orders[0][:args.batch_size] |
| diagnostic = next(iter(DataLoader(train_data, batch_size=args.batch_size, |
| sampler=FixedOrderSampler(diagnostic_indices), num_workers=0))) |
| model.train() |
| patch, landmarks, target_angles, teacher_vector, _ = (value.to(args.device) for value in diagnostic) |
| prediction = logits_to_angles(*model(patch, landmarks)) |
| student_vector = angles_to_vectors(prediction) |
| calibration_hard = SmoothAWLoss()(prediction, target_angles) |
| calibration_kd = (1.0 - (nn.functional.normalize(student_vector, dim=1) * |
| nn.functional.normalize(teacher_vector, dim=1)).sum(1)).mean() |
| hard_grad_norm, kd_grad_norm = gradient_norms(model, calibration_hard, calibration_kd) |
| effective_lambda = args.lambda_kd if args.lambda_kd is not None else ( |
| args.kd_gradient_ratio * hard_grad_norm / max(kd_grad_norm, 1e-12) |
| ) |
| config["effective_lambda_kd"] = effective_lambda |
| config["calibration_hard_gradient_norm"] = hard_grad_norm |
| config["calibration_pointkd_gradient_norm"] = kd_grad_norm |
| config["lambda_selection"] = "fixed_override" if args.lambda_kd is not None else "training_only_initial_gradient_norm_ratio" |
| train_data.close() |
| seed_everything(args.seed) |
| model.load_state_dict(torch.load(initial_path, map_location=args.device, weights_only=True), strict=True) |
| if args.arm == "shuffled_teacher_kd": |
| train_data.teacher_index_map = shuffled_mapping |
| (run_dir / "config.json").write_text(json.dumps(config, indent=2) + "\n", encoding="utf-8") |
|
|
| heldout_loader = DataLoader(heldout_data, batch_size=args.batch_size, shuffle=False, num_workers=args.workers) |
| optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=1e-2) |
| scheduler = torch.optim.lr_scheduler.OneCycleLR(optimizer, max_lr=args.lr, steps_per_epoch=math.ceil(len(train_data) / args.batch_size), epochs=args.epochs) |
| criterion = SmoothAWLoss() |
| log_path = run_dir / "epochs.jsonl" |
| with log_path.open("x", encoding="utf-8", newline="\n") as log: |
| for epoch, order in enumerate(orders, 1): |
| train_data.epoch = epoch |
| loader = DataLoader(train_data, batch_size=args.batch_size, sampler=FixedOrderSampler(order), num_workers=args.workers) |
| model.train(); sums = {"hard": 0.0, "kd": 0.0, "weighted_kd": 0.0, "gate": 0.0}; count = 0; grad_cos = None |
| train_values = {key: [] for key in ("error", "axis", "disagreement", "teacher_error", "pitch_s", "pitch_t", "yaw_s", "yaw_t")} |
| for patch, landmarks, target_angles, teacher_vector, teacher_error in loader: |
| patch, landmarks = patch.to(args.device), landmarks.to(args.device) |
| target_angles, teacher_vector, teacher_error = target_angles.to(args.device), teacher_vector.to(args.device), teacher_error.to(args.device) |
| optimizer.zero_grad(set_to_none=True) |
| prediction = logits_to_angles(*model(patch, landmarks)) |
| student_vector = angles_to_vectors(prediction) |
| hard = criterion(prediction, target_angles) |
| point = 1.0 - (nn.functional.normalize(student_vector, dim=1) * nn.functional.normalize(teacher_vector, dim=1)).sum(1) |
| gate = ((tau_bad - teacher_error) / max(tau_bad - tau_good, 1e-12)).clamp(0, 1) |
| weighted = point if args.arm in ("point_kd", "shuffled_teacher_kd") else gate * point if args.arm == "quality_gated_point_kd" else point * 0 |
| if grad_cos is None: grad_cos = gradient_cosine(model, hard, point.mean()) |
| loss = hard + effective_lambda * weighted.mean() |
| loss.backward(); optimizer.step(); scheduler.step() |
| batch = len(patch); count += batch |
| sums["hard"] += float(hard.detach()) * batch |
| sums["kd"] += float(point.mean().detach()) * batch |
| sums["weighted_kd"] += float(weighted.mean().detach()) * batch |
| sums["gate"] += float(gate.mean().detach()) * batch |
| with torch.no_grad(): |
| target_vector = angles_to_vectors(target_angles) |
| teacher_angles = torch.stack((torch.rad2deg(torch.asin((-teacher_vector[:, 1]).clamp(-1, 1))), torch.rad2deg(torch.atan2(-teacher_vector[:, 0], -teacher_vector[:, 2]))), 1) |
| batch_values = { |
| "error": angular_error(student_vector, target_vector), |
| "axis": (prediction - target_angles).abs().mean(1), |
| "disagreement": angular_error(student_vector, teacher_vector), |
| "teacher_error": angular_error(teacher_vector, target_vector), |
| "pitch_s": prediction[:, 0], "pitch_t": teacher_angles[:, 0], |
| "yaw_s": prediction[:, 1], "yaw_t": teacher_angles[:, 1], |
| } |
| for key, value in batch_values.items(): train_values[key].append(value.detach().cpu().numpy()) |
| train_values = {key: np.concatenate(value) for key, value in train_values.items()} |
| train_metrics = { |
| "train_student_3d_error_mean_deg": float(train_values["error"].mean()), |
| "train_student_axis_mae_deg": float(train_values["axis"].mean()), |
| "train_teacher_3d_error_mean_deg": float(train_values["teacher_error"].mean()), |
| "train_student_teacher_disagreement_mean_deg": float(train_values["disagreement"].mean()), |
| "train_pitch_pearson": correlation(train_values["pitch_s"], train_values["pitch_t"]), |
| "train_yaw_pearson": correlation(train_values["yaw_s"], train_values["yaw_t"]), |
| "train_pitch_spearman": spearman(train_values["pitch_s"], train_values["pitch_t"]), |
| "train_yaw_spearman": spearman(train_values["yaw_s"], train_values["yaw_t"]), |
| } |
| record = {"epoch": epoch, **{f"train_{k}_mean": v / count for k, v in sums.items()}, **train_metrics, "gradient_cosine_hard_vs_pointkd": grad_cos, **evaluate(model, heldout_loader, args.device)} |
| log.write(json.dumps(record, sort_keys=True) + "\n"); log.flush() |
| print(json.dumps({"arm": args.arm, "seed": args.seed, **record})) |
| final = json.loads(log_path.read_text(encoding="utf-8").splitlines()[-1]) |
| torch.save(model.state_dict(), run_dir / "final_state.pt") |
| summary = {"schema": "matched-pointkd-run-v1", "created_utc": datetime.now(timezone.utc).isoformat(), "held_out": args.held_out, "arm": args.arm, "seed": args.seed, "primary_result": final, "config_sha256": sha256(run_dir / "config.json"), "epoch_log_sha256": sha256(log_path), "final_state_sha256": sha256(run_dir / "final_state.pt")} |
| (run_dir / "summary.json").write_text(json.dumps(summary, indent=2) + "\n", encoding="utf-8") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|