"""Fit Gaussian time bands with synthetic time-only inputs as an attention control.""" from __future__ import annotations import argparse import csv import json import platform import random from datetime import datetime, timezone from pathlib import Path from typing import Any import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import numpy as np import torch from torch import Tensor, nn from .metrics import alignment_trajectory, normalized_attention_entropy from .models import SharedLatentTimeline, TextAnchoredCrossAttention from .types import MODALITIES, SequenceBatch GRID_SIZE = 50 HIDDEN_SIZE = 128 SIGMA = 0.10 SEED = 42 MODALITY_LENGTHS = {"text": GRID_SIZE, "audio": 256, "vision": 100} def _seed(seed: int) -> None: random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed) torch.backends.cudnn.deterministic = True torch.backends.cudnn.benchmark = False def _synthetic_sequence(length: int, device: torch.device) -> SequenceBatch: times = torch.linspace(0.0, 1.0, length, device=device).unsqueeze(0) # These features contain only the source's synthetic time coordinate. features = torch.stack( (times, times.square(), torch.ones_like(times)), dim=-1 ) valid = torch.ones_like(times, dtype=torch.bool) return SequenceBatch(features=features, times=times, valid=valid) def _targets( sequences: dict[str, SequenceBatch], method: str, device: torch.device ) -> dict[str, Tensor]: centers = (torch.arange(GRID_SIZE, device=device, dtype=torch.float32) + 0.5) / GRID_SIZE names = ("audio", "vision") if method == "M3" else MODALITIES targets: dict[str, Tensor] = {} for name in names: times = sequences[name].times logits = -0.5 * ((times[:, None, :] - centers[None, :, None]) / SIGMA).square() targets[name] = torch.softmax(logits, dim=-1) return targets def _loss(output: Any, targets: dict[str, Tensor]) -> Tensor: losses = [] for name, target in targets.items(): predicted = output.weights[name].clamp_min(1e-8) safe_target = target.clamp_min(1e-12) losses.append((safe_target * (safe_target.log() - predicted.log())).sum(-1).mean()) return torch.stack(losses).mean() def _attention_layers(model: nn.Module, method: str) -> list[nn.MultiheadAttention]: if method == "M3": return [model.audio_attention.attention, model.vision_attention.attention] return [model.attention[name].attention for name in MODALITIES] def _gradient_summary( model: nn.Module, method: str, loss: Tensor ) -> dict[str, float]: layers = _attention_layers(model, method) params = [layer.in_proj_weight for layer in layers] if method == "M4": params.append(model.slots) gradients = torch.autograd.grad(loss, params, retain_graph=True) q_sq = torch.zeros((), device=loss.device) k_sq = torch.zeros((), device=loss.device) for grad in gradients[: len(layers)]: q_sq += grad[:HIDDEN_SIZE].square().sum() k_sq += grad[HIDDEN_SIZE : 2 * HIDDEN_SIZE].square().sum() result = {"grad_WQ": float(q_sq.sqrt().item()), "grad_WK": float(k_sq.sqrt().item())} if method == "M4": result["grad_Z"] = float(gradients[-1].norm().item()) return result def _write_csv(path: Path, rows: list[dict[str, Any]]) -> None: if not rows: return path.parent.mkdir(parents=True, exist_ok=True) with path.open("w", newline="", encoding="utf-8-sig") as handle: fields = list(dict.fromkeys(key for row in rows for key in row)) writer = csv.DictWriter(handle, fieldnames=fields) writer.writeheader() writer.writerows(rows) def _save_plots( folder: Path, method: str, variant: str, output: Any, targets: dict[str, Tensor], sequences: dict[str, SequenceBatch], duration: float = 1.0, ) -> list[dict[str, Any]]: names = ("audio", "vision") if method == "M3" else MODALITIES rows: list[dict[str, Any]] = [] fig, axes = plt.subplots(len(names), 2, figsize=(11, 3.4 * len(names)), constrained_layout=True) if len(names) == 1: axes = np.asarray([axes]) fig_traj, ax_traj = plt.subplots(figsize=(8, 5), constrained_layout=True) grid = (np.arange(GRID_SIZE, dtype=np.float32) + 0.5) / GRID_SIZE for row_index, name in enumerate(names): weights = output.weights[name][0].detach().cpu().numpy() target = targets[name][0].detach().cpu().numpy() times = sequences[name].times[0].detach().cpu().numpy() entropy = normalized_attention_entropy( output.weights[name], sequences[name].valid ).mean().item() trajectory = alignment_trajectory( output.weights[name], sequences[name].times, torch.tensor([duration], device=output.weights[name].device) )[0] trajectory_np = trajectory.detach().cpu().numpy() span = float(trajectory_np[-1] - trajectory_np[0]) kl = float( (targets[name] * (targets[name].clamp_min(1e-12).log() - output.weights[name].clamp_min(1e-8).log())) .sum(-1) .mean() .item() ) rows.append( { "method": method, "variant": variant, "modality": name, "normalized_entropy": entropy, "trajectory_span": span, "mean_absolute_time_center_error": float( np.mean(np.abs(trajectory_np - grid)) ), "gaussian_target_kl": kl, } ) extent = (float(times[0]), float(times[-1]), 0.0, 1.0) image = axes[row_index, 0].imshow( weights, origin="lower", aspect="auto", interpolation="nearest", extent=extent, cmap="magma" ) axes[row_index, 0].set_title(f"{name}: learned A") axes[row_index, 0].set_xlabel("synthetic source time") axes[row_index, 0].set_ylabel("slot / K") fig.colorbar(image, ax=axes[row_index, 0], fraction=0.046, pad=0.04) image_target = axes[row_index, 1].imshow( target, origin="lower", aspect="auto", interpolation="nearest", extent=extent, cmap="magma" ) axes[row_index, 1].set_title(f"{name}: Gaussian target P") axes[row_index, 1].set_xlabel("synthetic source time") axes[row_index, 1].set_ylabel("slot / K") fig.colorbar(image_target, ax=axes[row_index, 1], fraction=0.046, pad=0.04) ax_traj.plot(grid, trajectory_np, label=f"{name} learned") ax_traj.plot([0, 1], [0, 1], "k:", label="uniform-time reference") ax_traj.set(xlim=(0, 1), ylim=(0, 1), xlabel="slot position", ylabel="expected source time") ax_traj.grid(alpha=0.2) ax_traj.legend() ax_traj.set_title(f"{method} {variant}: synthetic alignment trajectory") folder.mkdir(parents=True, exist_ok=True) fig.savefig(folder / f"{variant}_heatmap.png", dpi=160) fig_traj.savefig(folder / f"{variant}_trajectory.png", dpi=160) plt.close(fig) plt.close(fig_traj) return rows def _run_trial( method: str, variant: str, *, device: torch.device, steps: int, seed: int, output_dir: Path, ) -> tuple[list[dict[str, Any]], list[dict[str, Any]], dict[str, Any]]: _seed(seed) dimensions = {name: 3 for name in MODALITIES} if method == "M3": model: nn.Module = TextAnchoredCrossAttention( dimensions, grid_size=GRID_SIZE, hidden_size=HIDDEN_SIZE, heads=4, dropout=0.0 ) use_pe = False else: use_pe = variant == "M4_sinPE" model = SharedLatentTimeline( dimensions, grid_size=GRID_SIZE, hidden_size=HIDDEN_SIZE, heads=4, dropout=0.0, absolute_position_encoding=use_pe, ) model.to(device) sequences = { name: _synthetic_sequence(MODALITY_LENGTHS[name], device) for name in MODALITIES } targets = _targets(sequences, method, device) optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=0.0) history: list[dict[str, Any]] = [] output = None print(f"[{variant}] synthetic time-only inputs, steps={steps}, device={device}", flush=True) for step in range(1, steps + 1): model.train() output = model(sequences) loss = _loss(output, targets) if not torch.isfinite(loss): raise FloatingPointError(f"non-finite synthetic alignment loss for {variant} at step {step}") row: dict[str, Any] = {"method": method, "variant": variant, "step": step, "L_align": float(loss.item())} if step == 1 or step % 100 == 0 or step == steps: row.update(_gradient_summary(model, method, loss)) 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: print( f"[{variant} {step}/{steps}] KL={row['L_align']:.5f} " f"grad_Q/K/Z={row.get('grad_WQ', 0):.3g}/{row.get('grad_WK', 0):.3g}/" f"{row.get('grad_Z', float('nan')):.3g}", flush=True, ) assert output is not None model.eval() with torch.no_grad(): output = model(sequences) metric_rows = _save_plots( output_dir / variant, method, variant, output, targets, sequences ) history_rows = history checkpoint = output_dir / f"{variant}_checkpoint.pt" torch.save( { "method": method, "variant": variant, "seed": seed, "steps": steps, "absolute_position_encoding": use_pe, "model_state_dict": model.state_dict(), "synthetic_only": True, }, checkpoint, ) return metric_rows, history_rows, {"checkpoint": str(checkpoint), "final_loss": history[-1]["L_align"]} def run(args: argparse.Namespace) -> dict[str, Any]: output_dir = args.output_dir output_dir.mkdir(parents=True, exist_ok=True) 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") variants = [("M3", "M3"), ("M4", "M4_noPE"), ("M4", "M4_sinPE")] all_metrics: list[dict[str, Any]] = [] all_history: list[dict[str, Any]] = [] model_summaries: list[dict[str, Any]] = [] for method, variant in variants: metrics, history, summary = _run_trial( method, variant, device=device, steps=args.steps, seed=args.seed, output_dir=output_dir ) all_metrics.extend(metrics) all_history.extend(history) model_summaries.append({"method": method, "variant": variant, **summary}) _write_csv(output_dir / "metrics.csv", all_metrics) _write_csv(output_dir / "training_history.csv", all_history) manifest = { "created_utc": datetime.now(timezone.utc).isoformat(), "purpose": "Check whether the existing M3/M4 attention implementation can fit a time band when all inputs encode only synthetic time.", "sample_count": 1, "synthetic_only": True, "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, "feature_rule": "[t, t^2, 1] per source position; no extracted text/audio/video content is used", "sequence_lengths": MODALITY_LENGTHS, "optimizer": "AdamW", "learning_rate": 1e-3, "steps_per_variant": args.steps, "variants": model_summaries, "interpretation_limits": [ "This is an implementation/optimization control only; it does not evaluate real features or alignment accuracy.", "The time-only features deliberately provide source-position information that the real feature inputs may not contain.", ], } (output_dir / "run_manifest.json").write_text( json.dumps(manifest, ensure_ascii=False, indent=2, allow_nan=False), encoding="utf-8" ) print(f"[all done] wrote synthetic control to {output_dir}", 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=SEED) parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto") parser.add_argument( "--output-dir", type=Path, default=project / "outputs/alignment_debug/synthetic" ) return parser def main() -> None: args = build_parser().parse_args() run(args) if __name__ == "__main__": main()