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())