Files
modeling_zhaocui/deep_learning/Q1/q1/tsfa_experiment.py
T

1775 lines
80 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Grouped five-fold TSFA coarse-to-fine alignment experiment.
The M4 source-time model is frozen and supplies temporal candidate windows.
Only a content-only local semantic branch is trained. No emotion labels or
positional/time codes enter that branch.
"""
from __future__ import annotations
import argparse
import csv
import hashlib
import json
import math
import platform
import random
import shutil
import time
from collections import defaultdict
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Mapping, Sequence
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import torch
import torch.nn.functional as F
from torch import Tensor, nn
from .correspondence_eval import (
CorrespondenceProjection,
_cluster_bootstrap,
_fit_probe,
_sample_metrics,
_stack_ids,
_write_csv,
)
from .experiment_data import (
FeatureSample,
collate_feature_samples,
fit_feature_stats,
load_feature_samples,
)
from .m4_shared_latent_eval import (
ContentReconstructionProbe,
DECODER_TARGETS,
PAIRINGS,
SHIFTS,
_bootstrap_summary,
_content_pair_metrics,
_cycle_triangle_metrics,
_normalize_attention,
_pair_metrics,
_self_structure,
)
from .models import SharedLatentTimeline, TextAnchoredCrossAttention
from .types import MODALITIES
GRID_SIZE = 50
HIDDEN_SIZE = 128
HEADS = 4
BASELINE_VARIANTS = ("M3_noSourceTime", "M3_sourceTime", "M4_sourceTime")
TSFA_VARIANTS = ("TSFA-main", "TSFA-multiply", "TSFA-random", "TSFA-global")
ALL_METHODS = (*BASELINE_VARIANTS, *TSFA_VARIANTS)
HARD_NEGATIVE_OFFSETS = (-5, -3, -2, 2, 3, 5)
RECONSTRUCTION_METHODS = ("M4_sourceTime", "TSFA-main")
SHUFFLE_METHODS = ("M4_sourceTime", "TSFA-main")
class TSFASemanticBranch(nn.Module):
"""Content-only Text-to-Audio/Vision attention with no positional input."""
def __init__(self, dimension: int = HIDDEN_SIZE, match_dimension: int = 64) -> None:
super().__init__()
self.query = nn.Linear(dimension, dimension, bias=False)
self.keys = nn.ModuleDict(
{name: nn.Linear(dimension, dimension, bias=False) for name in ("audio", "vision")}
)
self.values = nn.ModuleDict(
{name: nn.Linear(dimension, dimension, bias=False) for name in ("audio", "vision")}
)
self.match = nn.ModuleDict(
{name: nn.Linear(dimension, match_dimension, bias=False) for name in MODALITIES}
)
def attend(
self,
text_content: Tensor,
source_content: Tensor,
modality: str,
candidate_mask: Tensor,
temporal_prior: Tensor | None = None,
) -> tuple[Tensor, Tensor]:
query = self.query(text_content)
keys = self.keys[modality](source_content)
scores = torch.bmm(query, keys.transpose(1, 2)) / math.sqrt(query.shape[-1])
if temporal_prior is not None:
scores = scores + temporal_prior.clamp_min(1e-8).log()
scores = scores.masked_fill(~candidate_mask, torch.finfo(scores.dtype).min)
weights = torch.softmax(scores, dim=-1)
values = self.values[modality](source_content)
return weights, torch.bmm(weights, values)
def match_embedding(self, modality: str, content: Tensor) -> Tensor:
return F.normalize(self.match[modality](content), dim=-1)
def _checkpoint_path(root: Path, variant: str, fold: int) -> Path:
if fold == 1:
return root / variant / "checkpoint.pt"
return root / f"fold_{fold:02d}" / variant / "checkpoint.pt"
def _instantiate_alignment_model(
variant: str, dimensions: Mapping[str, int]
) -> nn.Module:
if variant.startswith("M3_"):
return TextAnchoredCrossAttention(
dimensions,
grid_size=GRID_SIZE,
hidden_size=HIDDEN_SIZE,
heads=HEADS,
dropout=0.0,
source_time_encoding=variant == "M3_sourceTime",
)
return SharedLatentTimeline(
dimensions,
grid_size=GRID_SIZE,
hidden_size=HIDDEN_SIZE,
heads=HEADS,
dropout=0.0,
absolute_position_encoding=True,
source_time_encoding=True,
)
def _collect_fold_features(
*,
fold: int,
train_samples: Sequence[FeatureSample],
validation_samples: Sequence[FeatureSample],
feature_stats: Any,
checkpoint_root: Path,
device: torch.device,
batch_size: int,
) -> tuple[
dict[str, dict[str, dict[str, np.ndarray]]],
dict[str, dict[str, dict[str, np.ndarray]]],
dict[str, dict[str, np.ndarray]],
]:
"""Return baseline content/weights and frozen M4 temporal tensors."""
dimensions = {name: train_samples[0].features[name].shape[1] for name in MODALITIES}
all_samples = [*train_samples, *validation_samples]
all_ids = [sample.sample_id for sample in all_samples]
content_by_method: dict[str, dict[str, dict[str, np.ndarray]]] = {}
weights_by_method: dict[str, dict[str, dict[str, np.ndarray]]] = {}
temporal_by_id: dict[str, dict[str, np.ndarray]] = {}
for variant in BASELINE_VARIANTS:
checkpoint_path = _checkpoint_path(checkpoint_root, variant, fold)
if not checkpoint_path.is_file():
raise FileNotFoundError(f"missing {variant} fold {fold} checkpoint: {checkpoint_path}")
checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False)
if checkpoint.get("variant") != variant:
raise ValueError(f"checkpoint variant mismatch at {checkpoint_path}")
if set(checkpoint.get("train_sample_ids", [])) != {s.sample_id for s in train_samples}:
raise ValueError(f"checkpoint training IDs do not match fold {fold}: {variant}")
if set(checkpoint.get("validation_sample_ids", [])) != {s.sample_id for s in validation_samples}:
raise ValueError(f"checkpoint held-out IDs do not match fold {fold}: {variant}")
model = _instantiate_alignment_model(variant, dimensions).to(device)
model.load_state_dict(checkpoint["model_state_dict"], strict=True)
model.eval()
method_content: dict[str, dict[str, np.ndarray]] = {}
method_weights: dict[str, dict[str, np.ndarray]] = {}
with torch.no_grad():
for start in range(0, len(all_samples), batch_size):
batch_samples = all_samples[start : start + batch_size]
sequences, durations, _ = collate_feature_samples(batch_samples, feature_stats, device)
output = model(sequences, durations)
projected_sources = {
name: model.projections[name](sequences[name].features) for name in MODALITIES
}
for index, sample in enumerate(batch_samples):
one_content: dict[str, np.ndarray] = {}
one_weights: dict[str, np.ndarray] = {}
for name in MODALITIES:
length = len(sample.features[name])
weights = output.weights[name][index, :, :length]
if variant.startswith("M3_") and name == "text":
# M3's returned text query contains a timestamp PE
# residual in the source-time setting. Keep only
# attention-pooled text content for this probe.
pooled = torch.matmul(
weights, projected_sources[name][index, :length]
)
else:
# The actual attention output includes MHA value
# and output projections. Using A@pre-attention
# features would change the frozen baseline.
pooled = output.aligned[name][index]
one_content[name] = pooled.cpu().numpy().astype(np.float32, copy=False)
one_weights[name] = weights.cpu().numpy().astype(np.float32, copy=False)
method_content[sample.sample_id] = one_content
method_weights[sample.sample_id] = one_weights
if variant == "M4_sourceTime":
temporal_by_id[sample.sample_id] = {
"weights": one_weights,
"values": {
name: projected_sources[name][index, : len(sample.features[name])]
.cpu()
.numpy()
.astype(np.float32, copy=False)
for name in MODALITIES
},
"times": {
name: (sample.times[name] / max(sample.duration_s, 1e-8)).astype(
np.float32, copy=False
)
for name in MODALITIES
},
"valid": {name: sample.valid[name].copy() for name in MODALITIES},
"content": one_content,
}
content_by_method[variant] = method_content
weights_by_method[variant] = method_weights
del model, checkpoint
if device.type == "cuda":
torch.cuda.empty_cache()
if set(temporal_by_id) != set(all_ids):
raise ValueError("frozen M4 temporal features are incomplete")
return content_by_method, weights_by_method, temporal_by_id
def _stable_seed(seed: int, *parts: str | int) -> int:
payload = "|".join([str(seed), *(str(part) for part in parts)]).encode("utf-8")
return int.from_bytes(hashlib.sha256(payload).digest()[:8], "little") % (2**32)
def _random_centers(
record: Mapping[str, Any], modality: str, *, seed: int, draw: int
) -> np.ndarray:
valid = np.asarray(record["valid"][modality], dtype=bool)
times = np.asarray(record["times"][modality], dtype=np.float64)
if not valid.any():
raise ValueError(f"no valid time positions for random candidate window: {modality}")
rng = np.random.default_rng(_stable_seed(seed, record.get("sample_id", ""), modality, draw))
low, high = float(times[valid].min()), float(times[valid].max())
if high <= low:
return np.full(GRID_SIZE, low, dtype=np.float32)
return rng.uniform(low, high, size=GRID_SIZE).astype(np.float32)
def _candidate_mask(
times: Tensor,
valid: Tensor,
centers: Tensor,
*,
delta: float,
mode: str,
) -> tuple[Tensor, Tensor]:
if mode == "global":
mask = valid[:, None, :].expand(-1, GRID_SIZE, -1).clone()
else:
mask = (times[:, None, :] - centers[:, :, None]).abs() <= delta
mask &= valid[:, None, :]
empty = ~mask.any(dim=-1)
if empty.any():
distance = (times[:, None, :] - centers[:, :, None]).abs()
distance = distance.masked_fill(~valid[:, None, :], torch.inf)
nearest = distance.argmin(dim=-1)
batch_index, slot_index = empty.nonzero(as_tuple=True)
mask[batch_index, slot_index, nearest[batch_index, slot_index]] = True
return mask, empty
def _collate_temporal(
sample_ids: Sequence[str],
temporal_by_id: Mapping[str, Mapping[str, Any]],
device: torch.device,
) -> dict[str, Any]:
batch_size = len(sample_ids)
dimensions = {name: temporal_by_id[sample_ids[0]]["values"][name].shape[1] for name in MODALITIES}
tensors: dict[str, Any] = {
"text_content": torch.from_numpy(
np.stack([temporal_by_id[sample_id]["content"]["text"] for sample_id in sample_ids])
).to(device),
"values": {},
"times": {},
"valid": {},
"weights": {},
"centers": {},
}
for name in MODALITIES:
max_length = max(len(temporal_by_id[sample_id]["times"][name]) for sample_id in sample_ids)
values = torch.zeros(batch_size, max_length, dimensions[name], dtype=torch.float32, device=device)
times = torch.zeros(batch_size, max_length, dtype=torch.float32, device=device)
valid = torch.zeros(batch_size, max_length, dtype=torch.bool, device=device)
weights = torch.zeros(batch_size, GRID_SIZE, max_length, dtype=torch.float32, device=device)
for index, sample_id in enumerate(sample_ids):
record = temporal_by_id[sample_id]
length = len(record["times"][name])
values[index, :length] = torch.from_numpy(record["values"][name]).to(device)
times[index, :length] = torch.from_numpy(record["times"][name]).to(device)
valid[index, :length] = torch.from_numpy(record["valid"][name]).to(device)
weights[index, :, :length] = torch.from_numpy(record["weights"][name]).to(device)
centers = (weights * times[:, None, :]).sum(dim=-1)
tensors["values"][name] = values
tensors["times"][name] = times
tensors["valid"][name] = valid
tensors["weights"][name] = weights
tensors["centers"][name] = centers
return tensors
def _local_contrastive_loss(
branch: TSFASemanticBranch,
text_content: Tensor,
audio_content: Tensor,
vision_content: Tensor,
*,
temperature: float,
) -> Tensor:
offsets = HARD_NEGATIVE_OFFSETS
anchor_indices = torch.arange(5, GRID_SIZE - 5, device=text_content.device)
labels = torch.zeros(anchor_indices.numel(), dtype=torch.long, device=text_content.device)
def directed(query: Tensor, target: Tensor) -> Tensor:
candidates = torch.stack(
[target.index_select(1, anchor_indices + offset) for offset in (0, *offsets)], dim=2
)
anchors = query.index_select(1, anchor_indices)
logits = (anchors.unsqueeze(2) * candidates).sum(dim=-1) / temperature
return F.cross_entropy(logits.reshape(-1, logits.shape[-1]), labels.repeat(query.shape[0]))
text_embedding = branch.match_embedding("text", text_content)
audio_embedding = branch.match_embedding("audio", audio_content)
vision_embedding = branch.match_embedding("vision", vision_content)
loss_ta = (directed(text_embedding, audio_embedding) + directed(audio_embedding, text_embedding)) / 2
loss_tv = (directed(text_embedding, vision_embedding) + directed(vision_embedding, text_embedding)) / 2
return (loss_ta + loss_tv) / 2
def _fit_semantic_branch(
*,
fold: int,
train_samples: Sequence[FeatureSample],
temporal_by_id: Mapping[str, Mapping[str, Any]],
device: torch.device,
seed: int,
epochs: int,
batch_size: int,
learning_rate: float,
temperature: float,
delta: float,
) -> tuple[TSFASemanticBranch, list[dict[str, Any]]]:
fold_seed = seed + fold * 101
random.seed(fold_seed)
np.random.seed(fold_seed)
torch.manual_seed(fold_seed)
if device.type == "cuda":
torch.cuda.manual_seed_all(fold_seed)
branch = TSFASemanticBranch(HIDDEN_SIZE).to(device)
optimizer = torch.optim.AdamW(branch.parameters(), lr=learning_rate, weight_decay=1e-4)
rng = np.random.default_rng(fold_seed)
train_ids = [sample.sample_id for sample in train_samples]
history = []
branch.train()
for epoch in range(1, epochs + 1):
order = rng.permutation(len(train_ids))
losses = []
for start in range(0, len(order), batch_size):
ids = [train_ids[int(index)] for index in order[start : start + batch_size]]
batch = _collate_temporal(ids, temporal_by_id, device)
audio_mask, empty_audio = _candidate_mask(
batch["times"]["audio"],
batch["valid"]["audio"],
batch["centers"]["audio"],
delta=delta,
mode="local",
)
vision_mask, empty_vision = _candidate_mask(
batch["times"]["vision"],
batch["valid"]["vision"],
batch["centers"]["vision"],
delta=delta,
mode="local",
)
audio_weights, audio_content = branch.attend(
batch["text_content"], batch["values"]["audio"], "audio", audio_mask
)
vision_weights, vision_content = branch.attend(
batch["text_content"], batch["values"]["vision"], "vision", vision_mask
)
loss = _local_contrastive_loss(
branch,
batch["text_content"],
audio_content,
vision_content,
temperature=temperature,
)
if not torch.isfinite(loss):
raise FloatingPointError(f"non-finite TSFA local contrastive loss at fold {fold}, epoch {epoch}")
optimizer.zero_grad(set_to_none=True)
loss.backward()
nn.utils.clip_grad_norm_(branch.parameters(), 1.0)
optimizer.step()
losses.append(float(loss.detach().item()))
history.append(
{
"fold": fold,
"epoch": epoch,
"seed": fold_seed,
"train_loss": float(np.mean(losses)),
"candidate_empty_audio_rows": int(empty_audio.sum().item()),
"candidate_empty_vision_rows": int(empty_vision.sum().item()),
"attention_row_sum_error_audio": float((audio_weights.sum(-1) - 1).abs().max().item()),
"attention_row_sum_error_vision": float((vision_weights.sum(-1) - 1).abs().max().item()),
}
)
branch.eval()
return branch, history
def _candidate_statistics(
*,
method: str,
fold: int,
sample: FeatureSample,
modality: str,
candidate_mask: Tensor,
temporal_weights: Tensor,
times: Tensor,
centers: Tensor,
draw: int,
) -> dict[str, Any]:
mask = candidate_mask[0].to(torch.float32)
weights = temporal_weights[0]
times_row = times[0]
counts = mask.sum(dim=-1)
candidate_mean_times = (mask * times_row[None, :]).sum(dim=-1) / counts.clamp_min(1)
valid_count = int(sample.valid[modality].sum())
return {
"method": method,
"fold": fold,
"sample_id": sample.sample_id,
"video_id": sample.group_id,
"modality": modality,
"draw": draw,
"candidate_count_mean": float(counts.mean().item()),
"candidate_count_min": int(counts.min().item()),
"candidate_count_max": int(counts.max().item()),
"candidate_fraction_valid_mean": float((counts / max(valid_count, 1)).mean().item()),
"temporal_prior_mass_mean": float((weights * mask).sum(dim=-1).mean().item()),
"candidate_mean_center_error": float((candidate_mean_times - centers[0]).abs().mean().item()),
}
def _generate_tsfa_outputs(
*,
method: str,
fold: int,
sample_ids: Sequence[str],
samples_by_id: Mapping[str, FeatureSample],
temporal_by_id: Mapping[str, Mapping[str, Any]],
branch: TSFASemanticBranch,
device: torch.device,
delta: float,
seed: int,
draw: int = 0,
batch_size: int = 8,
) -> tuple[
dict[str, dict[str, np.ndarray]],
dict[str, dict[str, np.ndarray]],
list[dict[str, Any]],
]:
content_by_id: dict[str, dict[str, np.ndarray]] = {}
weights_by_id: dict[str, dict[str, np.ndarray]] = {}
candidate_rows: list[dict[str, Any]] = []
mode = "global" if method == "TSFA-global" else "local"
use_random_centers = method == "TSFA-random"
use_prior = method == "TSFA-multiply"
branch.eval()
with torch.no_grad():
for start in range(0, len(sample_ids), batch_size):
ids = list(sample_ids[start : start + batch_size])
batch = _collate_temporal(ids, temporal_by_id, device)
variant_masks: dict[str, Tensor] = {}
variant_centers: dict[str, Tensor] = {}
empty_rows: dict[str, Tensor] = {}
for modality in ("audio", "vision"):
centers = batch["centers"][modality]
if use_random_centers:
centers = torch.from_numpy(
np.stack(
[
_random_centers(
{**temporal_by_id[sample_id], "sample_id": sample_id},
modality,
seed=seed + fold * 1009,
draw=draw,
)
for sample_id in ids
]
)
).to(device)
candidate_mask, empty = _candidate_mask(
batch["times"][modality],
batch["valid"][modality],
centers,
delta=delta,
mode=mode,
)
variant_masks[modality] = candidate_mask
variant_centers[modality] = centers
empty_rows[modality] = empty
prior_audio = batch["weights"]["audio"] if use_prior else None
prior_vision = batch["weights"]["vision"] if use_prior else None
audio_weights, audio_content = branch.attend(
batch["text_content"],
batch["values"]["audio"],
"audio",
variant_masks["audio"],
temporal_prior=prior_audio,
)
vision_weights, vision_content = branch.attend(
batch["text_content"],
batch["values"]["vision"],
"vision",
variant_masks["vision"],
temporal_prior=prior_vision,
)
for index, sample_id in enumerate(ids):
sample = samples_by_id[sample_id]
record = temporal_by_id[sample_id]
content_by_id[sample_id] = {
"text": record["content"]["text"],
"audio": audio_content[index].cpu().numpy().astype(np.float32, copy=False),
"vision": vision_content[index].cpu().numpy().astype(np.float32, copy=False),
}
weights_by_id[sample_id] = {
"text": record["weights"]["text"],
"audio": audio_weights[index, :, : len(sample.features["audio"])]
.cpu()
.numpy()
.astype(np.float32, copy=False),
"vision": vision_weights[index, :, : len(sample.features["vision"])]
.cpu()
.numpy()
.astype(np.float32, copy=False),
}
for modality in ("audio", "vision"):
length = len(sample.features[modality])
candidate_rows.append(
_candidate_statistics(
method=method,
fold=fold,
sample=sample,
modality=modality,
candidate_mask=variant_masks[modality][index : index + 1, :, :length],
temporal_weights=batch["weights"][modality][index : index + 1, :, :length],
times=batch["times"][modality][index : index + 1, :length],
centers=variant_centers[modality][index : index + 1],
draw=draw,
)
)
return content_by_id, weights_by_id, candidate_rows
def _record_temporal_metrics(
*,
method: str,
fold: int,
sample: FeatureSample,
weights: Mapping[str, np.ndarray],
draw: int = 0,
) -> list[dict[str, Any]]:
normalized_times = {
name: np.asarray(sample.times[name], dtype=np.float64) / max(sample.duration_s, 1e-8)
for name in MODALITIES
}
# A hard TSFA window can leave valid source positions with exactly zero
# attention. Use the same supported positions in both directions so cycle
# and triangle maps have compatible shapes.
source_to_slot = {}
source_indices = {}
rows: list[dict[str, Any]] = []
for modality in MODALITIES:
source_to_slot[modality], source_indices[modality], _ = _normalize_attention(
weights[modality], sample.valid[modality]
)
_, self_metrics = _self_structure(weights[modality])
rows.append({"kind": "self", "modality": modality, **self_metrics})
pair_maps = {}
directions = (*PAIRINGS, *((right, left) for left, right in PAIRINGS))
for left, right in directions:
mapping = source_to_slot[left].T @ np.asarray(
weights[right][:, source_indices[right]], dtype=np.float64
)
mapping /= np.maximum(mapping.sum(axis=1, keepdims=True), 1e-12)
mapping = mapping.astype(np.float32)
pair_maps[f"{left}_{right}"] = (
mapping, source_indices[left], source_indices[right]
)
rows.append({
"kind": "pairwise",
"direction": f"{left}_to_{right}",
**_pair_metrics(
mapping,
normalized_times[left][source_indices[left]],
normalized_times[right][source_indices[right]],
),
})
cycle_rows, _ = _cycle_triangle_metrics(pair_maps, normalized_times)
for row in cycle_rows:
rows.append({"kind": "cycle_triangle", "cycle_kind": row["kind"], **{
key: value for key, value in row.items() if key != "kind"
}})
output = []
for row in rows:
output.append(
{
**row,
"method": method,
"fold": fold,
"sample_id": sample.sample_id,
"video_id": sample.group_id,
"draw": draw,
}
)
for modality in MODALITIES:
mu = np.asarray(weights[modality], dtype=np.float64) @ normalized_times[modality]
output.append(
{
"kind": "span",
"method": method,
"fold": fold,
"sample_id": sample.sample_id,
"video_id": sample.group_id,
"draw": draw,
"modality": modality,
"span_ratio": float(mu[-1] - mu[0]),
"span_abs": float(abs(mu[-1] - mu[0])),
}
)
return output
def _fit_method_probe(
*,
method: str,
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,
history_rows: list[dict[str, Any]],
checkpoint_store: dict[str, Any],
) -> tuple[list[dict[str, Any]], list[dict[str, Any]], dict[str, Tensor]]:
train_ids = [sample.sample_id for sample in train_samples]
validation_ids = [sample.sample_id for sample in validation_samples]
probe_seed = args.seed + fold * 101
projector, 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.probe_learning_rate,
temperature=args.probe_temperature,
)
history_rows.extend(
{"fold": fold, "method": method, "probe": "crossmodal_correspondence", "seed": probe_seed, **row}
for row in history
)
projector.eval()
with torch.no_grad():
validation = _stack_ids(validation_ids, content_by_id, device)
projected = projector(validation)
metrics: list[dict[str, Any]] = []
curves: list[dict[str, Any]] = []
for index, sample in enumerate(validation_samples):
one = {name: projected[name][index] for name in MODALITIES}
metrics.extend(
_sample_metrics(
method=method,
fold=fold,
sample=sample,
projected=one,
curve_rows=curves,
)
)
checkpoint_store[f"fold_{fold:02d}/{method}/correspondence_probe"] = {
"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 metrics, curves, projected
def _shuffle_control_rows(
*,
method: str,
fold: int,
validation_samples: Sequence[FeatureSample],
projected: Mapping[str, Tensor],
repeats: int,
seed: int,
device: torch.device,
) -> list[dict[str, Any]]:
rng = np.random.default_rng(seed + fold * 997)
rows = []
with torch.no_grad():
for index, sample in enumerate(validation_samples):
for repeat in range(repeats):
permuted = {
name: projected[name][index].index_select(
0, torch.as_tensor(rng.permutation(GRID_SIZE), device=device)
)
for name in MODALITIES
}
for left, right in PAIRINGS:
scores = (permuted[left] @ permuted[right].T).cpu().numpy()
rows.append(
_content_pair_metrics(
scores,
method=method,
fold=fold,
sample=sample,
pair=f"{left}_{right}",
control="independent_within_clip_permutation",
shuffle_id=repeat,
)
)
return rows
def _fit_and_score_decoder(
*,
method: str,
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,
history_rows: list[dict[str, Any]],
checkpoint_store: dict[str, Any],
) -> list[dict[str, Any]]:
seed = args.seed + fold * 211
torch.manual_seed(seed)
if device.type == "cuda":
torch.cuda.manual_seed_all(seed)
train = _stack_ids([sample.sample_id for sample in train_samples], content_by_id, device)
decoder = ContentReconstructionProbe(train["text"].shape[-1]).to(device)
optimizer = torch.optim.AdamW(decoder.parameters(), lr=args.decoder_learning_rate, weight_decay=1e-4)
rng = np.random.default_rng(seed)
decoder.train()
for epoch in range(1, args.decoder_epochs + 1):
order = rng.permutation(len(train_samples))
losses = []
for start in range(0, len(order), args.batch_size):
indexes = torch.as_tensor(order[start : start + args.batch_size], device=device)
batch = {name: value.index_select(0, indexes) for name, value in train.items()}
loss_rows = []
for target in MODALITIES:
left, right, _ = DECODER_TARGETS[target]
prediction = decoder(target, batch[left], batch[right])
loss_rows.append(F.smooth_l1_loss(prediction, batch[target]))
loss = torch.stack(loss_rows).mean()
optimizer.zero_grad(set_to_none=True)
loss.backward()
nn.utils.clip_grad_norm_(decoder.parameters(), 1.0)
optimizer.step()
losses.append(float(loss.detach().item()))
history_rows.append(
{"fold": fold, "method": method, "probe": "shifted_reconstruction", "seed": seed,
"epoch": epoch, "train_loss": float(np.mean(losses))}
)
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_prediction = decoder(target, left_eval, right_aligned)
shifted_prediction = decoder(target, left_shifted, right_shifted)
target_eval = target_values.index_select(0, target_idx)
aligned_mae = float((aligned_prediction - target_eval).abs().mean().item())
shifted_mae = float((shifted_prediction - target_eval).abs().mean().item())
rows.append(
{
"method": method,
"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}/{method}/reconstruction_probe"] = {
"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 _content_summary(
rows: Sequence[Mapping[str, Any]],
curves: Sequence[Mapping[str, Any]],
seed: int,
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
metric_names = (
"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_audio_to_text", "within_pm1_r1_audio_to_text",
"mase_slots_audio_to_text", "exact_r1_text_to_vision", "within_pm1_r1_text_to_vision",
"mase_slots_text_to_vision", "exact_r1_vision_to_text", "within_pm1_r1_vision_to_text",
"mase_slots_vision_to_text", "exact_r1_audio_to_vision", "within_pm1_r1_audio_to_vision",
"mase_slots_audio_to_vision", "exact_r1_vision_to_audio", "within_pm1_r1_vision_to_audio",
"mase_slots_vision_to_audio",
)
grouped: dict[tuple[str, str], list[Mapping[str, Any]]] = defaultdict(list)
for row in rows:
grouped[(str(row["method"]), str(row["pair"]))].append(row)
summary = []
for (method, pair), values in sorted(grouped.items()):
item: dict[str, Any] = {
"method": method,
"pair": pair,
"clip_count": len({row["sample_id"] for row in values}),
"draw_count": len(values) / max(len({row["sample_id"] for row in values}), 1),
}
for metric_index, metric in enumerate(metric_names):
if metric not in values[0] or values[0][metric] in ("", None):
continue
mean, low, high, groups = _cluster_bootstrap(
values,
metric,
seed=seed + metric_index + sum(ord(char) for char in method + pair),
repetitions=2000,
)
item[f"{metric}_video_macro_mean"] = mean
item[f"{metric}_ci95_low"] = low
item[f"{metric}_ci95_high"] = high
item["video_id_count"] = groups
summary.append(item)
curve_groups: dict[tuple[str, str, int], list[Mapping[str, Any]]] = defaultdict(list)
for row in curves:
curve_groups[(str(row["method"]), str(row["pair"]), int(row["delta"]))].append(row)
curve_summary = []
for (method, pair, delta), values in sorted(curve_groups.items()):
mean, low, high, groups = _cluster_bootstrap(
values,
"similarity",
seed=seed + delta + sum(ord(char) for char in method + pair),
repetitions=2000,
)
curve_summary.append(
{
"method": method,
"pair": pair,
"delta": delta,
"mean_similarity_video_macro": mean,
"ci95_low": low,
"ci95_high": high,
"video_id_count": groups,
}
)
return summary, curve_summary
def _temporal_summary(rows: Sequence[Mapping[str, Any]], seed: int) -> list[dict[str, Any]]:
specifications = (
("span", ("method", "modality"), ("span_ratio", "span_abs")),
("pairwise", ("method", "direction"), ("pairwise_time_mae", "pairwise_signed_lag", "pairwise_time_corr")),
("self", ("method", "modality"), ("near_similarity_offdiag", "far_similarity", "d_self_offdiag", "gram_target_error", "far_slot_leakage")),
("cycle_triangle", ("method", "cycle_kind"), ("cycle_band_error", "cycle_time_mae", "triangle_relative_error", "triangle_mean_absolute_residual")),
)
output = []
for index, (kind, group_keys, metrics) in enumerate(specifications):
selected = [row for row in rows if row.get("kind") == kind]
for item in _bootstrap_summary(selected, group_keys, metrics, seed=seed + index):
item["metric_family"] = kind
output.append(item)
return output
def _ablation_summary(
content_rows: Sequence[Mapping[str, Any]],
temporal_rows: Sequence[Mapping[str, Any]],
seed: int,
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
per_sample_content: dict[tuple[str, str, str], list[float]] = defaultdict(list)
for row in content_rows:
per_sample_content[(str(row["method"]), str(row["sample_id"]), str(row["video_id"]))].append(
float(row["matched_vs_shifted_auc"])
)
per_sample_time: dict[tuple[str, str, str], list[float]] = defaultdict(list)
for row in temporal_rows:
if row.get("kind") == "pairwise" and row.get("direction") in {
"text_to_audio", "text_to_vision", "audio_to_vision"
}:
per_sample_time[(str(row["method"]), str(row["sample_id"]), str(row["video_id"]))].append(
float(row["pairwise_time_mae"])
)
by_method_rows: list[dict[str, Any]] = []
for key, aucs in per_sample_content.items():
time_values = per_sample_time.get(key, [])
if not time_values:
continue
by_method_rows.append(
{
"method": key[0],
"sample_id": key[1],
"video_id": key[2],
"content_auc_mean": float(np.mean(aucs)),
"canonical_pairwise_time_mae_mean": float(np.mean(time_values)),
"temporal_quality_1_minus_mae": float(1 - np.mean(time_values)),
}
)
summary = []
for method in ALL_METHODS:
values = [row for row in by_method_rows if row["method"] == method]
if not values:
continue
result: dict[str, Any] = {
"method": method,
"clip_count": len(values),
"video_id_count": len({row["video_id"] for row in values}),
}
for index, metric in enumerate(("content_auc_mean", "canonical_pairwise_time_mae_mean", "temporal_quality_1_minus_mae")):
mean, low, high, groups = _cluster_bootstrap(
values, metric, seed=seed + index + sum(ord(char) for char in method), repetitions=2000
)
result[f"{metric}_video_macro_mean"] = mean
result[f"{metric}_ci95_low"] = low
result[f"{metric}_ci95_high"] = high
result["video_id_count"] = groups
summary.append(result)
return summary, by_method_rows
def _finalize_existing_outputs(output_dir: Path, seed: int) -> None:
"""Make paired held-out comparisons and a compact shareable result bundle."""
def read_csv(name: str) -> list[dict[str, str]]:
with (output_dir / name).open(newline="", encoding="utf-8-sig") as handle:
return list(csv.DictReader(handle))
def sha256_file(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
for chunk in iter(lambda: handle.read(1024 * 1024), b""):
digest.update(chunk)
return digest.hexdigest()
manifest_path = output_dir / "run_manifest.json"
manifest = json.loads(manifest_path.read_text(encoding="utf-8"))
paths = manifest["input_paths"]
feature_files = sorted(Path(paths["feature_dir"]).glob("*.npz"))
if len(feature_files) != manifest["sample_count"]:
raise ValueError("feature file count disagrees with TSFA sample count")
feature_digest = hashlib.sha256()
for feature_file in feature_files:
feature_digest.update(feature_file.name.encode("utf-8"))
feature_digest.update(sha256_file(feature_file).encode("ascii"))
checkpoint_root = Path(paths["checkpoint_root"])
manifest["input_sha256"] = {
"feature_manifest": sha256_file(Path(paths["feature_manifest"])),
"grouped_splits": sha256_file(Path(paths["grouped_splits"])),
"uv_lock": sha256_file(Path(__file__).resolve().parents[1] / "uv.lock"),
"features_combined": feature_digest.hexdigest(),
"feature_file_count": len(feature_files),
"frozen_alignment_checkpoints": {
f"fold_{fold:02d}/{variant}": sha256_file(_checkpoint_path(checkpoint_root, variant, fold))
for fold in range(1, 6)
for variant in BASELINE_VARIANTS
},
}
manifest_path.write_text(
json.dumps(manifest, ensure_ascii=False, indent=2, allow_nan=False), encoding="utf-8"
)
ablation_rows = read_csv("ablation_by_clip.csv")
content_rows = read_csv("content_metrics_by_clip.csv")
curve_rows = read_csv("shifted_similarity_by_clip.csv")
content_summary, curve_summary = _content_summary(content_rows, curve_rows, seed + 301)
_write_csv(output_dir / "content_metrics_summary.csv", content_summary)
_write_csv(output_dir / "shifted_similarity_summary.csv", curve_summary)
contrasts = (
("TSFA-main", "M4_sourceTime"),
("TSFA-main", "TSFA-random"),
("TSFA-main", "TSFA-global"),
("TSFA-multiply", "TSFA-main"),
("TSFA-main", "M3_noSourceTime"),
)
ablation_by_key = {
(row["method"], row["sample_id"]): row for row in ablation_rows
}
content_by_key: dict[tuple[str, str, str], list[float]] = defaultdict(list)
for row in content_rows:
content_by_key[(row["method"], row["sample_id"], row["pair"])].append(
float(row["matched_vs_shifted_auc"])
)
sample_ids = sorted({row["sample_id"] for row in ablation_rows})
paired_rows = []
metrics = (
("overall", "all_pairs", "content_auc_mean"),
("overall", "all_pairs", "canonical_pairwise_time_mae_mean"),
*(("content_pair", pair, "matched_vs_shifted_auc") for pair in (
"text_audio", "text_vision", "audio_vision"
)),
)
for contrast_index, (left_method, right_method) in enumerate(contrasts):
for metric_index, (family, pair, metric) in enumerate(metrics):
differences = []
for sample_id in sample_ids:
left_row = ablation_by_key[(left_method, sample_id)]
right_row = ablation_by_key[(right_method, sample_id)]
if left_row["video_id"] != right_row["video_id"]:
raise ValueError(f"paired video_id mismatch for {sample_id}")
if family == "overall":
left_value = float(left_row[metric])
right_value = float(right_row[metric])
else:
left_value = float(np.mean(content_by_key[(left_method, sample_id, pair)]))
right_value = float(np.mean(content_by_key[(right_method, sample_id, pair)]))
differences.append({
"video_id": left_row["video_id"],
"difference": left_value - right_value,
})
mean, low, high, video_count = _cluster_bootstrap(
differences,
"difference",
seed=seed + 307 + contrast_index * 101 + metric_index,
repetitions=2000,
)
paired_rows.append({
"left_method": left_method,
"right_method": right_method,
"metric_family": family,
"pair": pair,
"metric": metric,
"clip_count": len(differences),
"video_id_count": video_count,
"left_minus_right_video_macro_mean": mean,
"ci95_low": low,
"ci95_high": high,
})
_write_csv(output_dir / "paired_contrasts.csv", paired_rows)
bundle = output_dir / "report_bundle"
bundle.mkdir(parents=True, exist_ok=True)
names = (
"run_manifest.json", "tsfa_config.json", "ablation_summary.csv",
"paired_contrasts.csv", "content_metrics_summary.csv",
"temporal_metrics_summary.csv", "candidate_window_stats.csv",
"content_shuffle_summary.csv", "shifted_reconstruction_summary.csv",
"temporal_semantic_attention.png", "typical_sample_tsfa.png",
"temporal_content_tradeoff.png", "content_shuffle_control.png",
"shifted_similarity_curve.png", "shifted_reconstruction_gain.png",
)
for name in names:
shutil.copy2(output_dir / name, bundle / name)
(bundle / "README.md").write_text(
"# TSFA five-fold result bundle\n\n"
"Seven frozen-alignment comparisons on 100 held-out clips from 37 video IDs. "
"All Python runs use the project uv environment. Read `run_manifest.json` "
"for the complete protocol and limitations. `paired_contrasts.csv` reports "
"left minus right with video-ID-cluster bootstrap confidence intervals. "
"The full `outputs/tsfa/` directory retains per-clip CSVs and probe checkpoints.\n\n"
"The `alignment_only_*` summaries pool identical train-fold-standardized "
"raw source features through each method's alignment matrix. They isolate "
"source selection from native model value/output projections. "
"`tsfa_temporal_diagnostics_summary.csv` reports MVR, span, entropy, and "
"source coverage.\n",
encoding="utf-8",
)
def _time_edges(times: np.ndarray) -> np.ndarray:
values = np.asarray(times, dtype=np.float64)
if len(values) == 1:
return np.array([values[0] - 0.005, values[0] + 0.005])
centers = (values[:-1] + values[1:]) / 2
return np.concatenate(([values[0] - (centers[0] - values[0])], centers, [values[-1] + (values[-1] - centers[-1])]))
def _plot_temporal_semantic(
example: Mapping[str, Any], path: Path
) -> None:
fig, axes = plt.subplots(2, 2, figsize=(13, 8), constrained_layout=True)
matrices = (
("M4 temporal prior: Audio", "audio", example["temporal_weights"]["audio"]),
("M4 temporal prior: Vision", "vision", example["temporal_weights"]["vision"]),
("TSFA semantic: Text→Audio", "audio", example["tsfa_weights"]["audio"]),
("TSFA semantic: Text→Vision", "vision", example["tsfa_weights"]["vision"]),
)
image = None
slot_edges = np.arange(GRID_SIZE + 1, dtype=np.float64) - 0.5
for axis, (title, modality, matrix) in zip(axes.flat, matrices, strict=True):
mask = np.asarray(example["valid"][modality], dtype=bool)
times = np.asarray(example["times"][modality], dtype=np.float64)[mask]
image = axis.pcolormesh(
_time_edges(times), slot_edges, np.asarray(matrix)[:, mask], shading="flat", cmap="magma"
)
axis.set_title(title)
axis.set_xlabel(f"{modality.title()} normalized source time")
axis.set_ylabel("Shared slot i")
assert image is not None
fig.colorbar(image, ax=axes, fraction=0.03, pad=0.02, label="Attention probability")
fig.suptitle(f"Temporal prior and content-only semantic attention — {example['sample_id']}")
fig.savefig(path, dpi=180, bbox_inches="tight")
plt.close(fig)
def _plot_typical_trajectory(example: Mapping[str, Any], path: Path) -> None:
fig, axes = plt.subplots(1, 2, figsize=(12, 4.8), sharex=True, sharey=True, constrained_layout=True)
x = (np.arange(GRID_SIZE) + 0.5) / GRID_SIZE
for axis, modality in zip(axes, ("audio", "vision"), strict=True):
prior = example["temporal_weights"][modality] @ example["times"][modality]
semantic = example["tsfa_weights"][modality] @ example["times"][modality]
axis.plot(x, prior, label="M4 coarse center", color="#2878b5", linewidth=1.4)
axis.plot(x, semantic, label="TSFA content-selected time", color="#e1812c", linewidth=1.1)
axis.plot([0, 1], [0, 1], color="black", linestyle="--", linewidth=0.8)
axis.set_title(modality.title())
axis.set_xlabel("Shared slot time")
axis.grid(alpha=0.2)
axes[0].set_ylabel("Expected normalized source time")
axes[0].legend(frameon=False)
fig.suptitle(f"Temporal prior vs semantic selection — {example['sample_id']}")
fig.savefig(path, dpi=180, bbox_inches="tight")
plt.close(fig)
def _plot_tradeoff(rows: Sequence[Mapping[str, Any]], path: Path) -> None:
fig, axis = plt.subplots(figsize=(8, 6), constrained_layout=True)
colors = {
"M3_noSourceTime": "#777777",
"M3_sourceTime": "#a8a8a8",
"M4_sourceTime": "#2878b5",
"TSFA-main": "#d62728",
"TSFA-multiply": "#ff7f0e",
"TSFA-random": "#9467bd",
"TSFA-global": "#8c564b",
}
for row in rows:
axis.scatter(
row["temporal_quality_1_minus_mae_video_macro_mean"],
row["content_auc_mean_video_macro_mean"],
s=70,
color=colors.get(str(row["method"]), "black"),
)
axis.annotate(str(row["method"]),
(row["temporal_quality_1_minus_mae_video_macro_mean"], row["content_auc_mean_video_macro_mean"]),
xytext=(5, 5), textcoords="offset points", fontsize=8)
axis.axhline(0.5, color="black", linestyle="--", linewidth=0.8)
axis.set_xlabel("Temporal quality: 1 − mean pairwise normalized-time MAE")
axis.set_ylabel("Mean content-only matched-vs-shifted AUC")
axis.set_title("Temporal quality vs content correspondence (descriptive axes, no composite score)")
axis.grid(alpha=0.2)
fig.savefig(path, dpi=180, bbox_inches="tight")
plt.close(fig)
def _plot_shuffle(
normal_rows: Sequence[Mapping[str, Any]], shuffle_rows: Sequence[Mapping[str, Any]], path: Path
) -> None:
methods = SHUFFLE_METHODS
pairs = [f"{left}_{right}" for left, right in PAIRINGS]
fig, axes = plt.subplots(1, 3, figsize=(13, 4.5), sharey=True, constrained_layout=True)
colors = {"M4_sourceTime": "#2878b5", "TSFA-main": "#d62728"}
for axis, pair in zip(axes, pairs, strict=True):
x = np.arange(len(methods))
width = 0.32
for index, method in enumerate(methods):
normal = [row for row in normal_rows if row["method"] == method and row["pair"] == pair]
shuffled = [row for row in shuffle_rows if row["method"] == method and row["pair"] == pair]
mean_normal = _cluster_bootstrap(normal, "matched_vs_shifted_auc", seed=127 + index)[0]
mean_shuffle = _cluster_bootstrap(shuffled, "matched_vs_shifted_auc", seed=227 + index)[0]
axis.bar(index - width / 2, mean_normal, width, color=colors[method], label=f"{method}: aligned")
axis.bar(index + width / 2, mean_shuffle, width, color=colors[method], alpha=0.35, hatch="//",
label=f"{method}: shuffled")
axis.axhline(0.5, color="black", linestyle="--", linewidth=0.8)
axis.set_xticks(x, ["M4", "TSFA"])
axis.set_title(pair.replace("_", "–"))
axis.set_ylim(0.35, 0.8)
axis.grid(axis="y", alpha=0.2)
axes[0].set_ylabel("Matched-vs-shifted AUC")
handles, labels = axes[0].get_legend_handles_labels()
fig.legend(handles, labels, loc="outside lower center", ncol=4, frameon=False)
fig.suptitle("Content shuffle control")
fig.savefig(path, dpi=180, bbox_inches="tight")
plt.close(fig)
def _plot_shift_curve(rows: Sequence[Mapping[str, Any]], path: Path) -> None:
selected_methods = ("M4_sourceTime", "TSFA-main")
color_map = {"M4_sourceTime": "#2878b5", "TSFA-main": "#d62728"}
pairs = [f"{left}_{right}" for left, right in PAIRINGS]
fig, axes = plt.subplots(1, 3, figsize=(14, 4.5), sharey=True, constrained_layout=True)
for axis, pair in zip(axes, pairs, strict=True):
for method in selected_methods:
xs, means, lows, highs = [], [], [], []
for delta in range(-10, 11):
values = [row for row in rows if row["method"] == method and row["pair"] == pair and int(row["delta"]) == delta]
if not values:
continue
mean, low, high, _ = _cluster_bootstrap(
values, "similarity", seed=811 + delta + sum(ord(char) for char in method + pair), repetitions=1000
)
xs.append(delta)
means.append(mean)
lows.append(low)
highs.append(high)
axis.plot(xs, means, color=color_map[method], label=method, linewidth=1.4)
axis.fill_between(xs, lows, highs, color=color_map[method], alpha=0.12)
axis.axvline(0, color="black", linestyle="--", linewidth=0.8)
axis.set_title(pair.replace("_", "–"))
axis.set_xlabel("Slot shift Δ")
axis.grid(alpha=0.2)
axes[0].set_ylabel("Content-only cosine similarity")
axes[0].legend(frameon=False)
fig.suptitle("Content correspondence under temporal shift")
fig.savefig(path, dpi=180, bbox_inches="tight")
plt.close(fig)
def _plot_reconstruction(rows: Sequence[Mapping[str, Any]], path: Path) -> None:
colors = {"M4_sourceTime": "#2878b5", "TSFA-main": "#d62728"}
fig, axes = plt.subplots(1, 3, figsize=(14, 4.5), sharey=True, constrained_layout=True)
for axis, target in zip(axes, MODALITIES, strict=True):
for method in RECONSTRUCTION_METHODS:
xs, means, lows, highs = [], [], [], []
for distance in SHIFTS:
values = [
row for row in rows
if row["method"] == method and row["target_modality"] == target and int(row["abs_delta"]) == distance
]
if not values:
continue
mean, low, high, _ = _cluster_bootstrap(
values,
"gain_shift_minus_aligned",
seed=1009 + distance + sum(ord(char) for char in method + target),
repetitions=1000,
)
xs.append(distance)
means.append(mean)
lows.append(low)
highs.append(high)
axis.plot(xs, means, marker="o", color=colors[method], label=method)
axis.fill_between(xs, lows, highs, color=colors[method], alpha=0.14)
axis.axhline(0, color="black", linestyle="--", linewidth=0.8)
axis.set_title(f"Reconstruct {target}")
axis.set_xlabel("Absolute shift |Δ| (slots)")
axis.grid(alpha=0.2)
axes[0].set_ylabel("MAE(shifted) − MAE(aligned)")
axes[0].legend(frameon=False)
fig.suptitle("Aligned-vs-shifted reconstruction gain")
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)
samples_by_id = {sample.sample_id: sample for sample in samples}
splits = json.loads(args.splits.read_text(encoding="utf-8"))
if len(splits) != 5:
raise ValueError(f"TSFA requires the existing five grouped folds, found {len(splits)}")
heldout_ids = [sample_id for split in splits for sample_id in split["validation_sample_ids"]]
if len(heldout_ids) != len(set(heldout_ids)) or set(heldout_ids) != set(samples_by_id):
raise ValueError("five-fold validation partitions must cover every sample exactly once")
output_dir = args.output_dir
output_dir.mkdir(parents=True, exist_ok=True)
content_rows: list[dict[str, Any]] = []
curve_rows: list[dict[str, Any]] = []
temporal_rows: list[dict[str, Any]] = []
candidate_rows: list[dict[str, Any]] = []
shuffle_rows: list[dict[str, Any]] = []
reconstruction_rows: list[dict[str, Any]] = []
history_rows: list[dict[str, Any]] = []
checkpoint_store: dict[str, Any] = {}
fold_manifest = []
example_payload: dict[str, Any] | None = None
for split in splits:
fold = int(split["fold"])
train_samples = [samples_by_id[sample_id] for sample_id in split["train_sample_ids"]]
validation_samples = [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}: {sorted(train_groups & validation_groups)}")
feature_stats = fit_feature_stats(train_samples)
baseline_content, baseline_weights, temporal_by_id = _collect_fold_features(
fold=fold,
train_samples=train_samples,
validation_samples=validation_samples,
feature_stats=feature_stats,
checkpoint_root=args.checkpoint_root,
device=device,
batch_size=args.batch_size,
)
all_ids = [sample.sample_id for sample in [*train_samples, *validation_samples]]
all_ids_by_id = {sample.sample_id: sample for sample in [*train_samples, *validation_samples]}
branch, semantic_history = _fit_semantic_branch(
fold=fold,
train_samples=train_samples,
temporal_by_id=temporal_by_id,
device=device,
seed=args.seed,
epochs=args.semantic_epochs,
batch_size=args.batch_size,
learning_rate=args.semantic_learning_rate,
temperature=args.local_temperature,
delta=args.delta,
)
history_rows.extend(semantic_history)
checkpoint_store[f"fold_{fold:02d}/TSFA-semantic"] = {
"seed": args.seed + fold * 101,
"state_dict": {key: value.detach().cpu() for key, value in branch.state_dict().items()},
}
fold_content: dict[str, dict[str, dict[str, np.ndarray]]] = {
method: baseline_content[method] for method in BASELINE_VARIANTS
}
fold_weights: dict[str, dict[str, dict[str, np.ndarray]]] = {
method: baseline_weights[method] for method in BASELINE_VARIANTS
}
fold_candidate_rows: dict[str, list[dict[str, Any]]] = {}
for method in TSFA_VARIANTS:
generated_content, generated_weights, generated_candidates = _generate_tsfa_outputs(
method=method,
fold=fold,
sample_ids=all_ids,
samples_by_id=all_ids_by_id,
temporal_by_id=temporal_by_id,
branch=branch,
device=device,
delta=args.delta,
seed=args.seed,
draw=0,
batch_size=args.batch_size,
)
fold_content[method] = generated_content
fold_weights[method] = generated_weights
fold_candidate_rows[method] = generated_candidates
for method in TSFA_VARIANTS:
candidate_rows.extend(
row for row in fold_candidate_rows[method] if row["sample_id"] in {s.sample_id for s in validation_samples}
)
for method in ALL_METHODS:
method_metrics, method_curves, projected = _fit_method_probe(
method=method,
fold=fold,
train_samples=train_samples,
validation_samples=validation_samples,
content_by_id=fold_content[method],
device=device,
args=args,
history_rows=history_rows,
checkpoint_store=checkpoint_store,
)
content_rows.extend(method_metrics)
curve_rows.extend(method_curves)
for sample in validation_samples:
temporal_rows.extend(
_record_temporal_metrics(
method=method,
fold=fold,
sample=sample,
weights=fold_weights[method][sample.sample_id],
)
)
if method in SHUFFLE_METHODS:
shuffle_rows.extend(
_shuffle_control_rows(
method=method,
fold=fold,
validation_samples=validation_samples,
projected=projected,
repeats=args.shuffle_repeats,
seed=args.seed,
device=device,
)
)
if method in RECONSTRUCTION_METHODS:
reconstruction_rows.extend(
_fit_and_score_decoder(
method=method,
fold=fold,
train_samples=train_samples,
validation_samples=validation_samples,
content_by_id=fold_content[method],
device=device,
args=args,
history_rows=history_rows,
checkpoint_store=checkpoint_store,
)
)
if method == "TSFA-random":
for draw in range(1, args.random_window_repeats):
random_content, random_weights, random_candidate = _generate_tsfa_outputs(
method=method,
fold=fold,
sample_ids=[sample.sample_id for sample in validation_samples],
samples_by_id=samples_by_id,
temporal_by_id=temporal_by_id,
branch=branch,
device=device,
delta=args.delta,
seed=args.seed,
draw=draw,
batch_size=args.batch_size,
)
candidate_rows.extend(random_candidate)
random_metrics, random_curves, _ = _evaluate_fixed_projector(
method=method,
fold=fold,
validation_samples=validation_samples,
content_by_id=random_content,
projector_state=checkpoint_store[f"fold_{fold:02d}/{method}/correspondence_probe"]["state_dict"],
device=device,
)
content_rows.extend({**row, "draw": draw} for row in random_metrics)
curve_rows.extend({**row, "draw": draw} for row in random_curves)
for sample in validation_samples:
temporal_rows.extend(
_record_temporal_metrics(
method=method,
fold=fold,
sample=sample,
weights=random_weights[sample.sample_id],
draw=draw,
)
)
if example_payload is None:
example_id = args.example_id
if example_id not in samples_by_id or example_id not in fold_weights["M4_sourceTime"]:
example_id = validation_samples[0].sample_id
if example_id in fold_weights["M4_sourceTime"]:
example_payload = {
"sample_id": example_id,
"temporal_weights": fold_weights["M4_sourceTime"][example_id],
"tsfa_weights": fold_weights["TSFA-main"][example_id],
"times": temporal_by_id[example_id]["times"],
"valid": temporal_by_id[example_id]["valid"],
"fold": fold,
}
del projected
fold_manifest.append(
{
"fold": fold,
"train_count": len(train_samples),
"heldout_count": len(validation_samples),
"train_video_id_count": len(train_groups),
"heldout_video_id_count": len(validation_groups),
"video_id_overlap": sorted(train_groups & validation_groups),
"validation_sample_ids": [sample.sample_id for sample in validation_samples],
}
)
print(
f"[TSFA fold {fold}] frozen M3/M4 baselines loaded; semantic epochs={args.semantic_epochs}; "
f"train={len(train_samples)} heldout={len(validation_samples)}; "
f"content rows={len(content_rows)} temporal rows={len(temporal_rows)}",
flush=True,
)
del branch, baseline_content, baseline_weights, temporal_by_id, fold_content, fold_weights
if device.type == "cuda":
torch.cuda.empty_cache()
if example_payload is None:
raise RuntimeError("no held-out example was collected")
example_id = example_payload["sample_id"]
# The stored example maps need to be from the same fold as the sample.
example_fold = int(example_payload["fold"])
example_split = next(split for split in splits if int(split["fold"]) == example_fold)
example_train = [samples_by_id[sample_id] for sample_id in example_split["train_sample_ids"]]
example_val = [samples_by_id[sample_id] for sample_id in example_split["validation_sample_ids"]]
example_stats = fit_feature_stats(example_train)
_, _, example_temporal = _collect_fold_features(
fold=example_fold,
train_samples=example_train,
validation_samples=example_val,
feature_stats=example_stats,
checkpoint_root=args.checkpoint_root,
device=device,
batch_size=args.batch_size,
)
example_split_samples = {sample.sample_id: sample for sample in [*example_train, *example_val]}
example_fold_content, example_fold_weights, _ = _generate_tsfa_outputs(
method="TSFA-main",
fold=example_fold,
sample_ids=[example_id],
samples_by_id=samples_by_id,
temporal_by_id=example_temporal,
branch=_load_semantic_checkpoint(checkpoint_store, example_fold, device),
device=device,
delta=args.delta,
seed=args.seed,
draw=0,
batch_size=1,
)
example_payload["temporal_weights"] = example_temporal[example_id]["weights"]
example_payload["tsfa_weights"] = example_fold_weights[example_id]
example_payload["times"] = example_temporal[example_id]["times"]
example_payload["valid"] = example_temporal[example_id]["valid"]
np.savez_compressed(
output_dir / "typical_sample_tsfa_maps.npz",
sample_id=np.asarray(example_id),
**{f"A_tau_{name}": example_payload["temporal_weights"][name] for name in MODALITIES},
**{f"A_final_{name}": example_payload["tsfa_weights"][name] for name in MODALITIES},
**{f"tau_{name}": example_payload["times"][name] for name in MODALITIES},
**{f"valid_{name}": example_payload["valid"][name] for name in MODALITIES},
)
_plot_temporal_semantic(example_payload, output_dir / "temporal_semantic_attention.png")
_plot_typical_trajectory(example_payload, output_dir / "typical_sample_tsfa.png")
content_summary, curve_summary = _content_summary(content_rows, curve_rows, args.seed + 301)
temporal_summary = _temporal_summary(temporal_rows, args.seed + 302)
shuffle_summary = _bootstrap_summary(
shuffle_rows,
("method", "control", "pair"),
("matched_vs_shifted_auc", "same_minus_shifted_margin", "exact_r1_left_to_right", "within_pm1_r1_left_to_right"),
seed=args.seed + 303,
)
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 + 304,
)
candidate_summary = _bootstrap_summary(
candidate_rows,
("method", "modality"),
("candidate_count_mean", "candidate_fraction_valid_mean", "temporal_prior_mass_mean", "candidate_mean_center_error"),
seed=args.seed + 305,
)
ablation_summary, ablation_by_clip = _ablation_summary(content_rows, temporal_rows, args.seed + 306)
_write_csv(output_dir / "content_metrics_by_clip.csv", content_rows)
_write_csv(output_dir / "content_metrics_summary.csv", content_summary)
_write_csv(output_dir / "shifted_similarity_by_clip.csv", curve_rows)
_write_csv(output_dir / "shifted_similarity_summary.csv", curve_summary)
_write_csv(output_dir / "temporal_metrics_by_clip.csv", temporal_rows)
_write_csv(output_dir / "temporal_metrics_summary.csv", temporal_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 / "candidate_window_by_clip.csv", candidate_rows)
_write_csv(output_dir / "candidate_window_stats.csv", candidate_summary)
_write_csv(output_dir / "ablation_by_clip.csv", ablation_by_clip)
_write_csv(output_dir / "ablation_summary.csv", ablation_summary)
_write_csv(output_dir / "probe_training_history.csv", history_rows)
torch.save(checkpoint_store, output_dir / "probe_checkpoints.pt")
_plot_tradeoff(ablation_summary, output_dir / "temporal_content_tradeoff.png")
_plot_shuffle(content_rows, shuffle_rows, output_dir / "content_shuffle_control.png")
_plot_shift_curve(curve_rows, output_dir / "shifted_similarity_curve.png")
_plot_reconstruction(reconstruction_rows, output_dir / "shifted_reconstruction_gain.png")
manifest = {
"created_utc": datetime.now(timezone.utc).isoformat(),
"experiment": "TSFA: Temporal-Semantic Factorized Alignment",
"sample_count": len(samples),
"fold_count": len(splits),
"heldout_partition_count": len(heldout_ids),
"unique_video_id_count": len({sample.group_id for sample in samples}),
"folds": fold_manifest,
"alignment_models_retrained": False,
"temporal_branch": "frozen M4_sourceTime D5 checkpoint per grouped fold",
"semantic_branch_trained": True,
"emotion_labels_used_for_alignment": False,
"feature_extractors_changed": False,
"feature_encoder_note": "Existing BERT text and DeiT image features are reused unchanged.",
"example_sample_id": example_id,
"parameters": {
"seed": args.seed,
"grid_size": GRID_SIZE,
"hidden_size": HIDDEN_SIZE,
"heads_in_frozen_alignment": HEADS,
"delta_normalized_time": args.delta,
"hard_negative_offsets_slots": list(HARD_NEGATIVE_OFFSETS),
"semantic_epochs": args.semantic_epochs,
"semantic_learning_rate": args.semantic_learning_rate,
"local_contrastive_temperature": args.local_temperature,
"content_probe_epochs": args.probe_epochs,
"content_probe_temperature": args.probe_temperature,
"reconstruction_probe_epochs": args.decoder_epochs,
"shuffle_repeats": args.shuffle_repeats,
"random_window_repeats": args.random_window_repeats,
"window_ablation_status": "delta sweep and 80-percent attention-mass window deferred until the seven-method MVP is reviewed",
},
"methods": list(ALL_METHODS),
"protocol": {
"semantic_query_key_inputs": "Frozen M4 attention-pooled text value output is the semantic query input; M4 pre-attention projected Audio/Vision content is the semantic key/value input. No explicit source-time code, absolute PE, slot ID, or latent positional vector is passed to semantic Q/K. Temporal selection can still encode time indirectly.",
"baseline_content_values": "M4 and M3 Audio/Vision use each checkpoint's actual MHA attention output. M3 Text uses attention-pooled pre-attention text features to remove its source-time query residual.",
"temporal_role": "Frozen M4 attention supplies candidate masks; TSFA-multiply additionally multiplies semantic probabilities by M4 temporal attention",
"local_contrastive": "same-slot positives with within-clip slot-offset negatives at +/-2, +/-3, +/-5; symmetric Text-Audio and Text-Vision loss",
"random_candidate": "same delta-width windows with independently sampled centers at inference; 20 held-out random draws, one train draw for the train-only probe",
"global_candidate": "all valid source positions available to semantic attention",
"content_probe": "same train-only symmetric within-clip InfoNCE probe applied to each method; held-out negatives are positions more than two slots away",
"content_shuffle": "independently permute each modality's 50 projected content rows within each held-out clip; M4 and TSFA-main",
"reconstruction": "fold-train decoder predicts one aligned content representation from the other two; shift one partner by signed offsets +/-1,2,5,10 on identical target support",
"confidence_intervals": "95-percent video_id-cluster bootstrap; 2,000 repetitions for CSV summaries",
},
"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()),
"m3_m4_checkpoint_variants": list(BASELINE_VARIANTS),
},
"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": [
"Same-slot labels are a temporal training convention, not independent semantic ground truth.",
"Frozen M4 temporal attention was trained with a timestamp-derived Gaussian prior; time may enter content indirectly through candidate selection.",
"AUC, retrieval, and shifted similarity measure within-clip content matchability, not human event alignment.",
"Reconstruction predicts M4/TSFA pooled content values, not raw waveform or image pixels.",
"Five folds use one alignment-model seed; bootstrap intervals quantify video-group sampling uncertainty, not seed uncertainty.",
"Human event IoU/center error remains pending event annotation.",
],
}
(output_dir / "tsfa_config.json").write_text(
json.dumps(manifest["parameters"] | manifest["protocol"] | {"methods": list(ALL_METHODS)},
ensure_ascii=False, indent=2, allow_nan=False),
encoding="utf-8",
)
(output_dir / "run_manifest.json").write_text(
json.dumps(manifest, ensure_ascii=False, indent=2, allow_nan=False), encoding="utf-8"
)
_finalize_existing_outputs(output_dir, args.seed)
print(
f"[TSFA complete] samples={len(samples)} folds={len(splits)} "
f"elapsed={manifest['elapsed_seconds']:.1f}s output={output_dir}",
flush=True,
)
return manifest
def _evaluate_fixed_projector(
*,
method: str,
fold: int,
validation_samples: Sequence[FeatureSample],
content_by_id: Mapping[str, Mapping[str, np.ndarray]],
projector_state: Mapping[str, Tensor],
device: torch.device,
) -> tuple[list[dict[str, Any]], list[dict[str, Any]], dict[str, Tensor]]:
ids = [sample.sample_id for sample in validation_samples]
raw = _stack_ids(ids, content_by_id, device)
dimensions = {name: int(raw[name].shape[-1]) for name in MODALITIES}
projector = CorrespondenceProjection(dimensions).to(device)
projector.load_state_dict(projector_state, strict=True)
projector.eval()
with torch.no_grad():
projected = projector(raw)
metrics = []
curves = []
for index, sample in enumerate(validation_samples):
one = {name: projected[name][index] for name in MODALITIES}
metrics.extend(
_sample_metrics(method=method, fold=fold, sample=sample, projected=one, curve_rows=curves)
)
del projector
if device.type == "cuda":
torch.cuda.empty_cache()
return metrics, curves, projected
def _load_semantic_checkpoint(
store: Mapping[str, Any], fold: int, device: torch.device
) -> TSFASemanticBranch:
key = f"fold_{fold:02d}/TSFA-semantic"
branch = TSFASemanticBranch(HIDDEN_SIZE).to(device)
branch.load_state_dict(store[key]["state_dict"], strict=True)
branch.eval()
return branch
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("--seed", type=int, default=42)
parser.add_argument("--delta", type=float, default=0.10)
parser.add_argument("--semantic-epochs", type=int, default=40)
parser.add_argument("--semantic-learning-rate", type=float, default=1e-3)
parser.add_argument("--local-temperature", type=float, default=0.1)
parser.add_argument("--batch-size", type=int, default=8)
parser.add_argument("--probe-epochs", type=int, default=40)
parser.add_argument("--probe-learning-rate", type=float, default=1e-3)
parser.add_argument("--probe-temperature", type=float, default=0.1)
parser.add_argument("--decoder-epochs", type=int, default=40)
parser.add_argument("--decoder-learning-rate", type=float, default=1e-3)
parser.add_argument("--shuffle-repeats", type=int, default=20)
parser.add_argument("--random-window-repeats", type=int, default=20)
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/tsfa")
parser.add_argument("--finalize-existing", action="store_true", help="summarize an existing completed run without retraining")
return parser
def main() -> None:
args = build_parser().parse_args()
if args.finalize_existing:
_finalize_existing_outputs(args.output_dir, args.seed)
print(f"[TSFA finalize] output={args.output_dir}", flush=True)
return
if not 0 < args.delta <= 1:
raise ValueError("--delta must be in (0,1]")
if args.random_window_repeats < 1:
raise ValueError("--random-window-repeats must be at least 1")
run(args)
if __name__ == "__main__":
main()