"""Grouped held-out check for learned source-time positional features.""" from __future__ import annotations import argparse import csv import json import platform import time from datetime import datetime, timezone from pathlib import Path from typing import Any import numpy as np import torch from torch import nn from .alignment_debug import ( EXAMPLE_ID, GRID_SIZE, HEADS, HIDDEN_SIZE, LEARNING_RATE, _example_arrays, _gaussian_alignment_kl, _gaussian_targets, _gradient_norms, _metric_rows, _plot_example, _seed_everything, ) from .experiment_data import ( FeatureSample, FeatureStats, collate_feature_samples, fit_feature_stats, load_feature_samples, ) from .models import SharedLatentTimeline, TextAnchoredCrossAttention from .types import MODALITIES BATCH_SIZE = 8 DEFAULT_STEPS = 500 def _write_csv(path: Path, rows: list[dict[str, Any]]) -> None: if not rows: return path.parent.mkdir(parents=True, exist_ok=True) fields = list(dict.fromkeys(key for row in rows for key in row)) with path.open("w", newline="", encoding="utf-8-sig") as handle: writer = csv.DictWriter(handle, fieldnames=fields) writer.writeheader() writer.writerows(rows) def _make_model(method: str, dimensions: dict[str, int], source_time: bool) -> nn.Module: if method == "M3": return TextAnchoredCrossAttention( dimensions, grid_size=GRID_SIZE, hidden_size=HIDDEN_SIZE, heads=HEADS, dropout=0.0, source_time_encoding=source_time, ) return SharedLatentTimeline( dimensions, grid_size=GRID_SIZE, hidden_size=HIDDEN_SIZE, heads=HEADS, dropout=0.0, absolute_position_encoding=True, source_time_encoding=source_time, ) def _batches(samples: list[FeatureSample], rng: np.random.Generator): order = rng.permutation(len(samples)).tolist() for start in range(0, len(order), BATCH_SIZE): yield [samples[index] for index in order[start : start + BATCH_SIZE]] def _train_one( *, method: str, variant: str, source_time: bool, train_samples: list[FeatureSample], validation_samples: list[FeatureSample], stats: FeatureStats, example_id: str, output_dir: Path, device: torch.device, steps: int, seed: int, ) -> tuple[list[dict[str, Any]], list[dict[str, Any]], dict[str, Any]]: _seed_everything(seed) dimensions = { name: train_samples[0].features[name].shape[1] for name in MODALITIES } model = _make_model(method, dimensions, source_time).to(device) optimizer = torch.optim.AdamW(model.parameters(), lr=LEARNING_RATE, weight_decay=0.0) rng = np.random.default_rng(seed) history: list[dict[str, Any]] = [] print( f"[D5 {variant}] train={len(train_samples)} heldout={len(validation_samples)} " f"groups={len({s.group_id for s in train_samples})}/" f"{len({s.group_id for s in validation_samples})} steps={steps}", flush=True, ) model.train() for step in range(1, steps + 1): batch_indices = rng.choice( len(train_samples), size=min(BATCH_SIZE, len(train_samples)), replace=False ) batch_samples = [train_samples[int(index)] for index in batch_indices] sequences, durations, _ = collate_feature_samples(batch_samples, stats, device) output = model(sequences, durations) if source_time else model(sequences) targets = _gaussian_targets(method, output, sequences, durations) loss = _gaussian_alignment_kl(output, targets) if not torch.isfinite(loss): raise FloatingPointError(f"non-finite D5 loss for {variant} at step {step}") row: dict[str, Any] = { "experiment": "D5", "method": method, "variant": variant, "step": step, "L_align": float(loss.detach().item()), "source_time_encoding": source_time, } if step == 1 or step % 20 == 0 or step == steps: row.update(_gradient_norms(model, method, {"align": loss}, ("align",))) optimizer.zero_grad(set_to_none=True) loss.backward() nn.utils.clip_grad_norm_(model.parameters(), 2.0) optimizer.step() history.append(row) if step == 1 or step % 100 == 0 or step == steps: grad_z = row.get("grad_align_Z") grad_z_text = f"{grad_z:.3g}" if grad_z is not None else "NA" print( f"[D5 {variant} {step}/{steps}] KL={row['L_align']:.5f} " f"grad_Q/K/Z={row.get('grad_align_WQ', 0):.3g}/" f"{row.get('grad_align_WK', 0):.3g}/{grad_z_text}", flush=True, ) output_dir.mkdir(parents=True, exist_ok=True) _write_csv(output_dir / "history.csv", history) model.eval() metric_rows: list[dict[str, Any]] = [] heldout_arrays: dict[str, np.ndarray] | None = None with torch.no_grad(): eval_rng = np.random.default_rng(0) for batch_samples in _batches(validation_samples, eval_rng): sequences, durations, _ = collate_feature_samples(batch_samples, stats, device) output = model(sequences, durations) if source_time else model(sequences) for index, sample in enumerate(batch_samples): one_sequences = { name: type(sequences[name])( features=sequences[name].features[index : index + 1], times=sequences[name].times[index : index + 1], valid=sequences[name].valid[index : index + 1], ) for name in MODALITIES } one_output = type(output)( weights={name: output.weights[name][index : index + 1] for name in MODALITIES}, aligned={name: output.aligned[name][index : index + 1] for name in MODALITIES}, fallback_rows=output.fallback_rows, ) rows = _metric_rows( method, "D5", sample, one_output, one_sequences, durations[index : index + 1], stats=stats, device=device, ) for metric_row in rows: metric_row["variant"] = variant metric_row["source_time_encoding"] = source_time metric_rows.extend(rows) if sample.sample_id == example_id: heldout_arrays = _example_arrays( sample, one_output, one_sequences, method ) _write_csv(output_dir / "heldout_metrics.csv", metric_rows) if heldout_arrays is not None: heldout_sample = next(s for s in validation_samples if s.sample_id == example_id) np.savez_compressed(output_dir / "heldout_alignment.npz", **heldout_arrays) _plot_example(output_dir / variant, heldout_sample, heldout_arrays, method, "D5 held-out") checkpoint = output_dir / "checkpoint.pt" torch.save( { "experiment": "D5", "method": method, "variant": variant, "source_time_encoding": source_time, "absolute_position_encoding": method == "M4", "seed": seed, "steps": steps, "train_sample_ids": [sample.sample_id for sample in train_samples], "validation_sample_ids": [sample.sample_id for sample in validation_samples], "model_state_dict": model.state_dict(), }, checkpoint, ) training_summary = { "experiment": "D5", "method": method, "variant": variant, "source_time_encoding": source_time, "final_training_kl": history[-1]["L_align"], "heldout_sample_count": len(validation_samples), "checkpoint": str(checkpoint), } del model if device.type == "cuda": torch.cuda.empty_cache() return metric_rows, history, training_summary def run(args: argparse.Namespace) -> dict[str, Any]: started = time.time() if args.device == "auto": device = torch.device("cuda" if torch.cuda.is_available() else "cpu") else: device = torch.device(args.device) if device.type == "cuda" and not torch.cuda.is_available(): raise RuntimeError("CUDA was requested but is unavailable") samples = load_feature_samples(args.feature_dir, args.manifest) by_id = {sample.sample_id: sample for sample in samples} with args.splits.open("r", encoding="utf-8-sig") as handle: folds = json.load(handle) fold = next((item for item in folds if item["fold"] == args.fold), None) if fold is None: raise ValueError(f"fold {args.fold} is not present in {args.splits}") train_samples = [by_id[sample_id] for sample_id in fold["train_sample_ids"]] validation_samples = [by_id[sample_id] for sample_id in fold["validation_sample_ids"]] train_groups = {sample.group_id for sample in train_samples} validation_groups = {sample.group_id for sample in validation_samples} if train_groups & validation_groups: raise ValueError("train/validation video_id groups overlap") example_id = args.example_id or validation_samples[0].sample_id if example_id not in {sample.sample_id for sample in validation_samples}: raise ValueError(f"held-out example is not in validation fold: {example_id}") stats = fit_feature_stats(train_samples) output_root = args.output_dir output_root.mkdir(parents=True, exist_ok=True) specifications = ( ("M3", "M3_noSourceTime", False), ("M3", "M3_sourceTime", True), ("M4", "M4_noSourceTime", False), ("M4", "M4_sourceTime", True), ) all_metrics: list[dict[str, Any]] = [] all_history: list[dict[str, Any]] = [] summaries: list[dict[str, Any]] = [] for method, variant, source_time in specifications: metrics, history, summary = _train_one( method=method, variant=variant, source_time=source_time, train_samples=train_samples, validation_samples=validation_samples, stats=stats, example_id=example_id, output_dir=output_root / variant, device=device, steps=args.steps, seed=args.seed, ) all_metrics.extend(metrics) all_history.extend(history) summaries.append(summary) _write_csv(output_root / "per_sample_metrics.csv", all_metrics) _write_csv(output_root / "training_history.csv", all_history) aggregate_rows: list[dict[str, Any]] = [] metric_names = ( "mvr", "normalized_entropy", "c_row", "trajectory_span", "mean_absolute_time_center_error", "gaussian_target_kl", ) for variant, modality in sorted( {(row["variant"], row["modality"]) for row in all_metrics} ): rows = [row for row in all_metrics if row["variant"] == variant and row["modality"] == modality] aggregate: dict[str, Any] = { "variant": variant, "modality": modality, "sample_count": len(rows), } for name in metric_names: values = [float(row[name]) for row in rows if row.get(name) not in (None, "")] aggregate[f"mean_{name}"] = float(np.mean(values)) if values else "" aggregate_rows.append(aggregate) _write_csv(output_root / "heldout_summary.csv", aggregate_rows) _write_csv(output_root / "training_summary.csv", summaries) manifest = { "created_utc": datetime.now(timezone.utc).isoformat(), "experiment": "D5", "fold": args.fold, "heldout_example": example_id, "train_sample_count": len(train_samples), "heldout_sample_count": len(validation_samples), "train_video_ids": sorted(train_groups), "heldout_video_ids": sorted(validation_groups), "video_id_overlap": sorted(train_groups & validation_groups), "seed": args.seed, "device": str(device), "gpu_name": torch.cuda.get_device_name(device) if device.type == "cuda" else None, "python": platform.python_version(), "torch": torch.__version__, "grid_size": GRID_SIZE, "optimizer": "AdamW", "learning_rate": LEARNING_RATE, "steps_per_model": args.steps, "batch_size": BATCH_SIZE, "loss": "timestamp-derived Gaussian target KL only", "source_position_encoding": "fixed Fourier time code added to source key only; value remains projected content", "query_position_encoding": "M3 adds Fourier code at text-time centers; M4 uses fixed absolute sinusoidal slots plus the matching Fourier code at uniform slot centers", "normalization_fit_on_train_only": True, "variants": summaries, "elapsed_seconds": time.time() - started, "interpretation_limits": [ "This is one grouped video_id split and one seed; it is a focused held-out diagnostic, not a final method ranking.", "The Gaussian timestamp target is a weak temporal prior, not human alignment ground truth.", "The target supplies approximate time location; this experiment tests transfer of the time-conditioned attention mechanism, not semantic correctness by itself.", ], } (output_root / "run_manifest.json").write_text( json.dumps(manifest, ensure_ascii=False, indent=2, allow_nan=False), encoding="utf-8" ) print( f"[D5 done] train={len(train_samples)} heldout={len(validation_samples)} " f"video_id_groups={len(train_groups)}/{len(validation_groups)} " f"elapsed={manifest['elapsed_seconds']:.1f}s output={output_root}", flush=True, ) return manifest def build_parser() -> argparse.ArgumentParser: project = Path(__file__).resolve().parents[1] parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--steps", type=int, default=DEFAULT_STEPS) parser.add_argument("--seed", type=int, default=42) parser.add_argument("--fold", type=int, default=1) parser.add_argument("--example-id", type=str, default=None) parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto") parser.add_argument( "--feature-dir", type=Path, default=project / "outputs/q1_features/features" ) parser.add_argument("--manifest", type=Path, default=project / "outputs/audit/manifest.csv") parser.add_argument( "--splits", type=Path, default=project / "outputs/method_comparison/splits.json" ) parser.add_argument( "--output-dir", type=Path, default=project / "outputs/alignment_debug/heldout" ) return parser def main() -> None: args = build_parser().parse_args() run(args) if __name__ == "__main__": main()