"""Re-evaluate frozen M4 Shared Latent Timeline checkpoints structurally and functionally. No alignment model is trained here. The five grouped D5 M4_sourceTime checkpoints are evaluated with self-structure, induced pairwise maps, cycle and triangle consistency, content-only probes, content shuffling, and shifted reconstruction controls. """ from __future__ import annotations import argparse import csv import json import platform 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 sklearn.metrics import roc_auc_score from torch import Tensor, nn from .correspondence_eval import ( CorrespondenceProjection, _cluster_bootstrap, _fit_probe, _sample_metrics, _stack_ids, _write_csv, ) from .experiment_data import ( FeatureSample, FeatureStats, collate_feature_samples, fit_feature_stats, load_feature_samples, ) from .models import SharedLatentTimeline from .types import MODALITIES GRID_SIZE = 50 HIDDEN_SIZE = 128 HEADS = 4 PAIRINGS = (("text", "audio"), ("text", "vision"), ("audio", "vision")) STRUCTURE_RADIUS = 2 FAR_RADIUS = 10 SIGMA_SELF = 0.08 SIGMA_CYCLE = 0.05 SHIFTS = (1, 2, 5, 10) DECODER_TARGETS = { "audio": ("text", "vision", "vision"), "vision": ("text", "audio", "audio"), "text": ("audio", "vision", "vision"), } class ContentReconstructionProbe(nn.Module): """Small decoder that sees two content streams and no slot/position code.""" def __init__(self, dimension: int = HIDDEN_SIZE) -> None: super().__init__() self.decoders = nn.ModuleDict( { target: nn.Sequential( nn.Linear(dimension * 2, dimension * 2), nn.GELU(), nn.Linear(dimension * 2, dimension), ) for target in MODALITIES } ) def forward(self, target: str, left: Tensor, right: Tensor) -> Tensor: return self.decoders[target](torch.cat((left, right), dim=-1)) def _bootstrap_summary( rows: Sequence[Mapping[str, Any]], group_keys: Sequence[str], metric_names: Sequence[str], *, seed: int, repetitions: int = 2000, ) -> list[dict[str, Any]]: grouped: dict[tuple[Any, ...], list[Mapping[str, Any]]] = defaultdict(list) for row in rows: grouped[tuple(row[key] for key in group_keys)].append(row) output: list[dict[str, Any]] = [] for key, values in sorted(grouped.items(), key=lambda item: tuple(str(x) for x in item[0])): result: dict[str, Any] = dict(zip(group_keys, key)) result["clip_count"] = len({row.get("sample_id") for row in values}) result["video_id_count"] = len({str(row["video_id"]) for row in values}) for metric_index, metric in enumerate(metric_names): selected = [row for row in values if row.get(metric) not in (None, "")] if not selected: continue mean, low, high, group_count = _cluster_bootstrap( selected, metric, seed=seed + metric_index + sum(ord(char) for char in str(key)), repetitions=repetitions, ) result[f"{metric}_video_macro_mean"] = mean result[f"{metric}_ci95_low"] = low result[f"{metric}_ci95_high"] = high result["video_id_count"] = group_count output.append(result) return output def _normalize_attention( weights: np.ndarray, valid: np.ndarray ) -> tuple[np.ndarray, np.ndarray, np.ndarray]: """Return source-position-to-slot probabilities and supported source indices.""" matrix = np.asarray(weights, dtype=np.float64) valid_indices = np.flatnonzero(valid) column_mass = matrix[:, valid_indices].sum(axis=0) supported = column_mass > 1e-12 indices = valid_indices[supported] if len(indices) == 0: raise ValueError("attention has no source positions supported by any latent slot") normalized = matrix[:, indices] / column_mass[supported][None, :] return normalized, indices, column_mass def _pair_map( source_weights: np.ndarray, source_valid: np.ndarray, destination_weights: np.ndarray, destination_valid: np.ndarray, ) -> tuple[np.ndarray, np.ndarray, np.ndarray]: """Construct C^(source->destination) after column-normalizing source attention.""" source_to_slot, source_indices, _ = _normalize_attention(source_weights, source_valid) destination_indices = np.flatnonzero(destination_valid) slot_to_destination = np.asarray(destination_weights, dtype=np.float64)[:, destination_indices] mapping = source_to_slot.T @ slot_to_destination row_mass = mapping.sum(axis=1, keepdims=True) mapping = mapping / np.maximum(row_mass, 1e-12) return mapping.astype(np.float32), source_indices, destination_indices def _self_structure( weights: np.ndarray, *, sigma: float = SIGMA_SELF ) -> tuple[np.ndarray, dict[str, float]]: normalized = weights / np.maximum(np.linalg.norm(weights, axis=1, keepdims=True), 1e-12) gram = normalized @ normalized.T positions = (np.arange(gram.shape[0], dtype=np.float64) + 0.5) / gram.shape[0] distances = np.abs(positions[:, None] - positions[None, :]) near = distances <= STRUCTURE_RADIUS / GRID_SIZE far = distances >= FAR_RADIUS / GRID_SIZE far_leakage = distances > FAR_RADIUS / GRID_SIZE near_off_diagonal = near & ~np.eye(len(gram), dtype=bool) target = np.exp(-(distances**2) / (2 * sigma**2)) near_mean = float(gram[near].mean()) far_mean = float(gram[far].mean()) near_off_diagonal_mean = float(gram[near_off_diagonal].mean()) result = { # Keep the literal <= r definition from the task, and report an # off-diagonal version so the trivial unit diagonal cannot dominate. "near_similarity": near_mean, "near_similarity_offdiag": near_off_diagonal_mean, "far_similarity": far_mean, "d_self": float(near_mean - far_mean), "d_self_offdiag": float(near_off_diagonal_mean - far_mean), "gram_target_error": float(np.linalg.norm(gram - target) / max(np.linalg.norm(target), 1e-12)), "far_slot_leakage": float(gram[far_leakage].sum() / max(gram.sum(), 1e-12)), "c_row_offdiag": float(gram[~np.eye(len(gram), dtype=bool)].mean()), } return gram.astype(np.float32), result def _pair_metrics( mapping: np.ndarray, source_times: np.ndarray, destination_times: np.ndarray, ) -> dict[str, float]: row_sums = mapping.sum(axis=1) predicted = (mapping @ destination_times) / np.maximum(row_sums, 1e-12) error = predicted - source_times return { "pairwise_time_mae": float(np.abs(error).mean()), "pairwise_signed_lag": float(error.mean()), "pairwise_time_corr": float(np.corrcoef(source_times, predicted)[0, 1]) if len(source_times) > 1 and np.std(predicted) > 1e-12 and np.std(source_times) > 1e-12 else 0.0, "source_position_count": int(len(source_times)), "destination_position_count": int(len(destination_times)), } def _band_target(times: np.ndarray, sigma: float) -> np.ndarray: distances = times[:, None] - times[None, :] target = np.exp(-(distances**2) / (2 * sigma**2)) return target / np.maximum(target.sum(axis=1, keepdims=True), 1e-12) def _cycle_triangle_metrics( maps: Mapping[str, tuple[np.ndarray, np.ndarray, np.ndarray]], normalized_times: Mapping[str, np.ndarray], ) -> tuple[list[dict[str, Any]], dict[str, np.ndarray]]: cycle_rows = [] cycles: dict[str, np.ndarray] = {} for left, right in PAIRINGS: forward_key = f"{left}_{right}" reverse_key = f"{right}_{left}" forward, source_idx, destination_idx = maps[forward_key] reverse, reverse_source_idx, reverse_destination_idx = maps[reverse_key] if not np.array_equal(destination_idx, reverse_source_idx) or not np.array_equal( source_idx, reverse_destination_idx ): raise ValueError(f"pair map source supports disagree for {forward_key}/{reverse_key}") cycle = forward @ reverse times = normalized_times[left][source_idx] target = _band_target(times, SIGMA_CYCLE) pred_times = cycle @ times cycle_rows.append( { "kind": f"cycle_{left}_{right}_{left}", "cycle_band_error": float(np.linalg.norm(cycle - target) / max(np.linalg.norm(target), 1e-12)), "cycle_time_mae": float(np.abs(pred_times - times).mean()), } ) cycles[f"cycle_{left}_{right}_{left}"] = cycle.astype(np.float32) c_ta = maps["text_audio"][0] c_av = maps["audio_vision"][0] c_tv = maps["text_vision"][0] if c_ta.shape[1] != c_av.shape[0] or c_ta.shape[0] != c_tv.shape[0] or c_av.shape[1] != c_tv.shape[1]: raise ValueError("T-A, A-V, and T-V supports disagree; cannot calculate triangle consistency") path = c_ta @ c_av triangle_residual = path - c_tv cycles["triangle_TAV_residual"] = triangle_residual.astype(np.float32) cycle_rows.append( { "kind": "triangle_TAV", "triangle_relative_error": float( np.linalg.norm(triangle_residual) / max(np.linalg.norm(c_tv), 1e-12) ), "triangle_mean_absolute_residual": float(np.abs(triangle_residual).mean()), } ) return cycle_rows, cycles def _checkpoint_path(root: Path, fold: int) -> Path: if fold == 1: return root / "M4_sourceTime" / "checkpoint.pt" return root / f"fold_{fold:02d}" / "M4_sourceTime" / "checkpoint.pt" def _collect_fold( *, fold: int, train_samples: Sequence[FeatureSample], validation_samples: Sequence[FeatureSample], stats: FeatureStats, checkpoint_root: Path, device: torch.device, batch_size: int, ) -> tuple[dict[str, dict[str, np.ndarray]], dict[str, dict[str, np.ndarray]]]: checkpoint_path = _checkpoint_path(checkpoint_root, fold) if not checkpoint_path.is_file(): raise FileNotFoundError(f"missing M4_sourceTime checkpoint for fold {fold}: {checkpoint_path}") checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False) if checkpoint.get("variant") != "M4_sourceTime" or not checkpoint.get("source_time_encoding"): raise ValueError(f"checkpoint is not M4_sourceTime: {checkpoint_path}") if set(checkpoint.get("train_sample_ids", [])) != {sample.sample_id for sample in train_samples}: raise ValueError(f"M4 checkpoint training IDs do not match fold {fold}") if set(checkpoint.get("validation_sample_ids", [])) != {sample.sample_id for sample in validation_samples}: raise ValueError(f"M4 checkpoint validation IDs do not match fold {fold}") dimensions = {name: train_samples[0].features[name].shape[1] for name in MODALITIES} model = SharedLatentTimeline( dimensions, grid_size=GRID_SIZE, hidden_size=HIDDEN_SIZE, heads=HEADS, dropout=0.0, absolute_position_encoding=True, source_time_encoding=True, ).to(device) model.load_state_dict(checkpoint["model_state_dict"], strict=True) model.eval() weights_by_id: dict[str, dict[str, np.ndarray]] = {} content_by_id: dict[str, dict[str, np.ndarray]] = {} with torch.no_grad(): samples = [*train_samples, *validation_samples] for start in range(0, len(samples), batch_size): batch_samples = samples[start : start + batch_size] sequences, durations, _ = collate_feature_samples(batch_samples, stats, device) output = model(sequences, durations) for index, sample in enumerate(batch_samples): weights: dict[str, np.ndarray] = {} content: dict[str, np.ndarray] = {} for name in MODALITIES: length = len(sample.features[name]) weights[name] = output.weights[name][index, :, :length].detach().cpu().numpy().astype(np.float32) # M4's returned values are A^m V^m: content values pooled by attention. # Positional/query vectors are not concatenated into this representation. content[name] = output.aligned[name][index].detach().cpu().numpy().astype(np.float32) weights_by_id[sample.sample_id] = weights content_by_id[sample.sample_id] = content del model if device.type == "cuda": torch.cuda.empty_cache() return weights_by_id, content_by_id def _pairwise_for_sample( sample: FeatureSample, weights: Mapping[str, np.ndarray], ) -> tuple[dict[str, tuple[np.ndarray, np.ndarray, np.ndarray]], dict[str, np.ndarray], list[dict[str, Any]], dict[str, np.ndarray]]: normalized_times = { name: np.asarray(sample.times[name], dtype=np.float64) / max(sample.duration_s, 1e-8) for name in MODALITIES } pair_maps: dict[str, tuple[np.ndarray, np.ndarray, np.ndarray]] = {} pair_rows: list[dict[str, Any]] = [] self_grams: dict[str, np.ndarray] = {} for name in MODALITIES: gram, metrics = _self_structure(weights[name]) self_grams[name] = gram pair_rows.append( {"kind": "self", "modality": name, **metrics} ) directions = (*PAIRINGS, *((right, left) for left, right in PAIRINGS)) for left, right in directions: mapping, source_idx, destination_idx = _pair_map( weights[left], sample.valid[left], weights[right], sample.valid[right] ) key = f"{left}_{right}" pair_maps[key] = (mapping, source_idx, destination_idx) source_times = normalized_times[left][source_idx] destination_times = normalized_times[right][destination_idx] pair_rows.append( { "kind": "pairwise", "direction": f"{left}_to_{right}", **_pair_metrics(mapping, source_times, destination_times), } ) cycle_rows, cycle_maps = _cycle_triangle_metrics(pair_maps, normalized_times) for row in cycle_rows: row["kind_group"] = "cycle_triangle" pair_rows.append(row) return pair_maps, normalized_times, pair_rows, {**self_grams, **cycle_maps} def _fit_reconstruction_probe( train_ids: Sequence[str], content_by_id: Mapping[str, Mapping[str, np.ndarray]], *, device: torch.device, seed: int, epochs: int, batch_size: int, learning_rate: float, ) -> tuple[ContentReconstructionProbe, list[dict[str, Any]]]: torch.manual_seed(seed) if device.type == "cuda": torch.cuda.manual_seed_all(seed) train = _stack_ids(train_ids, content_by_id, device) model = ContentReconstructionProbe(train["text"].shape[-1]).to(device) optimizer = torch.optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=1e-4) rng = np.random.default_rng(seed) history = [] model.train() for epoch in range(1, epochs + 1): order = rng.permutation(len(train_ids)) losses = [] for start in range(0, len(order), batch_size): indexes = torch.as_tensor(order[start : start + batch_size], device=device) batch = {name: value.index_select(0, indexes) for name, value in train.items()} targets = [] for target in MODALITIES: left, right, _ = DECODER_TARGETS[target] prediction = model(target, batch[left], batch[right]) targets.append(F.smooth_l1_loss(prediction, batch[target])) loss = torch.stack(targets).mean() if not torch.isfinite(loss): raise FloatingPointError(f"non-finite content reconstruction loss at epoch {epoch}") optimizer.zero_grad(set_to_none=True) loss.backward() nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() losses.append(float(loss.detach().item())) history.append({"epoch": epoch, "train_loss": float(np.mean(losses))}) return model, history def _content_pair_metrics( scores: np.ndarray, *, method: str, fold: int, sample: FeatureSample, pair: str, control: str, shuffle_id: int | None = None, ) -> dict[str, Any]: k = scores.shape[0] diag = np.diag(scores) indexes = np.arange(k) negative_mask = np.abs(indexes[:, None] - indexes[None, :]) > 2 auc = roc_auc_score( np.r_[np.ones(k), np.zeros(int(negative_mask.sum()))], np.r_[diag, scores[negative_mask]], ) row_pred = scores.argmax(axis=1) col_pred = scores.T.argmax(axis=1) row_err = np.abs(row_pred - indexes) col_err = np.abs(col_pred - indexes) result: dict[str, Any] = { "method": method, "control": control, "fold": fold, "sample_id": sample.sample_id, "video_id": sample.group_id, "pair": pair, "same_time_similarity": float(diag.mean()), "shifted_far_similarity": float( np.mean( [scores[indexes[: k - d], indexes[d:]].mean() for d in range(3, 11)] + [scores[indexes[d:], indexes[: k - d]].mean() for d in range(3, 11)] ) ), "same_minus_shifted_margin": float( diag.mean() - np.mean( [scores[indexes[: k - d], indexes[d:]].mean() for d in range(3, 11)] + [scores[indexes[d:], indexes[: k - d]].mean() for d in range(3, 11)] ) ), "matched_vs_shifted_auc": float(auc), "exact_r1_left_to_right": float(np.mean(row_err == 0)), "within_pm1_r1_left_to_right": float(np.mean(row_err <= 1)), "mase_slots_left_to_right": float(row_err.mean()), "exact_r1_right_to_left": float(np.mean(col_err == 0)), "within_pm1_r1_right_to_left": float(np.mean(col_err <= 1)), "mase_slots_right_to_left": float(col_err.mean()), } if shuffle_id is not None: result["shuffle_id"] = shuffle_id return result def _fit_and_score_content( *, fold: int, train_samples: Sequence[FeatureSample], validation_samples: Sequence[FeatureSample], content_by_id: Mapping[str, Mapping[str, np.ndarray]], device: torch.device, args: argparse.Namespace, probe_seeds: dict[str, Any], ) -> tuple[list[dict[str, Any]], list[dict[str, Any]], list[dict[str, Any]], dict[str, Any]]: train_ids = [sample.sample_id for sample in train_samples] val_ids = [sample.sample_id for sample in validation_samples] probe_seed = args.seed + fold * 101 projector, projection_history = _fit_probe( train_ids, content_by_id, device=device, seed=probe_seed, epochs=args.probe_epochs, batch_size=args.batch_size, learning_rate=args.learning_rate, temperature=args.temperature, ) for row in projection_history: probe_seeds.setdefault("history", []).append( {"fold": fold, "probe": "content_projection", "seed": probe_seed, **row} ) projector.eval() with torch.no_grad(): validation = _stack_ids(val_ids, content_by_id, device) projected = projector(validation) normal_rows: list[dict[str, Any]] = [] curve_rows: list[dict[str, Any]] = [] for index, sample in enumerate(validation_samples): one = {name: projected[name][index] for name in MODALITIES} normal_rows.extend( _sample_metrics( method="M4_sourceTime_content", fold=fold, sample=sample, projected=one, curve_rows=curve_rows, ) ) rng = np.random.default_rng(probe_seed + 17) shuffle_rows = [] with torch.no_grad(): for shuffle_id in range(args.shuffle_repeats): for index, sample in enumerate(validation_samples): permuted = {} for name in MODALITIES: permutation = torch.as_tensor(rng.permutation(GRID_SIZE), device=device) permuted[name] = projected[name][index].index_select(0, permutation) for left, right in PAIRINGS: scores = (permuted[left] @ permuted[right].T).detach().cpu().numpy() shuffle_rows.append( _content_pair_metrics( scores, method="M4_sourceTime_content", fold=fold, sample=sample, pair=f"{left}_{right}", control="independent_within_clip_permutation", shuffle_id=shuffle_id, ) ) probe_seeds[f"fold_{fold:02d}/content_projection"] = { "seed": probe_seed, "state_dict": {key: value.detach().cpu() for key, value in projector.state_dict().items()}, } del projector if device.type == "cuda": torch.cuda.empty_cache() return normal_rows, curve_rows, shuffle_rows, probe_seeds def _fit_and_score_reconstruction( *, fold: int, train_samples: Sequence[FeatureSample], validation_samples: Sequence[FeatureSample], content_by_id: Mapping[str, Mapping[str, np.ndarray]], device: torch.device, args: argparse.Namespace, checkpoint_store: dict[str, Any], history_store: list[dict[str, Any]], ) -> list[dict[str, Any]]: seed = args.seed + fold * 211 decoder, history = _fit_reconstruction_probe( [sample.sample_id for sample in train_samples], content_by_id, device=device, seed=seed, epochs=args.decoder_epochs, batch_size=args.batch_size, learning_rate=args.decoder_learning_rate, ) history_store.extend( {"fold": fold, "probe": "content_reconstruction", "seed": seed, **row} for row in history ) decoder.eval() validation = _stack_ids([sample.sample_id for sample in validation_samples], content_by_id, device) rows = [] with torch.no_grad(): for target, (left, right, shifted_modality) in DECODER_TARGETS.items(): for sample_index, sample in enumerate(validation_samples): left_values = validation[left][sample_index] right_values = validation[right][sample_index] target_values = validation[target][sample_index] for delta in (*(-value for value in SHIFTS), *SHIFTS): source_indices = np.arange(max(0, -delta), min(GRID_SIZE, GRID_SIZE - delta)) shifted_indices = source_indices + delta target_idx = torch.as_tensor(source_indices, device=device) shifted_idx = torch.as_tensor(shifted_indices, device=device) left_eval = left_values.index_select(0, target_idx) right_aligned = right_values.index_select(0, target_idx) right_shifted = ( right_values.index_select(0, shifted_idx) if shifted_modality == right else right_aligned ) left_shifted = ( left_values.index_select(0, shifted_idx) if shifted_modality == left else left_eval ) aligned_pred = decoder(target, left_eval, right_aligned) shifted_pred = decoder(target, left_shifted, right_shifted) target_eval = target_values.index_select(0, target_idx) aligned_mae = float((aligned_pred - target_eval).abs().mean().item()) shifted_mae = float((shifted_pred - target_eval).abs().mean().item()) rows.append( { "method": "M4_sourceTime", "fold": fold, "sample_id": sample.sample_id, "video_id": sample.group_id, "target_modality": target, "shifted_modality": shifted_modality, "delta": delta, "abs_delta": abs(delta), "aligned_mae_same_support": aligned_mae, "shifted_mae": shifted_mae, "gain_shift_minus_aligned": shifted_mae - aligned_mae, } ) checkpoint_store[f"fold_{fold:02d}/content_reconstruction"] = { "seed": seed, "state_dict": {key: value.detach().cpu() for key, value in decoder.state_dict().items()}, } del decoder if device.type == "cuda": torch.cuda.empty_cache() return rows def _plot_self_gram(grams: Mapping[str, np.ndarray], path: Path) -> None: fig, axes = plt.subplots(1, 3, figsize=(13, 4.4), sharex=True, sharey=True, constrained_layout=True) for axis, name in zip(axes, MODALITIES, strict=True): image = axis.imshow(grams[name], origin="lower", aspect="equal", vmin=0, vmax=1, cmap="magma") axis.set_title(name.title()) axis.set_xlabel("Latent slot k") axis.set_xticks([0, 9, 19, 29, 39, 49]) axis.set_yticks([0, 9, 19, 29, 39, 49]) axes[0].set_ylabel("Latent slot i") fig.colorbar(image, ax=axes, fraction=0.025, pad=0.02, label="Cosine similarity of attention rows") fig.suptitle("M4 self-structure: pairwise similarity between latent slots") fig.savefig(path, dpi=180, bbox_inches="tight") plt.close(fig) def _time_cell_edges(times: np.ndarray) -> np.ndarray: """Convert ordered sample centers to physical-time cell boundaries.""" times = np.asarray(times, dtype=np.float64) if len(times) == 0: raise ValueError("cannot plot an empty time axis") if len(times) == 1: return np.array([times[0] - 0.005, times[0] + 0.005]) if np.any(np.diff(times) < 0): raise ValueError("feature times must be ordered for time-faithful heatmaps") midpoints = (times[:-1] + times[1:]) / 2 first = times[0] - (midpoints[0] - times[0]) last = times[-1] + (times[-1] - midpoints[-1]) return np.concatenate(([first], midpoints, [last])) def _plot_attention( weights: Mapping[str, np.ndarray], times: Mapping[str, np.ndarray], valid: Mapping[str, np.ndarray], path: Path, ) -> None: fig, axes = plt.subplots(1, 3, figsize=(15, 4.8), sharey=True, constrained_layout=True) slot_edges = np.arange(GRID_SIZE + 1, dtype=np.float64) - 0.5 for axis, name in zip(axes, MODALITIES, strict=True): mask = np.asarray(valid[name], dtype=bool) source_times = np.asarray(times[name], dtype=np.float64)[mask] matrix = np.asarray(weights[name], dtype=np.float64)[:, mask] image = axis.pcolormesh( _time_cell_edges(source_times), slot_edges, matrix, shading="flat", cmap="magma", ) axis.set_title(name.title()) axis.set_xlabel(f"{name.title()} normalized source time") axis.set_yticks([0, 9, 19, 29, 39, 49]) axes[0].set_ylabel("Shared latent slot") fig.colorbar(image, ax=axes, fraction=0.025, pad=0.02, label="Attention weight") fig.suptitle("M4 latent slots attending to each modality's source timeline") fig.savefig(path, dpi=180, bbox_inches="tight") plt.close(fig) def _plot_pairwise( maps: Mapping[str, tuple[np.ndarray, np.ndarray, np.ndarray]], times: Mapping[str, np.ndarray], path: Path, ) -> None: directions = (("text", "audio"), ("text", "vision"), ("audio", "vision")) fig, axes = plt.subplots(1, 3, figsize=(16, 4.8), constrained_layout=True) for axis, (left, right) in zip(axes, directions, strict=True): matrix, source_idx, destination_idx = maps[f"{left}_{right}"] source_times = times[left][source_idx] destination_times = times[right][destination_idx] image = axis.pcolormesh( _time_cell_edges(destination_times), _time_cell_edges(source_times), matrix, shading="flat", cmap="magma", ) axis.set_title(f"{left.title()} → {right.title()}") axis.set_xlabel(f"{right.title()} normalized time") axis.set_ylabel(f"{left.title()} normalized time") fig.colorbar(image, ax=axes, fraction=0.025, pad=0.02, label="Pairwise transition probability") fig.suptitle("M4 modality-to-modality maps induced through the shared latent timeline") fig.savefig(path, dpi=180, bbox_inches="tight") plt.close(fig) def _plot_pair_trajectories( maps: Mapping[str, tuple[np.ndarray, np.ndarray, np.ndarray]], times: Mapping[str, np.ndarray], path: Path, ) -> None: fig, axes = plt.subplots(1, 3, figsize=(15, 4.6), sharex=True, sharey=True) for axis, (left, right) in zip(axes, (("text", "audio"), ("text", "vision"), ("audio", "vision")), strict=True): matrix, source_idx, destination_idx = maps[f"{left}_{right}"] x = times[left][source_idx] y = matrix @ times[right][destination_idx] axis.plot(x, y, color="#2673b8", linewidth=1.2, marker=".", markersize=2) axis.plot([0, 1], [0, 1], color="black", linestyle="--", linewidth=1) axis.set_title(f"{left.title()} → {right.title()}") axis.set_xlabel(f"{left.title()} source time") axis.grid(alpha=0.2) axes[0].set_ylabel("Expected destination time") fig.suptitle("Pairwise temporal trajectories; diagonal indicates equal physical time") fig.tight_layout() fig.savefig(path, dpi=180, bbox_inches="tight") plt.close(fig) def _plot_cycles(cycles: Mapping[str, np.ndarray], path: Path) -> None: names = ("cycle_text_audio_text", "cycle_text_vision_text", "cycle_audio_vision_audio") fig, axes = plt.subplots(1, 4, figsize=(17, 4.4), constrained_layout=True) for axis, name in zip(axes[:3], names, strict=True): image = axis.imshow(cycles[name], origin="lower", aspect="auto", cmap="magma") axis.set_title(name.replace("cycle_", "").replace("_", "→")) axis.set_xlabel("Source position") axis.set_ylabel("Source position") residual = np.abs(cycles["triangle_TAV_residual"]) image = axes[3].imshow(residual, origin="lower", aspect="auto", cmap="viridis") axes[3].set_title("|T→A→V − T→V|") axes[3].set_xlabel("Vision position") axes[3].set_ylabel("Text position") fig.colorbar(image, ax=axes, fraction=0.024, pad=0.02, label="Return / path residual") fig.suptitle("Cycle transition maps and triangle-path residual") fig.savefig(path, dpi=180, bbox_inches="tight") plt.close(fig) def _plot_content_curve(rows: Sequence[Mapping[str, Any]], path: Path) -> None: colors = {"text_audio": "#2878b5", "text_vision": "#e1812c", "audio_vision": "#55a868"} fig, axes = plt.subplots(1, 3, figsize=(14, 4.3), sharey=True) for axis, pair in zip(axes, colors, strict=True): xs, means, lows, highs = [], [], [], [] for delta in range(-10, 11): selected = [row for row in rows if row["pair"] == pair and int(row["delta"]) == delta] if not selected: continue mean, low, high, _ = _cluster_bootstrap( selected, "similarity", seed=7301 + delta + sum(ord(char) for char in pair), repetitions=1000, ) xs.append(delta) means.append(mean) lows.append(low) highs.append(high) axis.plot(xs, means, color=colors[pair], linewidth=1.1, marker="o", markersize=3) axis.fill_between(xs, lows, highs, color=colors[pair], alpha=0.16) axis.axvline(0, color="black", linestyle="--", linewidth=0.8) axis.set_title(pair.replace("_", "–")) axis.set_xlabel("Temporal shift Δ (slots)") axis.grid(alpha=0.2) axes[0].set_ylabel("Content-only cosine similarity") fig.suptitle("M4 content-only same-slot vs shifted similarity (video_id bootstrap CI)") fig.tight_layout() fig.savefig(path, dpi=180, bbox_inches="tight") plt.close(fig) def _plot_shuffle(normal: Sequence[Mapping[str, Any]], shuffled: Sequence[Mapping[str, Any]], path: Path) -> None: pairs = [f"{left}_{right}" for left, right in PAIRINGS] normal_means = [] shuffle_means = [] for pair in pairs: nr = [row for row in normal if row["pair"] == pair] sr = [row for row in shuffled if row["pair"] == pair] normal_means.append(_cluster_bootstrap(nr, "matched_vs_shifted_auc", seed=921)[0]) shuffle_means.append(_cluster_bootstrap(sr, "matched_vs_shifted_auc", seed=922)[0]) positions = np.arange(len(pairs)) fig, axis = plt.subplots(figsize=(8, 4.6)) width = 0.34 axis.bar(positions - width / 2, normal_means, width, label="Content as aligned", color="#2878b5") axis.bar(positions + width / 2, shuffle_means, width, label="Independent slot shuffle", color="#c44e52") axis.axhline(0.5, color="black", linestyle="--", linewidth=0.9, label="Chance AUC") axis.set_xticks(positions, [pair.replace("_", "–") for pair in pairs]) axis.set_ylim(0.35, 0.8) axis.set_ylabel("Matched-vs-shifted AUC") axis.set_title("Content shuffle control") axis.legend(frameon=False) axis.grid(axis="y", alpha=0.2) fig.tight_layout() fig.savefig(path, dpi=180, bbox_inches="tight") plt.close(fig) def _plot_reconstruction_gain(rows: Sequence[Mapping[str, Any]], path: Path) -> None: targets = ("audio", "vision", "text") colors = {"audio": "#2878b5", "vision": "#e1812c", "text": "#55a868"} fig, axis = plt.subplots(figsize=(8, 5)) for target in targets: selected = [row for row in rows if row["target_modality"] == target] by_abs: dict[int, list[dict[str, Any]]] = defaultdict(list) for row in selected: by_abs[int(row["abs_delta"])].append(dict(row)) xs = sorted(by_abs) means, lows, highs = [], [], [] for distance in xs: mean, low, high, _ = _cluster_bootstrap( by_abs[distance], "gain_shift_minus_aligned", seed=991 + distance + sum(ord(char) for char in target), repetitions=1000, ) means.append(mean) lows.append(low) highs.append(high) axis.plot(xs, means, marker="o", color=colors[target], label=f"Reconstruct {target}") axis.fill_between(xs, lows, highs, color=colors[target], alpha=0.15) axis.axhline(0, color="black", linestyle="--", linewidth=0.9) axis.set_xlabel("Absolute temporal shift |Δ| (slots)") axis.set_ylabel("MAE(shifted) − MAE(aligned)") axis.set_title("Does the aligned partner help reconstruct the target content?") axis.legend(frameon=False) axis.grid(alpha=0.2) fig.tight_layout() fig.savefig(path, dpi=180, bbox_inches="tight") plt.close(fig) 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} splits = json.loads(args.splits.read_text(encoding="utf-8")) output_dir = args.output_dir output_dir.mkdir(parents=True, exist_ok=True) self_rows: list[dict[str, Any]] = [] pair_rows: list[dict[str, Any]] = [] cycle_triangle_rows: list[dict[str, Any]] = [] content_rows: list[dict[str, Any]] = [] content_curve_rows: list[dict[str, Any]] = [] shuffle_rows: list[dict[str, Any]] = [] reconstruction_rows: list[dict[str, Any]] = [] probe_history: list[dict[str, Any]] = [] probe_checkpoints: dict[str, Any] = {} fold_manifests = [] example_id = args.example_id example_arrays: dict[str, np.ndarray] | None = None example_times: dict[str, np.ndarray] | None = None example_pair_maps: dict[str, tuple[np.ndarray, np.ndarray, np.ndarray]] | None = None example_self_grams: dict[str, np.ndarray] | None = None example_cycles: dict[str, np.ndarray] | None = None example_weights: dict[str, np.ndarray] | None = None example_valid: dict[str, np.ndarray] | None = None for split in splits: fold = int(split["fold"]) train_samples = [by_id[sample_id] for sample_id in split["train_sample_ids"]] validation_samples = [by_id[sample_id] for sample_id in split["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(f"video_id leakage in fold {fold}") if example_id is not None and example_id not in {sample.sample_id for sample in validation_samples} and fold == 1: # The requested example can be left out; evaluation still covers all held-out clips. example_id = None stats = fit_feature_stats(train_samples) weights_by_id, content_by_id = _collect_fold( fold=fold, train_samples=train_samples, validation_samples=validation_samples, stats=stats, checkpoint_root=args.checkpoint_root, device=device, batch_size=args.batch_size, ) fold_manifests.append( { "fold": fold, "train_count": len(train_samples), "heldout_count": len(validation_samples), "train_video_ids": sorted(train_groups), "heldout_video_ids": sorted(validation_groups), "overlap": sorted(train_groups & validation_groups), } ) print( f"[M4 evaluation fold {fold}] train={len(train_samples)} heldout={len(validation_samples)} " f"video_ids={len(train_groups)}/{len(validation_groups)}", flush=True, ) for sample in validation_samples: maps, normalized_times, sample_rows, arrays = _pairwise_for_sample( sample, weights_by_id[sample.sample_id] ) for row in sample_rows: row.update({"method": "M4_sourceTime", "fold": fold, "sample_id": sample.sample_id, "video_id": sample.group_id}) if row["kind"] == "self": self_rows.append(row) elif row["kind"] == "pairwise": pair_rows.append(row) else: cycle_triangle_rows.append(row) if example_id == sample.sample_id or (example_id is None and sample.sample_id == "-3g5yACwYnA/13"): example_id = sample.sample_id example_pair_maps = maps example_times = normalized_times example_self_grams = {name: arrays[name] for name in MODALITIES} example_cycles = {name: arrays[name] for name in arrays if name.startswith("cycle_") or name.startswith("triangle_")} example_weights = weights_by_id[sample.sample_id] example_valid = {name: np.asarray(sample.valid[name], dtype=bool) for name in MODALITIES} example_arrays = {} for left, right in (*PAIRINGS, *((right, left) for left, right in PAIRINGS)): mapping, source_idx, destination_idx = maps[f"{left}_{right}"] example_arrays[f"C_{left}_{right}"] = mapping example_arrays[f"source_indices_{left}_{right}"] = source_idx example_arrays[f"destination_indices_{left}_{right}"] = destination_idx for name in MODALITIES: example_arrays[f"G_{name}"] = arrays[name] example_arrays[f"A_{name}"] = weights_by_id[sample.sample_id][name] example_arrays[f"tau_{name}"] = normalized_times[name] example_arrays[f"valid_{name}"] = np.asarray(sample.valid[name], dtype=bool) example_arrays.update(example_cycles) normal, curves, shuffled, probe_checkpoints = _fit_and_score_content( fold=fold, train_samples=train_samples, validation_samples=validation_samples, content_by_id=content_by_id, device=device, args=args, probe_seeds=probe_checkpoints, ) content_rows.extend(normal) content_curve_rows.extend(curves) shuffle_rows.extend(shuffled) probe_history.extend(probe_checkpoints.pop("history", [])) reconstruction_rows.extend( _fit_and_score_reconstruction( fold=fold, train_samples=train_samples, validation_samples=validation_samples, content_by_id=content_by_id, device=device, args=args, checkpoint_store=probe_checkpoints, history_store=probe_history, ) ) print( f"[M4 evaluation fold {fold}] content_probe={len(normal)} pair rows; " f"shuffle repeats={args.shuffle_repeats}; reconstruction rows={len(reconstruction_rows)}", flush=True, ) del weights_by_id, content_by_id if device.type == "cuda": torch.cuda.empty_cache() if example_arrays is not None: np.savez_compressed(output_dir / "example_m4_shared_latent_maps.npz", **example_arrays) assert example_pair_maps is not None and example_times is not None assert example_self_grams is not None and example_cycles is not None assert example_weights is not None and example_valid is not None _plot_attention(example_weights, example_times, example_valid, output_dir / "latent_attention_heatmaps.png") _plot_self_gram(example_self_grams, output_dir / "self_gram_heatmaps.png") _plot_pairwise(example_pair_maps, example_times, output_dir / "pairwise_heatmaps.png") _plot_pair_trajectories(example_pair_maps, example_times, output_dir / "pairwise_trajectories.png") _plot_cycles(example_cycles, output_dir / "cycle_triangle_maps.png") self_summary = _bootstrap_summary( self_rows, ("method", "modality"), ("near_similarity", "near_similarity_offdiag", "far_similarity", "d_self", "d_self_offdiag", "gram_target_error", "far_slot_leakage", "c_row_offdiag"), seed=args.seed, ) pair_summary = _bootstrap_summary( pair_rows, ("method", "direction"), ("pairwise_time_mae", "pairwise_signed_lag", "pairwise_time_corr"), seed=args.seed + 1, ) cycle_triangle_summary = _bootstrap_summary( cycle_triangle_rows, ("method", "kind"), ("cycle_band_error", "cycle_time_mae", "triangle_relative_error", "triangle_mean_absolute_residual"), seed=args.seed + 2, ) content_summary, content_curve_summary = _content_summary(content_rows, content_curve_rows, args.seed + 3) shuffle_summary = _bootstrap_summary( shuffle_rows, ("method", "control", "pair"), ("same_time_similarity", "same_minus_shifted_margin", "matched_vs_shifted_auc", "exact_r1_left_to_right", "within_pm1_r1_left_to_right", "mase_slots_left_to_right"), seed=args.seed + 4, ) reconstruction_summary = _bootstrap_summary( reconstruction_rows, ("method", "target_modality", "shifted_modality", "abs_delta"), ("aligned_mae_same_support", "shifted_mae", "gain_shift_minus_aligned"), seed=args.seed + 5, ) _write_csv(output_dir / "self_structure_by_clip.csv", self_rows) _write_csv(output_dir / "self_structure_summary.csv", self_summary) _write_csv(output_dir / "pairwise_by_clip.csv", pair_rows) _write_csv(output_dir / "pairwise_summary.csv", pair_summary) _write_csv(output_dir / "cycle_triangle_by_clip.csv", cycle_triangle_rows) _write_csv(output_dir / "cycle_triangle_summary.csv", cycle_triangle_summary) _write_csv(output_dir / "content_only_metrics_by_clip.csv", content_rows) _write_csv(output_dir / "content_only_metrics_summary.csv", content_summary) _write_csv(output_dir / "content_shift_curve_by_clip.csv", content_curve_rows) _write_csv(output_dir / "content_shift_curve_summary.csv", content_curve_summary) _write_csv(output_dir / "content_shuffle_by_clip.csv", shuffle_rows) _write_csv(output_dir / "content_shuffle_summary.csv", shuffle_summary) _write_csv(output_dir / "shifted_reconstruction_by_clip.csv", reconstruction_rows) _write_csv(output_dir / "shifted_reconstruction_summary.csv", reconstruction_summary) _write_csv(output_dir / "probe_training_history.csv", probe_history) torch.save(probe_checkpoints, output_dir / "content_probe_checkpoints.pt") _plot_content_curve(content_curve_rows, output_dir / "content_only_shifted_similarity.png") _plot_shuffle(content_rows, shuffle_rows, output_dir / "content_shuffle_control.png") _plot_reconstruction_gain(reconstruction_rows, output_dir / "shifted_reconstruction_gain.png") manifest = { "created_utc": datetime.now(timezone.utc).isoformat(), "experiment": "M4 Shared Latent Timeline re-evaluation", "alignment_model_retrained": False, "checkpoint_variant": "M4_sourceTime", "sample_count": len(samples), "fold_count": len(splits), "folds": fold_manifests, "example_sample_id": example_id, "parameters": { "seed": args.seed, "batch_size": args.batch_size, "grid_size": GRID_SIZE, "hidden_size": HIDDEN_SIZE, "heads": HEADS, "sigma_self_normalized_time": SIGMA_SELF, "near_radius_slots": STRUCTURE_RADIUS, "far_radius_slots": FAR_RADIUS, "sigma_cycle_normalized_time": SIGMA_CYCLE, "content_projection_dimension": 64, "content_probe_epochs": args.probe_epochs, "content_probe_learning_rate": args.learning_rate, "contrastive_temperature": args.temperature, "content_shuffle_repeats": args.shuffle_repeats, "reconstruction_probe_epochs": args.decoder_epochs, "reconstruction_probe_learning_rate": args.decoder_learning_rate, "reconstruction_shifts_slots": list(SHIFTS), "probe_seed_rule": "seed + fold * 101", "reconstruction_seed_rule": "seed + fold * 211", }, "input_paths": { "feature_dir": str(args.feature_dir.resolve()), "feature_manifest": str(args.manifest.resolve()), "grouped_splits": str(args.splits.resolve()), "checkpoint_root": str(args.checkpoint_root.resolve()), "alignment_checkpoints": [ str(_checkpoint_path(args.checkpoint_root, int(split["fold"])).resolve()) for split in splits ], }, "content_only_definition": "M4 output.aligned = attention weights times projected source values; no latent query vector or positional embedding is concatenated into the probe input", "functional_protocol": { "probe_training": "train-fold only, symmetric within-clip same-slot InfoNCE; no emotion labels", "shuffle": "independently permute each modality's 50 content rows within each held-out clip; keep row/slot index labels fixed; 20 permutations", "reconstruction": "train-fold decoder predicts one M4 value stream from the other two; at evaluation shift one partner stream by signed offsets +/-1,2,5,10 and compare MAE on identical valid target slots", "confidence_intervals": "95% cluster bootstrap over video_id groups, 2,000 repetitions for CSV summaries; 1,000 for figure ribbons", }, "device": str(device), "gpu_name": torch.cuda.get_device_name(device) if device.type == "cuda" else None, "python": platform.python_version(), "torch": torch.__version__, "elapsed_seconds": time.time() - started, "interpretation_limits": [ "Self Gram, pairwise maps, cycle, and triangle are structural consistency checks, not independent semantic ground truth.", "The content projection and reconstruction decoder are learned on training video groups and evaluated on held-out groups; their results measure transferable probe utility.", "The time-code M4 alignment checkpoint was trained with a timestamp-derived Gaussian prior, so its attention can still encode a position shortcut.", "Human event IoU and center-error validation remain unavailable until event intervals are annotated.", "The D5 alignment checkpoints use one alignment training seed; bootstrap intervals describe video-group sampling uncertainty, not seed uncertainty.", ], } (output_dir / "run_manifest.json").write_text( json.dumps(manifest, ensure_ascii=False, indent=2, allow_nan=False), encoding="utf-8" ) print( f"[M4 evaluation complete] samples={len(samples)} folds={len(splits)} " f"elapsed={manifest['elapsed_seconds']:.1f}s output={output_dir}", flush=True, ) return manifest def _content_summary( rows: Sequence[Mapping[str, Any]], curves: Sequence[Mapping[str, Any]], seed: int, ) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]: metrics = ( "same_time_similarity", "shifted_far_similarity", "same_minus_shifted_margin", "matched_vs_shifted_auc", "exact_r1_text_to_audio", "within_pm1_r1_text_to_audio", "mase_slots_text_to_audio", "exact_r1_text_to_vision", "within_pm1_r1_text_to_vision", "mase_slots_text_to_vision", "exact_r1_audio_to_vision", "within_pm1_r1_audio_to_vision", "mase_slots_audio_to_vision", ) grouped: dict[str, list[Mapping[str, Any]]] = defaultdict(list) for row in rows: grouped[str(row["pair"])].append(row) summary = [] for pair, values in grouped.items(): output: dict[str, Any] = { "method": "M4_sourceTime_content", "pair": pair, "clip_count": len(values), "video_id_count": len({str(row["video_id"]) for row in values}), } for i, metric in enumerate(metrics): if metric not in values[0]: continue mean, low, high, count = _cluster_bootstrap( values, metric, seed=seed + i + sum(ord(char) for char in pair) ) output[f"{metric}_video_macro_mean"] = mean output[f"{metric}_ci95_low"] = low output[f"{metric}_ci95_high"] = high output["video_id_count"] = count summary.append(output) curve_groups: dict[tuple[str, int], list[Mapping[str, Any]]] = defaultdict(list) for row in curves: curve_groups[(str(row["pair"]), int(row["delta"]))].append(row) curve_summary = [] for (pair, delta), values in sorted(curve_groups.items()): mean, low, high, count = _cluster_bootstrap( values, "similarity", seed=seed + delta + sum(ord(char) for char in pair) ) curve_summary.append( { "method": "M4_sourceTime_content", "pair": pair, "delta": delta, "mean_similarity_video_macro": mean, "ci95_low": low, "ci95_high": high, "video_id_count": count, } ) return summary, curve_summary def build_parser() -> argparse.ArgumentParser: project = Path(__file__).resolve().parents[1] parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto") parser.add_argument("--batch-size", type=int, default=8) parser.add_argument("--seed", type=int, default=42) parser.add_argument("--probe-epochs", type=int, default=40) parser.add_argument("--shuffle-repeats", type=int, default=20) parser.add_argument("--decoder-epochs", type=int, default=40) parser.add_argument("--learning-rate", type=float, default=1e-3) parser.add_argument("--temperature", type=float, default=0.1) parser.add_argument("--decoder-learning-rate", type=float, default=1e-3) parser.add_argument("--example-id", type=str, default="-3g5yACwYnA/13") 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( "--checkpoint-root", type=Path, default=project / "outputs/alignment_debug/heldout" ) parser.add_argument( "--output-dir", type=Path, default=project / "outputs/m4_shared_latent_eval" ) return parser def main() -> None: args = build_parser().parse_args() run(args) if __name__ == "__main__": main()