1775 lines
80 KiB
Python
1775 lines
80 KiB
Python
"""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()
|