"""Single-clip test of explicit source-time identities in learned alignment.""" 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, SIGMA, _example_arrays, _gaussian_alignment_kl, _gaussian_targets, _gradient_norms, _metric_rows, _plot_example, _seed_everything, ) from .experiment_data import fit_feature_stats, load_feature_samples, collate_feature_samples from .models import SharedLatentTimeline, TextAnchoredCrossAttention from .types import MODALITIES 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_encoding: 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_encoding, ) 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_encoding, ) def _run_trial( *, method: str, variant: str, source_time_encoding: bool, sample: Any, stats: Any, device: torch.device, steps: int, seed: int, output_dir: Path, ) -> tuple[dict[str, Any], list[dict[str, Any]]]: _seed_everything(seed) dimensions = {name: sample.features[name].shape[1] for name in MODALITIES} model = _make_model(method, dimensions, source_time_encoding).to(device) optimizer = torch.optim.AdamW(model.parameters(), lr=LEARNING_RATE, weight_decay=0.0) sequences, durations, _ = collate_feature_samples([sample], stats, device) history: list[dict[str, Any]] = [] print( f"[D4 {variant}] sample={sample.sample_id} steps={steps} " f"source_time_encoding={source_time_encoding} device={device}", flush=True, ) model.train() for step in range(1, steps + 1): output = model(sequences, durations) if source_time_encoding 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 D4 loss at {variant} step {step}") row: dict[str, Any] = { "experiment": "D4", "method": method, "variant": variant, "step": step, "L_align": float(loss.detach().item()), "source_time_encoding": source_time_encoding, } 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"[D4 {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}/" f"{grad_z_text}", flush=True, ) output_dir.mkdir(parents=True, exist_ok=True) _write_csv(output_dir / "history.csv", history) model.eval() with torch.no_grad(): output = model(sequences, durations) if source_time_encoding else model(sequences) metric_rows = _metric_rows( method, "D4", sample, output, sequences, durations, stats=stats, device=device, ) arrays = _example_arrays(sample, output, sequences, method) for row in metric_rows: row["variant"] = variant row["source_time_encoding"] = source_time_encoding _write_csv(output_dir / "metrics.csv", metric_rows) np.savez_compressed(output_dir / "example_alignment.npz", **arrays) _plot_example(output_dir / variant, sample, arrays, method, "D4") checkpoint = output_dir / "checkpoint.pt" torch.save( { "experiment": "D4", "method": method, "variant": variant, "source_time_encoding": source_time_encoding, "absolute_position_encoding": method == "M4", "seed": seed, "steps": steps, "model_state_dict": model.state_dict(), }, checkpoint, ) summary = { "experiment": "D4", "method": method, "variant": variant, "source_time_encoding": source_time_encoding, "final_training_kl": history[-1]["L_align"], "checkpoint": str(checkpoint), } del model if device.type == "cuda": torch.cuda.empty_cache() return summary, metric_rows def run(args: argparse.Namespace) -> dict[str, Any]: start = 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) sample_map = {sample.sample_id: sample for sample in samples} if EXAMPLE_ID not in sample_map: raise ValueError(f"diagnostic sample is missing: {EXAMPLE_ID}") sample = sample_map[EXAMPLE_ID] stats = fit_feature_stats([sample]) 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), ) summaries: list[dict[str, Any]] = [] all_metrics: list[dict[str, Any]] = [] for method, variant, use_source_time in specifications: summary, metrics = _run_trial( method=method, variant=variant, source_time_encoding=use_source_time, sample=sample, stats=stats, device=device, steps=args.steps, seed=args.seed, output_dir=output_root / variant, ) summaries.append(summary) all_metrics.extend(metrics) _write_csv(output_root / "summary.csv", summaries) _write_csv(output_root / "per_sample_metrics.csv", all_metrics) manifest = { "created_utc": datetime.now(timezone.utc).isoformat(), "experiment": "D4", "sample": sample.sample_id, "sample_count": 1, "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, "sigma_normalized_time": SIGMA, "optimizer": "AdamW", "learning_rate": LEARNING_RATE, "steps_per_model": args.steps, "source_position_encoding": "fixed Fourier time code added to source key only; value remains the projected real feature", "query_position_encoding": "M3 adds Fourier code at forced text-time centers; M4 keeps its fixed absolute sinusoidal slot code and adds the matching Fourier code at uniform slot centers", "modalities": ["audio", "vision"], "variants": summaries, "elapsed_seconds": time.time() - start, "interpretation_limits": [ "This is a one-sample overfit diagnostic, not a held-out accuracy result.", "The timestamp-derived Gaussian target is a weak temporal prior, not human alignment ground truth.", "The D4 variants test source/query time identity only; content features and downstream outputs still require separate evaluation.", ], } (output_root / "run_manifest.json").write_text( json.dumps(manifest, ensure_ascii=False, indent=2, allow_nan=False), encoding="utf-8" ) print(f"[D4 done] 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=1000) parser.add_argument("--seed", type=int, default=42) 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( "--output-dir", type=Path, default=project / "outputs/alignment_debug/source_time" ) return parser def main() -> None: args = build_parser().parse_args() run(args) if __name__ == "__main__": main()