1205 lines
53 KiB
Python
1205 lines
53 KiB
Python
"""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()
|