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

946 lines
39 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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())