Gaze-LIPE / experiments /train_matched_pointkd.py
thanhhuyvan's picture
Publish KD reproducibility investigation
178f61f
Raw
History Blame Contribute Delete
24.2 kB
"""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):
# Predictions are continuous, so exact ties are not expected in this diagnostic.
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"
# Calibrate objective scale using only the first deterministic training batch.
# Reloading the initial state and RNG afterward removes BatchNorm/dropout side effects.
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()