建立分批同步基线(基础文件)
This commit is contained in:
commit
7fc76aaafd
70 files changed
+18635
No files matched your search
@@ -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",
|
||||
]
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -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
|
||||
@@ -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()
|
||||
@@ -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())
|
||||
@@ -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())
|
||||
@@ -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()
|
||||
@@ -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
|
||||
@@ -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
@@ -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
@@ -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
|
||||
@@ -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
@@ -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")
|
||||
Reference in new issue
Block a user