from __future__ import annotations import argparse import csv import json import math import platform import random import time from collections import defaultdict from datetime import datetime, timezone from pathlib import Path from typing import Any, Mapping, Sequence import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import numpy as np import torch import torch.nn.functional as F from torch import Tensor, nn from .alignment import make_block_mask from .compare_methods import AlignmentReconstructor, _write_csv from .experiment_data import ( FeatureSample, FeatureStats, collate_feature_samples, fit_feature_stats, load_feature_samples, ) from .losses import ( cross_modal_contrastive_loss, temporal_span_loss, weak_temporal_band_loss, ) from .metrics import ( alignment_trajectory, attention_row_similarity, normalized_attention_entropy, monotonicity_violation_rate, ) from .models import SharedLatentTimeline, TextAnchoredCrossAttention from .types import MODALITIES EXAMPLE_ID = "-tPCytz4rww/12" SIGMA = 0.10 GRID_SIZE = 50 HIDDEN_SIZE = 128 HEADS = 4 SINGLE_SAMPLE_STEPS = 1000 FULL_DATA_STEPS = 500 BATCH_SIZE = 8 LEARNING_RATE = 1e-3 LOG_INTERVAL = 20 def _seed_everything(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 _model( method: str, dimensions: Mapping[str, int], *, absolute_position_encoding: bool = False, ) -> nn.Module: if method == "M3": return TextAnchoredCrossAttention( dimensions, grid_size=GRID_SIZE, hidden_size=HIDDEN_SIZE, heads=HEADS, dropout=0.0, ) if method == "M4": return SharedLatentTimeline( dimensions, grid_size=GRID_SIZE, hidden_size=HIDDEN_SIZE, heads=HEADS, dropout=0.0, absolute_position_encoding=absolute_position_encoding, ) raise ValueError(f"unknown method: {method}") def _centers( method: str, output: Any, sequences: Mapping[str, Any], durations: Tensor, ) -> tuple[dict[str, Tensor], Tensor]: batch_size = durations.shape[0] if method == "M3": text_centers = torch.bmm( output.weights["text"], sequences["text"].times.unsqueeze(-1) ).squeeze(-1) text_centers = text_centers / durations[:, None].clamp_min(1e-8) return {"audio": text_centers, "vision": text_centers}, text_centers centers = ( torch.arange(GRID_SIZE, dtype=durations.dtype, device=durations.device) + 0.5 ) / GRID_SIZE centers = centers.unsqueeze(0).expand(batch_size, -1) return {name: centers for name in MODALITIES}, centers def _gaussian_targets( method: str, output: Any, sequences: Mapping[str, Any], durations: Tensor, ) -> dict[str, Tensor]: centers, _ = _centers(method, output, sequences, durations) target_names = ("audio", "vision") if method == "M3" else MODALITIES targets: dict[str, Tensor] = {} for name in target_names: times = sequences[name].times / durations[:, None].clamp_min(1e-8) difference = (times[:, None, :] - centers[name][:, :, None]) / SIGMA logits = -0.5 * difference.square() logits = logits.masked_fill(~sequences[name].valid[:, None, :], -torch.inf) targets[name] = torch.softmax(logits, dim=-1) return targets def _gaussian_alignment_kl(output: Any, targets: Mapping[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) row_kl = (safe_target * (safe_target.log() - predicted.log())).sum(dim=-1) losses.append(row_kl.mean()) return torch.stack(losses).mean() def _training_components( experiment: str, method: str, output: Any, sequences: Mapping[str, Any], durations: Tensor, *, decoder: AlignmentReconstructor | None, block_generator: torch.Generator, ) -> tuple[Tensor, dict[str, Tensor], tuple[str, ...]]: centers, _ = _centers(method, output, sequences, durations) coverage_modalities = ("audio", "vision") if method == "M3" else MODALITIES span = temporal_span_loss( output, {name: sequences[name].times for name in MODALITIES}, durations, minimum_span=0.7, modalities=coverage_modalities, ) band = weak_temporal_band_loss( output, {name: sequences[name].times for name in MODALITIES}, durations, centers, margin=0.1, ) targets = _gaussian_targets(method, output, sequences, durations) align = _gaussian_alignment_kl(output, targets) reconstruction = align.new_zeros(()) contrastive = align.new_zeros(()) if experiment == "D0": total = 5.0 * span + 10.0 * band active = ("span", "band", "total") elif experiment in {"D1", "D2"}: total = align active = ("align", "total") elif experiment == "D3": if decoder is None: raise ValueError("D3 requires the reconstruction decoder") target_losses = [] batch_size, grid_size = output.aligned["text"].shape[:2] for target in MODALITIES: mask = make_block_mask( batch_size, grid_size, 0.2, output.aligned[target].device, generator=block_generator, ) prediction = decoder(target, output.aligned, mask) target_losses.append( F.smooth_l1_loss(prediction[mask], output.aligned[target][mask]) ) reconstruction = torch.stack(target_losses).mean() contrastive = cross_modal_contrastive_loss(output.aligned) total = align + reconstruction + contrastive active = ("align", "reconstruction", "contrastive", "total") else: raise ValueError(f"unknown experiment: {experiment}") return total, { "reconstruction": reconstruction, "contrastive": contrastive, "span": span, "band": band, "align": align, "total": total, }, active def _attention_projection_parameters(model: nn.Module, method: str) -> tuple[list[nn.Parameter], nn.Parameter | None]: if method == "M3": layers = (model.audio_attention.attention, model.vision_attention.attention) slot_parameter = None else: layers = tuple(model.attention[name].attention for name in MODALITIES) slot_parameter = model.slots parameters = [layer.in_proj_weight for layer in layers] return parameters, slot_parameter def _gradient_norms( model: nn.Module, method: str, component_losses: Mapping[str, Tensor], active_components: Sequence[str], ) -> dict[str, float | None]: qk_parameters, slot_parameter = _attention_projection_parameters(model, method) parameters = [*qk_parameters] if slot_parameter is not None: parameters.append(slot_parameter) hidden_size = HIDDEN_SIZE values: dict[str, float | None] = {} for component in active_components: loss = component_losses[component] gradients = torch.autograd.grad( loss, parameters, retain_graph=True, allow_unused=True, ) q_sq = torch.zeros((), device=loss.device) k_sq = torch.zeros((), device=loss.device) for gradient in gradients[: len(qk_parameters)]: if gradient is None: continue q_sq = q_sq + gradient[:hidden_size].square().sum() k_sq = k_sq + gradient[hidden_size : 2 * hidden_size].square().sum() values[f"grad_{component}_WQ"] = float(q_sq.sqrt().item()) values[f"grad_{component}_WK"] = float(k_sq.sqrt().item()) if slot_parameter is not None: slot_gradient = gradients[-1] values[f"grad_{component}_Z"] = ( float(slot_gradient.norm().item()) if slot_gradient is not None else 0.0 ) else: values[f"grad_{component}_Z"] = None return values def _metric_rows( method: str, experiment: str, sample: FeatureSample, output: Any, sequences: Mapping[str, Any], durations: Tensor, *, stats: FeatureStats, device: torch.device, ) -> list[dict[str, Any]]: targets = _gaussian_targets(method, output, sequences, durations) rows = [] for name in MODALITIES: weights = output.weights[name] valid = sequences[name].valid trajectory = alignment_trajectory(weights, sequences[name].times, durations) entropy = normalized_attention_entropy(weights, valid) centers, _ = _centers(method, output, sequences, durations) row: dict[str, Any] = { "experiment": experiment, "method": method, "sample_id": sample.sample_id, "modality": name, "mvr": float(monotonicity_violation_rate(trajectory).mean().item()), "normalized_entropy": float(entropy.mean().item()), "c_row": float(attention_row_similarity(weights).mean().item()), "c_far": float(attention_row_similarity(weights, min_separation=6).mean().item()), "expected_time_start": float(trajectory[0, 0].item()), "expected_time_end": float(trajectory[0, -1].item()), "trajectory_span": float((trajectory[0, -1] - trajectory[0, 0]).item()), "mean_absolute_time_center_error": float( (trajectory - centers.get(name, trajectory.new_full(trajectory.shape, float("nan")))) .abs() .mean() .item() ) if name in centers else None, "gaussian_target_kl": None, } if name in targets: target = targets[name].clamp_min(1e-12) prediction = weights.clamp_min(1e-8) row["gaussian_target_kl"] = float( (target * (target.log() - prediction.log())).sum(dim=-1).mean().item() ) rows.append(row) return rows def _example_arrays(sample: FeatureSample, output: Any, sequences: Mapping[str, Any], method: str) -> dict[str, np.ndarray]: targets = _gaussian_targets(method, output, sequences, torch.tensor([sample.duration_s], device=output.weights["text"].device)) arrays: dict[str, np.ndarray] = {"sample_id": np.asarray(sample.sample_id)} for name in MODALITIES: length = len(sample.times[name]) arrays[f"weights_{name}"] = output.weights[name][0, :, :length].detach().cpu().numpy() arrays[f"times_{name}_s"] = sample.times[name].astype(np.float32, copy=False) arrays[f"valid_{name}"] = sample.valid[name] trajectory = ( output.weights[name][0, :, :length] @ torch.as_tensor(sample.times[name], dtype=torch.float32, device=output.weights[name].device) / sample.duration_s ) arrays[f"trajectory_{name}"] = trajectory.detach().cpu().numpy() if name in targets: arrays[f"target_{name}"] = targets[name][0, :, :length].detach().cpu().numpy() return arrays def _plot_example( path_prefix: Path, sample: FeatureSample, arrays: Mapping[str, np.ndarray], method: str, experiment: str, ) -> None: modalities = ("audio", "vision") if method == "M3" else MODALITIES fig, axes = plt.subplots(len(modalities), 2, figsize=(12, 4 * len(modalities)), constrained_layout=True) if len(modalities) == 1: axes = np.asarray([axes]) for row, name in enumerate(modalities): times = arrays[f"times_{name}_s"] / max(sample.duration_s, 1e-8) extent = (float(times[0]), float(times[-1]), 0.0, 1.0) prediction = arrays[f"weights_{name}"] image = axes[row, 0].imshow( prediction, origin="lower", aspect="auto", interpolation="nearest", extent=extent, cmap="magma", ) axes[row, 0].set_title(f"{name}: learned A") axes[row, 0].set_xlabel("source time / clip duration") axes[row, 0].set_ylabel("slot index / K") fig.colorbar(image, ax=axes[row, 0], fraction=0.046, pad=0.04) target_key = f"target_{name}" if target_key in arrays: target = arrays[target_key] image_target = axes[row, 1].imshow( target, origin="lower", aspect="auto", interpolation="nearest", extent=extent, cmap="magma", ) axes[row, 1].set_title(f"{name}: Gaussian target P") fig.colorbar(image_target, ax=axes[row, 1], fraction=0.046, pad=0.04) else: axes[row, 1].imshow( np.zeros_like(prediction), origin="lower", aspect="auto", extent=extent, cmap="magma", ) axes[row, 1].set_title(f"{name}: target not used by M3") axes[row, 1].set_xlabel("source time / clip duration") axes[row, 1].set_ylabel("slot index / K") fig.suptitle(f"{experiment} {method} · {sample.sample_id}") fig.savefig(path_prefix.with_name(path_prefix.name + "_heatmap.png"), dpi=160) plt.close(fig) fig, ax = plt.subplots(figsize=(8, 5), constrained_layout=True) x = (np.arange(GRID_SIZE, dtype=np.float32) + 0.5) / GRID_SIZE for name in MODALITIES: trajectory = arrays[f"trajectory_{name}"] ax.plot(x, trajectory, label=f"{name} actual") if f"target_{name}" in arrays: target_weights = arrays[f"target_{name}"] target_time = target_weights @ arrays[f"times_{name}_s"] / max(sample.duration_s, 1e-8) ax.plot(x, target_time, linestyle="--", alpha=0.7, label=f"{name} target") ax.plot([0, 1], [0, 1], color="black", linestyle=":", alpha=0.6, label="uniform-time reference") ax.set(xlim=(0, 1), ylim=(0, 1), xlabel="shared slot position", ylabel="expected normalized source time") ax.grid(alpha=0.2) ax.legend(fontsize=8, ncol=2) ax.set_title(f"{experiment} {method} trajectory · {sample.sample_id}") fig.savefig(path_prefix.with_name(path_prefix.name + "_trajectory.png"), dpi=160) plt.close(fig) def _batches(samples: Sequence[FeatureSample], batch_size: int, 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 _evaluate_samples( method: str, experiment: str, model: nn.Module, samples: Sequence[FeatureSample], stats: FeatureStats, *, device: torch.device, batch_size: int, ) -> tuple[list[dict[str, Any]], dict[str, np.ndarray] | None]: rows: list[dict[str, Any]] = [] example_arrays: dict[str, np.ndarray] | None = None model.eval() rng = np.random.default_rng(0) with torch.no_grad(): for batch_samples in _batches(samples, batch_size, rng): sequences, durations, _ = collate_feature_samples(batch_samples, stats, device) output = 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, ) one_duration = durations[index : index + 1] rows.extend( _metric_rows( method, experiment, sample, one_output, one_sequences, one_duration, stats=stats, device=device, ) ) if sample.sample_id == EXAMPLE_ID: example_arrays = _example_arrays(sample, one_output, one_sequences, method) return rows, example_arrays def _run_trial( experiment: str, method: str, tag: str, train_samples: Sequence[FeatureSample], eval_samples: Sequence[FeatureSample], stats: FeatureStats, output_dir: Path, *, device: torch.device, seed: int, steps: int, absolute_position_encoding: bool, batch_size: int, include_content_losses: bool, ) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]: _seed_everything(seed) dimensions = {name: train_samples[0].features[name].shape[1] for name in MODALITIES} model = _model(method, dimensions, absolute_position_encoding=absolute_position_encoding).to(device) decoder = AlignmentReconstructor(HIDDEN_SIZE, dropout=0.0).to(device) if include_content_losses else None parameters = list(model.parameters()) + (list(decoder.parameters()) if decoder is not None else []) optimizer = torch.optim.AdamW(parameters, lr=LEARNING_RATE, weight_decay=0.0) block_generator = torch.Generator(device=device) block_generator.manual_seed(seed + 73) rng = np.random.default_rng(seed) order = np.arange(len(train_samples)) cursor = 0 history: list[dict[str, Any]] = [] log_every = LOG_INTERVAL active_names: tuple[str, ...] | None = None print( f"[{experiment} {tag}] samples={len(train_samples)} steps={steps} " f"sinusoidal_PE={absolute_position_encoding} lr={LEARNING_RATE}", flush=True, ) model.train() if decoder is not None: decoder.train() for step in range(1, steps + 1): if len(train_samples) == 1: batch_samples = [train_samples[0]] else: if cursor + batch_size > len(order): order = rng.permutation(len(train_samples)) cursor = 0 indices = order[cursor : cursor + batch_size] cursor += len(indices) batch_samples = [train_samples[int(index)] for index in indices] sequences, durations, _ = collate_feature_samples(batch_samples, stats, device) output = model(sequences) total, losses, active_components = _training_components( experiment, method, output, sequences, durations, decoder=decoder, block_generator=block_generator, ) active_names = active_components if not torch.isfinite(total): raise FloatingPointError(f"non-finite {experiment}/{tag} loss at step {step}") row: dict[str, Any] = { "experiment": experiment, "method": method, "variant": tag, "step": step, "epoch": math.ceil(step / max(1, math.ceil(len(train_samples) / batch_size))), "learning_rate": LEARNING_RATE, "absolute_position_encoding": absolute_position_encoding, "L_rec": float(losses["reconstruction"].detach().item()), "L_con": float(losses["contrastive"].detach().item()), "L_span": float(losses["span"].detach().item()), "L_band": float(losses["band"].detach().item()), "L_align": float(losses["align"].detach().item()), "L_total": float(total.detach().item()), } if step == 1 or step % log_every == 0 or step == steps: row.update(_gradient_norms(model, method, losses, active_components)) optimizer.zero_grad(set_to_none=True) total.backward() nn.utils.clip_grad_norm_(parameters, 2.0) optimizer.step() history.append(row) if step == 1 or step % log_every == 0 or step == steps: _write_csv(output_dir / f"{tag}_history.csv", history) if step % 100 == 0 or step == steps: diagnostic_component = "band" if experiment == "D0" else "align" print( f"[{experiment} {tag} step {step}/{steps}] " f"total={row['L_total']:.4f} align={row['L_align']:.4f} " f"span={row['L_span']:.4f} band={row['L_band']:.4f} " f"grad-{diagnostic_component}(Q/K/Z)=" f"{row.get(f'grad_{diagnostic_component}_WQ', float('nan')):.3g}/" f"{row.get(f'grad_{diagnostic_component}_WK', float('nan')):.3g}/" f"{row.get(f'grad_{diagnostic_component}_Z', float('nan')) if row.get(f'grad_{diagnostic_component}_Z') is not None else 'NA'}", flush=True, ) output_dir.mkdir(parents=True, exist_ok=True) _write_csv(output_dir / f"{tag}_history.csv", history) checkpoint = output_dir / f"{tag}_checkpoint.pt" torch.save( { "experiment": experiment, "method": method, "variant": tag, "seed": seed, "steps": steps, "absolute_position_encoding": absolute_position_encoding, "model_state_dict": model.state_dict(), "decoder_state_dict": decoder.state_dict() if decoder is not None else None, "history": history, }, checkpoint, ) metric_rows, example_arrays = _evaluate_samples( method, experiment, model, eval_samples, stats, device=device, batch_size=batch_size, ) _write_csv(output_dir / f"{tag}_metrics.csv", metric_rows) if example_arrays is not None: np.savez_compressed(output_dir / f"{tag}_example_alignment.npz", **example_arrays) example_sample = next(sample for sample in eval_samples if sample.sample_id == EXAMPLE_ID) _plot_example(output_dir / tag, example_sample, example_arrays, method, experiment) del model, decoder if device.type == "cuda": torch.cuda.empty_cache() return history, metric_rows def run(args: argparse.Namespace) -> dict[str, Any]: started = time.time() _seed_everything(args.seed) 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 not available") samples = load_feature_samples(args.feature_dir, args.manifest) if len(samples) != 100: raise ValueError(f"debug experiments expect the complete 100-sample set, found {len(samples)}") sample_by_id = {sample.sample_id: sample for sample in samples} if EXAMPLE_ID not in sample_by_id: raise ValueError(f"required debug sample is absent: {EXAMPLE_ID}") one_sample = sample_by_id[EXAMPLE_ID] output_root = args.output_dir output_root.mkdir(parents=True, exist_ok=True) full_stats = fit_feature_stats(samples) single_stats = fit_feature_stats([one_sample]) all_metrics: list[dict[str, Any]] = [] all_history_summaries: list[dict[str, Any]] = [] specifications = [ ("D0", "M3", "M3", [one_sample], [one_sample], single_stats, SINGLE_SAMPLE_STEPS, False, False), ("D0", "M4", "M4", [one_sample], [one_sample], single_stats, SINGLE_SAMPLE_STEPS, False, False), ("D1", "M3", "M3", [one_sample], [one_sample], single_stats, SINGLE_SAMPLE_STEPS, False, False), ("D1", "M4", "M4_noPE", [one_sample], [one_sample], single_stats, SINGLE_SAMPLE_STEPS, False, False), ("D1", "M4", "M4_sinPE", [one_sample], [one_sample], single_stats, SINGLE_SAMPLE_STEPS, True, False), ("D2", "M3", "M3", samples, samples, full_stats, FULL_DATA_STEPS, False, False), ("D2", "M4", "M4_sinPE", samples, samples, full_stats, FULL_DATA_STEPS, True, False), ("D3", "M3", "M3", samples, samples, full_stats, FULL_DATA_STEPS, False, True), ("D3", "M4", "M4_sinPE", samples, samples, full_stats, FULL_DATA_STEPS, True, True), ] for experiment, method, tag, train_set, eval_set, stats, steps, use_pe, content_losses in specifications: trial_dir = output_root / experiment trial_dir.mkdir(parents=True, exist_ok=True) history, metric_rows = _run_trial( experiment, method, tag, train_set, eval_set, stats, trial_dir, device=device, seed=args.seed, steps=steps, absolute_position_encoding=use_pe, batch_size=BATCH_SIZE, include_content_losses=content_losses, ) all_metrics.extend(metric_rows) selected_steps = [ row for row in history if any(key.startswith("grad_") and value not in (None, "") for key, value in row.items()) ] summary: dict[str, Any] = { "experiment": experiment, "method": method, "variant": tag, "steps": steps, "absolute_position_encoding": use_pe, "final_L_total": history[-1]["L_total"], "final_L_rec": history[-1]["L_rec"], "final_L_con": history[-1]["L_con"], "final_L_span": history[-1]["L_span"], "final_L_band": history[-1]["L_band"], "final_L_align": history[-1]["L_align"], } if selected_steps: for component in ("span", "band", "align", "reconstruction", "contrastive", "total"): for parameter in ("WQ", "WK", "Z"): key = f"grad_{component}_{parameter}" values = [float(row[key]) for row in selected_steps if row.get(key) not in (None, "")] if values: summary[f"mean_{key}"] = float(np.mean(values)) summary[f"final_{key}"] = values[-1] summary["final_metric_rows"] = len(metric_rows) all_history_summaries.append(summary) print( f"[done {experiment} {tag}] final total={summary['final_L_total']:.4f} " f"align={summary['final_L_align']:.4f}; metrics={len(metric_rows)}", flush=True, ) _write_csv(output_root / "debug_summary.csv", all_history_summaries) _write_csv(output_root / "per_sample_metrics.csv", all_metrics) manifest = { "created_utc": datetime.now(timezone.utc).isoformat(), "sample_count": len(samples), "diagnostic_sample": EXAMPLE_ID, "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__, "experiments": [ { "name": "D0", "scope": "single sample; span + barycenter band only", "steps_per_model": SINGLE_SAMPLE_STEPS, "loss": "5 * L_span + 10 * L_band", }, { "name": "D1", "scope": "single sample; Gaussian target KL only", "steps_per_model": SINGLE_SAMPLE_STEPS, "sigma_normalized_time": SIGMA, "M4_control": "no PE versus fixed sinusoidal PE", }, { "name": "D2", "scope": "all 100 clips; Gaussian target KL only; one seed; in-sample diagnostic", "steps_per_model": FULL_DATA_STEPS, "sigma_normalized_time": SIGMA, "M4_position_encoding": "fixed sinusoidal", }, { "name": "D3", "scope": "all 100 clips; Gaussian KL + masked reconstruction + contrastive; one seed; in-sample diagnostic", "steps_per_model": FULL_DATA_STEPS, "sigma_normalized_time": SIGMA, "M4_position_encoding": "fixed sinusoidal", }, ], "optimizer": "AdamW", "learning_rate": LEARNING_RATE, "dropout": 0.0, "batch_size": BATCH_SIZE, "gradient_logging_interval_steps": LOG_INTERVAL, "gradient_metrics": ["W_Q", "W_K", "M4 latent slots Z"], "features_changed": False, "M1_M2_changed": False, "elapsed_seconds": time.time() - started, "interpretation_limits": [ "D0 and D1 overfit one selected sample and diagnose optimization/representability only.", "D2 and D3 train and evaluate on the same 100 clips; they diagnose whether the target can be optimized, not generalization.", "Gaussian targets are weak temporal priors constructed from timestamps; they are not human alignment ground truth.", "M4 absolute sinusoidal encoding is enabled only for D1's PE control and D2/D3; prior M1-M4 and v2 results are unchanged.", ], } (output_root / "run_manifest.json").write_text( json.dumps(manifest, ensure_ascii=False, indent=2, allow_nan=False), encoding="utf-8" ) print( f"[all done] elapsed={manifest['elapsed_seconds']:.1f}s output={output_root}", flush=True, ) return manifest def build_parser() -> argparse.ArgumentParser: project_dir = Path(__file__).resolve().parents[1] parser = argparse.ArgumentParser( description="Diagnose gradient flow and temporal alignment learnability for M3/M4." ) parser.add_argument("--feature-dir", type=Path, default=project_dir / "outputs/q1_features/features") parser.add_argument("--manifest", type=Path, default=project_dir / "outputs/audit/manifest.csv") parser.add_argument("--output-dir", type=Path, default=project_dir / "outputs/alignment_debug") parser.add_argument("--device", default="auto", help="auto, cpu, or a torch device such as cuda:0") parser.add_argument("--seed", type=int, default=42) return parser def main() -> int: args = build_parser().parse_args() run(args) return 0 if __name__ == "__main__": raise SystemExit(main())