946 lines
39 KiB
Python
946 lines
39 KiB
Python
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())
|