建立分批同步基线(基础文件)

This commit is contained in:
gloamxun committed 2026-09-23 23:24:01 +08:00
commit 7fc76aaafd
70 files changed
+18635

No files matched your search

+14
View File
@@ -0,0 +1,14 @@
"""Core data structures and alignment methods for Q1."""
from .alignment import align_fixed_windows, align_forced_timestamps
from .models import SharedLatentTimeline, TextAnchoredCrossAttention
from .types import AlignmentOutput, SequenceBatch
__all__ = [
"AlignmentOutput",
"SequenceBatch",
"SharedLatentTimeline",
"TextAnchoredCrossAttention",
"align_fixed_windows",
"align_forced_timestamps",
]
+188
View File
@@ -0,0 +1,188 @@
from __future__ import annotations
from collections.abc import Sequence
import torch
from torch import Tensor
from .types import AlignmentOutput, MODALITIES, SequenceBatch
def uniform_time_intervals(durations: Tensor, grid_size: int) -> Tensor:
"""Return ``[B, K, 2]`` equal-duration windows in seconds."""
if durations.ndim != 1 or grid_size < 1:
raise ValueError("durations must be [B] and grid_size must be positive")
if bool((durations <= 0).any()) or not bool(torch.isfinite(durations).all()):
raise ValueError("durations must be finite and positive")
edges = torch.linspace(
0.0, 1.0, grid_size + 1, device=durations.device, dtype=durations.dtype
)[None, :] * durations[:, None]
return torch.stack((edges[:, :-1], edges[:, 1:]), dim=-1)
def word_intervals_to_grid(word_intervals: Tensor, grid_size: int) -> Tensor:
"""Resample ordered word spans into K consecutive text-order intervals.
Each grid slot covers an equal share of transcript word order. Its time
boundaries are interpolated from forced-alignment word boundaries, so long
and short words retain their actual duration on the audio/video timeline.
"""
if word_intervals.ndim != 2 or word_intervals.shape[1] != 2:
raise ValueError("word_intervals must have shape [word_count, 2]")
if word_intervals.shape[0] == 0 or grid_size < 1:
raise ValueError("at least one word interval and a positive grid size are required")
intervals = word_intervals.to(dtype=torch.float32)
if not bool(torch.isfinite(intervals).all()):
raise ValueError("word intervals must be finite")
if bool((intervals[:, 1] < intervals[:, 0]).any()):
raise ValueError("word interval end must not precede its start")
if bool((intervals[1:, 0] < intervals[:-1, 0]).any()):
raise ValueError("word intervals must be ordered by start time")
word_count = intervals.shape[0]
if word_count == 1:
boundaries = torch.cat((intervals[:1, 0], intervals[:1, 1]))
else:
between = (intervals[:-1, 1] + intervals[1:, 0]) / 2
boundaries = torch.cat((intervals[:1, 0], between, intervals[-1:, 1]))
boundaries = torch.cummax(boundaries, dim=0).values
positions = torch.linspace(
0, word_count, grid_size + 1, device=intervals.device, dtype=intervals.dtype
)
left = positions.floor().long().clamp(max=word_count)
right = (left + 1).clamp(max=word_count)
fraction = (positions - left.to(positions.dtype)).unsqueeze(-1)
time_edges = boundaries[left] + fraction.squeeze(-1) * (boundaries[right] - boundaries[left])
return torch.stack((time_edges[:-1], time_edges[1:]), dim=-1)
def index_alignment(valid: Tensor, grid_size: int) -> tuple[Tensor, int]:
"""Map equal ranges of valid sequence order to K grid slots."""
if valid.ndim != 2 or valid.dtype != torch.bool:
raise ValueError("valid must be a boolean [B, L] tensor")
if grid_size < 1:
raise ValueError("grid_size must be positive")
batch_size, length = valid.shape
result = torch.zeros(batch_size, grid_size, length, device=valid.device, dtype=torch.float32)
fallback_count = 0
for batch_index in range(batch_size):
positions = torch.nonzero(valid[batch_index], as_tuple=False).flatten()
count = positions.numel()
if count == 0:
raise ValueError("each sample must contain a valid position")
for grid_index in range(grid_size):
start = (grid_index * count) // grid_size
end = ((grid_index + 1) * count) // grid_size
if start == end:
source_index = min(int((grid_index + 0.5) * count / grid_size), count - 1)
result[batch_index, grid_index, positions[source_index]] = 1.0
fallback_count += 1
else:
chosen = positions[start:end]
result[batch_index, grid_index, chosen] = 1.0 / chosen.numel()
return result, fallback_count
def interval_alignment(times: Tensor, valid: Tensor, intervals: Tensor) -> tuple[Tensor, int]:
"""Create row-normalized interval membership weights with nearest-time fallback."""
if times.ndim != 2 or valid.shape != times.shape or valid.dtype != torch.bool:
raise ValueError("times and valid must have matching [B, L] shapes")
if intervals.ndim != 3 or intervals.shape[0] != times.shape[0] or intervals.shape[2] != 2:
raise ValueError("intervals must have shape [B, K, 2]")
batch_size, length = times.shape
grid_size = intervals.shape[1]
result = torch.zeros(batch_size, grid_size, length, device=times.device, dtype=torch.float32)
fallback_count = 0
for batch_index in range(batch_size):
valid_positions = torch.nonzero(valid[batch_index], as_tuple=False).flatten()
valid_times = times[batch_index, valid_positions]
for grid_index in range(grid_size):
start, end = intervals[batch_index, grid_index]
is_last = grid_index == grid_size - 1
in_window = (valid_times >= start) & (
(valid_times <= end) if is_last else (valid_times < end)
)
chosen = valid_positions[in_window]
if chosen.numel() > 0:
result[batch_index, grid_index, chosen] = 1.0 / chosen.numel()
else:
center = (start + end) / 2
nearest = valid_positions[torch.argmin((valid_times - center).abs())]
result[batch_index, grid_index, nearest] = 1.0
fallback_count += 1
return result, fallback_count
def _output_from_weights(
sequences: dict[str, SequenceBatch],
weights: dict[str, Tensor],
fallbacks: dict[str, int],
) -> AlignmentOutput:
aligned = {name: torch.bmm(weights[name].to(sequences[name].features.dtype), sequences[name].features)
for name in MODALITIES}
output = AlignmentOutput(weights=weights, aligned=aligned, fallback_rows=fallbacks)
output.validate({name: sequences[name].valid for name in MODALITIES})
return output
def align_fixed_windows(
sequences: dict[str, SequenceBatch], durations: Tensor, grid_size: int = 50
) -> AlignmentOutput:
"""M2: average each modality inside the same K equal-duration windows."""
intervals = uniform_time_intervals(durations, grid_size)
weights: dict[str, Tensor] = {}
fallbacks: dict[str, int] = {}
for name in MODALITIES:
weights[name], fallbacks[name] = interval_alignment(
sequences[name].times, sequences[name].valid, intervals
)
return _output_from_weights(sequences, weights, fallbacks)
def align_forced_timestamps(
sequences: dict[str, SequenceBatch],
word_intervals: Sequence[Tensor],
grid_size: int = 50,
) -> AlignmentOutput:
"""M1: use forced word timestamps to define text-ordered grid intervals."""
batch_size = sequences["text"].features.shape[0]
if len(word_intervals) != batch_size:
raise ValueError("provide one ordered word-interval array per sample")
device = sequences["text"].features.device
intervals = torch.stack(
[word_intervals_to_grid(spans.to(device), grid_size) for spans in word_intervals], dim=0
)
weights: dict[str, Tensor] = {}
fallbacks: dict[str, int] = {}
for name in MODALITIES:
weights[name], fallbacks[name] = interval_alignment(
sequences[name].times, sequences[name].valid, intervals
)
return _output_from_weights(sequences, weights, fallbacks)
def make_block_mask(
batch_size: int,
grid_size: int,
ratio: float,
device: torch.device | str,
generator: torch.Generator | None = None,
) -> Tensor:
"""Sample one continuous masked interval per sequence on the common grid."""
if not 0 < ratio < 1:
raise ValueError("ratio must be between zero and one")
if batch_size < 1 or grid_size < 2:
raise ValueError("batch_size must be positive and grid_size at least two")
block_length = min(max(1, round(grid_size * ratio)), grid_size - 1)
mask = torch.zeros(batch_size, grid_size, dtype=torch.bool, device=device)
starts = torch.randint(
0,
grid_size - block_length + 1,
(batch_size,),
device=device,
generator=generator,
)
offsets = torch.arange(block_length, device=device)
mask[torch.arange(batch_size, device=device)[:, None], starts[:, None] + offsets] = True
return mask
+776
View File
@@ -0,0 +1,776 @@
from __future__ import annotations
import argparse
import csv
import json
import math
import platform
import random
import time
from collections import defaultdict
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Mapping, Sequence
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import torch
import torch.nn.functional as F
from torch import Tensor, nn
from .alignment import make_block_mask
from .compare_methods import AlignmentReconstructor, _write_csv
from .experiment_data import (
FeatureSample,
FeatureStats,
collate_feature_samples,
fit_feature_stats,
load_feature_samples,
)
from .losses import (
cross_modal_contrastive_loss,
temporal_span_loss,
weak_temporal_band_loss,
)
from .metrics import (
alignment_trajectory,
attention_row_similarity,
normalized_attention_entropy,
monotonicity_violation_rate,
)
from .models import SharedLatentTimeline, TextAnchoredCrossAttention
from .types import MODALITIES
EXAMPLE_ID = "-tPCytz4rww/12"
SIGMA = 0.10
GRID_SIZE = 50
HIDDEN_SIZE = 128
HEADS = 4
SINGLE_SAMPLE_STEPS = 1000
FULL_DATA_STEPS = 500
BATCH_SIZE = 8
LEARNING_RATE = 1e-3
LOG_INTERVAL = 20
def _seed_everything(seed: int) -> None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)
torch.backends.cudnn.deterministic = True
torch.backends.cudnn.benchmark = False
def _model(
method: str,
dimensions: Mapping[str, int],
*,
absolute_position_encoding: bool = False,
) -> nn.Module:
if method == "M3":
return TextAnchoredCrossAttention(
dimensions,
grid_size=GRID_SIZE,
hidden_size=HIDDEN_SIZE,
heads=HEADS,
dropout=0.0,
)
if method == "M4":
return SharedLatentTimeline(
dimensions,
grid_size=GRID_SIZE,
hidden_size=HIDDEN_SIZE,
heads=HEADS,
dropout=0.0,
absolute_position_encoding=absolute_position_encoding,
)
raise ValueError(f"unknown method: {method}")
def _centers(
method: str,
output: Any,
sequences: Mapping[str, Any],
durations: Tensor,
) -> tuple[dict[str, Tensor], Tensor]:
batch_size = durations.shape[0]
if method == "M3":
text_centers = torch.bmm(
output.weights["text"], sequences["text"].times.unsqueeze(-1)
).squeeze(-1)
text_centers = text_centers / durations[:, None].clamp_min(1e-8)
return {"audio": text_centers, "vision": text_centers}, text_centers
centers = (
torch.arange(GRID_SIZE, dtype=durations.dtype, device=durations.device) + 0.5
) / GRID_SIZE
centers = centers.unsqueeze(0).expand(batch_size, -1)
return {name: centers for name in MODALITIES}, centers
def _gaussian_targets(
method: str,
output: Any,
sequences: Mapping[str, Any],
durations: Tensor,
) -> dict[str, Tensor]:
centers, _ = _centers(method, output, sequences, durations)
target_names = ("audio", "vision") if method == "M3" else MODALITIES
targets: dict[str, Tensor] = {}
for name in target_names:
times = sequences[name].times / durations[:, None].clamp_min(1e-8)
difference = (times[:, None, :] - centers[name][:, :, None]) / SIGMA
logits = -0.5 * difference.square()
logits = logits.masked_fill(~sequences[name].valid[:, None, :], -torch.inf)
targets[name] = torch.softmax(logits, dim=-1)
return targets
def _gaussian_alignment_kl(output: Any, targets: Mapping[str, Tensor]) -> Tensor:
losses = []
for name, target in targets.items():
predicted = output.weights[name].clamp_min(1e-8)
safe_target = target.clamp_min(1e-12)
row_kl = (safe_target * (safe_target.log() - predicted.log())).sum(dim=-1)
losses.append(row_kl.mean())
return torch.stack(losses).mean()
def _training_components(
experiment: str,
method: str,
output: Any,
sequences: Mapping[str, Any],
durations: Tensor,
*,
decoder: AlignmentReconstructor | None,
block_generator: torch.Generator,
) -> tuple[Tensor, dict[str, Tensor], tuple[str, ...]]:
centers, _ = _centers(method, output, sequences, durations)
coverage_modalities = ("audio", "vision") if method == "M3" else MODALITIES
span = temporal_span_loss(
output,
{name: sequences[name].times for name in MODALITIES},
durations,
minimum_span=0.7,
modalities=coverage_modalities,
)
band = weak_temporal_band_loss(
output,
{name: sequences[name].times for name in MODALITIES},
durations,
centers,
margin=0.1,
)
targets = _gaussian_targets(method, output, sequences, durations)
align = _gaussian_alignment_kl(output, targets)
reconstruction = align.new_zeros(())
contrastive = align.new_zeros(())
if experiment == "D0":
total = 5.0 * span + 10.0 * band
active = ("span", "band", "total")
elif experiment in {"D1", "D2"}:
total = align
active = ("align", "total")
elif experiment == "D3":
if decoder is None:
raise ValueError("D3 requires the reconstruction decoder")
target_losses = []
batch_size, grid_size = output.aligned["text"].shape[:2]
for target in MODALITIES:
mask = make_block_mask(
batch_size,
grid_size,
0.2,
output.aligned[target].device,
generator=block_generator,
)
prediction = decoder(target, output.aligned, mask)
target_losses.append(
F.smooth_l1_loss(prediction[mask], output.aligned[target][mask])
)
reconstruction = torch.stack(target_losses).mean()
contrastive = cross_modal_contrastive_loss(output.aligned)
total = align + reconstruction + contrastive
active = ("align", "reconstruction", "contrastive", "total")
else:
raise ValueError(f"unknown experiment: {experiment}")
return total, {
"reconstruction": reconstruction,
"contrastive": contrastive,
"span": span,
"band": band,
"align": align,
"total": total,
}, active
def _attention_projection_parameters(model: nn.Module, method: str) -> tuple[list[nn.Parameter], nn.Parameter | None]:
if method == "M3":
layers = (model.audio_attention.attention, model.vision_attention.attention)
slot_parameter = None
else:
layers = tuple(model.attention[name].attention for name in MODALITIES)
slot_parameter = model.slots
parameters = [layer.in_proj_weight for layer in layers]
return parameters, slot_parameter
def _gradient_norms(
model: nn.Module,
method: str,
component_losses: Mapping[str, Tensor],
active_components: Sequence[str],
) -> dict[str, float | None]:
qk_parameters, slot_parameter = _attention_projection_parameters(model, method)
parameters = [*qk_parameters]
if slot_parameter is not None:
parameters.append(slot_parameter)
hidden_size = HIDDEN_SIZE
values: dict[str, float | None] = {}
for component in active_components:
loss = component_losses[component]
gradients = torch.autograd.grad(
loss,
parameters,
retain_graph=True,
allow_unused=True,
)
q_sq = torch.zeros((), device=loss.device)
k_sq = torch.zeros((), device=loss.device)
for gradient in gradients[: len(qk_parameters)]:
if gradient is None:
continue
q_sq = q_sq + gradient[:hidden_size].square().sum()
k_sq = k_sq + gradient[hidden_size : 2 * hidden_size].square().sum()
values[f"grad_{component}_WQ"] = float(q_sq.sqrt().item())
values[f"grad_{component}_WK"] = float(k_sq.sqrt().item())
if slot_parameter is not None:
slot_gradient = gradients[-1]
values[f"grad_{component}_Z"] = (
float(slot_gradient.norm().item()) if slot_gradient is not None else 0.0
)
else:
values[f"grad_{component}_Z"] = None
return values
def _metric_rows(
method: str,
experiment: str,
sample: FeatureSample,
output: Any,
sequences: Mapping[str, Any],
durations: Tensor,
*,
stats: FeatureStats,
device: torch.device,
) -> list[dict[str, Any]]:
targets = _gaussian_targets(method, output, sequences, durations)
rows = []
for name in MODALITIES:
weights = output.weights[name]
valid = sequences[name].valid
trajectory = alignment_trajectory(weights, sequences[name].times, durations)
entropy = normalized_attention_entropy(weights, valid)
centers, _ = _centers(method, output, sequences, durations)
row: dict[str, Any] = {
"experiment": experiment,
"method": method,
"sample_id": sample.sample_id,
"modality": name,
"mvr": float(monotonicity_violation_rate(trajectory).mean().item()),
"normalized_entropy": float(entropy.mean().item()),
"c_row": float(attention_row_similarity(weights).mean().item()),
"c_far": float(attention_row_similarity(weights, min_separation=6).mean().item()),
"expected_time_start": float(trajectory[0, 0].item()),
"expected_time_end": float(trajectory[0, -1].item()),
"trajectory_span": float((trajectory[0, -1] - trajectory[0, 0]).item()),
"mean_absolute_time_center_error": float(
(trajectory - centers.get(name, trajectory.new_full(trajectory.shape, float("nan"))))
.abs()
.mean()
.item()
)
if name in centers
else None,
"gaussian_target_kl": None,
}
if name in targets:
target = targets[name].clamp_min(1e-12)
prediction = weights.clamp_min(1e-8)
row["gaussian_target_kl"] = float(
(target * (target.log() - prediction.log())).sum(dim=-1).mean().item()
)
rows.append(row)
return rows
def _example_arrays(sample: FeatureSample, output: Any, sequences: Mapping[str, Any], method: str) -> dict[str, np.ndarray]:
targets = _gaussian_targets(method, output, sequences, torch.tensor([sample.duration_s], device=output.weights["text"].device))
arrays: dict[str, np.ndarray] = {"sample_id": np.asarray(sample.sample_id)}
for name in MODALITIES:
length = len(sample.times[name])
arrays[f"weights_{name}"] = output.weights[name][0, :, :length].detach().cpu().numpy()
arrays[f"times_{name}_s"] = sample.times[name].astype(np.float32, copy=False)
arrays[f"valid_{name}"] = sample.valid[name]
trajectory = (
output.weights[name][0, :, :length]
@ torch.as_tensor(sample.times[name], dtype=torch.float32, device=output.weights[name].device)
/ sample.duration_s
)
arrays[f"trajectory_{name}"] = trajectory.detach().cpu().numpy()
if name in targets:
arrays[f"target_{name}"] = targets[name][0, :, :length].detach().cpu().numpy()
return arrays
def _plot_example(
path_prefix: Path,
sample: FeatureSample,
arrays: Mapping[str, np.ndarray],
method: str,
experiment: str,
) -> None:
modalities = ("audio", "vision") if method == "M3" else MODALITIES
fig, axes = plt.subplots(len(modalities), 2, figsize=(12, 4 * len(modalities)), constrained_layout=True)
if len(modalities) == 1:
axes = np.asarray([axes])
for row, name in enumerate(modalities):
times = arrays[f"times_{name}_s"] / max(sample.duration_s, 1e-8)
extent = (float(times[0]), float(times[-1]), 0.0, 1.0)
prediction = arrays[f"weights_{name}"]
image = axes[row, 0].imshow(
prediction,
origin="lower",
aspect="auto",
interpolation="nearest",
extent=extent,
cmap="magma",
)
axes[row, 0].set_title(f"{name}: learned A")
axes[row, 0].set_xlabel("source time / clip duration")
axes[row, 0].set_ylabel("slot index / K")
fig.colorbar(image, ax=axes[row, 0], fraction=0.046, pad=0.04)
target_key = f"target_{name}"
if target_key in arrays:
target = arrays[target_key]
image_target = axes[row, 1].imshow(
target,
origin="lower",
aspect="auto",
interpolation="nearest",
extent=extent,
cmap="magma",
)
axes[row, 1].set_title(f"{name}: Gaussian target P")
fig.colorbar(image_target, ax=axes[row, 1], fraction=0.046, pad=0.04)
else:
axes[row, 1].imshow(
np.zeros_like(prediction),
origin="lower",
aspect="auto",
extent=extent,
cmap="magma",
)
axes[row, 1].set_title(f"{name}: target not used by M3")
axes[row, 1].set_xlabel("source time / clip duration")
axes[row, 1].set_ylabel("slot index / K")
fig.suptitle(f"{experiment} {method} · {sample.sample_id}")
fig.savefig(path_prefix.with_name(path_prefix.name + "_heatmap.png"), dpi=160)
plt.close(fig)
fig, ax = plt.subplots(figsize=(8, 5), constrained_layout=True)
x = (np.arange(GRID_SIZE, dtype=np.float32) + 0.5) / GRID_SIZE
for name in MODALITIES:
trajectory = arrays[f"trajectory_{name}"]
ax.plot(x, trajectory, label=f"{name} actual")
if f"target_{name}" in arrays:
target_weights = arrays[f"target_{name}"]
target_time = target_weights @ arrays[f"times_{name}_s"] / max(sample.duration_s, 1e-8)
ax.plot(x, target_time, linestyle="--", alpha=0.7, label=f"{name} target")
ax.plot([0, 1], [0, 1], color="black", linestyle=":", alpha=0.6, label="uniform-time reference")
ax.set(xlim=(0, 1), ylim=(0, 1), xlabel="shared slot position", ylabel="expected normalized source time")
ax.grid(alpha=0.2)
ax.legend(fontsize=8, ncol=2)
ax.set_title(f"{experiment} {method} trajectory · {sample.sample_id}")
fig.savefig(path_prefix.with_name(path_prefix.name + "_trajectory.png"), dpi=160)
plt.close(fig)
def _batches(samples: Sequence[FeatureSample], batch_size: int, rng: np.random.Generator):
order = rng.permutation(len(samples)).tolist()
for start in range(0, len(order), batch_size):
yield [samples[index] for index in order[start : start + batch_size]]
def _evaluate_samples(
method: str,
experiment: str,
model: nn.Module,
samples: Sequence[FeatureSample],
stats: FeatureStats,
*,
device: torch.device,
batch_size: int,
) -> tuple[list[dict[str, Any]], dict[str, np.ndarray] | None]:
rows: list[dict[str, Any]] = []
example_arrays: dict[str, np.ndarray] | None = None
model.eval()
rng = np.random.default_rng(0)
with torch.no_grad():
for batch_samples in _batches(samples, batch_size, rng):
sequences, durations, _ = collate_feature_samples(batch_samples, stats, device)
output = model(sequences)
for index, sample in enumerate(batch_samples):
one_sequences = {
name: type(sequences[name])(
features=sequences[name].features[index : index + 1],
times=sequences[name].times[index : index + 1],
valid=sequences[name].valid[index : index + 1],
)
for name in MODALITIES
}
one_output = type(output)(
weights={name: output.weights[name][index : index + 1] for name in MODALITIES},
aligned={name: output.aligned[name][index : index + 1] for name in MODALITIES},
fallback_rows=output.fallback_rows,
)
one_duration = durations[index : index + 1]
rows.extend(
_metric_rows(
method,
experiment,
sample,
one_output,
one_sequences,
one_duration,
stats=stats,
device=device,
)
)
if sample.sample_id == EXAMPLE_ID:
example_arrays = _example_arrays(sample, one_output, one_sequences, method)
return rows, example_arrays
def _run_trial(
experiment: str,
method: str,
tag: str,
train_samples: Sequence[FeatureSample],
eval_samples: Sequence[FeatureSample],
stats: FeatureStats,
output_dir: Path,
*,
device: torch.device,
seed: int,
steps: int,
absolute_position_encoding: bool,
batch_size: int,
include_content_losses: bool,
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
_seed_everything(seed)
dimensions = {name: train_samples[0].features[name].shape[1] for name in MODALITIES}
model = _model(method, dimensions, absolute_position_encoding=absolute_position_encoding).to(device)
decoder = AlignmentReconstructor(HIDDEN_SIZE, dropout=0.0).to(device) if include_content_losses else None
parameters = list(model.parameters()) + (list(decoder.parameters()) if decoder is not None else [])
optimizer = torch.optim.AdamW(parameters, lr=LEARNING_RATE, weight_decay=0.0)
block_generator = torch.Generator(device=device)
block_generator.manual_seed(seed + 73)
rng = np.random.default_rng(seed)
order = np.arange(len(train_samples))
cursor = 0
history: list[dict[str, Any]] = []
log_every = LOG_INTERVAL
active_names: tuple[str, ...] | None = None
print(
f"[{experiment} {tag}] samples={len(train_samples)} steps={steps} "
f"sinusoidal_PE={absolute_position_encoding} lr={LEARNING_RATE}",
flush=True,
)
model.train()
if decoder is not None:
decoder.train()
for step in range(1, steps + 1):
if len(train_samples) == 1:
batch_samples = [train_samples[0]]
else:
if cursor + batch_size > len(order):
order = rng.permutation(len(train_samples))
cursor = 0
indices = order[cursor : cursor + batch_size]
cursor += len(indices)
batch_samples = [train_samples[int(index)] for index in indices]
sequences, durations, _ = collate_feature_samples(batch_samples, stats, device)
output = model(sequences)
total, losses, active_components = _training_components(
experiment,
method,
output,
sequences,
durations,
decoder=decoder,
block_generator=block_generator,
)
active_names = active_components
if not torch.isfinite(total):
raise FloatingPointError(f"non-finite {experiment}/{tag} loss at step {step}")
row: dict[str, Any] = {
"experiment": experiment,
"method": method,
"variant": tag,
"step": step,
"epoch": math.ceil(step / max(1, math.ceil(len(train_samples) / batch_size))),
"learning_rate": LEARNING_RATE,
"absolute_position_encoding": absolute_position_encoding,
"L_rec": float(losses["reconstruction"].detach().item()),
"L_con": float(losses["contrastive"].detach().item()),
"L_span": float(losses["span"].detach().item()),
"L_band": float(losses["band"].detach().item()),
"L_align": float(losses["align"].detach().item()),
"L_total": float(total.detach().item()),
}
if step == 1 or step % log_every == 0 or step == steps:
row.update(_gradient_norms(model, method, losses, active_components))
optimizer.zero_grad(set_to_none=True)
total.backward()
nn.utils.clip_grad_norm_(parameters, 2.0)
optimizer.step()
history.append(row)
if step == 1 or step % log_every == 0 or step == steps:
_write_csv(output_dir / f"{tag}_history.csv", history)
if step % 100 == 0 or step == steps:
diagnostic_component = "band" if experiment == "D0" else "align"
print(
f"[{experiment} {tag} step {step}/{steps}] "
f"total={row['L_total']:.4f} align={row['L_align']:.4f} "
f"span={row['L_span']:.4f} band={row['L_band']:.4f} "
f"grad-{diagnostic_component}(Q/K/Z)="
f"{row.get(f'grad_{diagnostic_component}_WQ', float('nan')):.3g}/"
f"{row.get(f'grad_{diagnostic_component}_WK', float('nan')):.3g}/"
f"{row.get(f'grad_{diagnostic_component}_Z', float('nan')) if row.get(f'grad_{diagnostic_component}_Z') is not None else 'NA'}",
flush=True,
)
output_dir.mkdir(parents=True, exist_ok=True)
_write_csv(output_dir / f"{tag}_history.csv", history)
checkpoint = output_dir / f"{tag}_checkpoint.pt"
torch.save(
{
"experiment": experiment,
"method": method,
"variant": tag,
"seed": seed,
"steps": steps,
"absolute_position_encoding": absolute_position_encoding,
"model_state_dict": model.state_dict(),
"decoder_state_dict": decoder.state_dict() if decoder is not None else None,
"history": history,
},
checkpoint,
)
metric_rows, example_arrays = _evaluate_samples(
method,
experiment,
model,
eval_samples,
stats,
device=device,
batch_size=batch_size,
)
_write_csv(output_dir / f"{tag}_metrics.csv", metric_rows)
if example_arrays is not None:
np.savez_compressed(output_dir / f"{tag}_example_alignment.npz", **example_arrays)
example_sample = next(sample for sample in eval_samples if sample.sample_id == EXAMPLE_ID)
_plot_example(output_dir / tag, example_sample, example_arrays, method, experiment)
del model, decoder
if device.type == "cuda":
torch.cuda.empty_cache()
return history, metric_rows
def run(args: argparse.Namespace) -> dict[str, Any]:
started = time.time()
_seed_everything(args.seed)
if args.device == "auto":
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
else:
device = torch.device(args.device)
if device.type == "cuda" and not torch.cuda.is_available():
raise RuntimeError("CUDA was requested but is not available")
samples = load_feature_samples(args.feature_dir, args.manifest)
if len(samples) != 100:
raise ValueError(f"debug experiments expect the complete 100-sample set, found {len(samples)}")
sample_by_id = {sample.sample_id: sample for sample in samples}
if EXAMPLE_ID not in sample_by_id:
raise ValueError(f"required debug sample is absent: {EXAMPLE_ID}")
one_sample = sample_by_id[EXAMPLE_ID]
output_root = args.output_dir
output_root.mkdir(parents=True, exist_ok=True)
full_stats = fit_feature_stats(samples)
single_stats = fit_feature_stats([one_sample])
all_metrics: list[dict[str, Any]] = []
all_history_summaries: list[dict[str, Any]] = []
specifications = [
("D0", "M3", "M3", [one_sample], [one_sample], single_stats, SINGLE_SAMPLE_STEPS, False, False),
("D0", "M4", "M4", [one_sample], [one_sample], single_stats, SINGLE_SAMPLE_STEPS, False, False),
("D1", "M3", "M3", [one_sample], [one_sample], single_stats, SINGLE_SAMPLE_STEPS, False, False),
("D1", "M4", "M4_noPE", [one_sample], [one_sample], single_stats, SINGLE_SAMPLE_STEPS, False, False),
("D1", "M4", "M4_sinPE", [one_sample], [one_sample], single_stats, SINGLE_SAMPLE_STEPS, True, False),
("D2", "M3", "M3", samples, samples, full_stats, FULL_DATA_STEPS, False, False),
("D2", "M4", "M4_sinPE", samples, samples, full_stats, FULL_DATA_STEPS, True, False),
("D3", "M3", "M3", samples, samples, full_stats, FULL_DATA_STEPS, False, True),
("D3", "M4", "M4_sinPE", samples, samples, full_stats, FULL_DATA_STEPS, True, True),
]
for experiment, method, tag, train_set, eval_set, stats, steps, use_pe, content_losses in specifications:
trial_dir = output_root / experiment
trial_dir.mkdir(parents=True, exist_ok=True)
history, metric_rows = _run_trial(
experiment,
method,
tag,
train_set,
eval_set,
stats,
trial_dir,
device=device,
seed=args.seed,
steps=steps,
absolute_position_encoding=use_pe,
batch_size=BATCH_SIZE,
include_content_losses=content_losses,
)
all_metrics.extend(metric_rows)
selected_steps = [
row
for row in history
if any(key.startswith("grad_") and value not in (None, "") for key, value in row.items())
]
summary: dict[str, Any] = {
"experiment": experiment,
"method": method,
"variant": tag,
"steps": steps,
"absolute_position_encoding": use_pe,
"final_L_total": history[-1]["L_total"],
"final_L_rec": history[-1]["L_rec"],
"final_L_con": history[-1]["L_con"],
"final_L_span": history[-1]["L_span"],
"final_L_band": history[-1]["L_band"],
"final_L_align": history[-1]["L_align"],
}
if selected_steps:
for component in ("span", "band", "align", "reconstruction", "contrastive", "total"):
for parameter in ("WQ", "WK", "Z"):
key = f"grad_{component}_{parameter}"
values = [float(row[key]) for row in selected_steps if row.get(key) not in (None, "")]
if values:
summary[f"mean_{key}"] = float(np.mean(values))
summary[f"final_{key}"] = values[-1]
summary["final_metric_rows"] = len(metric_rows)
all_history_summaries.append(summary)
print(
f"[done {experiment} {tag}] final total={summary['final_L_total']:.4f} "
f"align={summary['final_L_align']:.4f}; metrics={len(metric_rows)}",
flush=True,
)
_write_csv(output_root / "debug_summary.csv", all_history_summaries)
_write_csv(output_root / "per_sample_metrics.csv", all_metrics)
manifest = {
"created_utc": datetime.now(timezone.utc).isoformat(),
"sample_count": len(samples),
"diagnostic_sample": EXAMPLE_ID,
"seed": args.seed,
"device": str(device),
"gpu_name": torch.cuda.get_device_name(device) if device.type == "cuda" else None,
"python": platform.python_version(),
"torch": torch.__version__,
"experiments": [
{
"name": "D0",
"scope": "single sample; span + barycenter band only",
"steps_per_model": SINGLE_SAMPLE_STEPS,
"loss": "5 * L_span + 10 * L_band",
},
{
"name": "D1",
"scope": "single sample; Gaussian target KL only",
"steps_per_model": SINGLE_SAMPLE_STEPS,
"sigma_normalized_time": SIGMA,
"M4_control": "no PE versus fixed sinusoidal PE",
},
{
"name": "D2",
"scope": "all 100 clips; Gaussian target KL only; one seed; in-sample diagnostic",
"steps_per_model": FULL_DATA_STEPS,
"sigma_normalized_time": SIGMA,
"M4_position_encoding": "fixed sinusoidal",
},
{
"name": "D3",
"scope": "all 100 clips; Gaussian KL + masked reconstruction + contrastive; one seed; in-sample diagnostic",
"steps_per_model": FULL_DATA_STEPS,
"sigma_normalized_time": SIGMA,
"M4_position_encoding": "fixed sinusoidal",
},
],
"optimizer": "AdamW",
"learning_rate": LEARNING_RATE,
"dropout": 0.0,
"batch_size": BATCH_SIZE,
"gradient_logging_interval_steps": LOG_INTERVAL,
"gradient_metrics": ["W_Q", "W_K", "M4 latent slots Z"],
"features_changed": False,
"M1_M2_changed": False,
"elapsed_seconds": time.time() - started,
"interpretation_limits": [
"D0 and D1 overfit one selected sample and diagnose optimization/representability only.",
"D2 and D3 train and evaluate on the same 100 clips; they diagnose whether the target can be optimized, not generalization.",
"Gaussian targets are weak temporal priors constructed from timestamps; they are not human alignment ground truth.",
"M4 absolute sinusoidal encoding is enabled only for D1's PE control and D2/D3; prior M1-M4 and v2 results are unchanged.",
],
}
(output_root / "run_manifest.json").write_text(
json.dumps(manifest, ensure_ascii=False, indent=2, allow_nan=False), encoding="utf-8"
)
print(
f"[all done] elapsed={manifest['elapsed_seconds']:.1f}s output={output_root}",
flush=True,
)
return manifest
def build_parser() -> argparse.ArgumentParser:
project_dir = Path(__file__).resolve().parents[1]
parser = argparse.ArgumentParser(
description="Diagnose gradient flow and temporal alignment learnability for M3/M4."
)
parser.add_argument("--feature-dir", type=Path, default=project_dir / "outputs/q1_features/features")
parser.add_argument("--manifest", type=Path, default=project_dir / "outputs/audit/manifest.csv")
parser.add_argument("--output-dir", type=Path, default=project_dir / "outputs/alignment_debug")
parser.add_argument("--device", default="auto", help="auto, cpu, or a torch device such as cuda:0")
parser.add_argument("--seed", type=int, default=42)
return parser
def main() -> int:
args = build_parser().parse_args()
run(args)
return 0
if __name__ == "__main__":
raise SystemExit(main())
@@ -0,0 +1,382 @@
"""Grouped held-out check for learned source-time positional features."""
from __future__ import annotations
import argparse
import csv
import json
import platform
import time
from datetime import datetime, timezone
from pathlib import Path
from typing import Any
import numpy as np
import torch
from torch import nn
from .alignment_debug import (
EXAMPLE_ID,
GRID_SIZE,
HEADS,
HIDDEN_SIZE,
LEARNING_RATE,
_example_arrays,
_gaussian_alignment_kl,
_gaussian_targets,
_gradient_norms,
_metric_rows,
_plot_example,
_seed_everything,
)
from .experiment_data import (
FeatureSample,
FeatureStats,
collate_feature_samples,
fit_feature_stats,
load_feature_samples,
)
from .models import SharedLatentTimeline, TextAnchoredCrossAttention
from .types import MODALITIES
BATCH_SIZE = 8
DEFAULT_STEPS = 500
def _write_csv(path: Path, rows: list[dict[str, Any]]) -> None:
if not rows:
return
path.parent.mkdir(parents=True, exist_ok=True)
fields = list(dict.fromkeys(key for row in rows for key in row))
with path.open("w", newline="", encoding="utf-8-sig") as handle:
writer = csv.DictWriter(handle, fieldnames=fields)
writer.writeheader()
writer.writerows(rows)
def _make_model(method: str, dimensions: dict[str, int], source_time: bool) -> nn.Module:
if method == "M3":
return TextAnchoredCrossAttention(
dimensions,
grid_size=GRID_SIZE,
hidden_size=HIDDEN_SIZE,
heads=HEADS,
dropout=0.0,
source_time_encoding=source_time,
)
return SharedLatentTimeline(
dimensions,
grid_size=GRID_SIZE,
hidden_size=HIDDEN_SIZE,
heads=HEADS,
dropout=0.0,
absolute_position_encoding=True,
source_time_encoding=source_time,
)
def _batches(samples: list[FeatureSample], rng: np.random.Generator):
order = rng.permutation(len(samples)).tolist()
for start in range(0, len(order), BATCH_SIZE):
yield [samples[index] for index in order[start : start + BATCH_SIZE]]
def _train_one(
*,
method: str,
variant: str,
source_time: bool,
train_samples: list[FeatureSample],
validation_samples: list[FeatureSample],
stats: FeatureStats,
example_id: str,
output_dir: Path,
device: torch.device,
steps: int,
seed: int,
) -> tuple[list[dict[str, Any]], list[dict[str, Any]], dict[str, Any]]:
_seed_everything(seed)
dimensions = {
name: train_samples[0].features[name].shape[1] for name in MODALITIES
}
model = _make_model(method, dimensions, source_time).to(device)
optimizer = torch.optim.AdamW(model.parameters(), lr=LEARNING_RATE, weight_decay=0.0)
rng = np.random.default_rng(seed)
history: list[dict[str, Any]] = []
print(
f"[D5 {variant}] train={len(train_samples)} heldout={len(validation_samples)} "
f"groups={len({s.group_id for s in train_samples})}/"
f"{len({s.group_id for s in validation_samples})} steps={steps}",
flush=True,
)
model.train()
for step in range(1, steps + 1):
batch_indices = rng.choice(
len(train_samples), size=min(BATCH_SIZE, len(train_samples)), replace=False
)
batch_samples = [train_samples[int(index)] for index in batch_indices]
sequences, durations, _ = collate_feature_samples(batch_samples, stats, device)
output = model(sequences, durations) if source_time else model(sequences)
targets = _gaussian_targets(method, output, sequences, durations)
loss = _gaussian_alignment_kl(output, targets)
if not torch.isfinite(loss):
raise FloatingPointError(f"non-finite D5 loss for {variant} at step {step}")
row: dict[str, Any] = {
"experiment": "D5",
"method": method,
"variant": variant,
"step": step,
"L_align": float(loss.detach().item()),
"source_time_encoding": source_time,
}
if step == 1 or step % 20 == 0 or step == steps:
row.update(_gradient_norms(model, method, {"align": loss}, ("align",)))
optimizer.zero_grad(set_to_none=True)
loss.backward()
nn.utils.clip_grad_norm_(model.parameters(), 2.0)
optimizer.step()
history.append(row)
if step == 1 or step % 100 == 0 or step == steps:
grad_z = row.get("grad_align_Z")
grad_z_text = f"{grad_z:.3g}" if grad_z is not None else "NA"
print(
f"[D5 {variant} {step}/{steps}] KL={row['L_align']:.5f} "
f"grad_Q/K/Z={row.get('grad_align_WQ', 0):.3g}/"
f"{row.get('grad_align_WK', 0):.3g}/{grad_z_text}",
flush=True,
)
output_dir.mkdir(parents=True, exist_ok=True)
_write_csv(output_dir / "history.csv", history)
model.eval()
metric_rows: list[dict[str, Any]] = []
heldout_arrays: dict[str, np.ndarray] | None = None
with torch.no_grad():
eval_rng = np.random.default_rng(0)
for batch_samples in _batches(validation_samples, eval_rng):
sequences, durations, _ = collate_feature_samples(batch_samples, stats, device)
output = model(sequences, durations) if source_time else model(sequences)
for index, sample in enumerate(batch_samples):
one_sequences = {
name: type(sequences[name])(
features=sequences[name].features[index : index + 1],
times=sequences[name].times[index : index + 1],
valid=sequences[name].valid[index : index + 1],
)
for name in MODALITIES
}
one_output = type(output)(
weights={name: output.weights[name][index : index + 1] for name in MODALITIES},
aligned={name: output.aligned[name][index : index + 1] for name in MODALITIES},
fallback_rows=output.fallback_rows,
)
rows = _metric_rows(
method,
"D5",
sample,
one_output,
one_sequences,
durations[index : index + 1],
stats=stats,
device=device,
)
for metric_row in rows:
metric_row["variant"] = variant
metric_row["source_time_encoding"] = source_time
metric_rows.extend(rows)
if sample.sample_id == example_id:
heldout_arrays = _example_arrays(
sample, one_output, one_sequences, method
)
_write_csv(output_dir / "heldout_metrics.csv", metric_rows)
if heldout_arrays is not None:
heldout_sample = next(s for s in validation_samples if s.sample_id == example_id)
np.savez_compressed(output_dir / "heldout_alignment.npz", **heldout_arrays)
_plot_example(output_dir / variant, heldout_sample, heldout_arrays, method, "D5 held-out")
checkpoint = output_dir / "checkpoint.pt"
torch.save(
{
"experiment": "D5",
"method": method,
"variant": variant,
"source_time_encoding": source_time,
"absolute_position_encoding": method == "M4",
"seed": seed,
"steps": steps,
"train_sample_ids": [sample.sample_id for sample in train_samples],
"validation_sample_ids": [sample.sample_id for sample in validation_samples],
"model_state_dict": model.state_dict(),
},
checkpoint,
)
training_summary = {
"experiment": "D5",
"method": method,
"variant": variant,
"source_time_encoding": source_time,
"final_training_kl": history[-1]["L_align"],
"heldout_sample_count": len(validation_samples),
"checkpoint": str(checkpoint),
}
del model
if device.type == "cuda":
torch.cuda.empty_cache()
return metric_rows, history, training_summary
def run(args: argparse.Namespace) -> dict[str, Any]:
started = time.time()
if args.device == "auto":
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
else:
device = torch.device(args.device)
if device.type == "cuda" and not torch.cuda.is_available():
raise RuntimeError("CUDA was requested but is unavailable")
samples = load_feature_samples(args.feature_dir, args.manifest)
by_id = {sample.sample_id: sample for sample in samples}
with args.splits.open("r", encoding="utf-8-sig") as handle:
folds = json.load(handle)
fold = next((item for item in folds if item["fold"] == args.fold), None)
if fold is None:
raise ValueError(f"fold {args.fold} is not present in {args.splits}")
train_samples = [by_id[sample_id] for sample_id in fold["train_sample_ids"]]
validation_samples = [by_id[sample_id] for sample_id in fold["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("train/validation video_id groups overlap")
example_id = args.example_id or validation_samples[0].sample_id
if example_id not in {sample.sample_id for sample in validation_samples}:
raise ValueError(f"held-out example is not in validation fold: {example_id}")
stats = fit_feature_stats(train_samples)
output_root = args.output_dir
output_root.mkdir(parents=True, exist_ok=True)
specifications = (
("M3", "M3_noSourceTime", False),
("M3", "M3_sourceTime", True),
("M4", "M4_noSourceTime", False),
("M4", "M4_sourceTime", True),
)
all_metrics: list[dict[str, Any]] = []
all_history: list[dict[str, Any]] = []
summaries: list[dict[str, Any]] = []
for method, variant, source_time in specifications:
metrics, history, summary = _train_one(
method=method,
variant=variant,
source_time=source_time,
train_samples=train_samples,
validation_samples=validation_samples,
stats=stats,
example_id=example_id,
output_dir=output_root / variant,
device=device,
steps=args.steps,
seed=args.seed,
)
all_metrics.extend(metrics)
all_history.extend(history)
summaries.append(summary)
_write_csv(output_root / "per_sample_metrics.csv", all_metrics)
_write_csv(output_root / "training_history.csv", all_history)
aggregate_rows: list[dict[str, Any]] = []
metric_names = (
"mvr",
"normalized_entropy",
"c_row",
"trajectory_span",
"mean_absolute_time_center_error",
"gaussian_target_kl",
)
for variant, modality in sorted(
{(row["variant"], row["modality"]) for row in all_metrics}
):
rows = [row for row in all_metrics if row["variant"] == variant and row["modality"] == modality]
aggregate: dict[str, Any] = {
"variant": variant,
"modality": modality,
"sample_count": len(rows),
}
for name in metric_names:
values = [float(row[name]) for row in rows if row.get(name) not in (None, "")]
aggregate[f"mean_{name}"] = float(np.mean(values)) if values else ""
aggregate_rows.append(aggregate)
_write_csv(output_root / "heldout_summary.csv", aggregate_rows)
_write_csv(output_root / "training_summary.csv", summaries)
manifest = {
"created_utc": datetime.now(timezone.utc).isoformat(),
"experiment": "D5",
"fold": args.fold,
"heldout_example": example_id,
"train_sample_count": len(train_samples),
"heldout_sample_count": len(validation_samples),
"train_video_ids": sorted(train_groups),
"heldout_video_ids": sorted(validation_groups),
"video_id_overlap": sorted(train_groups & validation_groups),
"seed": args.seed,
"device": str(device),
"gpu_name": torch.cuda.get_device_name(device) if device.type == "cuda" else None,
"python": platform.python_version(),
"torch": torch.__version__,
"grid_size": GRID_SIZE,
"optimizer": "AdamW",
"learning_rate": LEARNING_RATE,
"steps_per_model": args.steps,
"batch_size": BATCH_SIZE,
"loss": "timestamp-derived Gaussian target KL only",
"source_position_encoding": "fixed Fourier time code added to source key only; value remains projected content",
"query_position_encoding": "M3 adds Fourier code at text-time centers; M4 uses fixed absolute sinusoidal slots plus the matching Fourier code at uniform slot centers",
"normalization_fit_on_train_only": True,
"variants": summaries,
"elapsed_seconds": time.time() - started,
"interpretation_limits": [
"This is one grouped video_id split and one seed; it is a focused held-out diagnostic, not a final method ranking.",
"The Gaussian timestamp target is a weak temporal prior, not human alignment ground truth.",
"The target supplies approximate time location; this experiment tests transfer of the time-conditioned attention mechanism, not semantic correctness by itself.",
],
}
(output_root / "run_manifest.json").write_text(
json.dumps(manifest, ensure_ascii=False, indent=2, allow_nan=False), encoding="utf-8"
)
print(
f"[D5 done] train={len(train_samples)} heldout={len(validation_samples)} "
f"video_id_groups={len(train_groups)}/{len(validation_groups)} "
f"elapsed={manifest['elapsed_seconds']:.1f}s output={output_root}",
flush=True,
)
return manifest
def build_parser() -> argparse.ArgumentParser:
project = Path(__file__).resolve().parents[1]
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--steps", type=int, default=DEFAULT_STEPS)
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--fold", type=int, default=1)
parser.add_argument("--example-id", type=str, default=None)
parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto")
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(
"--output-dir", type=Path, default=project / "outputs/alignment_debug/heldout"
)
return parser
def main() -> None:
args = build_parser().parse_args()
run(args)
if __name__ == "__main__":
main()
@@ -0,0 +1,273 @@
"""Single-clip test of explicit source-time identities in learned alignment."""
from __future__ import annotations
import argparse
import csv
import json
import platform
import time
from datetime import datetime, timezone
from pathlib import Path
from typing import Any
import numpy as np
import torch
from torch import nn
from .alignment_debug import (
EXAMPLE_ID,
GRID_SIZE,
HEADS,
HIDDEN_SIZE,
LEARNING_RATE,
SIGMA,
_example_arrays,
_gaussian_alignment_kl,
_gaussian_targets,
_gradient_norms,
_metric_rows,
_plot_example,
_seed_everything,
)
from .experiment_data import fit_feature_stats, load_feature_samples, collate_feature_samples
from .models import SharedLatentTimeline, TextAnchoredCrossAttention
from .types import MODALITIES
def _write_csv(path: Path, rows: list[dict[str, Any]]) -> None:
if not rows:
return
path.parent.mkdir(parents=True, exist_ok=True)
fields = list(dict.fromkeys(key for row in rows for key in row))
with path.open("w", newline="", encoding="utf-8-sig") as handle:
writer = csv.DictWriter(handle, fieldnames=fields)
writer.writeheader()
writer.writerows(rows)
def _make_model(method: str, dimensions: dict[str, int], source_time_encoding: bool) -> nn.Module:
if method == "M3":
return TextAnchoredCrossAttention(
dimensions,
grid_size=GRID_SIZE,
hidden_size=HIDDEN_SIZE,
heads=HEADS,
dropout=0.0,
source_time_encoding=source_time_encoding,
)
return SharedLatentTimeline(
dimensions,
grid_size=GRID_SIZE,
hidden_size=HIDDEN_SIZE,
heads=HEADS,
dropout=0.0,
absolute_position_encoding=True,
source_time_encoding=source_time_encoding,
)
def _run_trial(
*,
method: str,
variant: str,
source_time_encoding: bool,
sample: Any,
stats: Any,
device: torch.device,
steps: int,
seed: int,
output_dir: Path,
) -> tuple[dict[str, Any], list[dict[str, Any]]]:
_seed_everything(seed)
dimensions = {name: sample.features[name].shape[1] for name in MODALITIES}
model = _make_model(method, dimensions, source_time_encoding).to(device)
optimizer = torch.optim.AdamW(model.parameters(), lr=LEARNING_RATE, weight_decay=0.0)
sequences, durations, _ = collate_feature_samples([sample], stats, device)
history: list[dict[str, Any]] = []
print(
f"[D4 {variant}] sample={sample.sample_id} steps={steps} "
f"source_time_encoding={source_time_encoding} device={device}",
flush=True,
)
model.train()
for step in range(1, steps + 1):
output = model(sequences, durations) if source_time_encoding else model(sequences)
targets = _gaussian_targets(method, output, sequences, durations)
loss = _gaussian_alignment_kl(output, targets)
if not torch.isfinite(loss):
raise FloatingPointError(f"non-finite D4 loss at {variant} step {step}")
row: dict[str, Any] = {
"experiment": "D4",
"method": method,
"variant": variant,
"step": step,
"L_align": float(loss.detach().item()),
"source_time_encoding": source_time_encoding,
}
if step == 1 or step % 20 == 0 or step == steps:
row.update(_gradient_norms(model, method, {"align": loss}, ("align",)))
optimizer.zero_grad(set_to_none=True)
loss.backward()
nn.utils.clip_grad_norm_(model.parameters(), 2.0)
optimizer.step()
history.append(row)
if step == 1 or step % 100 == 0 or step == steps:
grad_z = row.get("grad_align_Z")
grad_z_text = f"{grad_z:.3g}" if grad_z is not None else "NA"
print(
f"[D4 {variant} {step}/{steps}] KL={row['L_align']:.5f} "
f"grad_Q/K/Z={row.get('grad_align_WQ', 0):.3g}/"
f"{row.get('grad_align_WK', 0):.3g}/"
f"{grad_z_text}",
flush=True,
)
output_dir.mkdir(parents=True, exist_ok=True)
_write_csv(output_dir / "history.csv", history)
model.eval()
with torch.no_grad():
output = model(sequences, durations) if source_time_encoding else model(sequences)
metric_rows = _metric_rows(
method,
"D4",
sample,
output,
sequences,
durations,
stats=stats,
device=device,
)
arrays = _example_arrays(sample, output, sequences, method)
for row in metric_rows:
row["variant"] = variant
row["source_time_encoding"] = source_time_encoding
_write_csv(output_dir / "metrics.csv", metric_rows)
np.savez_compressed(output_dir / "example_alignment.npz", **arrays)
_plot_example(output_dir / variant, sample, arrays, method, "D4")
checkpoint = output_dir / "checkpoint.pt"
torch.save(
{
"experiment": "D4",
"method": method,
"variant": variant,
"source_time_encoding": source_time_encoding,
"absolute_position_encoding": method == "M4",
"seed": seed,
"steps": steps,
"model_state_dict": model.state_dict(),
},
checkpoint,
)
summary = {
"experiment": "D4",
"method": method,
"variant": variant,
"source_time_encoding": source_time_encoding,
"final_training_kl": history[-1]["L_align"],
"checkpoint": str(checkpoint),
}
del model
if device.type == "cuda":
torch.cuda.empty_cache()
return summary, metric_rows
def run(args: argparse.Namespace) -> dict[str, Any]:
start = 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)
sample_map = {sample.sample_id: sample for sample in samples}
if EXAMPLE_ID not in sample_map:
raise ValueError(f"diagnostic sample is missing: {EXAMPLE_ID}")
sample = sample_map[EXAMPLE_ID]
stats = fit_feature_stats([sample])
output_root = args.output_dir
output_root.mkdir(parents=True, exist_ok=True)
specifications = (
("M3", "M3_noSourceTime", False),
("M3", "M3_sourceTime", True),
("M4", "M4_noSourceTime", False),
("M4", "M4_sourceTime", True),
)
summaries: list[dict[str, Any]] = []
all_metrics: list[dict[str, Any]] = []
for method, variant, use_source_time in specifications:
summary, metrics = _run_trial(
method=method,
variant=variant,
source_time_encoding=use_source_time,
sample=sample,
stats=stats,
device=device,
steps=args.steps,
seed=args.seed,
output_dir=output_root / variant,
)
summaries.append(summary)
all_metrics.extend(metrics)
_write_csv(output_root / "summary.csv", summaries)
_write_csv(output_root / "per_sample_metrics.csv", all_metrics)
manifest = {
"created_utc": datetime.now(timezone.utc).isoformat(),
"experiment": "D4",
"sample": sample.sample_id,
"sample_count": 1,
"seed": args.seed,
"device": str(device),
"gpu_name": torch.cuda.get_device_name(device) if device.type == "cuda" else None,
"python": platform.python_version(),
"torch": torch.__version__,
"grid_size": GRID_SIZE,
"sigma_normalized_time": SIGMA,
"optimizer": "AdamW",
"learning_rate": LEARNING_RATE,
"steps_per_model": args.steps,
"source_position_encoding": "fixed Fourier time code added to source key only; value remains the projected real feature",
"query_position_encoding": "M3 adds Fourier code at forced text-time centers; M4 keeps its fixed absolute sinusoidal slot code and adds the matching Fourier code at uniform slot centers",
"modalities": ["audio", "vision"],
"variants": summaries,
"elapsed_seconds": time.time() - start,
"interpretation_limits": [
"This is a one-sample overfit diagnostic, not a held-out accuracy result.",
"The timestamp-derived Gaussian target is a weak temporal prior, not human alignment ground truth.",
"The D4 variants test source/query time identity only; content features and downstream outputs still require separate evaluation.",
],
}
(output_root / "run_manifest.json").write_text(
json.dumps(manifest, ensure_ascii=False, indent=2, allow_nan=False), encoding="utf-8"
)
print(f"[D4 done] elapsed={manifest['elapsed_seconds']:.1f}s output={output_root}", flush=True)
return manifest
def build_parser() -> argparse.ArgumentParser:
project = Path(__file__).resolve().parents[1]
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--steps", type=int, default=1000)
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto")
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(
"--output-dir", type=Path, default=project / "outputs/alignment_debug/source_time"
)
return parser
def main() -> None:
args = build_parser().parse_args()
run(args)
if __name__ == "__main__":
main()
+246
View File
@@ -0,0 +1,246 @@
from __future__ import annotations
import argparse
import csv
import json
import subprocess
from collections import Counter
from fractions import Fraction
from pathlib import Path
from typing import Any
from openpyxl import load_workbook
def _identifier(value: Any) -> str:
if value is None:
return ""
if isinstance(value, float) and value.is_integer():
return str(int(value))
return str(value).strip()
def _read_labels(path: Path) -> list[dict[str, Any]]:
workbook = load_workbook(path, read_only=True, data_only=True)
sheet = workbook.active
rows = sheet.iter_rows(values_only=True)
header = next(rows, None)
if header is None:
raise ValueError(f"empty workbook: {path}")
names = [str(value).strip().lower() if value is not None else "" for value in header]
required = ("video_id", "clip_id", "text", "label", "annotation")
missing = set(required) - set(names)
if missing:
raise ValueError(f"missing required label columns: {sorted(missing)}")
indexes = {name: names.index(name) for name in required}
records = []
for values in rows:
if not values or all(value is None for value in values):
continue
record = {name: values[index] if index < len(values) else None for name, index in indexes.items()}
record["video_id"] = _identifier(record["video_id"])
record["clip_id"] = _identifier(record["clip_id"])
record["text"] = "" if record["text"] is None else str(record["text"]).strip()
record["annotation"] = "" if record["annotation"] is None else str(record["annotation"]).strip()
if record["label"] is not None:
try:
record["label"] = float(record["label"])
except (TypeError, ValueError):
pass
records.append(record)
workbook.close()
return records
def _probe_video(path: Path) -> dict[str, Any]:
result = subprocess.run(
[
"ffprobe", "-v", "error", "-show_entries",
"format=duration:stream=codec_type,codec_name,width,height,avg_frame_rate,r_frame_rate,nb_frames,sample_rate,channels:frame=media_type,best_effort_timestamp_time,pkt_duration_time",
"-show_frames", "-of", "json", str(path),
],
check=True,
capture_output=True,
text=True,
)
payload = json.loads(result.stdout)
streams = payload.get("streams", [])
video = next((stream for stream in streams if stream.get("codec_type") == "video"), {})
audio = next((stream for stream in streams if stream.get("codec_type") == "audio"), {})
fps_text = video.get("avg_frame_rate") or video.get("r_frame_rate") or "0/1"
try:
fps = float(Fraction(fps_text))
except (ValueError, ZeroDivisionError):
fps = 0.0
container_duration = float(payload.get("format", {}).get("duration", 0.0))
frame_times: dict[str, list[tuple[float, float]]] = {"video": [], "audio": []}
for frame in payload.get("frames", []):
media_type = frame.get("media_type")
timestamp = frame.get("best_effort_timestamp_time")
if media_type not in frame_times or timestamp is None:
continue
try:
start = float(timestamp)
packet_duration = float(frame.get("pkt_duration_time", 0.0) or 0.0)
except (TypeError, ValueError):
continue
frame_times[media_type].append((start, packet_duration))
video_times = frame_times["video"]
audio_times = frame_times["audio"]
fallback_frame_duration = 1.0 / fps if fps > 0 else 0.0
video_start = min((item[0] for item in video_times), default=0.0)
video_end = max((start + (packet_duration or fallback_frame_duration) for start, packet_duration in video_times), default=container_duration)
audio_start = min((item[0] for item in audio_times), default=0.0)
audio_end = max((start + packet_duration for start, packet_duration in audio_times), default=container_duration)
decoded_duration = max(0.0, video_end - video_start)
audio_duration = max(0.0, audio_end - audio_start)
return {
"duration_s": decoded_duration or container_duration,
"container_duration_s": container_duration,
"audio_duration_s": audio_duration,
"video_timeline_start_s": video_start,
"video_timeline_end_s": video_end,
"audio_timeline_start_s": audio_start,
"audio_timeline_end_s": audio_end,
"video_codec": video.get("codec_name", ""),
"width": video.get("width", ""),
"height": video.get("height", ""),
"fps": fps,
"video_frames": len(video_times),
"container_video_frames": video.get("nb_frames", ""),
"audio_codec": audio.get("codec_name", ""),
"audio_sample_rate": audio.get("sample_rate", ""),
"audio_channels": audio.get("channels", ""),
"has_audio": bool(audio),
}
def audit_dataset(video_root: Path, label_file: Path, output_dir: Path, expected_count: int = 100) -> dict[str, Any]:
records = _read_labels(label_file)
videos: dict[tuple[str, str], Path] = {}
duplicate_video_files: list[str] = []
for video in sorted(video_root.rglob("*.mp4")):
key = (video.parent.name, video.stem)
if key in videos:
duplicate_video_files.append(str(video))
else:
videos[key] = video
keys = [(str(row["video_id"]), str(row["clip_id"])) for row in records]
counts = Counter(keys)
duplicate_rows = [list(key) for key, count in counts.items() if count > 1]
matched_keys = set(keys) & set(videos)
missing_keys = [key for key in keys if key not in videos]
extra_keys = sorted(set(videos) - set(keys))
class_mismatches = []
probe_errors: list[dict[str, str]] = []
probe_by_key: dict[tuple[str, str], dict[str, Any]] = {}
for key in matched_keys:
try:
probe_by_key[key] = _probe_video(videos[key])
except (OSError, subprocess.CalledProcessError, json.JSONDecodeError, ValueError) as error:
probe_errors.append({"video_id": key[0], "clip_id": key[1], "error": str(error)})
output_dir.mkdir(parents=True, exist_ok=True)
manifest_path = output_dir / "manifest.csv"
with manifest_path.open("w", encoding="utf-8-sig", newline="") as file:
writer = csv.DictWriter(
file,
fieldnames=(
"video_id", "clip_id", "group_id", "text", "label", "annotation", "video_path",
"video_exists", "duration_s", "container_duration_s", "audio_duration_s",
"video_timeline_start_s", "video_timeline_end_s", "audio_timeline_start_s",
"audio_timeline_end_s", "video_codec", "width", "height", "fps",
"video_frames", "container_video_frames", "audio_codec", "audio_sample_rate",
"audio_channels", "has_audio",
),
)
writer.writeheader()
for row, key in zip(records, keys):
label = row["label"]
annotation = str(row["annotation"]).strip().lower()
expected_class = "negative" if isinstance(label, (int, float)) and label < 0 else (
"positive" if isinstance(label, (int, float)) and label > 0 else "neutral"
)
if annotation in {"negative", "neutral", "positive"} and annotation != expected_class:
class_mismatches.append({"video_id": key[0], "clip_id": key[1], "label": label, "annotation": annotation})
path = videos.get(key)
probe = probe_by_key.get(key, {})
writer.writerow({
"video_id": key[0],
"clip_id": key[1],
"group_id": key[0],
"text": row["text"],
"label": row["label"],
"annotation": row["annotation"],
"video_path": str(path.relative_to(video_root)) if path else "",
"video_exists": bool(path),
**probe,
})
durations = [info["duration_s"] for info in probe_by_key.values() if info["duration_s"] > 0]
duration_out_of_range = [
{"video_id": key[0], "clip_id": key[1], "duration_s": info["duration_s"]}
for key, info in probe_by_key.items()
if not (2.648 <= info["duration_s"] <= 34.567)
]
audio_missing = [list(key) for key, info in probe_by_key.items() if not info["has_audio"]]
summary = {
"video_root": str(video_root),
"label_file": str(label_file),
"expected_count": expected_count,
"label_rows": len(records),
"unique_video_clip_pairs": len(set(keys)),
"unique_video_ids": len({key[0] for key in keys}),
"video_files": len(videos),
"matched_samples": len(matched_keys),
"coverage_rate": len(matched_keys) / max(len(records), 1),
"duration_min_s": min(durations) if durations else None,
"duration_max_s": max(durations) if durations else None,
"stated_duration_range_s": [2.648, 34.567],
"duration_out_of_range": duration_out_of_range,
"audio_stream_missing": audio_missing,
"media_probe_errors": probe_errors,
"missing_video_pairs": [list(key) for key in missing_keys],
"unlisted_video_pairs": [list(key) for key in extra_keys],
"duplicate_label_pairs": duplicate_rows,
"duplicate_video_files": duplicate_video_files,
"label_class_mismatches": class_mismatches,
"manifest": str(manifest_path),
}
summary["coverage_status"] = (
"complete_with_metadata_warnings" if duration_out_of_range else "complete"
)
summary_path = output_dir / "coverage_summary.json"
summary_path.write_text(json.dumps(summary, ensure_ascii=False, indent=2), encoding="utf-8")
return summary
def main() -> int:
project_dir = Path(__file__).resolve().parents[1]
repo_dir = project_dir.parent
default_data = repo_dir / "E题数据" / "附件1-数据集原始多模态样本" / "MOSEI数据集部分原始视频-100条"
parser = argparse.ArgumentParser(description="Audit the 100 raw-video Q1 samples and export a manifest.")
parser.add_argument("--video-root", type=Path, default=default_data)
parser.add_argument("--labels", type=Path, default=default_data / "label-100.xlsx")
parser.add_argument("--output-dir", type=Path, default=project_dir / "outputs" / "audit")
parser.add_argument("--expected-count", type=int, default=100)
args = parser.parse_args()
summary = audit_dataset(args.video_root, args.labels, args.output_dir, args.expected_count)
print(json.dumps(summary, ensure_ascii=False, indent=2))
complete = (
summary["label_rows"] == args.expected_count
and summary["unique_video_clip_pairs"] == args.expected_count
and summary["matched_samples"] == args.expected_count
and not summary["duplicate_label_pairs"]
and not summary["label_class_mismatches"]
and not summary["audio_stream_missing"]
and not summary["media_probe_errors"]
)
return 0 if complete else 1
if __name__ == "__main__":
raise SystemExit(main())
+945
View File
@@ -0,0 +1,945 @@
from __future__ import annotations
import argparse
import csv
import json
import math
import platform
import random
import statistics
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
from sklearn.model_selection import GroupKFold
from torch import Tensor, nn
from .alignment import align_fixed_windows, align_forced_timestamps, make_block_mask
from .experiment_data import (
FeatureSample,
FeatureStats,
collate_feature_samples,
fit_feature_stats,
load_feature_samples,
)
from .experiment_probes import (
RetrievalProjection,
run_frozen_emotion_probe,
run_reconstruction_probe,
run_retrieval_probe,
)
from .metrics import (
attention_row_similarity,
alignment_trajectory,
attention_width80,
monotonicity_violation_rate,
normalized_attention_entropy,
)
from .models import SharedLatentTimeline, TextAnchoredCrossAttention
from .types import AlignmentOutput, MODALITIES
class AlignmentReconstructor(nn.Module):
"""Shared M3/M4 training decoder: predict one stream from the other two."""
def __init__(self, hidden_size: int = 128, dropout: float = 0.1) -> None:
super().__init__()
self.decoders = nn.ModuleDict(
{
target: nn.Sequential(
nn.Linear(hidden_size * 2, hidden_size),
nn.GELU(),
nn.Dropout(dropout),
nn.Linear(hidden_size, hidden_size),
)
for target in MODALITIES
}
)
def forward(self, target: str, sources: Mapping[str, Tensor], mask: Tensor) -> Tensor:
values = [
sources[name].masked_fill(mask.unsqueeze(-1), 0.0)
for name in MODALITIES
if name != target
]
return self.decoders[target](torch.cat(values, dim=-1))
def _seed_everything(seed: int) -> None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)
if hasattr(torch.backends, "cudnn"):
torch.backends.cudnn.deterministic = True
torch.backends.cudnn.benchmark = False
def _batches(
samples: Sequence[FeatureSample], batch_size: int, *, shuffle: bool, rng: np.random.Generator
) -> list[list[FeatureSample]]:
if shuffle:
order = rng.permutation(len(samples)).tolist()
else:
order = list(range(len(samples)))
return [[samples[index] for index in order[start : start + batch_size]]
for start in range(0, len(order), batch_size)]
def _training_objective(
output: AlignmentOutput,
sequences: Mapping[str, Any],
durations: Tensor,
decoder: AlignmentReconstructor,
generator: torch.Generator,
*,
method: str = "M3",
loss_variant: str = "v1",
) -> tuple[Tensor, dict[str, Tensor]]:
reconstruction_terms = []
batch_size, grid_size = output.aligned["text"].shape[:2]
for target in MODALITIES:
mask = make_block_mask(
batch_size,
grid_size,
0.2,
output.aligned[target].device,
generator=generator,
)
prediction = decoder(target, output.aligned, mask)
reconstruction_terms.append(
nn.functional.smooth_l1_loss(prediction[mask], output.aligned[target][mask])
)
reconstruction = torch.stack(reconstruction_terms).mean()
times = {name: sequences[name].times for name in MODALITIES}
variant_weights = {
"v1": (0.0, 0.0, 0.0),
"v2_a": (5.0, 0.0, 0.0),
"v2_b": (5.0, 0.5, 0.0),
"v2_c": (5.0, 0.5, 10.0),
}
if loss_variant not in variant_weights:
raise ValueError(f"unknown loss variant: {loss_variant}")
lambda_span, lambda_div, lambda_band = variant_weights[loss_variant]
grid_size = output.weights["text"].shape[1]
if method == "M3":
text_reference = torch.bmm(
output.weights["text"], sequences["text"].times.unsqueeze(-1)
).squeeze(-1)
text_reference = text_reference / durations[:, None].clamp_min(1e-8)
band_targets = {"audio": text_reference, "vision": text_reference}
diversity_modalities = ("audio", "vision")
else:
centers = (
torch.arange(grid_size, dtype=durations.dtype, device=durations.device) + 0.5
) / grid_size
reference = centers.unsqueeze(0).expand(durations.shape[0], -1)
band_targets = {name: reference for name in MODALITIES}
diversity_modalities = MODALITIES
from .losses import alignment_training_loss
return alignment_training_loss(
output,
times,
durations,
reconstruction,
lambda_rec=1.0,
lambda_con=1.0,
lambda_mono=0.1,
lambda_span=lambda_span,
lambda_div=lambda_div,
lambda_band=lambda_band,
epsilon=0.02,
minimum_span=0.7,
coverage_modalities=diversity_modalities,
diversity_modalities=diversity_modalities,
diversity_min_separation=6,
band_targets=band_targets,
band_margin=0.1,
)
def _make_learned_model(
method: str,
dimensions: Mapping[str, int],
grid_size: int,
hidden_size: int,
heads: int,
dropout: float,
) -> nn.Module:
if method == "M3":
return TextAnchoredCrossAttention(
dimensions,
grid_size=grid_size,
hidden_size=hidden_size,
heads=heads,
dropout=dropout,
)
if method == "M4":
return SharedLatentTimeline(
dimensions,
grid_size=grid_size,
hidden_size=hidden_size,
heads=heads,
dropout=dropout,
)
raise ValueError(f"unknown learned method: {method}")
def _fit_learned_model(
method: str,
train_samples: Sequence[FeatureSample],
val_samples: Sequence[FeatureSample],
stats: FeatureStats,
*,
device: torch.device,
seed: int,
grid_size: int,
hidden_size: int,
heads: int,
dropout: float,
batch_size: int,
max_epochs: int,
patience: int,
learning_rate: float,
checkpoint_path: Path,
loss_variant: str = "v1",
) -> tuple[nn.Module, dict[str, Any]]:
_seed_everything(seed)
dimensions = {name: train_samples[0].features[name].shape[1] for name in MODALITIES}
model = _make_learned_model(method, dimensions, grid_size, hidden_size, heads, dropout).to(device)
decoder = AlignmentReconstructor(hidden_size, dropout).to(device)
optimizer = torch.optim.AdamW(
[*model.parameters(), *decoder.parameters()], lr=learning_rate, weight_decay=1e-4
)
rng = np.random.default_rng(seed)
train_mask_generator = torch.Generator(device=device)
train_mask_generator.manual_seed(seed + 31)
history: list[dict[str, float]] = []
best_loss = math.inf
best_epoch = 0
best_model: dict[str, Tensor] | None = None
best_decoder: dict[str, Tensor] | None = None
patience_used = 0
for epoch in range(1, max_epochs + 1):
model.train()
decoder.train()
train_total = 0.0
train_count = 0
for batch_samples in _batches(train_samples, batch_size, shuffle=True, rng=rng):
sequences, durations, _ = collate_feature_samples(batch_samples, stats, device)
output = model(sequences)
total, _ = _training_objective(
output,
sequences,
durations,
decoder,
train_mask_generator,
method=method,
loss_variant=loss_variant,
)
if not torch.isfinite(total):
raise FloatingPointError(f"non-finite {method} objective at epoch {epoch}")
optimizer.zero_grad(set_to_none=True)
total.backward()
nn.utils.clip_grad_norm_([*model.parameters(), *decoder.parameters()], 1.0)
optimizer.step()
train_total += float(total.detach().item()) * len(batch_samples)
train_count += len(batch_samples)
model.eval()
decoder.eval()
val_generator = torch.Generator(device=device)
val_generator.manual_seed(seed + 99991)
val_total = 0.0
val_count = 0
val_metric_sums: dict[str, float] = defaultdict(float)
with torch.no_grad():
for batch_samples in _batches(val_samples, batch_size, shuffle=False, rng=rng):
sequences, durations, _ = collate_feature_samples(batch_samples, stats, device)
output = model(sequences)
total, parts = _training_objective(
output,
sequences,
durations,
decoder,
val_generator,
method=method,
loss_variant=loss_variant,
)
count = len(batch_samples)
val_total += float(total.item()) * count
val_count += count
for key, value in parts.items():
val_metric_sums[key] += float(value.item()) * count
for name in MODALITIES:
val_metric_sums[f"c_row_{name}"] += float(
attention_row_similarity(output.weights[name]).mean().item()
) * count
val_metric_sums[f"c_far_{name}"] += float(
attention_row_similarity(output.weights[name], min_separation=6)
.mean()
.item()
) * count
val_mean = val_total / max(val_count, 1)
history.append(
{
"epoch": float(epoch),
"train_total": train_total / max(train_count, 1),
"validation_total": val_mean,
**{
f"validation_{key}": value / max(val_count, 1)
for key, value in val_metric_sums.items()
},
}
)
if val_mean < best_loss - 1e-6:
best_loss = val_mean
best_epoch = epoch
best_model = {key: value.detach().cpu().clone() for key, value in model.state_dict().items()}
best_decoder = {
key: value.detach().cpu().clone() for key, value in decoder.state_dict().items()
}
patience_used = 0
else:
patience_used += 1
if patience_used >= patience:
break
if best_model is None or best_decoder is None:
raise RuntimeError(f"{method} training did not produce a finite checkpoint")
model.load_state_dict(best_model)
decoder.load_state_dict(best_decoder)
model.eval()
decoder.eval()
checkpoint_path.parent.mkdir(parents=True, exist_ok=True)
torch.save(
{
"method": method,
"seed": seed,
"loss_variant": loss_variant,
"best_epoch": best_epoch,
"best_validation_objective": best_loss,
"model_state_dict": best_model,
"training_decoder_state_dict": best_decoder,
"history": history,
},
checkpoint_path,
)
return model, {
"best_epoch": best_epoch,
"best_validation_objective": best_loss,
"history": history,
"best_validation_metrics": history[best_epoch - 1],
"checkpoint": str(checkpoint_path),
}
def _baseline_output(
method: str,
sequences: Mapping[str, Any],
durations: Tensor,
word_intervals: Sequence[Tensor],
grid_size: int,
) -> AlignmentOutput:
if method == "M1":
return align_forced_timestamps(sequences, word_intervals, grid_size)
if method == "M2":
return align_fixed_windows(sequences, durations, grid_size)
raise ValueError(f"not a fixed baseline: {method}")
def _alignment_rows(
method: str,
seed_label: str,
fold: int,
samples: Sequence[FeatureSample],
output: AlignmentOutput,
device: torch.device,
epsilon: float,
) -> list[dict[str, Any]]:
rows = []
for index, sample in enumerate(samples):
for name in MODALITIES:
length = len(sample.times[name])
weights = output.weights[name][index : index + 1, :, :length]
times = torch.as_tensor(sample.times[name], dtype=torch.float32, device=device)[None]
valid = torch.as_tensor(sample.valid[name], dtype=torch.bool, device=device)[None]
duration = torch.tensor([sample.duration_s], dtype=torch.float32, device=device)
trajectory = alignment_trajectory(weights, times, duration)
mvr = monotonicity_violation_rate(trajectory, epsilon)
entropy = normalized_attention_entropy(weights, valid)
width = attention_width80(weights)
rows.append(
{
"method": method,
"seed": seed_label,
"fold": fold,
"sample_id": sample.sample_id,
"modality": name,
"mvr": float(mvr[0].item()),
"normalized_entropy": float(entropy.mean().item()),
"width80_source_positions": float(width.float().mean().item()),
"c_row": float(attention_row_similarity(weights).mean().item()),
"c_far": float(
attention_row_similarity(weights, min_separation=6).mean().item()
),
"expected_time_start_s": float(trajectory[0, 0].item() * sample.duration_s),
"expected_time_end_s": float(trajectory[0, -1].item() * sample.duration_s),
"trajectory_span_fraction": float((trajectory[0, -1] - trajectory[0, 0]).item()),
}
)
return rows
def _collect_representations(
method: str,
samples: Sequence[FeatureSample],
stats: FeatureStats,
*,
device: torch.device,
grid_size: int,
batch_size: int,
model: nn.Module | None = None,
) -> tuple[dict[str, dict[str, np.ndarray]], dict[str, dict[str, np.ndarray]]]:
aligned_by_id: dict[str, dict[str, np.ndarray]] = {}
weights_by_id: dict[str, dict[str, np.ndarray]] = {}
rng = np.random.default_rng(0)
if model is not None:
model.eval()
with torch.no_grad():
for batch_samples in _batches(samples, batch_size, shuffle=False, rng=rng):
sequences, durations, intervals = collate_feature_samples(batch_samples, stats, device)
if method in {"M1", "M2"}:
output = _baseline_output(method, sequences, durations, intervals, grid_size)
elif model is not None:
output = model(sequences)
else:
raise ValueError(f"a trained model is required for {method}")
for index, sample in enumerate(batch_samples):
aligned: dict[str, np.ndarray] = {}
weights: dict[str, np.ndarray] = {}
for name in MODALITIES:
length = len(sample.features[name])
matrix = output.weights[name][index, :, :length]
source = sequences[name].features[index, :length]
pooled = matrix.to(source.dtype) @ source
aligned[name] = pooled.detach().cpu().numpy().astype(np.float32, copy=False)
weights[name] = matrix.detach().cpu().numpy().astype(np.float32, copy=False)
aligned_by_id[sample.sample_id] = aligned
weights_by_id[sample.sample_id] = weights
return aligned_by_id, weights_by_id
def _save_alignment(
output_dir: Path,
method: str,
seed_label: str,
sample: FeatureSample,
weights: Mapping[str, np.ndarray],
) -> None:
path = output_dir / "alignments" / method / f"seed_{seed_label}" / (
sample.sample_id.replace("/", "__") + ".npz"
)
path.parent.mkdir(parents=True, exist_ok=True)
values: dict[str, np.ndarray] = {"sample_id": np.asarray(sample.sample_id)}
for name in MODALITIES:
trajectory = weights[name] @ sample.times[name] / max(sample.duration_s, 1e-8)
values[f"weights_{name}"] = weights[name]
values[f"times_{name}_s"] = sample.times[name].astype(np.float32, copy=False)
values[f"valid_{name}"] = sample.valid[name]
values[f"trajectory_{name}"] = trajectory.astype(np.float32, copy=False)
np.savez_compressed(path, **values)
def _write_csv(path: Path, rows: Sequence[Mapping[str, Any]]) -> None:
if not rows:
return
path.parent.mkdir(parents=True, exist_ok=True)
columns = list(dict.fromkeys(key for row in rows for key in row))
with path.open("w", encoding="utf-8-sig", newline="") as file:
writer = csv.DictWriter(file, fieldnames=columns, extrasaction="ignore")
writer.writeheader()
writer.writerows(rows)
def _group_summary(
rows: Sequence[Mapping[str, Any]], group_columns: Sequence[str], metric_columns: Sequence[str]
) -> list[dict[str, Any]]:
groups: dict[tuple[Any, ...], list[Mapping[str, Any]]] = defaultdict(list)
for row in rows:
groups[tuple(row[column] for column in group_columns)].append(row)
output = []
for key, values in groups.items():
summary: dict[str, Any] = dict(zip(group_columns, key))
summary["n"] = len(values)
for metric in metric_columns:
numbers = [float(row[metric]) for row in values if row.get(metric) not in (None, "")]
numbers = [value for value in numbers if math.isfinite(value)]
if numbers:
summary[f"{metric}_mean"] = statistics.fmean(numbers)
summary[f"{metric}_std"] = statistics.stdev(numbers) if len(numbers) > 1 else 0.0
output.append(summary)
return output
def _comparison_table(summaries: Mapping[str, Sequence[Mapping[str, Any]]]) -> list[dict[str, Any]]:
"""Make one compact, multi-metric table without inventing a composite score."""
alignment = {(row["method"], row["modality"]): row for row in summaries["alignment"]}
retrieval = {(row["method"], row["direction"]): row for row in summaries["retrieval"]}
reconstruction = {
(row["method"], row["target_modality"]): row for row in summaries["reconstruction"]
}
emotion = {row["method"]: row for row in summaries["emotion"]}
table: list[dict[str, Any]] = []
for method in ("M1", "M2", "M3", "M4"):
row: dict[str, Any] = {"method": method}
for modality in MODALITIES:
metrics = alignment[(method, modality)]
row[f"mvr_{modality}"] = metrics.get("mvr_mean")
row[f"entropy_{modality}"] = metrics.get("normalized_entropy_mean")
row[f"trajectory_span_{modality}"] = metrics.get("trajectory_span_fraction_mean")
for direction in ("text_to_audio", "text_to_vision"):
metrics = retrieval[(method, direction)]
row[f"r_at_1_{direction}"] = metrics.get("r_at_1_mean")
row[f"r_at_5_{direction}"] = metrics.get("r_at_5_mean")
for modality in MODALITIES:
metrics = reconstruction[(method, modality)]
row[f"reconstruction_mae_{modality}"] = metrics.get("mae_standardized_mean")
for metric in ("accuracy", "macro_f1", "mae", "pearson"):
row[f"emotion_{metric}"] = emotion[method].get(f"{metric}_mean")
table.append(row)
return table
def _make_figures(
output_dir: Path,
example_id: str,
example_sample: FeatureSample,
example_weights: Mapping[str, Mapping[str, np.ndarray]],
grid_size: int,
) -> None:
method_order = ("M1", "M2", "M3", "M4")
fig, axes = plt.subplots(4, 3, figsize=(15, 13), constrained_layout=True)
for row, method in enumerate(method_order):
if method not in example_weights:
continue
for col, name in enumerate(MODALITIES):
ax = axes[row, col]
matrix = example_weights[method][name]
image = ax.imshow(matrix, origin="lower", aspect="auto", interpolation="nearest", cmap="magma")
ax.set_title(f"{method} · {name}")
ax.set_xlabel("source position")
ax.set_ylabel("shared grid slot")
ax.set_yticks(np.linspace(0, grid_size - 1, 5, dtype=int))
fig.colorbar(image, ax=ax, fraction=0.046, pad=0.04)
fig.suptitle(f"Alignment matrices on held-out sample {example_id}")
fig.savefig(output_dir / "typical_alignment_heatmaps.png", dpi=170)
plt.close(fig)
fig, axes = plt.subplots(2, 2, figsize=(13, 9), constrained_layout=True)
x = (np.arange(grid_size, dtype=np.float32) + 0.5) / grid_size
for ax, method in zip(axes.flat, method_order):
for name in MODALITIES:
matrix = example_weights[method][name]
duration = example_sample.duration_s
trajectory = matrix @ example_sample.times[name] / max(duration, 1e-8)
ax.plot(x, trajectory, label=name)
ax.plot([0, 1], [0, 1], linestyle="--", color="black", alpha=0.5, label="uniform-time reference")
ax.set_title(method)
ax.set_xlabel("shared-grid position")
ax.set_ylabel("expected source time / clip duration")
ax.set_xlim(0, 1)
ax.set_ylim(0, 1)
ax.grid(alpha=0.2)
axes[0, 0].legend(fontsize=8)
fig.suptitle(f"Alignment trajectories on held-out sample {example_id}")
fig.savefig(output_dir / "typical_alignment_trajectories.png", dpi=170)
plt.close(fig)
def run(args: argparse.Namespace) -> dict[str, Any]:
start_time = time.time()
_seed_everything(args.seeds[0])
if args.device == "auto":
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
else:
device = torch.device(args.device)
if device.type == "cuda" and not torch.cuda.is_available():
raise RuntimeError("CUDA was requested but is not available in this WSL environment")
samples = load_feature_samples(args.feature_dir, args.manifest)
if len(samples) != args.expected_samples:
raise ValueError(f"expected {args.expected_samples} samples, found {len(samples)}")
groups = [sample.group_id for sample in samples]
if len(set(groups)) < args.folds:
raise ValueError("fewer video_id groups than requested folds")
split_iter = GroupKFold(n_splits=args.folds).split(np.zeros(len(samples)), groups=groups)
splits = [(train.tolist(), val.tolist()) for train, val in split_iter]
splits_json = [
{
"fold": fold + 1,
"train_sample_ids": [samples[index].sample_id for index in train],
"validation_sample_ids": [samples[index].sample_id for index in val],
"train_video_ids": sorted({samples[index].group_id for index in train}),
"validation_video_ids": sorted({samples[index].group_id for index in val}),
}
for fold, (train, val) in enumerate(splits)
]
args.output_dir.mkdir(parents=True, exist_ok=True)
print(
f"[start] samples={len(samples)} groups={len(set(groups))} folds={args.folds} "
f"seeds={args.seeds} device={device}",
flush=True,
)
(args.output_dir / "splits.json").write_text(
json.dumps(splits_json, ensure_ascii=False, indent=2), encoding="utf-8"
)
alignment_rows: list[dict[str, Any]] = []
retrieval_rows: list[dict[str, Any]] = []
reconstruction_rows: list[dict[str, Any]] = []
emotion_rows: list[dict[str, Any]] = []
training_rows: list[dict[str, Any]] = []
examples: dict[str, dict[str, Mapping[str, np.ndarray]]] = defaultdict(dict)
sample_by_id = {sample.sample_id: sample for sample in samples}
preferred_example = args.example_id if args.example_id in sample_by_id else samples[0].sample_id
preferred_seed = args.seeds[0]
example_sample = sample_by_id[preferred_example]
for fold_index, (train_indices, val_indices) in enumerate(splits, start=1):
train_samples = [samples[index] for index in train_indices]
val_samples = [samples[index] for index in val_indices]
stats = fit_feature_stats(train_samples)
fold_output = args.output_dir / f"fold_{fold_index:02d}"
print(
f"[fold {fold_index}/{args.folds}] train={len(train_samples)} validation={len(val_samples)} "
f"train_video_ids={len({sample.group_id for sample in train_samples})} "
f"validation_video_ids={len({sample.group_id for sample in val_samples})}",
flush=True,
)
# M1 and M2 are deterministic methods with no learned alignment loss.
for method in ("M1", "M2"):
print(f"[fold {fold_index}] evaluate {method} and train identical probes", flush=True)
train_aligned, _ = _collect_representations(
method,
train_samples,
stats,
device=device,
grid_size=args.grid_size,
batch_size=args.batch_size,
)
val_aligned, val_weights = _collect_representations(
method,
val_samples,
stats,
device=device,
grid_size=args.grid_size,
batch_size=args.batch_size,
)
combined = {**train_aligned, **val_aligned}
rng = np.random.default_rng(args.seeds[0] + fold_index)
for batch_samples in _batches(val_samples, args.batch_size, shuffle=False, rng=rng):
sequences, durations, intervals = collate_feature_samples(batch_samples, stats, device)
output = _baseline_output(method, sequences, durations, intervals, args.grid_size)
alignment_rows.extend(
_alignment_rows(method, "fixed", fold_index, batch_samples, output, device, args.mvr_epsilon)
)
for sample in val_samples:
_save_alignment(args.output_dir, method, "fixed", sample, val_weights[sample.sample_id])
if sample.sample_id == preferred_example:
examples[sample.sample_id][method] = val_weights[sample.sample_id]
probe_seed = args.seeds[0] + fold_index * 100
retrieval_rows.extend(
{
"method": method,
"seed": "fixed",
"fold": fold_index,
**row,
}
for row in run_retrieval_probe(
[sample.sample_id for sample in train_samples],
[sample.sample_id for sample in val_samples],
combined,
device=device,
seed=probe_seed,
epochs=args.retrieval_probe_epochs,
batch_size=args.probe_batch_size,
)
)
reconstruction_rows.extend(
{"method": method, "seed": "fixed", "fold": fold_index, **row}
for row in run_reconstruction_probe(
[sample.sample_id for sample in train_samples],
[sample.sample_id for sample in val_samples],
combined,
device=device,
seed=probe_seed + 1,
ratio=args.mask_ratio,
epochs=args.reconstruction_probe_epochs,
batch_size=args.probe_batch_size,
)
)
emotion_rows.append(
{
"method": method,
"seed": "fixed",
"fold": fold_index,
**run_frozen_emotion_probe(train_samples, val_samples, combined),
}
)
# M3/M4 share one unsupervised objective, split, and training budget.
for seed in args.seeds:
for method in ("M3", "M4"):
print(f"[fold {fold_index}] train {method}, seed={seed}", flush=True)
checkpoint = fold_output / f"seed_{seed}" / f"{method}.pt"
model, training_info = _fit_learned_model(
method,
train_samples,
val_samples,
stats,
device=device,
seed=seed + fold_index * 1009,
grid_size=args.grid_size,
hidden_size=args.hidden_size,
heads=args.heads,
dropout=args.dropout,
batch_size=args.batch_size,
max_epochs=args.epochs,
patience=args.patience,
learning_rate=args.learning_rate,
checkpoint_path=checkpoint,
)
training_rows.append(
{
"method": method,
"seed": seed,
"fold": fold_index,
"best_epoch": training_info["best_epoch"],
"best_validation_objective": training_info["best_validation_objective"],
"checkpoint": training_info["checkpoint"],
}
)
print(
f"[fold {fold_index}] {method}, seed={seed} best_epoch={training_info['best_epoch']} "
f"val_objective={training_info['best_validation_objective']:.5f}; running frozen probes",
flush=True,
)
train_aligned, _ = _collect_representations(
method,
train_samples,
stats,
device=device,
grid_size=args.grid_size,
batch_size=args.batch_size,
model=model,
)
val_aligned, val_weights = _collect_representations(
method,
val_samples,
stats,
device=device,
grid_size=args.grid_size,
batch_size=args.batch_size,
model=model,
)
combined = {**train_aligned, **val_aligned}
rng = np.random.default_rng(seed + fold_index)
for batch_samples in _batches(val_samples, args.batch_size, shuffle=False, rng=rng):
sequences, durations, _ = collate_feature_samples(batch_samples, stats, device)
output = model(sequences)
alignment_rows.extend(
_alignment_rows(method, str(seed), fold_index, batch_samples, output, device, args.mvr_epsilon)
)
for sample in val_samples:
_save_alignment(args.output_dir, method, str(seed), sample, val_weights[sample.sample_id])
if sample.sample_id == preferred_example and seed == preferred_seed:
examples[sample.sample_id][method] = val_weights[sample.sample_id]
probe_seed = seed + fold_index * 100 + (3 if method == "M3" else 7)
retrieval_rows.extend(
{
"method": method,
"seed": seed,
"fold": fold_index,
**row,
}
for row in run_retrieval_probe(
[sample.sample_id for sample in train_samples],
[sample.sample_id for sample in val_samples],
combined,
device=device,
seed=probe_seed,
epochs=args.retrieval_probe_epochs,
batch_size=args.probe_batch_size,
)
)
reconstruction_rows.extend(
{"method": method, "seed": seed, "fold": fold_index, **row}
for row in run_reconstruction_probe(
[sample.sample_id for sample in train_samples],
[sample.sample_id for sample in val_samples],
combined,
device=device,
seed=probe_seed + 1,
ratio=args.mask_ratio,
epochs=args.reconstruction_probe_epochs,
batch_size=args.probe_batch_size,
)
)
emotion_rows.append(
{
"method": method,
"seed": seed,
"fold": fold_index,
**run_frozen_emotion_probe(train_samples, val_samples, combined),
}
)
del model
if device.type == "cuda":
torch.cuda.empty_cache()
_write_csv(args.output_dir / "alignment_metrics.csv", alignment_rows)
_write_csv(args.output_dir / "retrieval_probe_metrics.csv", retrieval_rows)
_write_csv(args.output_dir / "reconstruction_probe_metrics.csv", reconstruction_rows)
_write_csv(args.output_dir / "frozen_emotion_probe_metrics.csv", emotion_rows)
_write_csv(args.output_dir / "training_summary.csv", training_rows)
summaries = {
"alignment": _group_summary(
alignment_rows,
("method", "modality"),
("mvr", "normalized_entropy", "width80_source_positions", "expected_time_start_s", "expected_time_end_s", "trajectory_span_fraction"),
),
"retrieval": _group_summary(
retrieval_rows,
("method", "direction"),
("r_at_1", "r_at_5", "mrr"),
),
"reconstruction": _group_summary(
reconstruction_rows,
("method", "target_modality"),
("mae_standardized", "smooth_l1_standardized"),
),
"emotion": _group_summary(
emotion_rows,
("method",),
("accuracy", "macro_f1", "mae", "pearson"),
),
}
(args.output_dir / "summary.json").write_text(
json.dumps(summaries, ensure_ascii=False, indent=2, allow_nan=False), encoding="utf-8"
)
for name, rows in summaries.items():
_write_csv(args.output_dir / f"{name}_summary.csv", rows)
comparison_table = _comparison_table(summaries)
_write_csv(args.output_dir / "comparison_summary.csv", comparison_table)
if preferred_example in examples and set(examples[preferred_example]) == {"M1", "M2", "M3", "M4"}:
_make_figures(args.output_dir, preferred_example, example_sample, examples[preferred_example], args.grid_size)
manifest = {
"created_utc": datetime.now(timezone.utc).isoformat(),
"sample_count": len(samples),
"group_count": len(set(groups)),
"folds": args.folds,
"seeds": args.seeds,
"device": str(device),
"gpu_name": torch.cuda.get_device_name(device) if device.type == "cuda" else None,
"python": platform.python_version(),
"torch": torch.__version__,
"parameters": {
"grid_size": args.grid_size,
"hidden_size": args.hidden_size,
"heads": args.heads,
"dropout": args.dropout,
"batch_size": args.batch_size,
"epochs_max": args.epochs,
"early_stopping_patience": args.patience,
"learning_rate": args.learning_rate,
"mask_ratio": args.mask_ratio,
"retrieval_probe_epochs": args.retrieval_probe_epochs,
"reconstruction_probe_epochs": args.reconstruction_probe_epochs,
"mvr_epsilon": args.mvr_epsilon,
},
"objective": "masked reconstruction + cross-modal contrastive + temporal monotonicity; emotion labels unused",
"split_rule": "GroupKFold by group_id/video_id",
"example_sample_id": preferred_example,
"elapsed_seconds": time.time() - start_time,
"interpretation_limits": [
"Grid-index retrieval is a representation-consistency probe, not independent temporal ground truth.",
"Masked reconstruction uses a decoder trained on the training fold and reports standardized-feature errors.",
"The emotion probe is a small-sample downstream utility check, not a claim of generalization to MOSEI.",
"No human event timestamps are available, so human IoU/MATE is not reported.",
],
}
(args.output_dir / "run_manifest.json").write_text(
json.dumps(manifest, ensure_ascii=False, indent=2), encoding="utf-8"
)
(args.output_dir / "README.md").write_text(
"# Q1 method comparison\n\n"
"This folder contains grouped cross-validation results for M1–M4. M1/M2 are fixed rules; M3/M4 are trained without emotion labels. "
"All learned models and probes use training-fold-only feature normalization, and folds are grouped by `video_id`.\n\n"
"`alignment_metrics.csv` reports expected-time trajectories, monotonicity, attention entropy, and width diagnostics. "
"`retrieval_probe_metrics.csv` uses a separately trained linear projection probe; its grid-index positives are not independent temporal ground truth. "
"`reconstruction_probe_metrics.csv` reports held-out masked reconstruction error in training-fold standardized feature units. "
"`frozen_emotion_probe_metrics.csv` is a small-sample downstream utility check.\n\n"
"No human event-time annotation is present, so the results cannot establish direct human alignment accuracy. "
"See `run_manifest.json` for parameters, seeds, device, and interpretation limits.\n",
encoding="utf-8",
)
print(
f"[done] elapsed_seconds={manifest['elapsed_seconds']:.1f} output={args.output_dir}",
flush=True,
)
return manifest
def build_parser() -> argparse.ArgumentParser:
project_dir = Path(__file__).resolve().parents[1]
parser = argparse.ArgumentParser(description="Compare Q1 M1-M4 alignment methods with grouped CV.")
parser.add_argument("--feature-dir", type=Path, default=project_dir / "outputs/q1_features/features")
parser.add_argument("--manifest", type=Path, default=project_dir / "outputs/audit/manifest.csv")
parser.add_argument("--output-dir", type=Path, default=project_dir / "outputs/method_comparison")
parser.add_argument("--device", default="auto", help="auto, cpu, or a torch device such as cuda:0")
parser.add_argument("--expected-samples", type=int, default=100)
parser.add_argument("--folds", type=int, default=5)
parser.add_argument("--seeds", type=int, nargs="+", default=[42, 3407, 2026])
parser.add_argument("--grid-size", type=int, default=50)
parser.add_argument("--hidden-size", type=int, default=128)
parser.add_argument("--heads", type=int, default=4)
parser.add_argument("--dropout", type=float, default=0.1)
parser.add_argument("--batch-size", type=int, default=8)
parser.add_argument("--probe-batch-size", type=int, default=16)
parser.add_argument("--epochs", type=int, default=50)
parser.add_argument("--patience", type=int, default=8)
parser.add_argument("--learning-rate", type=float, default=1e-4)
parser.add_argument("--mask-ratio", type=float, default=0.2)
parser.add_argument("--retrieval-probe-epochs", type=int, default=20)
parser.add_argument("--reconstruction-probe-epochs", type=int, default=25)
parser.add_argument("--mvr-epsilon", type=float, default=0.02)
parser.add_argument("--example-id", default="-tPCytz4rww/12")
return parser
def main() -> int:
args = build_parser().parse_args()
manifest = run(args)
print(json.dumps(manifest, ensure_ascii=False, indent=2))
return 0
if __name__ == "__main__":
raise SystemExit(main())
+773
View File
@@ -0,0 +1,773 @@
"""Evaluate same-time cross-modal correspondence against within-clip shifts.
The alignment model is frozen. A small train-only linear projection probe maps
the aligned raw features into a shared space using diagonal-versus-off-diagonal
InfoNCE within each training clip. All reported correspondence scores are then
computed on held-out video_id groups.
"""
from __future__ import annotations
import argparse
import csv
import json
import platform
import random
import time
from collections import defaultdict
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Mapping, Sequence
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import torch
import torch.nn.functional as F
from sklearn.metrics import roc_auc_score
from torch import Tensor, nn
from .compare_methods import _collect_representations
from .experiment_data import (
FeatureSample,
FeatureStats,
collate_feature_samples,
fit_feature_stats,
load_feature_samples,
)
from .models import SharedLatentTimeline, TextAnchoredCrossAttention
from .types import MODALITIES
GRID_SIZE = 50
HIDDEN_SIZE = 128
HEADS = 4
OUTPUT_SIZE = 64
PAIRINGS = (("text", "audio"), ("text", "vision"), ("audio", "vision"))
METHODS = ("M1", "M2", "M3_noSourceTime", "M3_sourceTime", "M4_noSourceTime", "M4_sourceTime")
CURVE_MAX_SHIFT = 10
NEGATIVE_RADIUS = 2
class CorrespondenceProjection(nn.Module):
"""Equal-output-size modality heads used only as a frozen-feature probe."""
def __init__(self, dimensions: Mapping[str, int], output_size: int = OUTPUT_SIZE) -> None:
super().__init__()
self.projections = nn.ModuleDict(
{name: nn.Linear(dimensions[name], output_size, bias=False) for name in MODALITIES}
)
def forward(self, values: Mapping[str, Tensor]) -> dict[str, Tensor]:
return {
name: F.normalize(self.projections[name](values[name]), dim=-1)
for name in MODALITIES
}
def _write_csv(path: Path, rows: Sequence[Mapping[str, Any]]) -> None:
if not rows:
return
path.parent.mkdir(parents=True, exist_ok=True)
fields = list(dict.fromkeys(key for row in rows for key in row))
with path.open("w", newline="", encoding="utf-8-sig") as handle:
writer = csv.DictWriter(handle, fieldnames=fields, extrasaction="ignore")
writer.writeheader()
writer.writerows(rows)
def _batches(samples: Sequence[FeatureSample], batch_size: int):
for start in range(0, len(samples), batch_size):
yield list(samples[start : start + batch_size])
def _make_model(method: str, dimensions: Mapping[str, int], source_time: bool) -> nn.Module:
if method == "M3":
return TextAnchoredCrossAttention(
dimensions,
grid_size=GRID_SIZE,
hidden_size=HIDDEN_SIZE,
heads=HEADS,
dropout=0.0,
source_time_encoding=source_time,
)
return SharedLatentTimeline(
dimensions,
grid_size=GRID_SIZE,
hidden_size=HIDDEN_SIZE,
heads=HEADS,
dropout=0.0,
absolute_position_encoding=True,
source_time_encoding=source_time,
)
def _checkpoint_path(root: Path, fold: int, variant: str) -> Path:
if fold == 1:
return root / variant / "checkpoint.pt"
return root / f"fold_{fold:02d}" / variant / "checkpoint.pt"
def _collect_learned_representations(
*,
method_name: str,
fold: int,
train_samples: Sequence[FeatureSample],
validation_samples: Sequence[FeatureSample],
stats: FeatureStats,
checkpoint_root: Path,
device: torch.device,
batch_size: int,
) -> dict[str, dict[str, np.ndarray]]:
method = method_name[:2]
source_time = method_name.endswith("_sourceTime")
checkpoint_path = _checkpoint_path(checkpoint_root, fold, method_name)
if not checkpoint_path.is_file():
raise FileNotFoundError(
f"missing {method_name} fold {fold} checkpoint: {checkpoint_path}; "
"run q1.alignment_heldout_debug for this fold first"
)
checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False)
if checkpoint.get("variant") != method_name:
raise ValueError(f"checkpoint variant does not match {method_name}: {checkpoint_path}")
if bool(checkpoint.get("source_time_encoding")) != source_time:
raise ValueError(f"checkpoint source-time setting does not match {method_name}: {checkpoint_path}")
expected_train = {sample.sample_id for sample in train_samples}
expected_validation = {sample.sample_id for sample in validation_samples}
if set(checkpoint.get("train_sample_ids", [])) != expected_train:
raise ValueError(f"checkpoint training IDs do not match split for {method_name}, fold {fold}")
if set(checkpoint.get("validation_sample_ids", [])) != expected_validation:
raise ValueError(f"checkpoint validation IDs do not match split for {method_name}, fold {fold}")
dimensions = {name: train_samples[0].features[name].shape[1] for name in MODALITIES}
model = _make_model(method, dimensions, source_time).to(device)
model.load_state_dict(checkpoint["model_state_dict"], strict=True)
model.eval()
result: dict[str, dict[str, np.ndarray]] = {}
with torch.no_grad():
for batch_samples in _batches([*train_samples, *validation_samples], batch_size):
sequences, durations, _ = collate_feature_samples(batch_samples, stats, device)
output = model(sequences, durations) if source_time else model(sequences)
for index, sample in enumerate(batch_samples):
sample_values: dict[str, np.ndarray] = {}
for name in MODALITIES:
length = len(sample.features[name])
weights = output.weights[name][index, :, :length]
source = sequences[name].features[index, :length]
pooled = weights.to(source.dtype) @ source
if pooled.shape[0] != GRID_SIZE:
raise ValueError(f"unexpected grid size for {sample.sample_id}/{name}")
sample_values[name] = pooled.detach().cpu().numpy().astype(np.float32, copy=False)
result[sample.sample_id] = sample_values
del model
if device.type == "cuda":
torch.cuda.empty_cache()
return result
def _stack_ids(
sample_ids: Sequence[str],
aligned_by_id: Mapping[str, Mapping[str, np.ndarray]],
device: torch.device,
) -> dict[str, Tensor]:
return {
name: torch.as_tensor(
np.stack([aligned_by_id[sample_id][name] for sample_id in sample_ids]),
dtype=torch.float32,
device=device,
)
for name in MODALITIES
}
def _within_clip_infonce(
projected: Mapping[str, Tensor], temperature: float
) -> Tensor:
"""Symmetric diagonal-vs-shifted loss; negatives stay inside each clip."""
batch_size, grid_size = projected["text"].shape[:2]
target = torch.arange(grid_size, device=projected["text"].device).repeat(batch_size)
losses = []
for left_index, left_name in enumerate(MODALITIES):
for right_name in MODALITIES[left_index + 1 :]:
scores = torch.einsum(
"bkd,bld->bkl", projected[left_name], projected[right_name]
) / temperature
forward = F.cross_entropy(scores.reshape(-1, grid_size), target)
reverse = F.cross_entropy(scores.transpose(1, 2).reshape(-1, grid_size), target)
losses.append((forward + reverse) / 2)
return torch.stack(losses).mean()
def _fit_probe(
train_ids: Sequence[str],
aligned_by_id: Mapping[str, Mapping[str, np.ndarray]],
*,
device: torch.device,
seed: int,
epochs: int,
batch_size: int,
learning_rate: float,
temperature: float,
) -> tuple[CorrespondenceProjection, list[dict[str, Any]]]:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
if device.type == "cuda":
torch.cuda.manual_seed_all(seed)
train = _stack_ids(train_ids, aligned_by_id, device)
dimensions = {name: int(train[name].shape[-1]) for name in MODALITIES}
model = CorrespondenceProjection(dimensions).to(device)
optimizer = torch.optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=1e-4)
rng = np.random.default_rng(seed)
history: list[dict[str, Any]] = []
model.train()
for epoch in range(1, epochs + 1):
order = rng.permutation(len(train_ids))
loss_sum = 0.0
batches = 0
for start in range(0, len(order), batch_size):
indexes = torch.as_tensor(order[start : start + batch_size], device=device)
batch = {name: value.index_select(0, indexes) for name, value in train.items()}
loss = _within_clip_infonce(model(batch), temperature)
if not torch.isfinite(loss):
raise FloatingPointError(f"non-finite correspondence probe loss at epoch {epoch}")
optimizer.zero_grad(set_to_none=True)
loss.backward()
nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
loss_sum += float(loss.detach().item())
batches += 1
history.append({"epoch": epoch, "train_loss": loss_sum / max(batches, 1)})
return model, history
def _shifted_scores(scores: np.ndarray, delta: int) -> np.ndarray:
grid_size = scores.shape[0]
if delta >= 0:
indexes = np.arange(0, grid_size - delta)
return scores[indexes, indexes + delta]
indexes = np.arange(-delta, grid_size)
return scores[indexes, indexes + delta]
def _sample_metrics(
*,
method: str,
fold: int,
sample: FeatureSample,
projected: Mapping[str, Tensor],
curve_rows: list[dict[str, Any]],
) -> list[dict[str, Any]]:
rows = []
for left_name, right_name in PAIRINGS:
left = projected[left_name]
right = projected[right_name]
score = (left @ right.T).detach().cpu().numpy()
grid_size = score.shape[0]
if score.shape != (GRID_SIZE, GRID_SIZE):
raise ValueError(f"expected {GRID_SIZE}x{GRID_SIZE} score matrix")
diagonal = np.diag(score)
negative_mask = np.abs(np.arange(grid_size)[:, None] - np.arange(grid_size)[None, :]) > NEGATIVE_RADIUS
auc = float(
roc_auc_score(
np.concatenate((np.ones(grid_size), np.zeros(int(negative_mask.sum())))),
np.concatenate((diagonal, score[negative_mask])),
)
)
far_shifts = [
_shifted_scores(score, delta).mean()
for delta in range(-CURVE_MAX_SHIFT, CURVE_MAX_SHIFT + 1)
if abs(delta) > NEGATIVE_RADIUS
]
row: dict[str, Any] = {
"method": method,
"fold": fold,
"sample_id": sample.sample_id,
"video_id": sample.group_id,
"pair": f"{left_name}_{right_name}",
"same_time_similarity": float(diagonal.mean()),
"shifted_far_similarity": float(np.mean(far_shifts)),
"same_minus_shifted_margin": float(diagonal.mean() - np.mean(far_shifts)),
"matched_vs_shifted_auc": auc,
}
for direction, directed_score in (
(f"{left_name}_to_{right_name}", score),
(f"{right_name}_to_{left_name}", score.T),
):
prediction = directed_score.argmax(axis=1)
error = np.abs(prediction - np.arange(grid_size))
row[f"exact_r1_{direction}"] = float(np.mean(error == 0))
row[f"within_pm1_r1_{direction}"] = float(np.mean(error <= 1))
row[f"mase_slots_{direction}"] = float(np.mean(error))
for delta in range(-CURVE_MAX_SHIFT, CURVE_MAX_SHIFT + 1):
curve_rows.append(
{
"method": method,
"fold": fold,
"sample_id": sample.sample_id,
"video_id": sample.group_id,
"pair": f"{left_name}_{right_name}",
"delta": delta,
"similarity": float(_shifted_scores(score, delta).mean()),
}
)
rows.append(row)
return rows
def _cluster_bootstrap(
rows: Sequence[Mapping[str, Any]],
metric: str,
*,
seed: int,
repetitions: int = 2000,
) -> tuple[float, float, float, int]:
grouped: dict[str, list[float]] = defaultdict(list)
for row in rows:
value = float(row[metric])
if np.isfinite(value):
grouped[str(row["video_id"])].append(value)
groups = np.asarray([np.mean(values) for values in grouped.values()], dtype=np.float64)
if groups.size == 0:
return float("nan"), float("nan"), float("nan"), 0
mean = float(groups.mean())
if groups.size == 1:
return mean, mean, mean, 1
rng = np.random.default_rng(seed)
indexes = rng.integers(0, groups.size, size=(repetitions, groups.size))
boot = groups[indexes].mean(axis=1)
low, high = np.quantile(boot, [0.025, 0.975])
return mean, float(low), float(high), int(groups.size)
def _summaries(
metric_rows: Sequence[Mapping[str, Any]],
curve_rows: Sequence[Mapping[str, Any]],
*,
seed: int,
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
metrics = (
"same_time_similarity",
"shifted_far_similarity",
"same_minus_shifted_margin",
"matched_vs_shifted_auc",
"exact_r1_text_to_audio",
"within_pm1_r1_text_to_audio",
"mase_slots_text_to_audio",
"exact_r1_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 metric_rows:
grouped[(str(row["method"]), str(row["pair"]))].append(row)
summary_rows = []
for (method, pair), rows in sorted(grouped.items()):
result: dict[str, Any] = {
"method": method,
"pair": pair,
"clip_count": len(rows),
"video_id_count": len({str(row["video_id"]) for row in rows}),
}
for metric_index, metric in enumerate(metrics):
if metric not in rows[0]:
continue
mean, low, high, _ = _cluster_bootstrap(
rows,
metric,
seed=seed + metric_index + sum(ord(character) for character in method + pair),
)
result[f"{metric}_mean"] = mean
result[f"{metric}_ci95_low"] = low
result[f"{metric}_ci95_high"] = high
curve_for_group = [
row for row in curve_rows if row["method"] == method and row["pair"] == pair
]
curve_by_delta: dict[int, list[dict[str, Any]]] = defaultdict(list)
for row in curve_for_group:
curve_by_delta[int(row["delta"])].append(dict(row))
mean_curve = {
delta: _cluster_bootstrap(
values,
"similarity",
seed=seed + delta + sum(ord(character) for character in method + pair),
)[0]
for delta, values in curve_by_delta.items()
}
if mean_curve:
result["peak_delta"] = max(mean_curve, key=mean_curve.get)
result["peak_similarity"] = mean_curve[result["peak_delta"]]
summary_rows.append(result)
curve_summary = []
curve_groups: dict[tuple[str, str, int], list[Mapping[str, Any]]] = defaultdict(list)
for row in curve_rows:
curve_groups[(str(row["method"]), str(row["pair"]), int(row["delta"]))].append(row)
for (method, pair, delta), rows in sorted(curve_groups.items()):
mean, low, high, group_count = _cluster_bootstrap(
rows,
"similarity",
seed=seed + delta + sum(ord(character) for character in method + pair),
)
curve_summary.append(
{
"method": method,
"pair": pair,
"delta": delta,
"mean_similarity": mean,
"ci95_low": low,
"ci95_high": high,
"video_id_count": group_count,
}
)
return summary_rows, curve_summary
def _summarize_time_localization(
checkpoint_root: Path, fold_count: int, *, seed: int
) -> list[dict[str, Any]]:
grouped: dict[tuple[str, str], list[dict[str, Any]]] = defaultdict(list)
for fold in range(1, fold_count + 1):
metrics_path = (
checkpoint_root / "per_sample_metrics.csv"
if fold == 1
else checkpoint_root / f"fold_{fold:02d}" / "per_sample_metrics.csv"
)
if not metrics_path.is_file():
raise FileNotFoundError(f"missing D5 per-sample metrics for fold {fold}: {metrics_path}")
with metrics_path.open("r", newline="", encoding="utf-8-sig") as handle:
for source in csv.DictReader(handle):
sample_id = source["sample_id"]
row: dict[str, Any] = {
"video_id": sample_id.split("/", 1)[0],
"sample_id": sample_id,
"variant": source["variant"],
"modality": source["modality"],
}
for metric in (
"normalized_entropy",
"trajectory_span",
"gaussian_target_kl",
"mean_absolute_time_center_error",
):
raw = source.get(metric, "")
if raw not in (None, ""):
value = float(raw)
if np.isfinite(value):
row[metric] = value
grouped[(row["variant"], row["modality"])].append(row)
metric_names = (
"normalized_entropy",
"trajectory_span",
"gaussian_target_kl",
"mean_absolute_time_center_error",
)
results = []
for (variant, modality), rows in sorted(grouped.items()):
result: dict[str, Any] = {
"variant": variant,
"modality": modality,
"clip_count": len(rows),
"video_id_count": len({row["video_id"] for row in rows}),
}
for index, metric in enumerate(metric_names):
values = [row for row in rows if metric in row]
if not values:
continue
mean, low, high, _ = _cluster_bootstrap(
values,
metric,
seed=seed + index + sum(ord(character) for character in variant + modality),
)
result[f"{metric}_video_macro_mean"] = mean
result[f"{metric}_ci95_low"] = low
result[f"{metric}_ci95_high"] = high
results.append(result)
return results
def _plot_curve(curve_rows: Sequence[Mapping[str, Any]], output_path: Path) -> None:
pairs = tuple(f"{left}_{right}" for left, right in PAIRINGS)
colors = {
"M1": "#555555",
"M2": "#9a9a9a",
"M3_noSourceTime": "#2878b5",
"M3_sourceTime": "#f08a24",
"M4_noSourceTime": "#55a868",
"M4_sourceTime": "#c44e52",
}
fig, axes = plt.subplots(1, 3, figsize=(17, 5), sharey=True)
for axis, pair in zip(axes, pairs, strict=True):
for method in METHODS:
rows = [row for row in curve_rows if row["pair"] == pair and row["method"] == method]
if not rows:
continue
deltas = sorted({int(row["delta"]) for row in rows})
xs, means, lows, highs = [], [], [], []
for delta in deltas:
subset = [row for row in rows if int(row["delta"]) == delta]
mean, low, high, _ = _cluster_bootstrap(
subset,
"similarity",
seed=9821 + delta + sum(ord(character) for character in method + pair),
repetitions=1000,
)
xs.append(delta)
means.append(mean)
lows.append(low)
highs.append(high)
axis.plot(xs, means, marker="o", markersize=3, linewidth=1.6, label=method, color=colors[method])
axis.fill_between(xs, lows, highs, color=colors[method], alpha=0.10, linewidth=0)
axis.axvline(0, color="black", linestyle="--", linewidth=0.9, alpha=0.6)
axis.set_title(pair.replace("_", "–"))
axis.set_xlabel("Temporal shift Δ (slots)")
axis.grid(alpha=0.2)
axes[0].set_ylabel("Cross-modal cosine similarity")
handles, labels = axes[-1].get_legend_handles_labels()
fig.legend(handles, labels, loc="lower center", ncol=3, frameon=False, bbox_to_anchor=(0.5, -0.02))
fig.suptitle("Same-slot vs shifted cross-modal similarity (95% video_id bootstrap CI)", y=1.02)
fig.tight_layout(rect=(0, 0.08, 1, 0.98))
output_path.parent.mkdir(parents=True, exist_ok=True)
fig.savefig(output_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")
loaded_samples = load_feature_samples(args.feature_dir, args.manifest)
by_id = {sample.sample_id: sample for sample in loaded_samples}
split_rows = json.loads(args.splits.read_text(encoding="utf-8"))
output_dir = args.output_dir
output_dir.mkdir(parents=True, exist_ok=True)
all_metric_rows: list[dict[str, Any]] = []
all_curve_rows: list[dict[str, Any]] = []
all_history_rows: list[dict[str, Any]] = []
saved_probes: dict[str, Any] = {}
fold_manifests = []
for split in split_rows:
fold = int(split["fold"])
train_samples = [by_id[sample_id] for sample_id in split["train_sample_ids"]]
validation_samples = [by_id[sample_id] for sample_id in split["validation_sample_ids"]]
train_groups = {sample.group_id for sample in train_samples}
validation_groups = {sample.group_id for sample in validation_samples}
if train_groups & validation_groups:
raise ValueError(f"video_id leakage in fold {fold}")
stats = fit_feature_stats(train_samples)
print(
f"[correspondence fold {fold}] train={len(train_samples)} heldout={len(validation_samples)} "
f"video_ids={len(train_groups)}/{len(validation_groups)}",
flush=True,
)
fold_manifests.append(
{
"fold": fold,
"train_count": len(train_samples),
"heldout_count": len(validation_samples),
"train_video_ids": sorted(train_groups),
"heldout_video_ids": sorted(validation_groups),
"overlap": sorted(train_groups & validation_groups),
}
)
for method_index, method_name in enumerate(METHODS):
if method_name in {"M1", "M2"}:
aligned_train, _ = _collect_representations(
method_name,
train_samples,
stats,
device=device,
grid_size=GRID_SIZE,
batch_size=args.batch_size,
)
aligned_validation, _ = _collect_representations(
method_name,
validation_samples,
stats,
device=device,
grid_size=GRID_SIZE,
batch_size=args.batch_size,
)
aligned = {**aligned_train, **aligned_validation}
else:
aligned = _collect_learned_representations(
method_name=method_name,
fold=fold,
train_samples=train_samples,
validation_samples=validation_samples,
stats=stats,
checkpoint_root=args.checkpoint_root,
device=device,
batch_size=args.batch_size,
)
probe_seed = args.seed + fold * 101 + method_index
probe, history = _fit_probe(
[sample.sample_id for sample in train_samples],
aligned,
device=device,
seed=probe_seed,
epochs=args.epochs,
batch_size=args.batch_size,
learning_rate=args.learning_rate,
temperature=args.temperature,
)
for row in history:
all_history_rows.append(
{"fold": fold, "method": method_name, "seed": probe_seed, **row}
)
probe.eval()
with torch.no_grad():
validation = _stack_ids(
[sample.sample_id for sample in validation_samples], aligned, device
)
projected = probe(validation)
for sample_index, sample in enumerate(validation_samples):
one = {name: projected[name][sample_index] for name in MODALITIES}
all_metric_rows.extend(
_sample_metrics(
method=method_name,
fold=fold,
sample=sample,
projected=one,
curve_rows=all_curve_rows,
)
)
saved_probes[f"fold_{fold:02d}/{method_name}"] = {
"input_dimensions": {
name: int(aligned[train_samples[0].sample_id][name].shape[-1])
for name in MODALITIES
},
"state_dict": {key: value.detach().cpu() for key, value in probe.state_dict().items()},
"seed": probe_seed,
}
print(
f"[correspondence fold {fold} {method_name}] "
f"probe_loss={history[-1]['train_loss']:.4f}",
flush=True,
)
del probe
if device.type == "cuda":
torch.cuda.empty_cache()
summary_rows, curve_summary_rows = _summaries(
all_metric_rows, all_curve_rows, seed=args.seed
)
time_localization_rows = _summarize_time_localization(
args.checkpoint_root, len(split_rows), seed=args.seed
)
_write_csv(output_dir / "correspondence_clip_metrics.csv", all_metric_rows)
_write_csv(output_dir / "shifted_similarity_by_clip.csv", all_curve_rows)
_write_csv(output_dir / "metric_summary.csv", summary_rows)
_write_csv(output_dir / "shift_curve_summary.csv", curve_summary_rows)
_write_csv(output_dir / "time_localization_summary.csv", time_localization_rows)
_write_csv(output_dir / "probe_training_history.csv", all_history_rows)
_plot_curve(all_curve_rows, output_dir / "shifted_similarity_curve.png")
torch.save(saved_probes, output_dir / "probe_checkpoints.pt")
manifest = {
"created_utc": datetime.now(timezone.utc).isoformat(),
"experiment": "Q1-Correspondence Evaluation",
"sample_count": len(loaded_samples),
"fold_count": len(split_rows),
"folds": fold_manifests,
"methods": list(METHODS),
"grid_size": GRID_SIZE,
"probe": {
"type": "one bias-free linear projection per modality, L2 normalized",
"output_dimension": OUTPUT_SIZE,
"training_objective": "symmetric within-clip diagonal-vs-off-diagonal InfoNCE across all three modality pairs",
"negative_source": "same training clip, different grid slots",
"epochs": args.epochs,
"batch_size": args.batch_size,
"learning_rate": args.learning_rate,
"temperature": args.temperature,
"emotion_labels_used": False,
"heldout_groups_used_for_probe_training_or_selection": False,
},
"evaluation": {
"shift_curve": f"mean cosine sim(z_i^m, z_(i+delta)^n), delta=-{CURVE_MAX_SHIFT}..{CURVE_MAX_SHIFT}",
"temporal_retrieval": "candidate slots restricted to the same held-out clip; report exact R@1, within +/-1 R@1, and MASE slots in both directions",
"matched_vs_shifted_auc": f"per-clip ROC AUC; positives are same-slot pairs, negatives are same-clip pairs with |i-j|>{NEGATIVE_RADIUS}",
"confidence_intervals": "95% percentile bootstrap resampling video_id groups, not clips or slots; 2,000 repetitions for CSV summaries and 1,000 for plot ribbons",
"random_retrieval_reference": {
"exact_r1": 1.0 / GRID_SIZE,
"within_pm1_r1": (3.0 * GRID_SIZE - 2.0) / GRID_SIZE**2,
"expected_mase_slots": (GRID_SIZE**2 - 1) / (3.0 * GRID_SIZE),
},
},
"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": [
"The projection is a supervised diagnostic probe for same-slot matchability, not an alignment model or independent ground truth.",
"A strong result means train-video same-time matchability transfers to held-out video_id groups; it does not establish semantic equivalence between modalities.",
"M3/M4 D5 training used timestamp-derived Gaussian targets; this evaluation tests whether their frozen pooled content representations carry transferable same-time correspondence beyond that temporal target.",
"The effective sample size for confidence intervals is the number of video_id groups, not the number of slots.",
],
}
(output_dir / "run_manifest.json").write_text(
json.dumps(manifest, ensure_ascii=False, indent=2, allow_nan=False), encoding="utf-8"
)
print(
f"[correspondence done] clips={len(all_metric_rows)} "
f"elapsed={manifest['elapsed_seconds']:.1f}s output={output_dir}",
flush=True,
)
return manifest
def build_parser() -> argparse.ArgumentParser:
project = Path(__file__).resolve().parents[1]
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto")
parser.add_argument("--batch-size", type=int, default=8)
parser.add_argument("--epochs", type=int, default=40)
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--learning-rate", type=float, default=1e-3)
parser.add_argument("--temperature", type=float, default=0.1)
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/correspondence_eval"
)
return parser
def main() -> None:
args = build_parser().parse_args()
run(args)
if __name__ == "__main__":
main()
+159
View File
@@ -0,0 +1,159 @@
from __future__ import annotations
import csv
from dataclasses import dataclass
from pathlib import Path
from typing import Mapping, Sequence
import numpy as np
import torch
from torch import Tensor
from .types import MODALITIES, SequenceBatch
@dataclass(frozen=True)
class FeatureSample:
sample_id: str
group_id: str
duration_s: float
sentiment: float
polarity: int
word_intervals: np.ndarray
features: Mapping[str, np.ndarray]
times: Mapping[str, np.ndarray]
valid: Mapping[str, np.ndarray]
@dataclass(frozen=True)
class FeatureStats:
mean: Mapping[str, np.ndarray]
scale: Mapping[str, np.ndarray]
def load_feature_samples(feature_dir: Path, manifest_path: Path) -> list[FeatureSample]:
"""Read the extracted NPZ files and their source-time/label manifest."""
with manifest_path.open("r", encoding="utf-8-sig", newline="") as file:
rows = list(csv.DictReader(file))
if not rows:
raise ValueError(f"sample manifest is empty: {manifest_path}")
samples: list[FeatureSample] = []
seen: set[str] = set()
for row in rows:
video_id = row["video_id"]
clip_id = row["clip_id"]
sample_id = f"{video_id}/{clip_id}"
if sample_id in seen:
raise ValueError(f"duplicate sample in manifest: {sample_id}")
seen.add(sample_id)
path = feature_dir / f"{video_id}__{clip_id}.npz"
if not path.is_file():
raise FileNotFoundError(f"feature file missing for {sample_id}: {path}")
with np.load(path, allow_pickle=False) as archive:
features = {
"text": np.asarray(archive["text_features"], dtype=np.float32),
"audio": np.asarray(archive["audio_features"], dtype=np.float32),
"vision": np.asarray(archive["vision_features"], dtype=np.float32),
}
times = {
"text": np.asarray(archive["word_intervals_s"], dtype=np.float32).mean(axis=1),
"audio": np.asarray(archive["audio_times_s"], dtype=np.float32),
"vision": np.asarray(archive["vision_times_s"], dtype=np.float32),
}
valid = {
"text": np.ones(features["text"].shape[0], dtype=np.bool_),
"audio": np.asarray(archive["audio_valid"], dtype=np.bool_),
"vision": np.asarray(archive["vision_valid"], dtype=np.bool_),
}
word_intervals = np.asarray(archive["word_intervals_s"], dtype=np.float32)
for name in MODALITIES:
if features[name].ndim != 2 or times[name].shape != (features[name].shape[0],):
raise ValueError(f"invalid {name} feature/timestamp shape in {sample_id}")
if valid[name].shape != times[name].shape or not valid[name].any():
raise ValueError(f"{sample_id} has no valid {name} sequence positions")
if not np.isfinite(features[name][valid[name]]).all():
raise ValueError(f"non-finite valid {name} values in {sample_id}")
if not np.isfinite(times[name][valid[name]]).all():
raise ValueError(f"non-finite valid {name} timestamps in {sample_id}")
annotation = row["annotation"].strip().lower()
polarity_by_name = {"negative": 0, "neutral": 1, "positive": 2}
if annotation not in polarity_by_name:
raise ValueError(f"unknown polarity label {annotation!r} in {sample_id}")
samples.append(
FeatureSample(
sample_id=sample_id,
group_id=row.get("group_id") or video_id,
duration_s=float(row["duration_s"]),
sentiment=float(row["label"]),
polarity=polarity_by_name[annotation],
word_intervals=word_intervals,
features=features,
times=times,
valid=valid,
)
)
return samples
def fit_feature_stats(samples: Sequence[FeatureSample]) -> FeatureStats:
"""Fit per-modality z-score parameters using only the training fold."""
if not samples:
raise ValueError("cannot fit feature statistics on an empty sample list")
means: dict[str, np.ndarray] = {}
scales: dict[str, np.ndarray] = {}
for name in MODALITIES:
values = np.concatenate(
[sample.features[name][sample.valid[name]] for sample in samples], axis=0
).astype(np.float64, copy=False)
mean = values.mean(axis=0)
scale = values.std(axis=0)
scale[scale < 1e-6] = 1.0
means[name] = mean.astype(np.float32)
scales[name] = scale.astype(np.float32)
return FeatureStats(mean=means, scale=scales)
def standardized_features(sample: FeatureSample, stats: FeatureStats) -> dict[str, np.ndarray]:
return {
name: ((sample.features[name] - stats.mean[name]) / stats.scale[name]).astype(
np.float32, copy=False
)
for name in MODALITIES
}
def collate_feature_samples(
samples: Sequence[FeatureSample],
stats: FeatureStats,
device: torch.device,
) -> tuple[dict[str, SequenceBatch], Tensor, list[Tensor]]:
"""Pad one variable-length batch in memory; padding is masked and never saved."""
if not samples:
raise ValueError("cannot collate an empty sample list")
batch_size = len(samples)
sequences: dict[str, SequenceBatch] = {}
for name in MODALITIES:
lengths = [sample.features[name].shape[0] for sample in samples]
max_length = max(lengths)
dimension = samples[0].features[name].shape[1]
feature_batch = torch.zeros(batch_size, max_length, dimension, dtype=torch.float32)
time_batch = torch.zeros(batch_size, max_length, dtype=torch.float32)
valid_batch = torch.zeros(batch_size, max_length, dtype=torch.bool)
for index, sample in enumerate(samples):
features = standardized_features(sample, stats)[name]
length = len(features)
feature_batch[index, :length] = torch.from_numpy(features)
time_batch[index, :length] = torch.from_numpy(sample.times[name])
valid_batch[index, :length] = torch.from_numpy(sample.valid[name])
sequences[name] = SequenceBatch(
features=feature_batch.to(device),
times=time_batch.to(device),
valid=valid_batch.to(device),
)
durations = torch.tensor([sample.duration_s for sample in samples], dtype=torch.float32, device=device)
word_intervals = [torch.from_numpy(sample.word_intervals).to(device) for sample in samples]
return sequences, durations, word_intervals
+462
View File
@@ -0,0 +1,462 @@
from __future__ import annotations
from collections.abc import Mapping, Sequence
import numpy as np
import torch
import torch.nn.functional as F
from sklearn.dummy import DummyClassifier
from sklearn.linear_model import LogisticRegression, Ridge
from sklearn.metrics import accuracy_score, f1_score, mean_absolute_error
from sklearn.preprocessing import StandardScaler
from torch import Tensor, nn
from .alignment import make_block_mask
from .losses import cross_modal_contrastive_loss
from .metrics import retrieval_metrics
from .types import MODALITIES
class RetrievalProjection(nn.Module):
"""Equal-capacity linear heads for the cross-modal retrieval probe."""
def __init__(self, dimensions: Mapping[str, int], output_size: int = 128) -> None:
super().__init__()
self.projections = nn.ModuleDict(
{name: nn.Linear(dimensions[name], output_size, bias=False) for name in MODALITIES}
)
def forward(self, aligned: Mapping[str, Tensor]) -> dict[str, Tensor]:
return {
name: F.normalize(self.projections[name](aligned[name]), dim=-1)
for name in MODALITIES
}
class MaskedCrossModalDecoder(nn.Module):
"""Reconstruct one missing modality from the other aligned streams."""
def __init__(self, dimensions: Mapping[str, int], hidden_size: int = 256) -> None:
super().__init__()
self.dimensions = dict(dimensions)
input_size = sum(dimensions.values()) + len(MODALITIES)
self.decoders = nn.ModuleDict(
{
target: nn.Sequential(
nn.Linear(input_size, hidden_size),
nn.GELU(),
nn.Dropout(0.1),
nn.Linear(hidden_size, dimensions[target]),
)
for target in MODALITIES
}
)
def forward(self, aligned: Mapping[str, Tensor], target: str, mask: Tensor) -> Tensor:
values = []
availability = []
for name in MODALITIES:
present = torch.ones_like(mask, dtype=aligned[name].dtype)
source = aligned[name]
if name == target:
present = (~mask).to(source.dtype)
source = source.masked_fill(mask.unsqueeze(-1), 0.0)
values.append(source)
availability.append(present.unsqueeze(-1))
inputs = torch.cat((*values, *availability), dim=-1)
return self.decoders[target](inputs)
def _stack_aligned(
sample_ids: Sequence[str], aligned_by_id: Mapping[str, Mapping[str, np.ndarray]], device: torch.device
) -> dict[str, Tensor]:
return {
name: torch.as_tensor(
np.stack([aligned_by_id[sample_id][name] for sample_id in sample_ids]),
dtype=torch.float32,
device=device,
)
for name in MODALITIES
}
def run_retrieval_probe(
train_ids: Sequence[str],
val_ids: Sequence[str],
aligned_by_id: Mapping[str, Mapping[str, np.ndarray]],
*,
device: torch.device,
seed: int,
epochs: int = 20,
batch_size: int = 16,
learning_rate: float = 1e-3,
temperature: float = 0.1,
) -> list[dict[str, float | str]]:
"""Fit modality projections on training clips, then score held-out retrieval.
The positive is a matching common-grid index within a clip. This measures
representation consistency; it is not independent temporal ground truth.
"""
if not train_ids or not val_ids:
raise ValueError("retrieval probe needs non-empty train and validation sets")
device_gen = torch.Generator(device=device)
device_gen.manual_seed(seed)
torch.manual_seed(seed)
train = _stack_aligned(train_ids, aligned_by_id, device)
validation = _stack_aligned(val_ids, aligned_by_id, device)
dimensions = {name: int(train[name].shape[-1]) for name in MODALITIES}
model = RetrievalProjection(dimensions).to(device)
optimizer = torch.optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=1e-4)
rng = np.random.default_rng(seed)
model.train()
for _ in range(epochs):
order = rng.permutation(len(train_ids))
for start in range(0, len(order), batch_size):
indices = torch.as_tensor(order[start : start + batch_size], device=device)
batch = {name: value.index_select(0, indices) for name, value in train.items()}
projected = model(batch)
loss = cross_modal_contrastive_loss(projected, temperature=temperature)
optimizer.zero_grad(set_to_none=True)
loss.backward()
nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
model.eval()
with torch.no_grad():
projected = model(validation)
directions = (("text", "audio"), ("audio", "text"), ("text", "vision"),
("vision", "text"), ("audio", "vision"), ("vision", "audio"))
rows: list[dict[str, float | str]] = []
for query_name, target_name in directions:
values = retrieval_metrics(projected[query_name], projected[target_name])
rows.append({"direction": f"{query_name}_to_{target_name}", **values})
return rows
def run_within_clip_temporal_retrieval_probe(
train_ids: Sequence[str],
val_ids: Sequence[str],
aligned_by_id: Mapping[str, Mapping[str, np.ndarray]],
*,
device: torch.device,
seed: int,
epochs: int = 20,
batch_size: int = 16,
learning_rate: float = 1e-3,
temperature: float = 0.1,
tolerance: int = 1,
top_k: int = 3,
) -> list[dict[str, float | str]]:
"""Fit train-only cross-modal projections, then retrieve slots within each clip.
Unlike global grid retrieval, each query's candidates are restricted to
the target modality from that same held-out clip. A result is correct if
its slot is within ``tolerance`` of the query slot.
"""
if not train_ids or not val_ids:
raise ValueError("temporal retrieval needs non-empty train and validation sets")
if tolerance < 0 or top_k < 1:
raise ValueError("tolerance must be non-negative and top_k positive")
torch.manual_seed(seed)
train = _stack_aligned(train_ids, aligned_by_id, device)
validation = _stack_aligned(val_ids, aligned_by_id, device)
dimensions = {name: int(train[name].shape[-1]) for name in MODALITIES}
model = RetrievalProjection(dimensions).to(device)
optimizer = torch.optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=1e-4)
rng = np.random.default_rng(seed)
model.train()
for _ in range(epochs):
order = rng.permutation(len(train_ids))
for start in range(0, len(order), batch_size):
indices = torch.as_tensor(order[start : start + batch_size], device=device)
batch = {name: value.index_select(0, indices) for name, value in train.items()}
projected = model(batch)
loss = cross_modal_contrastive_loss(projected, temperature=temperature)
optimizer.zero_grad(set_to_none=True)
loss.backward()
nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
model.eval()
with torch.no_grad():
projected = model(validation)
directions = (
("text", "audio"),
("audio", "text"),
("text", "vision"),
("vision", "text"),
("audio", "vision"),
("vision", "audio"),
)
rows: list[dict[str, float | str]] = []
for query_name, target_name in directions:
distances_top1 = []
hits_top1 = []
hits_topk = []
for clip_index in range(len(val_ids)):
query = projected[query_name][clip_index]
target = projected[target_name][clip_index]
scores = query @ target.T
count = scores.shape[0]
k = min(top_k, count)
candidates = scores.topk(k=k, dim=-1).indices
slots = torch.arange(count, device=device)[:, None]
distances = (candidates - slots).abs()
distances_top1.append(distances[:, 0].float())
hits_top1.append((distances[:, 0] <= tolerance).float())
hits_topk.append((distances <= tolerance).any(dim=-1).float())
top1_distance = torch.cat(distances_top1)
rows.append(
{
"direction": f"{query_name}_to_{target_name}",
"r_at_1": float(torch.cat(hits_top1).mean().item()),
"r_at_3": float(torch.cat(hits_topk).mean().item()),
"mase_slots": float(top1_distance.mean().item()),
"exact_r_at_1": float((top1_distance == 0).float().mean().item()),
"tolerance_slots": float(tolerance),
"candidate_slots_per_clip": float(projected[query_name].shape[1]),
"queries": float(len(val_ids) * projected[query_name].shape[1]),
}
)
return rows
def _fixed_block_mask(
count: int, grid_size: int, ratio: float, device: torch.device, salt: int
) -> Tensor:
block = min(max(1, round(grid_size * ratio)), grid_size - 1)
starts = torch.tensor(
[(index * 17 + salt * 13) % (grid_size - block + 1) for index in range(count)],
dtype=torch.long,
device=device,
)
offsets = torch.arange(block, device=device)
mask = torch.zeros(count, grid_size, dtype=torch.bool, device=device)
mask[torch.arange(count, device=device)[:, None], starts[:, None] + offsets] = True
return mask
def run_reconstruction_probe(
train_ids: Sequence[str],
val_ids: Sequence[str],
aligned_by_id: Mapping[str, Mapping[str, np.ndarray]],
*,
device: torch.device,
seed: int,
ratio: float = 0.2,
epochs: int = 25,
batch_size: int = 16,
learning_rate: float = 1e-3,
) -> list[dict[str, float | str]]:
"""Train the same decoder family on frozen alignments and score held-out clips."""
if not train_ids or not val_ids:
raise ValueError("reconstruction probe needs non-empty train and validation sets")
torch.manual_seed(seed)
generator = torch.Generator(device=device)
generator.manual_seed(seed)
train = _stack_aligned(train_ids, aligned_by_id, device)
validation = _stack_aligned(val_ids, aligned_by_id, device)
dimensions = {name: int(train[name].shape[-1]) for name in MODALITIES}
model = MaskedCrossModalDecoder(dimensions).to(device)
optimizer = torch.optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=1e-4)
rng = np.random.default_rng(seed)
model.train()
for _ in range(epochs):
order = rng.permutation(len(train_ids))
for start in range(0, len(order), batch_size):
indices = torch.as_tensor(order[start : start + batch_size], device=device)
batch = {name: value.index_select(0, indices) for name, value in train.items()}
target_losses = []
for target in MODALITIES:
mask = make_block_mask(
len(indices),
batch[target].shape[1],
ratio,
device,
generator=generator,
)
prediction = model(batch, target, mask)
target_losses.append(F.smooth_l1_loss(prediction[mask], batch[target][mask]))
loss = torch.stack(target_losses).mean()
optimizer.zero_grad(set_to_none=True)
loss.backward()
nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
model.eval()
rows: list[dict[str, float | str]] = []
with torch.no_grad():
for target_index, target in enumerate(MODALITIES):
mask = _fixed_block_mask(
len(val_ids), validation[target].shape[1], ratio, device, target_index
)
prediction = model(validation, target, mask)
residual = (prediction[mask] - validation[target][mask]).abs()
smooth = F.smooth_l1_loss(prediction[mask], validation[target][mask])
rows.append(
{
"target_modality": target,
"mask_ratio": ratio,
"mae_standardized": float(residual.mean().item()),
"smooth_l1_standardized": float(smooth.item()),
"masked_values": int(residual.numel()),
}
)
return rows
def run_shuffled_alignment_reconstruction_probe(
train_ids: Sequence[str],
val_ids: Sequence[str],
aligned_by_id: Mapping[str, Mapping[str, np.ndarray]],
*,
device: torch.device,
seed: int,
ratio: float = 0.2,
epochs: int = 25,
batch_size: int = 16,
learning_rate: float = 1e-3,
shuffle_repeats: int = 5,
) -> list[dict[str, float | str]]:
"""Compare normal reconstruction with cross-modal slots shuffled at eval.
The decoder is trained once on aligned training-fold representations.
For the control, the two non-target modalities share a random slot
permutation within each validation clip; the target stream and target
values remain in their original order. Thus the metric isolates how much
correctly matched cross-modal slots help the frozen decoder.
"""
if not train_ids or not val_ids:
raise ValueError("reconstruction needs non-empty train and validation sets")
if shuffle_repeats < 1:
raise ValueError("shuffle_repeats must be positive")
torch.manual_seed(seed)
generator = torch.Generator(device=device)
generator.manual_seed(seed)
train = _stack_aligned(train_ids, aligned_by_id, device)
validation = _stack_aligned(val_ids, aligned_by_id, device)
dimensions = {name: int(train[name].shape[-1]) for name in MODALITIES}
model = MaskedCrossModalDecoder(dimensions).to(device)
optimizer = torch.optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=1e-4)
rng = np.random.default_rng(seed)
model.train()
for _ in range(epochs):
order = rng.permutation(len(train_ids))
for start in range(0, len(order), batch_size):
indices = torch.as_tensor(order[start : start + batch_size], device=device)
batch = {name: value.index_select(0, indices) for name, value in train.items()}
target_losses = []
for target in MODALITIES:
mask = make_block_mask(
len(indices), batch[target].shape[1], ratio, device, generator=generator
)
prediction = model(batch, target, mask)
target_losses.append(F.smooth_l1_loss(prediction[mask], batch[target][mask]))
loss = torch.stack(target_losses).mean()
optimizer.zero_grad(set_to_none=True)
loss.backward()
nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
model.eval()
rows: list[dict[str, float | str]] = []
with torch.no_grad():
for target_index, target in enumerate(MODALITIES):
mask = _fixed_block_mask(
len(val_ids), validation[target].shape[1], ratio, device, target_index
)
aligned_prediction = model(validation, target, mask)
aligned_error = (aligned_prediction[mask] - validation[target][mask]).abs().mean()
shuffle_errors = []
permutation_rng = np.random.default_rng(seed + 1709 + target_index)
slot_count = validation[target].shape[1]
for _ in range(shuffle_repeats):
permutations = np.stack(
[permutation_rng.permutation(slot_count) for _ in val_ids]
)
permutation_tensor = torch.as_tensor(permutations, dtype=torch.long, device=device)
shuffled = {}
for name, values in validation.items():
if name == target:
shuffled[name] = values
else:
gather_indices = permutation_tensor.unsqueeze(-1).expand_as(values)
shuffled[name] = values.gather(1, gather_indices)
shuffled_prediction = model(shuffled, target, mask)
shuffled_error = (
shuffled_prediction[mask] - validation[target][mask]
).abs().mean()
shuffle_errors.append(float(shuffled_error.item()))
shuffled_mean = float(np.mean(shuffle_errors))
rows.append(
{
"target_modality": target,
"mask_ratio": ratio,
"mae_aligned": float(aligned_error.item()),
"mae_shuffled_mean": shuffled_mean,
"mae_shuffled_std": float(np.std(shuffle_errors, ddof=1))
if shuffle_repeats > 1
else 0.0,
"gain_align": shuffled_mean - float(aligned_error.item()),
"shuffle_repeats": float(shuffle_repeats),
"masked_values": int(mask.sum().item() * dimensions[target]),
}
)
return rows
def _emotion_features(
sample_ids: Sequence[str], aligned_by_id: Mapping[str, Mapping[str, np.ndarray]], bins: int = 5
) -> np.ndarray:
outputs = []
for sample_id in sample_ids:
combined = np.concatenate([aligned_by_id[sample_id][name] for name in MODALITIES], axis=-1)
segments = np.array_split(combined, bins, axis=0)
outputs.append(np.concatenate([segment.mean(axis=0) for segment in segments]))
return np.stack(outputs).astype(np.float32, copy=False)
def run_frozen_emotion_probe(
train_samples: Sequence[object],
val_samples: Sequence[object],
aligned_by_id: Mapping[str, Mapping[str, np.ndarray]],
) -> dict[str, float | int]:
"""Evaluate a regularized five-bin linear probe on frozen aligned features."""
train_ids = [sample.sample_id for sample in train_samples]
val_ids = [sample.sample_id for sample in val_samples]
x_train = _emotion_features(train_ids, aligned_by_id)
x_val = _emotion_features(val_ids, aligned_by_id)
y_train = np.asarray([sample.polarity for sample in train_samples], dtype=np.int64)
y_val = np.asarray([sample.polarity for sample in val_samples], dtype=np.int64)
target_train = np.asarray([sample.sentiment for sample in train_samples], dtype=np.float64)
target_val = np.asarray([sample.sentiment for sample in val_samples], dtype=np.float64)
scaler = StandardScaler()
x_train = scaler.fit_transform(x_train)
x_val = scaler.transform(x_val)
if np.unique(y_train).size > 1:
classifier = LogisticRegression(
C=0.1, class_weight="balanced", max_iter=2000, solver="lbfgs", random_state=0
)
else:
classifier = DummyClassifier(strategy="most_frequent")
classifier.fit(x_train, y_train)
predicted_class = classifier.predict(x_val)
regressor = Ridge(alpha=10.0)
regressor.fit(x_train, target_train)
predicted_score = regressor.predict(x_val)
if len(target_val) > 1 and np.std(predicted_score) > 0 and np.std(target_val) > 0:
pearson = float(np.corrcoef(predicted_score, target_val)[0, 1])
else:
pearson = float("nan")
return {
"accuracy": float(accuracy_score(y_val, predicted_class)),
"macro_f1": float(f1_score(y_val, predicted_class, average="macro", zero_division=0)),
"mae": float(mean_absolute_error(target_val, predicted_score)),
"pearson": pearson,
"n_train": len(train_samples),
"n_validation": len(val_samples),
}
File diff suppressed because it is too large. Load diff
+175
View File
@@ -0,0 +1,175 @@
from __future__ import annotations
from collections.abc import Mapping
import torch
import torch.nn.functional as F
from torch import Tensor
from .metrics import attention_row_similarity
from .types import AlignmentOutput, MODALITIES
def temporal_monotonicity_loss(
output: AlignmentOutput,
times: Mapping[str, Tensor],
durations: Tensor,
epsilon: float = 0.02,
) -> Tensor:
"""Penalize backward motion on the normalized clip timeline."""
if epsilon < 0:
raise ValueError("epsilon must be non-negative")
losses = []
for name in MODALITIES:
mu = torch.bmm(output.weights[name], times[name].unsqueeze(-1)).squeeze(-1)
mu = mu / durations[:, None].clamp_min(torch.finfo(mu.dtype).eps)
backward = F.relu(mu[:, :-1] - mu[:, 1:] - epsilon)
losses.append(backward.square().mean())
return torch.stack(losses).mean()
def cross_modal_contrastive_loss(
aligned: Mapping[str, Tensor], temperature: float = 0.1
) -> Tensor:
"""Symmetric in-batch InfoNCE over same-sample, same-grid-slot positives."""
if temperature <= 0:
raise ValueError("temperature must be positive")
if set(aligned) != set(MODALITIES):
raise ValueError(f"aligned must contain exactly {MODALITIES}")
pair_losses = []
for left_index, left_name in enumerate(MODALITIES):
for right_name in MODALITIES[left_index + 1 :]:
left = F.normalize(aligned[left_name].flatten(0, 1), dim=-1)
right = F.normalize(aligned[right_name].flatten(0, 1), dim=-1)
if left.shape != right.shape:
raise ValueError("contrastive representations must share [B, K, D]")
logits = left @ right.T / temperature
labels = torch.arange(logits.shape[0], device=logits.device)
pair_losses.append(
(F.cross_entropy(logits, labels) + F.cross_entropy(logits.T, labels)) / 2
)
return torch.stack(pair_losses).mean()
def masked_reconstruction_loss(prediction: Tensor, target: Tensor, mask: Tensor) -> Tensor:
"""Smooth-L1 loss over masked grid slots only."""
if prediction.shape != target.shape:
raise ValueError("prediction and target must have the same shape")
if mask.shape != target.shape[:2] or mask.dtype != torch.bool:
raise ValueError("mask must be boolean with shape [B, K]")
if not bool(mask.any()):
raise ValueError("mask must select at least one target slot")
element_loss = F.smooth_l1_loss(prediction, target, reduction="none")
return element_loss[mask].mean()
def temporal_span_loss(
output: AlignmentOutput,
times: Mapping[str, Tensor],
durations: Tensor,
minimum_span: float = 0.7,
modalities: tuple[str, ...] = MODALITIES,
) -> Tensor:
"""Penalize an expected-time path that does not cover enough of a clip."""
if not 0.0 <= minimum_span <= 1.0:
raise ValueError("minimum_span must be in [0, 1]")
if not modalities or any(name not in MODALITIES for name in modalities):
raise ValueError("modalities must be a non-empty subset of MODALITIES")
losses = []
for name in modalities:
mu = torch.bmm(output.weights[name], times[name].unsqueeze(-1)).squeeze(-1)
mu = mu / durations[:, None].clamp_min(torch.finfo(mu.dtype).eps)
span = mu[:, -1] - mu[:, 0]
losses.append(F.relu(minimum_span - span).square().mean())
return torch.stack(losses).mean()
def attention_diversity_loss(
output: AlignmentOutput,
modalities: tuple[str, ...] = MODALITIES,
min_separation: int = 6,
) -> Tensor:
"""Penalize similar attention rows for slots far apart on the grid."""
if min_separation < 1:
raise ValueError("min_separation must be at least one")
if not modalities or any(name not in MODALITIES for name in modalities):
raise ValueError("modalities must be a non-empty subset of MODALITIES")
return torch.stack(
[attention_row_similarity(output.weights[name], min_separation).mean() for name in modalities]
).mean()
def weak_temporal_band_loss(
output: AlignmentOutput,
times: Mapping[str, Tensor],
durations: Tensor,
targets: Mapping[str, Tensor],
margin: float = 0.1,
) -> Tensor:
"""Allow soft alignment while keeping expected times near weak slot anchors."""
if margin < 0:
raise ValueError("margin must be non-negative")
if not targets:
return next(iter(output.weights.values())).sum() * 0.0
losses = []
for name, target in targets.items():
if name not in MODALITIES:
raise ValueError(f"unknown modality in temporal targets: {name}")
mu = torch.bmm(output.weights[name], times[name].unsqueeze(-1)).squeeze(-1)
mu = mu / durations[:, None].clamp_min(torch.finfo(mu.dtype).eps)
if target.shape != mu.shape:
raise ValueError(f"band target for {name} must have shape {tuple(mu.shape)}")
distance = (mu - target).abs()
losses.append(F.relu(distance - margin).square().mean())
return torch.stack(losses).mean()
def alignment_training_loss(
output: AlignmentOutput,
times: Mapping[str, Tensor],
durations: Tensor,
reconstruction: Tensor,
*,
lambda_rec: float = 1.0,
lambda_con: float = 1.0,
lambda_mono: float = 0.1,
lambda_span: float = 0.0,
lambda_div: float = 0.0,
lambda_band: float = 0.0,
epsilon: float = 0.02,
minimum_span: float = 0.7,
coverage_modalities: tuple[str, ...] = MODALITIES,
diversity_modalities: tuple[str, ...] = MODALITIES,
diversity_min_separation: int = 6,
band_targets: Mapping[str, Tensor] | None = None,
band_margin: float = 0.1,
) -> tuple[Tensor, dict[str, Tensor]]:
"""Shared M3/M4 training objective; emotion labels are deliberately unused."""
contrastive = cross_modal_contrastive_loss(output.aligned)
monotonicity = temporal_monotonicity_loss(output, times, durations, epsilon)
span = temporal_span_loss(
output, times, durations, minimum_span, modalities=coverage_modalities
)
diversity = attention_diversity_loss(
output, diversity_modalities, min_separation=diversity_min_separation
)
band = weak_temporal_band_loss(
output, times, durations, band_targets or {}, margin=band_margin
)
total = (
lambda_rec * reconstruction
+ lambda_con * contrastive
+ lambda_mono * monotonicity
+ lambda_span * span
+ lambda_div * diversity
+ lambda_band * band
)
return total, {
"reconstruction": reconstruction,
"contrastive": contrastive,
"monotonicity": monotonicity,
"span": span,
"diversity": diversity,
"band": band,
"total": total,
}
File diff suppressed because it is too large. Load diff
+146
View File
@@ -0,0 +1,146 @@
from __future__ import annotations
import numpy as np
import torch
import torch.nn.functional as F
from torch import Tensor
from .types import MODALITIES
def alignment_trajectory(weights: Tensor, times: Tensor, durations: Tensor) -> Tensor:
"""Return expected normalized source time at each common-grid position."""
if weights.ndim != 3 or times.shape != (weights.shape[0], weights.shape[2]):
raise ValueError("weights [B,K,L] and times [B,L] must agree")
if durations.shape != (weights.shape[0],):
raise ValueError("durations must have shape [B]")
expected_seconds = torch.bmm(weights, times.unsqueeze(-1)).squeeze(-1)
return expected_seconds / durations[:, None].clamp_min(torch.finfo(expected_seconds.dtype).eps)
def monotonicity_violation_rate(trajectory: Tensor, epsilon: float = 0.02) -> Tensor:
"""Per-sample fraction of adjacent grid pairs that move backwards by epsilon."""
if trajectory.ndim != 2 or trajectory.shape[1] < 2:
raise ValueError("trajectory must have shape [B, K] with K >= 2")
if epsilon < 0:
raise ValueError("epsilon must be non-negative")
return ((trajectory[:, :-1] - trajectory[:, 1:]) > epsilon).float().mean(dim=1)
def normalized_attention_entropy(weights: Tensor, valid: Tensor) -> Tensor:
"""Per-row entropy normalized by the number of valid source positions."""
if weights.ndim != 3 or valid.shape != (weights.shape[0], weights.shape[2]):
raise ValueError("weights [B,K,L] and valid [B,L] must agree")
safe = weights.clamp_min(torch.finfo(weights.dtype).tiny)
entropy = -(weights * safe.log()).sum(dim=-1)
counts = valid.sum(dim=-1).clamp_min(1)
denominator = counts.float().log().clamp_min(torch.finfo(torch.float32).eps)
normalized = entropy / denominator[:, None]
return torch.where(counts[:, None] > 1, normalized, torch.zeros_like(normalized))
def attention_row_similarity(weights: Tensor, min_separation: int = 1) -> Tensor:
"""Mean cosine similarity between attention rows, per sample.
``min_separation=1`` compares every distinct pair (the C_row collapse
score). A value of 6 compares only pairs more than five slots apart,
matching the training diversity loss. Scores near one mean that slots
attend to nearly the same source distribution.
"""
if weights.ndim != 3:
raise ValueError("weights must have shape [B, K, L]")
if min_separation < 1:
raise ValueError("min_separation must be at least one")
batch_size, grid_size, _ = weights.shape
if grid_size <= min_separation:
return torch.zeros(batch_size, dtype=weights.dtype, device=weights.device)
normalized = F.normalize(weights, p=2, dim=-1, eps=1e-12)
similarities = torch.bmm(normalized, normalized.transpose(1, 2))
pair_mask = torch.triu(
torch.ones(grid_size, grid_size, dtype=torch.bool, device=weights.device),
diagonal=min_separation,
)
return similarities[:, pair_mask].mean(dim=-1)
def attention_width80(weights: Tensor, threshold: float = 0.8) -> Tensor:
"""Shortest contiguous source-index span containing the requested mass.
Returns integer widths with shape ``[B, K]``. This is a diagnostic, not a
score to maximize or minimize on its own.
"""
if weights.ndim != 3 or not 0 < threshold <= 1:
raise ValueError("weights must be [B,K,L] and threshold in (0, 1]")
rows = weights.detach().to(device="cpu", dtype=torch.float64).numpy()
widths = np.empty(rows.shape[:2], dtype=np.int64)
for batch_index in range(rows.shape[0]):
for grid_index in range(rows.shape[1]):
row = rows[batch_index, grid_index]
left = 0
mass = 0.0
best = len(row)
for right, value in enumerate(row):
mass += float(value)
while left <= right and mass - float(row[left]) >= threshold:
mass -= float(row[left])
left += 1
if mass + 1e-12 >= threshold:
best = min(best, right - left + 1)
widths[batch_index, grid_index] = best
return torch.from_numpy(widths)
def retrieval_metrics(query: Tensor, target: Tensor, chunk_size: int = 256) -> dict[str, float]:
"""Grid-slot retrieval; same flattened sample/slot index is the positive.
Use only as a representation-consistency probe. It is not independent
temporal ground truth; report human-labeled temporal scores separately.
"""
if query.ndim != 3 or target.ndim != 3 or query.shape != target.shape:
raise ValueError("query and target must have matching [N, K, D] shapes")
if query.shape[0] * query.shape[1] < 1:
raise ValueError("retrieval requires at least one grid position")
q = F.normalize(query.flatten(0, 1), dim=-1)
t = F.normalize(target.flatten(0, 1), dim=-1)
total = q.shape[0]
ranks = torch.empty(total, dtype=torch.long, device=q.device)
for start in range(0, total, chunk_size):
stop = min(start + chunk_size, total)
scores = q[start:stop] @ t.T
positives = scores[torch.arange(stop - start, device=q.device),
torch.arange(start, stop, device=q.device)]
ranks[start:stop] = 1 + (scores > positives[:, None]).sum(dim=1)
ranks_f = ranks.float()
return {
"r_at_1": float((ranks <= 1).float().mean().item()),
"r_at_5": float((ranks <= min(5, total)).float().mean().item()),
"mrr": float((1.0 / ranks_f).mean().item()),
"queries": float(total),
}
def summarize_alignment(
weights: dict[str, Tensor],
times: dict[str, Tensor],
valid: dict[str, Tensor],
durations: Tensor,
epsilon: float = 0.02,
) -> dict[str, dict[str, float]]:
"""Produce sample-aggregated E1-E3 diagnostics for each modality."""
summary: dict[str, dict[str, float]] = {}
for name in MODALITIES:
trajectory = alignment_trajectory(weights[name], times[name], durations)
mvr = monotonicity_violation_rate(trajectory, epsilon)
entropy = normalized_attention_entropy(weights[name], valid[name])
width = attention_width80(weights[name])
summary[name] = {
"mvr": float(mvr.mean().item()),
"normalized_entropy": float(entropy.mean().item()),
"width80_indices": float(width.float().mean().item()),
"mean_time_start": float(trajectory[:, 0].mean().item()),
"mean_time_end": float(trajectory[:, -1].mean().item()),
"trajectory_span_fraction": float(
(trajectory[:, -1] - trajectory[:, 0]).mean().item()
),
}
return summary
+247
View File
@@ -0,0 +1,247 @@
from __future__ import annotations
from collections.abc import Mapping
import torch
from torch import Tensor, nn
from .alignment import index_alignment
from .types import AlignmentOutput, MODALITIES, SequenceBatch
def _sinusoidal_position_encoding(length: int, dimension: int) -> Tensor:
"""Build a fixed, small-amplitude sinusoidal code for ordered slots."""
positions = torch.arange(length, dtype=torch.float32).unsqueeze(1)
frequencies = torch.exp(
torch.arange(0, dimension, 2, dtype=torch.float32)
* (-torch.log(torch.tensor(10000.0)) / dimension)
)
encoding = torch.zeros(length, dimension, dtype=torch.float32)
encoding[:, 0::2] = torch.sin(positions * frequencies)
odd_width = encoding[:, 1::2].shape[1]
if odd_width:
encoding[:, 1::2] = torch.cos(positions * frequencies[:odd_width])
return encoding * (dimension**-0.5)
def _temporal_position_encoding(times: Tensor, dimension: int) -> Tensor:
"""Encode normalized source times with a fixed Fourier feature bank."""
half_width = (dimension + 1) // 2
frequencies = torch.logspace(
0.0,
1.6989700043360187,
steps=half_width,
device=times.device,
dtype=times.dtype,
)
angles = (2.0 * torch.pi) * times.unsqueeze(-1) * frequencies
encoding = torch.empty(*times.shape, dimension, device=times.device, dtype=times.dtype)
encoding[..., 0::2] = torch.sin(angles)
if dimension > 1:
encoding[..., 1::2] = torch.cos(angles[..., : encoding[..., 1::2].shape[-1]])
return encoding * (dimension**-0.25)
def _validate_inputs(sequences: Mapping[str, SequenceBatch], grid_size: int) -> int:
if set(sequences) != set(MODALITIES):
raise ValueError(f"sequences must contain exactly {MODALITIES}")
batch_sizes = {sequences[name].features.shape[0] for name in MODALITIES}
if len(batch_sizes) != 1:
raise ValueError("all modalities must have the same batch size")
if grid_size < 1:
raise ValueError("grid_size must be positive")
return batch_sizes.pop()
class _CrossAttention(nn.Module):
def __init__(self, dimension: int, heads: int, dropout: float) -> None:
super().__init__()
if dimension % heads != 0:
raise ValueError("dimension must be divisible by heads")
self.attention = nn.MultiheadAttention(
# Keep the returned alignment matrix row-stochastic during training.
# PyTorch applies attention dropout to returned weights when it is
# nonzero, which breaks the shared AlignmentOutput contract.
embed_dim=dimension, num_heads=heads, dropout=0.0, batch_first=True
)
self.input_dropout = nn.Dropout(dropout)
def forward(
self,
query: Tensor,
source: Tensor,
valid: Tensor,
*,
source_position: Tensor | None = None,
) -> tuple[Tensor, Tensor]:
key = source if source_position is None else source + source_position
values, weights = self.attention(
self.input_dropout(query),
self.input_dropout(key),
self.input_dropout(source),
key_padding_mask=~valid,
need_weights=True,
average_attn_weights=True,
)
return values, weights
class TextAnchoredCrossAttention(nn.Module):
"""M3: transcript-order text slots query Audio and Vision sequences."""
def __init__(
self,
dimensions: Mapping[str, int],
grid_size: int = 50,
hidden_size: int = 128,
heads: int = 4,
dropout: float = 0.1,
source_time_encoding: bool = False,
) -> None:
super().__init__()
if set(dimensions) != set(MODALITIES):
raise ValueError(f"dimensions must contain exactly {MODALITIES}")
self.grid_size = grid_size
self.source_time_encoding = source_time_encoding
self.hidden_size = hidden_size
self.projections = nn.ModuleDict(
{name: nn.Linear(dimensions[name], hidden_size) for name in MODALITIES}
)
self.audio_attention = _CrossAttention(hidden_size, heads, dropout)
self.vision_attention = _CrossAttention(hidden_size, heads, dropout)
def forward(
self, sequences: Mapping[str, SequenceBatch], durations: Tensor | None = None
) -> AlignmentOutput:
_validate_inputs(sequences, self.grid_size)
if self.source_time_encoding:
if durations is None or durations.shape != (sequences["text"].features.shape[0],):
raise ValueError("durations with shape [B] are required for source time encoding")
durations = durations.to(device=sequences["text"].times.device).clamp_min(1e-8)
projected = {
name: self.projections[name](sequences[name].features) for name in MODALITIES
}
text_weights, text_fallbacks = index_alignment(sequences["text"].valid, self.grid_size)
text_query = torch.bmm(text_weights.to(projected["text"].dtype), projected["text"])
source_positions: dict[str, Tensor] = {}
if self.source_time_encoding:
assert durations is not None
text_centers = torch.bmm(
text_weights.to(sequences["text"].times.dtype),
sequences["text"].times.unsqueeze(-1),
).squeeze(-1) / durations[:, None]
text_query = text_query + _temporal_position_encoding(
text_centers, self.hidden_size
).to(text_query.dtype)
for name in ("audio", "vision"):
normalized_times = sequences[name].times / durations[:, None]
source_positions[name] = _temporal_position_encoding(
normalized_times, self.hidden_size
).to(projected[name].dtype)
audio_values, audio_weights = self.audio_attention(
text_query,
projected["audio"],
sequences["audio"].valid,
source_position=source_positions.get("audio"),
)
vision_values, vision_weights = self.vision_attention(
text_query,
projected["vision"],
sequences["vision"].valid,
source_position=source_positions.get("vision"),
)
output = AlignmentOutput(
weights={"text": text_weights, "audio": audio_weights, "vision": vision_weights},
aligned={"text": text_query, "audio": audio_values, "vision": vision_values},
fallback_rows={"text": text_fallbacks, "audio": 0, "vision": 0},
)
output.validate({name: sequences[name].valid for name in MODALITIES})
return output
class SharedLatentTimeline(nn.Module):
"""M4: K learned shared slots attend independently to all three modalities."""
def __init__(
self,
dimensions: Mapping[str, int],
grid_size: int = 50,
hidden_size: int = 128,
heads: int = 4,
dropout: float = 0.1,
absolute_position_encoding: bool = False,
source_time_encoding: bool = False,
) -> None:
super().__init__()
if set(dimensions) != set(MODALITIES):
raise ValueError(f"dimensions must contain exactly {MODALITIES}")
if hidden_size % heads != 0:
raise ValueError("hidden_size must be divisible by heads")
self.grid_size = grid_size
self.hidden_size = hidden_size
self.absolute_position_encoding = absolute_position_encoding
self.source_time_encoding = source_time_encoding
self.register_buffer(
"sinusoidal_positions",
_sinusoidal_position_encoding(grid_size, hidden_size),
persistent=False,
)
self.projections = nn.ModuleDict(
{name: nn.Linear(dimensions[name], hidden_size) for name in MODALITIES}
)
self.slots = nn.Parameter(torch.empty(grid_size, hidden_size))
nn.init.normal_(self.slots, mean=0.0, std=hidden_size**-0.5)
self.attention = nn.ModuleDict(
{name: _CrossAttention(hidden_size, heads, dropout) for name in MODALITIES}
)
def forward(
self, sequences: Mapping[str, SequenceBatch], durations: Tensor | None = None
) -> AlignmentOutput:
batch_size = _validate_inputs(sequences, self.grid_size)
if self.source_time_encoding:
if durations is None or durations.shape != (batch_size,):
raise ValueError("durations with shape [B] are required for source time encoding")
durations = durations.to(device=sequences["text"].times.device).clamp_min(1e-8)
latent_queries = self.slots.unsqueeze(0).expand(batch_size, -1, -1)
if self.absolute_position_encoding:
latent_queries = latent_queries + self.sinusoidal_positions.unsqueeze(0)
if self.source_time_encoding:
centers = (
torch.arange(
self.grid_size,
dtype=sequences["text"].times.dtype,
device=sequences["text"].times.device,
)
+ 0.5
) / self.grid_size
centers = centers.unsqueeze(0).expand(batch_size, -1)
latent_queries = latent_queries + _temporal_position_encoding(
centers, self.hidden_size
).to(latent_queries.dtype)
weights: dict[str, Tensor] = {}
aligned: dict[str, Tensor] = {}
for name in MODALITIES:
source = self.projections[name](sequences[name].features)
source_position = None
if self.source_time_encoding:
assert durations is not None
normalized_times = sequences[name].times / durations[:, None]
source_position = _temporal_position_encoding(
normalized_times, self.hidden_size
).to(source.dtype)
aligned[name], weights[name] = self.attention[name](
latent_queries,
source,
sequences[name].valid,
source_position=source_position,
)
output = AlignmentOutput(
weights=weights,
aligned=aligned,
fallback_rows={name: 0 for name in MODALITIES},
)
output.validate({name: sequences[name].valid for name in MODALITIES})
return output
@@ -0,0 +1,194 @@
"""Summarize the saved D0-D3 alignment-diagnostic metrics without retraining."""
from __future__ import annotations
import argparse
import csv
from collections import defaultdict
from pathlib import Path
import shutil
from statistics import mean
from typing import Any
METRICS = (
"mvr",
"normalized_entropy",
"c_row",
"trajectory_span",
"mean_absolute_time_center_error",
"gaussian_target_kl",
)
def merge_trial_metrics(run_dir: Path, destination: Path) -> None:
"""Rebuild the combined file with a variant label for PE controls."""
rows: list[dict[str, str]] = []
fields: list[str] | None = None
for experiment in ("D0", "D1", "D2", "D3"):
for path in sorted((run_dir / experiment).glob("*_metrics.csv")):
variant = path.name.removesuffix("_metrics.csv")
with path.open(newline="", encoding="utf-8-sig") as handle:
reader = csv.DictReader(handle)
if fields is None:
fields = ["experiment", "method", "variant"] + [
name for name in (reader.fieldnames or [])
if name not in {"experiment", "method", "variant"}
]
for row in reader:
row["variant"] = variant
rows.append(row)
if not fields:
raise FileNotFoundError(f"no per-trial metric CSV files found under {run_dir}")
destination.parent.mkdir(parents=True, exist_ok=True)
with destination.open("w", newline="", encoding="utf-8-sig") as handle:
writer = csv.DictWriter(handle, fieldnames=fields)
writer.writeheader()
writer.writerows(rows)
def summarize(source: Path, destination: Path) -> list[dict[str, Any]]:
groups: dict[tuple[str, str, str], list[dict[str, str]]] = defaultdict(list)
with source.open(newline="", encoding="utf-8-sig") as handle:
for row in csv.DictReader(handle):
if row["experiment"] in {"D2", "D3"}:
groups[(row["experiment"], row["method"], row["modality"])].append(row)
summaries: list[dict[str, Any]] = []
for (experiment, method, modality), rows in sorted(groups.items()):
summary: dict[str, Any] = {
"experiment": experiment,
"method": method,
"modality": modality,
"sample_count": len(rows),
}
for metric in METRICS:
values = [float(row[metric]) for row in rows if row.get(metric, "") != ""]
summary[f"mean_{metric}"] = mean(values) if values else ""
summary[f"n_{metric}"] = len(values)
summaries.append(summary)
destination.parent.mkdir(parents=True, exist_ok=True)
with destination.open("w", newline="", encoding="utf-8-sig") as handle:
writer = csv.DictWriter(handle, fieldnames=list(summaries[0]))
writer.writeheader()
writer.writerows(summaries)
return summaries
def build_report_bundle(run_dir: Path) -> Path:
"""Copy compact diagnostic evidence, leaving large checkpoints local."""
bundle = run_dir / "report_bundle"
bundle.mkdir(parents=True, exist_ok=True)
for name in (
"debug_summary.csv",
"metric_summary.csv",
"per_sample_metrics.csv",
"run_manifest.json",
"experiment.log",
):
source = run_dir / name
if source.exists():
shutil.copy2(source, bundle / name)
for experiment in ("D0", "D1", "D2", "D3"):
source_dir = run_dir / experiment
target_dir = bundle / experiment
target_dir.mkdir(exist_ok=True)
for pattern in ("*_history.csv", "*_metrics.csv", "*_heatmap.png", "*_trajectory.png"):
for source in source_dir.glob(pattern):
shutil.copy2(source, target_dir / source.name)
synthetic_source = run_dir / "synthetic"
synthetic_target = bundle / "synthetic"
synthetic_target.mkdir(exist_ok=True)
for name in ("metrics.csv", "training_history.csv", "run_manifest.json"):
source = synthetic_source / name
if source.exists():
shutil.copy2(source, synthetic_target / name)
for source in synthetic_source.rglob("*.png"):
target = synthetic_target / source.relative_to(synthetic_source)
target.parent.mkdir(parents=True, exist_ok=True)
shutil.copy2(source, target)
source_time_dir = run_dir / "source_time"
source_time_target = bundle / "source_time"
source_time_target.mkdir(exist_ok=True)
for name in ("summary.csv", "per_sample_metrics.csv", "run_manifest.json"):
source = source_time_dir / name
if source.exists():
shutil.copy2(source, source_time_target / name)
for variant_dir in source_time_dir.glob("M*"):
target_dir = source_time_target / variant_dir.name
target_dir.mkdir(exist_ok=True)
for pattern in ("history.csv", "metrics.csv", "*_heatmap.png", "*_trajectory.png"):
for source in variant_dir.glob(pattern):
shutil.copy2(source, target_dir / source.name)
heldout_dir = run_dir / "heldout"
heldout_target = bundle / "heldout"
heldout_target.mkdir(exist_ok=True)
for name in (
"heldout_summary.csv",
"per_sample_metrics.csv",
"training_summary.csv",
"training_history.csv",
"run_manifest.json",
):
source = heldout_dir / name
if source.exists():
shutil.copy2(source, heldout_target / name)
for variant_dir in heldout_dir.glob("M*"):
target_dir = heldout_target / variant_dir.name
target_dir.mkdir(exist_ok=True)
for pattern in ("history.csv", "heldout_metrics.csv", "*_heatmap.png", "*_trajectory.png"):
for source in variant_dir.glob(pattern):
shutil.copy2(source, target_dir / source.name)
readme = """# Q1 M3/M4 可学习性诊断结果包
本包包含 D0–D3 诊断的汇总表、运行清单、日志、各试验的损失历史、逐样本指标和代表性图像。PyTorch 检查点与逐样本 `.npz` 注意力矩阵留在上级 `outputs/alignment_debug/`,因此结果包较轻。
- D0/D1:单样本过拟合诊断。
- D1-S:输入仅为合成时间坐标;M3/M4 都能把 Gaussian 目标拟合至约 1e-4 KL。
- D2/D3:100 条样本上的同集训练/评价诊断,单个随机种子,不用于声称泛化。
- D4:真实单样本加入显式时间 key/query 特征后,Audio 注意力形成局部时间带。
- D5:按 `video_id` 留出 20 条样本;M4 的 Audio/Vision 时间带能迁移,M3 有改善但仍未充分贴近目标。
- Gaussian 时间目标是根据源时间戳生成的弱先验,不是人工对齐真值。
- M3 的 Gaussian KL 取 Audio/Vision 平均,M4 取 Text/Audio/Vision 平均;KL 数值仅用于各自的优化诊断,不可作为方法排名。
- D1 真实特征上 Text/Vision 能拟合,Audio 仍失败;合成对照成功,提示真实源特征的显式时间身份值得优先验证。
- 全数据 Q/K 梯度仍非零,但学习到的注意力没有稳定形成局部时间带;加入重构和对比目标后 Audio 塌缩更明显。
- 环境:Fedora WSL,`uv`,NVIDIA GeForce RTX 5070 Ti,Python 3.14.7,PyTorch 2.14.0+cu130。
详细解释见项目根目录的 `RESULTS.md`。
"""
(bundle / "README.md").write_text(readme, encoding="utf-8")
return bundle
def main() -> None:
project = Path(__file__).resolve().parents[1]
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument(
"--input",
type=Path,
default=project / "outputs/alignment_debug/per_sample_metrics.csv",
)
parser.add_argument(
"--output",
type=Path,
default=project / "outputs/alignment_debug/metric_summary.csv",
)
args = parser.parse_args()
merge_trial_metrics(args.input.parent, args.input)
for row in summarize(args.input, args.output):
values = " ".join(
f"{name}={row[f'mean_{name}']:.4f}"
for name in METRICS
if row[f"mean_{name}"] != ""
)
print(
f"{row['experiment']} {row['method']} {row['modality']} "
f"n={row['sample_count']} {values}"
)
print(f"Wrote {args.output}")
print(f"Built report bundle at {build_report_bundle(args.input.parent)}")
if __name__ == "__main__":
main()
@@ -0,0 +1,338 @@
"""Fit Gaussian time bands with synthetic time-only inputs as an attention control."""
from __future__ import annotations
import argparse
import csv
import json
import platform
import random
from datetime import datetime, timezone
from pathlib import Path
from typing import Any
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import torch
from torch import Tensor, nn
from .metrics import alignment_trajectory, normalized_attention_entropy
from .models import SharedLatentTimeline, TextAnchoredCrossAttention
from .types import MODALITIES, SequenceBatch
GRID_SIZE = 50
HIDDEN_SIZE = 128
SIGMA = 0.10
SEED = 42
MODALITY_LENGTHS = {"text": GRID_SIZE, "audio": 256, "vision": 100}
def _seed(seed: int) -> None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)
torch.backends.cudnn.deterministic = True
torch.backends.cudnn.benchmark = False
def _synthetic_sequence(length: int, device: torch.device) -> SequenceBatch:
times = torch.linspace(0.0, 1.0, length, device=device).unsqueeze(0)
# These features contain only the source's synthetic time coordinate.
features = torch.stack(
(times, times.square(), torch.ones_like(times)), dim=-1
)
valid = torch.ones_like(times, dtype=torch.bool)
return SequenceBatch(features=features, times=times, valid=valid)
def _targets(
sequences: dict[str, SequenceBatch], method: str, device: torch.device
) -> dict[str, Tensor]:
centers = (torch.arange(GRID_SIZE, device=device, dtype=torch.float32) + 0.5) / GRID_SIZE
names = ("audio", "vision") if method == "M3" else MODALITIES
targets: dict[str, Tensor] = {}
for name in names:
times = sequences[name].times
logits = -0.5 * ((times[:, None, :] - centers[None, :, None]) / SIGMA).square()
targets[name] = torch.softmax(logits, dim=-1)
return targets
def _loss(output: Any, targets: dict[str, Tensor]) -> Tensor:
losses = []
for name, target in targets.items():
predicted = output.weights[name].clamp_min(1e-8)
safe_target = target.clamp_min(1e-12)
losses.append((safe_target * (safe_target.log() - predicted.log())).sum(-1).mean())
return torch.stack(losses).mean()
def _attention_layers(model: nn.Module, method: str) -> list[nn.MultiheadAttention]:
if method == "M3":
return [model.audio_attention.attention, model.vision_attention.attention]
return [model.attention[name].attention for name in MODALITIES]
def _gradient_summary(
model: nn.Module, method: str, loss: Tensor
) -> dict[str, float]:
layers = _attention_layers(model, method)
params = [layer.in_proj_weight for layer in layers]
if method == "M4":
params.append(model.slots)
gradients = torch.autograd.grad(loss, params, retain_graph=True)
q_sq = torch.zeros((), device=loss.device)
k_sq = torch.zeros((), device=loss.device)
for grad in gradients[: len(layers)]:
q_sq += grad[:HIDDEN_SIZE].square().sum()
k_sq += grad[HIDDEN_SIZE : 2 * HIDDEN_SIZE].square().sum()
result = {"grad_WQ": float(q_sq.sqrt().item()), "grad_WK": float(k_sq.sqrt().item())}
if method == "M4":
result["grad_Z"] = float(gradients[-1].norm().item())
return result
def _write_csv(path: Path, rows: list[dict[str, Any]]) -> None:
if not rows:
return
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("w", newline="", encoding="utf-8-sig") as handle:
fields = list(dict.fromkeys(key for row in rows for key in row))
writer = csv.DictWriter(handle, fieldnames=fields)
writer.writeheader()
writer.writerows(rows)
def _save_plots(
folder: Path,
method: str,
variant: str,
output: Any,
targets: dict[str, Tensor],
sequences: dict[str, SequenceBatch],
duration: float = 1.0,
) -> list[dict[str, Any]]:
names = ("audio", "vision") if method == "M3" else MODALITIES
rows: list[dict[str, Any]] = []
fig, axes = plt.subplots(len(names), 2, figsize=(11, 3.4 * len(names)), constrained_layout=True)
if len(names) == 1:
axes = np.asarray([axes])
fig_traj, ax_traj = plt.subplots(figsize=(8, 5), constrained_layout=True)
grid = (np.arange(GRID_SIZE, dtype=np.float32) + 0.5) / GRID_SIZE
for row_index, name in enumerate(names):
weights = output.weights[name][0].detach().cpu().numpy()
target = targets[name][0].detach().cpu().numpy()
times = sequences[name].times[0].detach().cpu().numpy()
entropy = normalized_attention_entropy(
output.weights[name], sequences[name].valid
).mean().item()
trajectory = alignment_trajectory(
output.weights[name], sequences[name].times, torch.tensor([duration], device=output.weights[name].device)
)[0]
trajectory_np = trajectory.detach().cpu().numpy()
span = float(trajectory_np[-1] - trajectory_np[0])
kl = float(
(targets[name] * (targets[name].clamp_min(1e-12).log() - output.weights[name].clamp_min(1e-8).log()))
.sum(-1)
.mean()
.item()
)
rows.append(
{
"method": method,
"variant": variant,
"modality": name,
"normalized_entropy": entropy,
"trajectory_span": span,
"mean_absolute_time_center_error": float(
np.mean(np.abs(trajectory_np - grid))
),
"gaussian_target_kl": kl,
}
)
extent = (float(times[0]), float(times[-1]), 0.0, 1.0)
image = axes[row_index, 0].imshow(
weights, origin="lower", aspect="auto", interpolation="nearest", extent=extent, cmap="magma"
)
axes[row_index, 0].set_title(f"{name}: learned A")
axes[row_index, 0].set_xlabel("synthetic source time")
axes[row_index, 0].set_ylabel("slot / K")
fig.colorbar(image, ax=axes[row_index, 0], fraction=0.046, pad=0.04)
image_target = axes[row_index, 1].imshow(
target, origin="lower", aspect="auto", interpolation="nearest", extent=extent, cmap="magma"
)
axes[row_index, 1].set_title(f"{name}: Gaussian target P")
axes[row_index, 1].set_xlabel("synthetic source time")
axes[row_index, 1].set_ylabel("slot / K")
fig.colorbar(image_target, ax=axes[row_index, 1], fraction=0.046, pad=0.04)
ax_traj.plot(grid, trajectory_np, label=f"{name} learned")
ax_traj.plot([0, 1], [0, 1], "k:", label="uniform-time reference")
ax_traj.set(xlim=(0, 1), ylim=(0, 1), xlabel="slot position", ylabel="expected source time")
ax_traj.grid(alpha=0.2)
ax_traj.legend()
ax_traj.set_title(f"{method} {variant}: synthetic alignment trajectory")
folder.mkdir(parents=True, exist_ok=True)
fig.savefig(folder / f"{variant}_heatmap.png", dpi=160)
fig_traj.savefig(folder / f"{variant}_trajectory.png", dpi=160)
plt.close(fig)
plt.close(fig_traj)
return rows
def _run_trial(
method: str,
variant: str,
*,
device: torch.device,
steps: int,
seed: int,
output_dir: Path,
) -> tuple[list[dict[str, Any]], list[dict[str, Any]], dict[str, Any]]:
_seed(seed)
dimensions = {name: 3 for name in MODALITIES}
if method == "M3":
model: nn.Module = TextAnchoredCrossAttention(
dimensions, grid_size=GRID_SIZE, hidden_size=HIDDEN_SIZE, heads=4, dropout=0.0
)
use_pe = False
else:
use_pe = variant == "M4_sinPE"
model = SharedLatentTimeline(
dimensions,
grid_size=GRID_SIZE,
hidden_size=HIDDEN_SIZE,
heads=4,
dropout=0.0,
absolute_position_encoding=use_pe,
)
model.to(device)
sequences = {
name: _synthetic_sequence(MODALITY_LENGTHS[name], device) for name in MODALITIES
}
targets = _targets(sequences, method, device)
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=0.0)
history: list[dict[str, Any]] = []
output = None
print(f"[{variant}] synthetic time-only inputs, steps={steps}, device={device}", flush=True)
for step in range(1, steps + 1):
model.train()
output = model(sequences)
loss = _loss(output, targets)
if not torch.isfinite(loss):
raise FloatingPointError(f"non-finite synthetic alignment loss for {variant} at step {step}")
row: dict[str, Any] = {"method": method, "variant": variant, "step": step, "L_align": float(loss.item())}
if step == 1 or step % 100 == 0 or step == steps:
row.update(_gradient_summary(model, method, loss))
optimizer.zero_grad(set_to_none=True)
loss.backward()
nn.utils.clip_grad_norm_(model.parameters(), 2.0)
optimizer.step()
history.append(row)
if step == 1 or step % 100 == 0 or step == steps:
print(
f"[{variant} {step}/{steps}] KL={row['L_align']:.5f} "
f"grad_Q/K/Z={row.get('grad_WQ', 0):.3g}/{row.get('grad_WK', 0):.3g}/"
f"{row.get('grad_Z', float('nan')):.3g}",
flush=True,
)
assert output is not None
model.eval()
with torch.no_grad():
output = model(sequences)
metric_rows = _save_plots(
output_dir / variant, method, variant, output, targets, sequences
)
history_rows = history
checkpoint = output_dir / f"{variant}_checkpoint.pt"
torch.save(
{
"method": method,
"variant": variant,
"seed": seed,
"steps": steps,
"absolute_position_encoding": use_pe,
"model_state_dict": model.state_dict(),
"synthetic_only": True,
},
checkpoint,
)
return metric_rows, history_rows, {"checkpoint": str(checkpoint), "final_loss": history[-1]["L_align"]}
def run(args: argparse.Namespace) -> dict[str, Any]:
output_dir = args.output_dir
output_dir.mkdir(parents=True, exist_ok=True)
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")
variants = [("M3", "M3"), ("M4", "M4_noPE"), ("M4", "M4_sinPE")]
all_metrics: list[dict[str, Any]] = []
all_history: list[dict[str, Any]] = []
model_summaries: list[dict[str, Any]] = []
for method, variant in variants:
metrics, history, summary = _run_trial(
method, variant, device=device, steps=args.steps, seed=args.seed, output_dir=output_dir
)
all_metrics.extend(metrics)
all_history.extend(history)
model_summaries.append({"method": method, "variant": variant, **summary})
_write_csv(output_dir / "metrics.csv", all_metrics)
_write_csv(output_dir / "training_history.csv", all_history)
manifest = {
"created_utc": datetime.now(timezone.utc).isoformat(),
"purpose": "Check whether the existing M3/M4 attention implementation can fit a time band when all inputs encode only synthetic time.",
"sample_count": 1,
"synthetic_only": True,
"seed": args.seed,
"device": str(device),
"gpu_name": torch.cuda.get_device_name(device) if device.type == "cuda" else None,
"python": platform.python_version(),
"torch": torch.__version__,
"grid_size": GRID_SIZE,
"sigma_normalized_time": SIGMA,
"feature_rule": "[t, t^2, 1] per source position; no extracted text/audio/video content is used",
"sequence_lengths": MODALITY_LENGTHS,
"optimizer": "AdamW",
"learning_rate": 1e-3,
"steps_per_variant": args.steps,
"variants": model_summaries,
"interpretation_limits": [
"This is an implementation/optimization control only; it does not evaluate real features or alignment accuracy.",
"The time-only features deliberately provide source-position information that the real feature inputs may not contain.",
],
}
(output_dir / "run_manifest.json").write_text(
json.dumps(manifest, ensure_ascii=False, indent=2, allow_nan=False), encoding="utf-8"
)
print(f"[all done] wrote synthetic control to {output_dir}", flush=True)
return manifest
def build_parser() -> argparse.ArgumentParser:
project = Path(__file__).resolve().parents[1]
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--steps", type=int, default=1000)
parser.add_argument("--seed", type=int, default=SEED)
parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto")
parser.add_argument(
"--output-dir", type=Path, default=project / "outputs/alignment_debug/synthetic"
)
return parser
def main() -> None:
args = build_parser().parse_args()
run(args)
if __name__ == "__main__":
main()
@@ -0,0 +1,647 @@
from __future__ import annotations
import argparse
import json
import math
import platform
import statistics
import time
from collections import defaultdict
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Mapping, Sequence
import numpy as np
import torch
from sklearn.model_selection import GroupKFold
from .compare_methods import (
_alignment_rows,
_batches,
_baseline_output,
_collect_representations,
_fit_learned_model,
_make_figures,
_save_alignment,
_seed_everything,
_write_csv,
)
from .experiment_data import (
FeatureSample,
collate_feature_samples,
fit_feature_stats,
load_feature_samples,
)
from .experiment_probes import (
run_shuffled_alignment_reconstruction_probe,
run_within_clip_temporal_retrieval_probe,
)
from .types import MODALITIES
VARIANTS = ("v2_a", "v2_b", "v2_c")
METRIC_COLUMNS = {
"alignment": (
"mvr",
"normalized_entropy",
"width80_source_positions",
"trajectory_span_fraction",
"c_row",
"c_far",
),
"retrieval": ("r_at_1", "r_at_3", "mase_slots", "exact_r_at_1"),
"reconstruction": (
"mae_aligned",
"mae_shuffled_mean",
"mae_shuffled_std",
"gain_align",
),
}
def _summarize(
rows: Sequence[Mapping[str, Any]], group_keys: Sequence[str], metrics: Sequence[str]
) -> list[dict[str, Any]]:
grouped: dict[tuple[Any, ...], list[Mapping[str, Any]]] = defaultdict(list)
for row in rows:
grouped[tuple(row[key] for key in group_keys)].append(row)
results = []
for key, values in grouped.items():
summary: dict[str, Any] = dict(zip(group_keys, key))
summary["n"] = len(values)
for metric in metrics:
numbers = [float(row[metric]) for row in values if row.get(metric) not in (None, "")]
numbers = [number for number in numbers if math.isfinite(number)]
if numbers:
summary[f"{metric}_mean"] = statistics.fmean(numbers)
summary[f"{metric}_std"] = statistics.stdev(numbers) if len(numbers) > 1 else 0.0
results.append(summary)
return results
def _comparison_table(summaries: Mapping[str, Sequence[Mapping[str, Any]]]) -> list[dict[str, Any]]:
alignment = {(row["method"], row["modality"]): row for row in summaries["alignment"]}
retrieval = {(row["method"], row["direction"]): row for row in summaries["retrieval"]}
reconstruction = {
(row["method"], row["target_modality"]): row for row in summaries["reconstruction"]
}
rows = []
for method in ("M1", "M2", "M3", "M4"):
row: dict[str, Any] = {"method": method}
for modality in MODALITIES:
metrics = alignment[(method, modality)]
for key in ("mvr", "normalized_entropy", "trajectory_span_fraction", "c_row", "c_far"):
row[f"{key}_{modality}"] = metrics.get(f"{key}_mean")
for direction in ("text_to_audio", "text_to_vision", "audio_to_vision"):
metrics = retrieval[(method, direction)]
for key in ("r_at_1", "r_at_3", "mase_slots", "exact_r_at_1"):
row[f"{key}_{direction}"] = metrics.get(f"{key}_mean")
for modality in MODALITIES:
metrics = reconstruction[(method, modality)]
for key in ("mae_aligned", "mae_shuffled_mean", "gain_align"):
row[f"{key}_{modality}"] = metrics.get(f"{key}_mean")
rows.append(row)
return rows
def _folds_from_baseline(
samples: Sequence[FeatureSample], baseline_splits: Path, requested_folds: int
) -> list[tuple[list[FeatureSample], list[FeatureSample], dict[str, Any]]]:
by_id = {sample.sample_id: sample for sample in samples}
if baseline_splits.is_file():
split_rows = json.loads(baseline_splits.read_text(encoding="utf-8"))
else:
groups = [sample.group_id for sample in samples]
splitter = GroupKFold(n_splits=requested_folds)
split_rows = []
for fold, (train, validation) in enumerate(
splitter.split(np.zeros(len(samples)), groups=groups), start=1
):
split_rows.append(
{
"fold": fold,
"train_sample_ids": [samples[index].sample_id for index in train],
"validation_sample_ids": [samples[index].sample_id for index in validation],
"train_video_ids": sorted({samples[index].group_id for index in train}),
"validation_video_ids": sorted(
{samples[index].group_id for index in validation}
),
}
)
if len(split_rows) != requested_folds:
raise ValueError(
f"baseline split file has {len(split_rows)} folds; expected {requested_folds}"
)
folds = []
seen_validation: list[str] = []
for row in split_rows:
train_ids = row["train_sample_ids"]
validation_ids = row["validation_sample_ids"]
if set(train_ids) & set(validation_ids):
raise ValueError(f"sample leakage in fold {row['fold']}")
train_samples = [by_id[sample_id] for sample_id in train_ids]
val_samples = [by_id[sample_id] for sample_id in validation_ids]
train_groups = {sample.group_id for sample in train_samples}
val_groups = {sample.group_id for sample in val_samples}
if train_groups & val_groups:
raise ValueError(f"video_id leakage in fold {row['fold']}")
seen_validation.extend(validation_ids)
folds.append((train_samples, val_samples, row))
if len(seen_validation) != len(samples) or set(seen_validation) != set(by_id):
raise ValueError("baseline folds do not cover the current complete sample manifest")
return folds
def _evaluate_fixed_methods(
args: argparse.Namespace,
samples: Sequence[FeatureSample],
folds: Sequence[tuple[list[FeatureSample], list[FeatureSample], dict[str, Any]]],
device: torch.device,
output_dir: Path,
) -> tuple[dict[str, list[dict[str, Any]]], dict[str, dict[str, Mapping[str, np.ndarray]]]]:
output_dir.mkdir(parents=True, exist_ok=True)
rows: dict[str, list[dict[str, Any]]] = {
"alignment": [],
"retrieval": [],
"reconstruction": [],
}
example_weights: dict[str, dict[str, Mapping[str, np.ndarray]]] = defaultdict(dict)
sample_by_id = {sample.sample_id: sample for sample in samples}
example_id = args.example_id if args.example_id in sample_by_id else samples[0].sample_id
for fold_index, (train_samples, val_samples, split_row) in enumerate(folds, start=1):
stats = fit_feature_stats(train_samples)
print(
f"[fixed fold {fold_index}/{len(folds)}] train={len(train_samples)} "
f"validation={len(val_samples)}; evaluating unchanged M1/M2",
flush=True,
)
for method in ("M1", "M2"):
train_aligned, _ = _collect_representations(
method,
train_samples,
stats,
device=device,
grid_size=args.grid_size,
batch_size=args.batch_size,
)
val_aligned, val_weights = _collect_representations(
method,
val_samples,
stats,
device=device,
grid_size=args.grid_size,
batch_size=args.batch_size,
)
combined = {**train_aligned, **val_aligned}
for batch_samples in _batches(
val_samples,
args.batch_size,
shuffle=False,
rng=np.random.default_rng(args.seeds[0] + fold_index),
):
sequences, durations, intervals = collate_feature_samples(
batch_samples, stats, device
)
output = _baseline_output(method, sequences, durations, intervals, args.grid_size)
rows["alignment"].extend(
_alignment_rows(
method,
"fixed",
fold_index,
batch_samples,
output,
device,
args.mvr_epsilon,
)
)
for sample in val_samples:
_save_alignment(output_dir, method, "fixed", sample, val_weights[sample.sample_id])
if sample.sample_id == example_id:
example_weights[sample.sample_id][method] = val_weights[sample.sample_id]
probe_seed = args.seeds[0] + fold_index * 100
rows["retrieval"].extend(
{"method": method, "seed": "fixed", "fold": fold_index, **row}
for row in run_within_clip_temporal_retrieval_probe(
[sample.sample_id for sample in train_samples],
[sample.sample_id for sample in val_samples],
combined,
device=device,
seed=probe_seed,
epochs=args.retrieval_probe_epochs,
batch_size=args.probe_batch_size,
tolerance=args.retrieval_tolerance,
top_k=3,
)
)
rows["reconstruction"].extend(
{"method": method, "seed": "fixed", "fold": fold_index, **row}
for row in run_shuffled_alignment_reconstruction_probe(
[sample.sample_id for sample in train_samples],
[sample.sample_id for sample in val_samples],
combined,
device=device,
seed=probe_seed + 1,
ratio=args.mask_ratio,
epochs=args.reconstruction_probe_epochs,
batch_size=args.probe_batch_size,
shuffle_repeats=args.shuffle_repeats,
)
)
for key, values in rows.items():
_write_csv(output_dir / f"{key}_metrics.csv", values)
(output_dir / "README.md").write_text(
"# Fixed M1/M2 reference probes\n\n"
"M1 and M2 are recomputed from the same saved video-grouped folds and unchanged. "
"The within-clip retrieval projection is fitted on each training fold; held-out candidates "
"come only from the same clip. The reconstruction decoder is trained on aligned training "
"representations, then compared with a control that shuffles the two non-target streams.\n",
encoding="utf-8",
)
return rows, example_weights
def _write_stage_summary(
stage_dir: Path,
fixed_rows: Mapping[str, Sequence[dict[str, Any]]],
learned_rows: Mapping[str, Sequence[dict[str, Any]]],
) -> dict[str, list[dict[str, Any]]]:
all_rows = {
key: [*fixed_rows[key], *learned_rows[key]]
for key in ("alignment", "retrieval", "reconstruction")
}
group_columns = {
"alignment": ("method", "modality"),
"retrieval": ("method", "direction"),
"reconstruction": ("method", "target_modality"),
}
summaries = {
key: _summarize(rows, group_columns[key], METRIC_COLUMNS[key])
for key, rows in all_rows.items()
}
for key, rows in all_rows.items():
_write_csv(stage_dir / f"{key}_metrics_with_fixed.csv", rows)
_write_csv(stage_dir / f"{key}_summary.csv", summaries[key])
_write_csv(stage_dir / "comparison_summary.csv", _comparison_table(summaries))
(stage_dir / "summary.json").write_text(
json.dumps(summaries, ensure_ascii=False, indent=2, allow_nan=False), encoding="utf-8"
)
return summaries
def _run_variant(
variant: str,
args: argparse.Namespace,
samples: Sequence[FeatureSample],
folds: Sequence[tuple[list[FeatureSample], list[FeatureSample], dict[str, Any]]],
fixed_rows: Mapping[str, Sequence[dict[str, Any]]],
fixed_examples: Mapping[str, Mapping[str, Mapping[str, np.ndarray]]],
device: torch.device,
) -> dict[str, Any]:
start_time = time.time()
stage_dir = args.output_dir / variant
stage_dir.mkdir(parents=True, exist_ok=True)
(stage_dir / "splits.json").write_text(
json.dumps([row for _, _, row in folds], ensure_ascii=False, indent=2), encoding="utf-8"
)
learned_rows: dict[str, list[dict[str, Any]]] = {
"alignment": [],
"retrieval": [],
"reconstruction": [],
}
training_summary: list[dict[str, Any]] = []
training_history: list[dict[str, Any]] = []
example_weights: dict[str, dict[str, Mapping[str, np.ndarray]]] = defaultdict(dict)
sample_by_id = {sample.sample_id: sample for sample in samples}
example_id = args.example_id if args.example_id in sample_by_id else samples[0].sample_id
example_sample = sample_by_id[example_id]
if example_id in fixed_examples:
example_weights[example_id].update(fixed_examples[example_id])
for fold_index, (train_samples, val_samples, _) in enumerate(folds, start=1):
stats = fit_feature_stats(train_samples)
fold_dir = stage_dir / f"fold_{fold_index:02d}"
print(
f"[{variant} fold {fold_index}/{len(folds)}] training M3/M4 with "
f"{len(train_samples)} train and {len(val_samples)} validation clips",
flush=True,
)
for seed in args.seeds:
for method in ("M3", "M4"):
checkpoint = fold_dir / f"seed_{seed}" / f"{method}.pt"
model, info = _fit_learned_model(
method,
train_samples,
val_samples,
stats,
device=device,
seed=seed + fold_index * 1009,
grid_size=args.grid_size,
hidden_size=args.hidden_size,
heads=args.heads,
dropout=args.dropout,
batch_size=args.batch_size,
max_epochs=args.epochs,
patience=args.patience,
learning_rate=args.learning_rate,
checkpoint_path=checkpoint,
loss_variant=variant,
)
training_summary.append(
{
"loss_variant": variant,
"method": method,
"seed": seed,
"fold": fold_index,
"best_epoch": info["best_epoch"],
"best_validation_objective": info["best_validation_objective"],
**{
key: value
for key, value in info["best_validation_metrics"].items()
if key not in {"epoch", "train_total"}
},
"checkpoint": info["checkpoint"],
}
)
training_history.extend(
{
"loss_variant": variant,
"method": method,
"seed": seed,
"fold": fold_index,
**epoch,
}
for epoch in info["history"]
)
print(
f"[{variant} fold {fold_index}] {method} seed={seed} "
f"best_epoch={info['best_epoch']} "
f"val={info['best_validation_objective']:.4f} "
f"C_row(audio/vision)="
f"{info['best_validation_metrics']['validation_c_row_audio']:.3f}/"
f"{info['best_validation_metrics']['validation_c_row_vision']:.3f}",
flush=True,
)
train_aligned, _ = _collect_representations(
method,
train_samples,
stats,
device=device,
grid_size=args.grid_size,
batch_size=args.batch_size,
model=model,
)
val_aligned, val_weights = _collect_representations(
method,
val_samples,
stats,
device=device,
grid_size=args.grid_size,
batch_size=args.batch_size,
model=model,
)
combined = {**train_aligned, **val_aligned}
for batch_samples in _batches(
val_samples,
args.batch_size,
shuffle=False,
rng=np.random.default_rng(seed + fold_index),
):
sequences, durations, _ = collate_feature_samples(
batch_samples, stats, device
)
output = model(sequences)
learned_rows["alignment"].extend(
_alignment_rows(
method,
str(seed),
fold_index,
batch_samples,
output,
device,
args.mvr_epsilon,
)
)
for sample in val_samples:
_save_alignment(stage_dir, method, str(seed), sample, val_weights[sample.sample_id])
if sample.sample_id == example_id and seed == args.seeds[0]:
example_weights[sample.sample_id][method] = val_weights[sample.sample_id]
probe_seed = seed + fold_index * 100 + (3 if method == "M3" else 7)
learned_rows["retrieval"].extend(
{
"method": method,
"seed": seed,
"fold": fold_index,
**row,
}
for row in run_within_clip_temporal_retrieval_probe(
[sample.sample_id for sample in train_samples],
[sample.sample_id for sample in val_samples],
combined,
device=device,
seed=probe_seed,
epochs=args.retrieval_probe_epochs,
batch_size=args.probe_batch_size,
tolerance=args.retrieval_tolerance,
top_k=3,
)
)
learned_rows["reconstruction"].extend(
{
"method": method,
"seed": seed,
"fold": fold_index,
**row,
}
for row in run_shuffled_alignment_reconstruction_probe(
[sample.sample_id for sample in train_samples],
[sample.sample_id for sample in val_samples],
combined,
device=device,
seed=probe_seed + 1,
ratio=args.mask_ratio,
epochs=args.reconstruction_probe_epochs,
batch_size=args.probe_batch_size,
shuffle_repeats=args.shuffle_repeats,
)
)
del model
if device.type == "cuda":
torch.cuda.empty_cache()
_write_csv(stage_dir / "training_summary.csv", training_summary)
_write_csv(stage_dir / "training_history.csv", training_history)
for key, rows in learned_rows.items():
_write_csv(stage_dir / f"{key}_metrics_learned_only.csv", rows)
summaries = _write_stage_summary(stage_dir, fixed_rows, learned_rows)
_write_csv(stage_dir / "training_summary.csv", training_summary)
_write_csv(stage_dir / "training_history.csv", training_history)
if example_id in example_weights and set(example_weights[example_id]) == {"M1", "M2", "M3", "M4"}:
_make_figures(stage_dir, example_id, example_sample, example_weights[example_id], args.grid_size)
manifest = {
"variant": variant,
"loss_coefficients": {
"lambda_reconstruction": 1.0,
"lambda_contrastive": 1.0,
"lambda_monotonicity": 0.1,
"lambda_span": 5.0,
"lambda_diversity": 0.5 if variant in {"v2_b", "v2_c"} else 0.0,
"lambda_band": 10.0 if variant == "v2_c" else 0.0,
"coverage_floor": 0.7,
"diversity_slot_separation": 6,
"band_margin": 0.1,
},
"sample_count": len(samples),
"video_group_folds": len(folds),
"seeds": args.seeds,
"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() - start_time,
"probes": {
"within_clip_retrieval_tolerance_slots": args.retrieval_tolerance,
"retrieval_top_k": 3,
"shuffled_reconstruction_repeats": args.shuffle_repeats,
"masked_block_ratio": args.mask_ratio,
},
"limits": [
"Retrieval projections are fitted on training-fold grid-slot positives; test candidates are restricted to the same held-out clip.",
"The reconstruction control shuffles the two non-target modality slot streams and preserves the target stream.",
"No human event timestamps are available, so temporal probes do not replace manual annotation.",
"Slot-regularization losses impose weak temporal structure and must be interpreted alongside the unregularized M1/M2 reference.",
],
}
(stage_dir / "run_manifest.json").write_text(
json.dumps(manifest, ensure_ascii=False, indent=2, allow_nan=False), encoding="utf-8"
)
(stage_dir / "README.md").write_text(
f"# {variant} M3/M4 alignment variant\n\n"
"M1/M2 rows in the comparison files are fixed references recomputed on the original grouped folds. "
"Only M3/M4 training losses changed. See `training_history.csv` for per-epoch loss components and row-collapse scores. "
"`retrieval_metrics_learned_only.csv` restricts candidates to the same clip. "
"`reconstruction_metrics_learned_only.csv` contrasts aligned and shuffled non-target streams. "
"The experiment does not include RoPE, Gaussian bias, or latent-length changes.\n",
encoding="utf-8",
)
print(
f"[{variant} done] elapsed={manifest['elapsed_seconds']:.1f}s "
f"output={stage_dir}; summary rows={sum(len(rows) for rows in summaries.values())}",
flush=True,
)
return manifest
def run(args: argparse.Namespace) -> dict[str, Any]:
start_time = time.time()
_seed_everything(args.seeds[0])
if args.device == "auto":
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
else:
device = torch.device(args.device)
if device.type == "cuda" and not torch.cuda.is_available():
raise RuntimeError("CUDA was requested but is not available")
samples = load_feature_samples(args.feature_dir, args.manifest)
if len(samples) != args.expected_samples:
raise ValueError(f"expected {args.expected_samples} samples, found {len(samples)}")
folds = _folds_from_baseline(
samples, args.baseline_dir / "splits.json", args.folds
)
args.output_dir.mkdir(parents=True, exist_ok=True)
print(
f"[start] samples={len(samples)} folds={len(folds)} seeds={args.seeds} "
f"device={device} variants={','.join(VARIANTS)}",
flush=True,
)
fixed_rows, fixed_examples = _evaluate_fixed_methods(
args,
samples,
folds,
device,
args.output_dir / "fixed_baselines",
)
stage_manifests = []
for variant in VARIANTS:
stage_manifests.append(
_run_variant(
variant,
args,
samples,
folds,
fixed_rows,
fixed_examples,
device,
)
)
manifest = {
"created_utc": datetime.now(timezone.utc).isoformat(),
"sample_count": len(samples),
"group_count": len({sample.group_id for sample in samples}),
"folds": len(folds),
"seeds": args.seeds,
"device": str(device),
"gpu_name": torch.cuda.get_device_name(device) if device.type == "cuda" else None,
"python": platform.python_version(),
"torch": torch.__version__,
"feature_dir": str(args.feature_dir),
"baseline_dir": str(args.baseline_dir),
"variants": stage_manifests,
"elapsed_seconds": time.time() - start_time,
"fixed_methods_unchanged": ["M1", "M2"],
"feature_extraction_changed": False,
}
(args.output_dir / "run_manifest.json").write_text(
json.dumps(manifest, ensure_ascii=False, indent=2, allow_nan=False), encoding="utf-8"
)
print(
f"[all done] elapsed={manifest['elapsed_seconds']:.1f}s output={args.output_dir}",
flush=True,
)
return manifest
def build_parser() -> argparse.ArgumentParser:
project_dir = Path(__file__).resolve().parents[1]
parser = argparse.ArgumentParser(
description="Train staged M3/M4 alignment-loss variants and stronger temporal probes."
)
parser.add_argument("--feature-dir", type=Path, default=project_dir / "outputs/q1_features/features")
parser.add_argument("--manifest", type=Path, default=project_dir / "outputs/audit/manifest.csv")
parser.add_argument("--baseline-dir", type=Path, default=project_dir / "outputs/method_comparison")
parser.add_argument("--output-dir", type=Path, default=project_dir / "outputs/alignment_v2")
parser.add_argument("--device", default="auto", help="auto, cpu, or a torch device such as cuda:0")
parser.add_argument("--expected-samples", type=int, default=100)
parser.add_argument("--folds", type=int, default=5)
parser.add_argument("--seeds", type=int, nargs="+", default=[42, 3407, 2026])
parser.add_argument("--grid-size", type=int, default=50)
parser.add_argument("--hidden-size", type=int, default=128)
parser.add_argument("--heads", type=int, default=4)
parser.add_argument("--dropout", type=float, default=0.1)
parser.add_argument("--batch-size", type=int, default=8)
parser.add_argument("--probe-batch-size", type=int, default=16)
parser.add_argument("--epochs", type=int, default=50)
parser.add_argument("--patience", type=int, default=8)
parser.add_argument("--learning-rate", type=float, default=1e-4)
parser.add_argument("--mask-ratio", type=float, default=0.2)
parser.add_argument("--retrieval-probe-epochs", type=int, default=20)
parser.add_argument("--reconstruction-probe-epochs", type=int, default=25)
parser.add_argument("--retrieval-tolerance", type=int, default=1)
parser.add_argument("--shuffle-repeats", type=int, default=5)
parser.add_argument("--mvr-epsilon", type=float, default=0.02)
parser.add_argument("--example-id", default="-tPCytz4rww/12")
return parser
def main() -> int:
args = build_parser().parse_args()
manifest = run(args)
print(json.dumps(manifest, ensure_ascii=False, indent=2))
return 0
if __name__ == "__main__":
raise SystemExit(main())
@@ -0,0 +1,353 @@
"""Evaluate TSFA attention maps using identical raw source content for every method.
This probe pools the same train-fold-standardized BERT, audio, and DeiT features
with each method's alignment matrix. It excludes native model value/output
projections from the representation being scored.
"""
from __future__ import annotations
import csv
import json
import shutil
import time
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Mapping, Sequence
import numpy as np
import torch
from .correspondence_eval import _cluster_bootstrap, _write_csv
from .experiment_data import (
FeatureSample,
collate_feature_samples,
fit_feature_stats,
load_feature_samples,
)
from .m4_shared_latent_eval import _bootstrap_summary
from .tsfa_experiment import (
ALL_METHODS,
BASELINE_VARIANTS,
GRID_SIZE,
TSFA_VARIANTS,
_ablation_summary,
_collect_fold_features,
_content_summary,
_evaluate_fixed_projector,
_fit_method_probe,
_generate_tsfa_outputs,
_load_semantic_checkpoint,
build_parser,
)
from .types import MODALITIES
def _pool_same_raw_features(
samples: Sequence[FeatureSample],
feature_stats: Any,
weights_by_id: Mapping[str, Mapping[str, np.ndarray]],
device: torch.device,
batch_size: int,
) -> dict[str, dict[str, np.ndarray]]:
content: dict[str, dict[str, np.ndarray]] = {}
with torch.no_grad():
for start in range(0, len(samples), batch_size):
batch_samples = list(samples[start : start + batch_size])
sequences, _, _ = collate_feature_samples(batch_samples, feature_stats, device)
for index, sample in enumerate(batch_samples):
content[sample.sample_id] = {}
for modality in MODALITIES:
length = len(sample.features[modality])
weights = torch.as_tensor(
weights_by_id[sample.sample_id][modality],
dtype=torch.float32,
device=device,
)
if weights.shape != (GRID_SIZE, length):
raise ValueError(f"unexpected alignment shape for {sample.sample_id}/{modality}")
source = sequences[modality].features[index, :length]
content[sample.sample_id][modality] = (
(weights @ source).cpu().numpy().astype(np.float32, copy=False)
)
return content
def _paired_summary(rows: Sequence[Mapping[str, Any]], seed: int) -> list[dict[str, Any]]:
by_key = {(row["method"], row["sample_id"]): row for row in rows}
sample_ids = sorted({row["sample_id"] for row in rows})
contrasts = (
("TSFA-main", "M4_sourceTime"),
("TSFA-main", "TSFA-random"),
("TSFA-main", "TSFA-global"),
("TSFA-multiply", "TSFA-main"),
("TSFA-main", "M3_noSourceTime"),
)
output = []
for contrast_index, (left, right) in enumerate(contrasts):
for metric_index, metric in enumerate((
"content_auc_mean", "canonical_pairwise_time_mae_mean"
)):
differences = []
for sample_id in sample_ids:
left_row = by_key[(left, sample_id)]
right_row = by_key[(right, sample_id)]
if left_row["video_id"] != right_row["video_id"]:
raise ValueError(f"video group mismatch for {sample_id}")
differences.append({
"video_id": left_row["video_id"],
"difference": float(left_row[metric]) - float(right_row[metric]),
})
mean, low, high, groups = _cluster_bootstrap(
differences,
"difference",
seed=seed + contrast_index * 101 + metric_index,
repetitions=2000,
)
output.append({
"left_method": left,
"right_method": right,
"metric": metric,
"clip_count": len(differences),
"video_id_count": groups,
"left_minus_right_video_macro_mean": mean,
"ci95_low": low,
"ci95_high": high,
})
return output
def _temporal_diagnostic_rows(
method: str,
fold: int,
sample: FeatureSample,
weights: Mapping[str, np.ndarray],
draw: int = 0,
) -> list[dict[str, Any]]:
output = []
for modality in MODALITIES:
valid = np.asarray(sample.valid[modality], dtype=bool)
attention = np.asarray(weights[modality], dtype=np.float64)[:, valid]
times = np.asarray(sample.times[modality], dtype=np.float64)[valid] / max(sample.duration_s, 1e-8)
centers = attention @ times
backward = centers[:-1] - centers[1:]
entropy = -(attention * np.log(np.maximum(attention, 1e-12))).sum(axis=1)
output.append({
"method": method,
"fold": fold,
"sample_id": sample.sample_id,
"video_id": sample.group_id,
"draw": draw,
"modality": modality,
"mvr_epsilon_0_01": float(np.mean(backward > 0.01)),
"mvr_epsilon_0_02": float(np.mean(backward > 0.02)),
"mvr_epsilon_0_05": float(np.mean(backward > 0.05)),
"time_span_ratio": float(centers[-1] - centers[0]),
"normalized_entropy_mean": float(entropy.mean() / max(np.log(attention.shape[1]), 1e-12)),
"source_coverage_rate": float(np.mean(attention.sum(axis=0) > 1e-12)),
})
return output
def run(args: Any) -> None:
started = time.time()
device = torch.device(
("cuda" if torch.cuda.is_available() else "cpu") if args.device == "auto" else args.device
)
if device.type == "cuda" and not torch.cuda.is_available():
raise RuntimeError("CUDA 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("expected five grouped folds")
store = torch.load(args.output_dir / "probe_checkpoints.pt", map_location="cpu", weights_only=False)
content_rows: list[dict[str, Any]] = []
curve_rows: list[dict[str, Any]] = []
temporal_diagnostic_rows: list[dict[str, Any]] = []
history_rows: list[dict[str, Any]] = []
probe_store: dict[str, Any] = {}
heldout_ids = []
for split in splits:
fold = int(split["fold"])
train = [samples_by_id[sample_id] for sample_id in split["train_sample_ids"]]
validation = [samples_by_id[sample_id] for sample_id in split["validation_sample_ids"]]
if {sample.group_id for sample in train} & {sample.group_id for sample in validation}:
raise ValueError(f"video_id leakage in fold {fold}")
heldout_ids.extend(sample.sample_id for sample in validation)
feature_stats = fit_feature_stats(train)
_, baseline_weights, temporal_by_id = _collect_fold_features(
fold=fold,
train_samples=train,
validation_samples=validation,
feature_stats=feature_stats,
checkpoint_root=args.checkpoint_root,
device=device,
batch_size=args.batch_size,
)
branch = _load_semantic_checkpoint(store, fold, device)
all_samples = [*train, *validation]
all_ids = [sample.sample_id for sample in all_samples]
fold_weights = {method: baseline_weights[method] for method in BASELINE_VARIANTS}
for method in TSFA_VARIANTS:
_, weights, _ = _generate_tsfa_outputs(
method=method,
fold=fold,
sample_ids=all_ids,
samples_by_id=samples_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_weights[method] = weights
for method in ALL_METHODS:
pooled = _pool_same_raw_features(
all_samples, feature_stats, fold_weights[method], device, args.batch_size
)
metrics, curves, _ = _fit_method_probe(
method=method,
fold=fold,
train_samples=train,
validation_samples=validation,
content_by_id=pooled,
device=device,
args=args,
history_rows=history_rows,
checkpoint_store=probe_store,
)
content_rows.extend(metrics)
curve_rows.extend(curves)
for sample in validation:
temporal_diagnostic_rows.extend(_temporal_diagnostic_rows(
method, fold, sample, fold_weights[method][sample.sample_id]
))
if method == "TSFA-random":
for draw in range(1, args.random_window_repeats):
_, random_weights, _ = _generate_tsfa_outputs(
method=method,
fold=fold,
sample_ids=[sample.sample_id for sample in validation],
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,
)
random_pooled = _pool_same_raw_features(
validation, feature_stats, random_weights, device, args.batch_size
)
repeated_metrics, repeated_curves, _ = _evaluate_fixed_projector(
method=method,
fold=fold,
validation_samples=validation,
content_by_id=random_pooled,
projector_state=probe_store[
f"fold_{fold:02d}/{method}/correspondence_probe"
]["state_dict"],
device=device,
)
content_rows.extend({**row, "draw": draw} for row in repeated_metrics)
curve_rows.extend({**row, "draw": draw} for row in repeated_curves)
for sample in validation:
temporal_diagnostic_rows.extend(_temporal_diagnostic_rows(
method, fold, sample, random_weights[sample.sample_id], draw
))
print(f"[TSFA alignment-only fold {fold}] heldout={len(validation)}", flush=True)
del branch, baseline_weights, temporal_by_id, fold_weights
if device.type == "cuda":
torch.cuda.empty_cache()
if len(heldout_ids) != 100 or len(set(heldout_ids)) != 100:
raise ValueError("held-out fold coverage is not exactly 100 distinct samples")
output_dir = args.output_dir
content_summary, curve_summary = _content_summary(content_rows, curve_rows, args.seed + 901)
with (output_dir / "temporal_metrics_by_clip.csv").open(
newline="", encoding="utf-8-sig"
) as handle:
temporal_rows = list(csv.DictReader(handle))
ablation_summary, ablation_by_clip = _ablation_summary(
content_rows, temporal_rows, args.seed + 902
)
paired = _paired_summary(ablation_by_clip, args.seed + 903)
temporal_diagnostic_summary = _bootstrap_summary(
temporal_diagnostic_rows,
("method", "modality"),
(
"mvr_epsilon_0_01", "mvr_epsilon_0_02", "mvr_epsilon_0_05",
"time_span_ratio", "normalized_entropy_mean", "source_coverage_rate",
),
seed=args.seed + 904,
)
_write_csv(output_dir / "alignment_only_content_by_clip.csv", content_rows)
_write_csv(output_dir / "alignment_only_content_summary.csv", content_summary)
_write_csv(output_dir / "alignment_only_shift_curve_summary.csv", curve_summary)
_write_csv(output_dir / "alignment_only_ablation_by_clip.csv", ablation_by_clip)
_write_csv(output_dir / "alignment_only_ablation_summary.csv", ablation_summary)
_write_csv(output_dir / "alignment_only_paired_contrasts.csv", paired)
_write_csv(output_dir / "alignment_only_probe_training_history.csv", history_rows)
_write_csv(output_dir / "tsfa_temporal_diagnostics_by_clip.csv", temporal_diagnostic_rows)
_write_csv(output_dir / "tsfa_temporal_diagnostics_summary.csv", temporal_diagnostic_summary)
torch.save(probe_store, output_dir / "alignment_only_probe_checkpoints.pt")
run_manifest = {
"created_utc": datetime.now(timezone.utc).isoformat(),
"experiment": "TSFA shared raw-content alignment-only probe",
"sample_count": len(samples),
"heldout_count": len(heldout_ids),
"video_id_count": len({sample.group_id for sample in samples}),
"fold_count": len(splits),
"seed": args.seed,
"delta": args.delta,
"probe_epochs": args.probe_epochs,
"random_window_repeats": args.random_window_repeats,
"feature_protocol": "Within each fold, feature normalization is fitted on 80 training clips. For every method and modality, the same normalized raw source feature matrix is pooled using that method's A^m; native M3/M4/TSFA value and output projections are excluded from scored features.",
"probe_protocol": "One 64-dimensional linear projector per modality, trained on 80 training clips with the same within-clip InfoNCE protocol and identical fold seed across methods. Scores use only 20 held-out clips per fold.",
"interpretation_limit": "A same-slot positive is a timestamp/slot convention, not independently annotated semantic ground truth. Source-time attention can still encode time through selected raw values.",
"device": str(device),
"elapsed_seconds": time.time() - started,
}
(output_dir / "alignment_only_run_manifest.json").write_text(
json.dumps(run_manifest, ensure_ascii=False, indent=2), encoding="utf-8"
)
bundle = output_dir / "report_bundle"
for name in (
"alignment_only_content_summary.csv", "alignment_only_ablation_summary.csv",
"alignment_only_paired_contrasts.csv", "alignment_only_run_manifest.json",
"tsfa_temporal_diagnostics_summary.csv",
):
shutil.copy2(output_dir / name, bundle / name)
bundle_readme = bundle / "README.md"
note = (
"\nThe `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.\n"
)
current = bundle_readme.read_text(encoding="utf-8")
if "The `alignment_only_*` summaries" not in current:
bundle_readme.write_text(current + note, encoding="utf-8")
print(
f"[TSFA alignment-only complete] samples={len(heldout_ids)} "
f"elapsed={run_manifest['elapsed_seconds']:.1f}s output={output_dir}",
flush=True,
)
def main() -> None:
args = build_parser().parse_args()
if args.finalize_existing:
raise ValueError("--finalize-existing belongs to q1.tsfa_experiment")
if not 0 < args.delta <= 1:
raise ValueError("--delta must be in (0,1]")
run(args)
if __name__ == "__main__":
main()
File diff suppressed because it is too large. Load diff
+86
View File
@@ -0,0 +1,86 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Mapping
import torch
from torch import Tensor
MODALITIES = ("text", "audio", "vision")
@dataclass
class SequenceBatch:
"""A padded batch of timed feature sequences.
``features`` has shape ``[batch, length, dimension]``; ``times`` and
``valid`` have shape ``[batch, length]``. Times are seconds from the start
of each clip. Padding positions must be false in ``valid``.
"""
features: Tensor
times: Tensor
valid: Tensor
def __post_init__(self) -> None:
if self.features.ndim != 3:
raise ValueError("features must have shape [batch, length, dimension]")
expected = self.features.shape[:2]
if tuple(self.times.shape) != tuple(expected):
raise ValueError("times must match the batch and sequence dimensions")
if tuple(self.valid.shape) != tuple(expected):
raise ValueError("valid must match the batch and sequence dimensions")
if self.valid.dtype != torch.bool:
raise TypeError("valid must be a boolean tensor")
if self.features.device != self.times.device or self.features.device != self.valid.device:
raise ValueError("features, times, and valid must be on the same device")
if not bool(self.valid.any(dim=1).all()):
raise ValueError("every sample must contain at least one valid feature")
if not bool(torch.isfinite(self.features[self.valid]).all()):
raise ValueError("valid features must be finite")
if not bool(torch.isfinite(self.times[self.valid]).all()):
raise ValueError("valid timestamps must be finite")
@dataclass
class AlignmentOutput:
"""Unified alignment result for Text, Audio, and Vision.
Each ``weights[m]`` is a row-stochastic matrix with shape ``[B, K, L_m]``.
``aligned[m]`` is the resulting common-grid representation ``[B, K, D_m]``.
"""
weights: Mapping[str, Tensor]
aligned: Mapping[str, Tensor]
fallback_rows: Mapping[str, int] | None = None
def validate(self, valid: Mapping[str, Tensor] | None = None, atol: float = 1e-4) -> None:
if set(self.weights) != set(MODALITIES) or set(self.aligned) != set(MODALITIES):
raise ValueError(f"weights and aligned must contain exactly {MODALITIES}")
batch_size: int | None = None
grid_size: int | None = None
for modality in MODALITIES:
matrix = self.weights[modality]
values = self.aligned[modality]
if matrix.ndim != 3 or values.ndim != 3:
raise ValueError(f"{modality} alignment and values must be rank 3")
if matrix.shape[:2] != values.shape[:2]:
raise ValueError(f"{modality} weights and aligned values disagree on [B, K]")
if batch_size is None:
batch_size, grid_size = matrix.shape[:2]
elif matrix.shape[0] != batch_size or matrix.shape[1] != grid_size:
raise ValueError("all modalities must use the same batch and grid sizes")
if not bool(torch.isfinite(matrix).all()) or bool((matrix < -atol).any()):
raise ValueError(f"{modality} weights must be finite and non-negative")
if not torch.allclose(
matrix.sum(dim=-1),
torch.ones_like(matrix.sum(dim=-1)),
atol=atol,
rtol=0,
):
raise ValueError(f"{modality} alignment rows must sum to one")
if valid is not None:
allowed = valid[modality][:, None, :]
if bool((matrix.masked_select(~allowed.expand_as(matrix)).abs() > atol).any()):
raise ValueError(f"{modality} alignment assigns weight to padding")