建立分批同步基线(基础文件)
This commit is contained in:
@@ -0,0 +1,945 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import csv
|
||||
import json
|
||||
import math
|
||||
import platform
|
||||
import random
|
||||
import statistics
|
||||
import time
|
||||
from collections import defaultdict
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Any, Mapping, Sequence
|
||||
|
||||
import matplotlib
|
||||
|
||||
matplotlib.use("Agg")
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
import torch
|
||||
from sklearn.model_selection import GroupKFold
|
||||
from torch import Tensor, nn
|
||||
|
||||
from .alignment import align_fixed_windows, align_forced_timestamps, make_block_mask
|
||||
from .experiment_data import (
|
||||
FeatureSample,
|
||||
FeatureStats,
|
||||
collate_feature_samples,
|
||||
fit_feature_stats,
|
||||
load_feature_samples,
|
||||
)
|
||||
from .experiment_probes import (
|
||||
RetrievalProjection,
|
||||
run_frozen_emotion_probe,
|
||||
run_reconstruction_probe,
|
||||
run_retrieval_probe,
|
||||
)
|
||||
from .metrics import (
|
||||
attention_row_similarity,
|
||||
alignment_trajectory,
|
||||
attention_width80,
|
||||
monotonicity_violation_rate,
|
||||
normalized_attention_entropy,
|
||||
)
|
||||
from .models import SharedLatentTimeline, TextAnchoredCrossAttention
|
||||
from .types import AlignmentOutput, MODALITIES
|
||||
|
||||
|
||||
class AlignmentReconstructor(nn.Module):
|
||||
"""Shared M3/M4 training decoder: predict one stream from the other two."""
|
||||
|
||||
def __init__(self, hidden_size: int = 128, dropout: float = 0.1) -> None:
|
||||
super().__init__()
|
||||
self.decoders = nn.ModuleDict(
|
||||
{
|
||||
target: nn.Sequential(
|
||||
nn.Linear(hidden_size * 2, hidden_size),
|
||||
nn.GELU(),
|
||||
nn.Dropout(dropout),
|
||||
nn.Linear(hidden_size, hidden_size),
|
||||
)
|
||||
for target in MODALITIES
|
||||
}
|
||||
)
|
||||
|
||||
def forward(self, target: str, sources: Mapping[str, Tensor], mask: Tensor) -> Tensor:
|
||||
values = [
|
||||
sources[name].masked_fill(mask.unsqueeze(-1), 0.0)
|
||||
for name in MODALITIES
|
||||
if name != target
|
||||
]
|
||||
return self.decoders[target](torch.cat(values, dim=-1))
|
||||
|
||||
|
||||
def _seed_everything(seed: int) -> None:
|
||||
random.seed(seed)
|
||||
np.random.seed(seed)
|
||||
torch.manual_seed(seed)
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.manual_seed_all(seed)
|
||||
if hasattr(torch.backends, "cudnn"):
|
||||
torch.backends.cudnn.deterministic = True
|
||||
torch.backends.cudnn.benchmark = False
|
||||
|
||||
|
||||
def _batches(
|
||||
samples: Sequence[FeatureSample], batch_size: int, *, shuffle: bool, rng: np.random.Generator
|
||||
) -> list[list[FeatureSample]]:
|
||||
if shuffle:
|
||||
order = rng.permutation(len(samples)).tolist()
|
||||
else:
|
||||
order = list(range(len(samples)))
|
||||
return [[samples[index] for index in order[start : start + batch_size]]
|
||||
for start in range(0, len(order), batch_size)]
|
||||
|
||||
|
||||
def _training_objective(
|
||||
output: AlignmentOutput,
|
||||
sequences: Mapping[str, Any],
|
||||
durations: Tensor,
|
||||
decoder: AlignmentReconstructor,
|
||||
generator: torch.Generator,
|
||||
*,
|
||||
method: str = "M3",
|
||||
loss_variant: str = "v1",
|
||||
) -> tuple[Tensor, dict[str, Tensor]]:
|
||||
reconstruction_terms = []
|
||||
batch_size, grid_size = output.aligned["text"].shape[:2]
|
||||
for target in MODALITIES:
|
||||
mask = make_block_mask(
|
||||
batch_size,
|
||||
grid_size,
|
||||
0.2,
|
||||
output.aligned[target].device,
|
||||
generator=generator,
|
||||
)
|
||||
prediction = decoder(target, output.aligned, mask)
|
||||
reconstruction_terms.append(
|
||||
nn.functional.smooth_l1_loss(prediction[mask], output.aligned[target][mask])
|
||||
)
|
||||
reconstruction = torch.stack(reconstruction_terms).mean()
|
||||
times = {name: sequences[name].times for name in MODALITIES}
|
||||
variant_weights = {
|
||||
"v1": (0.0, 0.0, 0.0),
|
||||
"v2_a": (5.0, 0.0, 0.0),
|
||||
"v2_b": (5.0, 0.5, 0.0),
|
||||
"v2_c": (5.0, 0.5, 10.0),
|
||||
}
|
||||
if loss_variant not in variant_weights:
|
||||
raise ValueError(f"unknown loss variant: {loss_variant}")
|
||||
lambda_span, lambda_div, lambda_band = variant_weights[loss_variant]
|
||||
grid_size = output.weights["text"].shape[1]
|
||||
if method == "M3":
|
||||
text_reference = torch.bmm(
|
||||
output.weights["text"], sequences["text"].times.unsqueeze(-1)
|
||||
).squeeze(-1)
|
||||
text_reference = text_reference / durations[:, None].clamp_min(1e-8)
|
||||
band_targets = {"audio": text_reference, "vision": text_reference}
|
||||
diversity_modalities = ("audio", "vision")
|
||||
else:
|
||||
centers = (
|
||||
torch.arange(grid_size, dtype=durations.dtype, device=durations.device) + 0.5
|
||||
) / grid_size
|
||||
reference = centers.unsqueeze(0).expand(durations.shape[0], -1)
|
||||
band_targets = {name: reference for name in MODALITIES}
|
||||
diversity_modalities = MODALITIES
|
||||
from .losses import alignment_training_loss
|
||||
|
||||
return alignment_training_loss(
|
||||
output,
|
||||
times,
|
||||
durations,
|
||||
reconstruction,
|
||||
lambda_rec=1.0,
|
||||
lambda_con=1.0,
|
||||
lambda_mono=0.1,
|
||||
lambda_span=lambda_span,
|
||||
lambda_div=lambda_div,
|
||||
lambda_band=lambda_band,
|
||||
epsilon=0.02,
|
||||
minimum_span=0.7,
|
||||
coverage_modalities=diversity_modalities,
|
||||
diversity_modalities=diversity_modalities,
|
||||
diversity_min_separation=6,
|
||||
band_targets=band_targets,
|
||||
band_margin=0.1,
|
||||
)
|
||||
|
||||
|
||||
def _make_learned_model(
|
||||
method: str,
|
||||
dimensions: Mapping[str, int],
|
||||
grid_size: int,
|
||||
hidden_size: int,
|
||||
heads: int,
|
||||
dropout: float,
|
||||
) -> nn.Module:
|
||||
if method == "M3":
|
||||
return TextAnchoredCrossAttention(
|
||||
dimensions,
|
||||
grid_size=grid_size,
|
||||
hidden_size=hidden_size,
|
||||
heads=heads,
|
||||
dropout=dropout,
|
||||
)
|
||||
if method == "M4":
|
||||
return SharedLatentTimeline(
|
||||
dimensions,
|
||||
grid_size=grid_size,
|
||||
hidden_size=hidden_size,
|
||||
heads=heads,
|
||||
dropout=dropout,
|
||||
)
|
||||
raise ValueError(f"unknown learned method: {method}")
|
||||
|
||||
|
||||
def _fit_learned_model(
|
||||
method: str,
|
||||
train_samples: Sequence[FeatureSample],
|
||||
val_samples: Sequence[FeatureSample],
|
||||
stats: FeatureStats,
|
||||
*,
|
||||
device: torch.device,
|
||||
seed: int,
|
||||
grid_size: int,
|
||||
hidden_size: int,
|
||||
heads: int,
|
||||
dropout: float,
|
||||
batch_size: int,
|
||||
max_epochs: int,
|
||||
patience: int,
|
||||
learning_rate: float,
|
||||
checkpoint_path: Path,
|
||||
loss_variant: str = "v1",
|
||||
) -> tuple[nn.Module, dict[str, Any]]:
|
||||
_seed_everything(seed)
|
||||
dimensions = {name: train_samples[0].features[name].shape[1] for name in MODALITIES}
|
||||
model = _make_learned_model(method, dimensions, grid_size, hidden_size, heads, dropout).to(device)
|
||||
decoder = AlignmentReconstructor(hidden_size, dropout).to(device)
|
||||
optimizer = torch.optim.AdamW(
|
||||
[*model.parameters(), *decoder.parameters()], lr=learning_rate, weight_decay=1e-4
|
||||
)
|
||||
rng = np.random.default_rng(seed)
|
||||
train_mask_generator = torch.Generator(device=device)
|
||||
train_mask_generator.manual_seed(seed + 31)
|
||||
history: list[dict[str, float]] = []
|
||||
best_loss = math.inf
|
||||
best_epoch = 0
|
||||
best_model: dict[str, Tensor] | None = None
|
||||
best_decoder: dict[str, Tensor] | None = None
|
||||
patience_used = 0
|
||||
|
||||
for epoch in range(1, max_epochs + 1):
|
||||
model.train()
|
||||
decoder.train()
|
||||
train_total = 0.0
|
||||
train_count = 0
|
||||
for batch_samples in _batches(train_samples, batch_size, shuffle=True, rng=rng):
|
||||
sequences, durations, _ = collate_feature_samples(batch_samples, stats, device)
|
||||
output = model(sequences)
|
||||
total, _ = _training_objective(
|
||||
output,
|
||||
sequences,
|
||||
durations,
|
||||
decoder,
|
||||
train_mask_generator,
|
||||
method=method,
|
||||
loss_variant=loss_variant,
|
||||
)
|
||||
if not torch.isfinite(total):
|
||||
raise FloatingPointError(f"non-finite {method} objective at epoch {epoch}")
|
||||
optimizer.zero_grad(set_to_none=True)
|
||||
total.backward()
|
||||
nn.utils.clip_grad_norm_([*model.parameters(), *decoder.parameters()], 1.0)
|
||||
optimizer.step()
|
||||
train_total += float(total.detach().item()) * len(batch_samples)
|
||||
train_count += len(batch_samples)
|
||||
|
||||
model.eval()
|
||||
decoder.eval()
|
||||
val_generator = torch.Generator(device=device)
|
||||
val_generator.manual_seed(seed + 99991)
|
||||
val_total = 0.0
|
||||
val_count = 0
|
||||
val_metric_sums: dict[str, float] = defaultdict(float)
|
||||
with torch.no_grad():
|
||||
for batch_samples in _batches(val_samples, batch_size, shuffle=False, rng=rng):
|
||||
sequences, durations, _ = collate_feature_samples(batch_samples, stats, device)
|
||||
output = model(sequences)
|
||||
total, parts = _training_objective(
|
||||
output,
|
||||
sequences,
|
||||
durations,
|
||||
decoder,
|
||||
val_generator,
|
||||
method=method,
|
||||
loss_variant=loss_variant,
|
||||
)
|
||||
count = len(batch_samples)
|
||||
val_total += float(total.item()) * count
|
||||
val_count += count
|
||||
for key, value in parts.items():
|
||||
val_metric_sums[key] += float(value.item()) * count
|
||||
for name in MODALITIES:
|
||||
val_metric_sums[f"c_row_{name}"] += float(
|
||||
attention_row_similarity(output.weights[name]).mean().item()
|
||||
) * count
|
||||
val_metric_sums[f"c_far_{name}"] += float(
|
||||
attention_row_similarity(output.weights[name], min_separation=6)
|
||||
.mean()
|
||||
.item()
|
||||
) * count
|
||||
val_mean = val_total / max(val_count, 1)
|
||||
history.append(
|
||||
{
|
||||
"epoch": float(epoch),
|
||||
"train_total": train_total / max(train_count, 1),
|
||||
"validation_total": val_mean,
|
||||
**{
|
||||
f"validation_{key}": value / max(val_count, 1)
|
||||
for key, value in val_metric_sums.items()
|
||||
},
|
||||
}
|
||||
)
|
||||
if val_mean < best_loss - 1e-6:
|
||||
best_loss = val_mean
|
||||
best_epoch = epoch
|
||||
best_model = {key: value.detach().cpu().clone() for key, value in model.state_dict().items()}
|
||||
best_decoder = {
|
||||
key: value.detach().cpu().clone() for key, value in decoder.state_dict().items()
|
||||
}
|
||||
patience_used = 0
|
||||
else:
|
||||
patience_used += 1
|
||||
if patience_used >= patience:
|
||||
break
|
||||
|
||||
if best_model is None or best_decoder is None:
|
||||
raise RuntimeError(f"{method} training did not produce a finite checkpoint")
|
||||
model.load_state_dict(best_model)
|
||||
decoder.load_state_dict(best_decoder)
|
||||
model.eval()
|
||||
decoder.eval()
|
||||
checkpoint_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
torch.save(
|
||||
{
|
||||
"method": method,
|
||||
"seed": seed,
|
||||
"loss_variant": loss_variant,
|
||||
"best_epoch": best_epoch,
|
||||
"best_validation_objective": best_loss,
|
||||
"model_state_dict": best_model,
|
||||
"training_decoder_state_dict": best_decoder,
|
||||
"history": history,
|
||||
},
|
||||
checkpoint_path,
|
||||
)
|
||||
return model, {
|
||||
"best_epoch": best_epoch,
|
||||
"best_validation_objective": best_loss,
|
||||
"history": history,
|
||||
"best_validation_metrics": history[best_epoch - 1],
|
||||
"checkpoint": str(checkpoint_path),
|
||||
}
|
||||
|
||||
|
||||
def _baseline_output(
|
||||
method: str,
|
||||
sequences: Mapping[str, Any],
|
||||
durations: Tensor,
|
||||
word_intervals: Sequence[Tensor],
|
||||
grid_size: int,
|
||||
) -> AlignmentOutput:
|
||||
if method == "M1":
|
||||
return align_forced_timestamps(sequences, word_intervals, grid_size)
|
||||
if method == "M2":
|
||||
return align_fixed_windows(sequences, durations, grid_size)
|
||||
raise ValueError(f"not a fixed baseline: {method}")
|
||||
|
||||
|
||||
def _alignment_rows(
|
||||
method: str,
|
||||
seed_label: str,
|
||||
fold: int,
|
||||
samples: Sequence[FeatureSample],
|
||||
output: AlignmentOutput,
|
||||
device: torch.device,
|
||||
epsilon: float,
|
||||
) -> list[dict[str, Any]]:
|
||||
rows = []
|
||||
for index, sample in enumerate(samples):
|
||||
for name in MODALITIES:
|
||||
length = len(sample.times[name])
|
||||
weights = output.weights[name][index : index + 1, :, :length]
|
||||
times = torch.as_tensor(sample.times[name], dtype=torch.float32, device=device)[None]
|
||||
valid = torch.as_tensor(sample.valid[name], dtype=torch.bool, device=device)[None]
|
||||
duration = torch.tensor([sample.duration_s], dtype=torch.float32, device=device)
|
||||
trajectory = alignment_trajectory(weights, times, duration)
|
||||
mvr = monotonicity_violation_rate(trajectory, epsilon)
|
||||
entropy = normalized_attention_entropy(weights, valid)
|
||||
width = attention_width80(weights)
|
||||
rows.append(
|
||||
{
|
||||
"method": method,
|
||||
"seed": seed_label,
|
||||
"fold": fold,
|
||||
"sample_id": sample.sample_id,
|
||||
"modality": name,
|
||||
"mvr": float(mvr[0].item()),
|
||||
"normalized_entropy": float(entropy.mean().item()),
|
||||
"width80_source_positions": float(width.float().mean().item()),
|
||||
"c_row": float(attention_row_similarity(weights).mean().item()),
|
||||
"c_far": float(
|
||||
attention_row_similarity(weights, min_separation=6).mean().item()
|
||||
),
|
||||
"expected_time_start_s": float(trajectory[0, 0].item() * sample.duration_s),
|
||||
"expected_time_end_s": float(trajectory[0, -1].item() * sample.duration_s),
|
||||
"trajectory_span_fraction": float((trajectory[0, -1] - trajectory[0, 0]).item()),
|
||||
}
|
||||
)
|
||||
return rows
|
||||
|
||||
|
||||
def _collect_representations(
|
||||
method: str,
|
||||
samples: Sequence[FeatureSample],
|
||||
stats: FeatureStats,
|
||||
*,
|
||||
device: torch.device,
|
||||
grid_size: int,
|
||||
batch_size: int,
|
||||
model: nn.Module | None = None,
|
||||
) -> tuple[dict[str, dict[str, np.ndarray]], dict[str, dict[str, np.ndarray]]]:
|
||||
aligned_by_id: dict[str, dict[str, np.ndarray]] = {}
|
||||
weights_by_id: dict[str, dict[str, np.ndarray]] = {}
|
||||
rng = np.random.default_rng(0)
|
||||
if model is not None:
|
||||
model.eval()
|
||||
with torch.no_grad():
|
||||
for batch_samples in _batches(samples, batch_size, shuffle=False, rng=rng):
|
||||
sequences, durations, intervals = collate_feature_samples(batch_samples, stats, device)
|
||||
if method in {"M1", "M2"}:
|
||||
output = _baseline_output(method, sequences, durations, intervals, grid_size)
|
||||
elif model is not None:
|
||||
output = model(sequences)
|
||||
else:
|
||||
raise ValueError(f"a trained model is required for {method}")
|
||||
for index, sample in enumerate(batch_samples):
|
||||
aligned: dict[str, np.ndarray] = {}
|
||||
weights: dict[str, np.ndarray] = {}
|
||||
for name in MODALITIES:
|
||||
length = len(sample.features[name])
|
||||
matrix = output.weights[name][index, :, :length]
|
||||
source = sequences[name].features[index, :length]
|
||||
pooled = matrix.to(source.dtype) @ source
|
||||
aligned[name] = pooled.detach().cpu().numpy().astype(np.float32, copy=False)
|
||||
weights[name] = matrix.detach().cpu().numpy().astype(np.float32, copy=False)
|
||||
aligned_by_id[sample.sample_id] = aligned
|
||||
weights_by_id[sample.sample_id] = weights
|
||||
return aligned_by_id, weights_by_id
|
||||
|
||||
|
||||
def _save_alignment(
|
||||
output_dir: Path,
|
||||
method: str,
|
||||
seed_label: str,
|
||||
sample: FeatureSample,
|
||||
weights: Mapping[str, np.ndarray],
|
||||
) -> None:
|
||||
path = output_dir / "alignments" / method / f"seed_{seed_label}" / (
|
||||
sample.sample_id.replace("/", "__") + ".npz"
|
||||
)
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
values: dict[str, np.ndarray] = {"sample_id": np.asarray(sample.sample_id)}
|
||||
for name in MODALITIES:
|
||||
trajectory = weights[name] @ sample.times[name] / max(sample.duration_s, 1e-8)
|
||||
values[f"weights_{name}"] = weights[name]
|
||||
values[f"times_{name}_s"] = sample.times[name].astype(np.float32, copy=False)
|
||||
values[f"valid_{name}"] = sample.valid[name]
|
||||
values[f"trajectory_{name}"] = trajectory.astype(np.float32, copy=False)
|
||||
np.savez_compressed(path, **values)
|
||||
|
||||
|
||||
def _write_csv(path: Path, rows: Sequence[Mapping[str, Any]]) -> None:
|
||||
if not rows:
|
||||
return
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
columns = list(dict.fromkeys(key for row in rows for key in row))
|
||||
with path.open("w", encoding="utf-8-sig", newline="") as file:
|
||||
writer = csv.DictWriter(file, fieldnames=columns, extrasaction="ignore")
|
||||
writer.writeheader()
|
||||
writer.writerows(rows)
|
||||
|
||||
|
||||
def _group_summary(
|
||||
rows: Sequence[Mapping[str, Any]], group_columns: Sequence[str], metric_columns: Sequence[str]
|
||||
) -> list[dict[str, Any]]:
|
||||
groups: dict[tuple[Any, ...], list[Mapping[str, Any]]] = defaultdict(list)
|
||||
for row in rows:
|
||||
groups[tuple(row[column] for column in group_columns)].append(row)
|
||||
output = []
|
||||
for key, values in groups.items():
|
||||
summary: dict[str, Any] = dict(zip(group_columns, key))
|
||||
summary["n"] = len(values)
|
||||
for metric in metric_columns:
|
||||
numbers = [float(row[metric]) for row in values if row.get(metric) not in (None, "")]
|
||||
numbers = [value for value in numbers if math.isfinite(value)]
|
||||
if numbers:
|
||||
summary[f"{metric}_mean"] = statistics.fmean(numbers)
|
||||
summary[f"{metric}_std"] = statistics.stdev(numbers) if len(numbers) > 1 else 0.0
|
||||
output.append(summary)
|
||||
return output
|
||||
|
||||
|
||||
def _comparison_table(summaries: Mapping[str, Sequence[Mapping[str, Any]]]) -> list[dict[str, Any]]:
|
||||
"""Make one compact, multi-metric table without inventing a composite score."""
|
||||
alignment = {(row["method"], row["modality"]): row for row in summaries["alignment"]}
|
||||
retrieval = {(row["method"], row["direction"]): row for row in summaries["retrieval"]}
|
||||
reconstruction = {
|
||||
(row["method"], row["target_modality"]): row for row in summaries["reconstruction"]
|
||||
}
|
||||
emotion = {row["method"]: row for row in summaries["emotion"]}
|
||||
table: list[dict[str, Any]] = []
|
||||
for method in ("M1", "M2", "M3", "M4"):
|
||||
row: dict[str, Any] = {"method": method}
|
||||
for modality in MODALITIES:
|
||||
metrics = alignment[(method, modality)]
|
||||
row[f"mvr_{modality}"] = metrics.get("mvr_mean")
|
||||
row[f"entropy_{modality}"] = metrics.get("normalized_entropy_mean")
|
||||
row[f"trajectory_span_{modality}"] = metrics.get("trajectory_span_fraction_mean")
|
||||
for direction in ("text_to_audio", "text_to_vision"):
|
||||
metrics = retrieval[(method, direction)]
|
||||
row[f"r_at_1_{direction}"] = metrics.get("r_at_1_mean")
|
||||
row[f"r_at_5_{direction}"] = metrics.get("r_at_5_mean")
|
||||
for modality in MODALITIES:
|
||||
metrics = reconstruction[(method, modality)]
|
||||
row[f"reconstruction_mae_{modality}"] = metrics.get("mae_standardized_mean")
|
||||
for metric in ("accuracy", "macro_f1", "mae", "pearson"):
|
||||
row[f"emotion_{metric}"] = emotion[method].get(f"{metric}_mean")
|
||||
table.append(row)
|
||||
return table
|
||||
|
||||
|
||||
def _make_figures(
|
||||
output_dir: Path,
|
||||
example_id: str,
|
||||
example_sample: FeatureSample,
|
||||
example_weights: Mapping[str, Mapping[str, np.ndarray]],
|
||||
grid_size: int,
|
||||
) -> None:
|
||||
method_order = ("M1", "M2", "M3", "M4")
|
||||
fig, axes = plt.subplots(4, 3, figsize=(15, 13), constrained_layout=True)
|
||||
for row, method in enumerate(method_order):
|
||||
if method not in example_weights:
|
||||
continue
|
||||
for col, name in enumerate(MODALITIES):
|
||||
ax = axes[row, col]
|
||||
matrix = example_weights[method][name]
|
||||
image = ax.imshow(matrix, origin="lower", aspect="auto", interpolation="nearest", cmap="magma")
|
||||
ax.set_title(f"{method} · {name}")
|
||||
ax.set_xlabel("source position")
|
||||
ax.set_ylabel("shared grid slot")
|
||||
ax.set_yticks(np.linspace(0, grid_size - 1, 5, dtype=int))
|
||||
fig.colorbar(image, ax=ax, fraction=0.046, pad=0.04)
|
||||
fig.suptitle(f"Alignment matrices on held-out sample {example_id}")
|
||||
fig.savefig(output_dir / "typical_alignment_heatmaps.png", dpi=170)
|
||||
plt.close(fig)
|
||||
|
||||
fig, axes = plt.subplots(2, 2, figsize=(13, 9), constrained_layout=True)
|
||||
x = (np.arange(grid_size, dtype=np.float32) + 0.5) / grid_size
|
||||
for ax, method in zip(axes.flat, method_order):
|
||||
for name in MODALITIES:
|
||||
matrix = example_weights[method][name]
|
||||
duration = example_sample.duration_s
|
||||
trajectory = matrix @ example_sample.times[name] / max(duration, 1e-8)
|
||||
ax.plot(x, trajectory, label=name)
|
||||
ax.plot([0, 1], [0, 1], linestyle="--", color="black", alpha=0.5, label="uniform-time reference")
|
||||
ax.set_title(method)
|
||||
ax.set_xlabel("shared-grid position")
|
||||
ax.set_ylabel("expected source time / clip duration")
|
||||
ax.set_xlim(0, 1)
|
||||
ax.set_ylim(0, 1)
|
||||
ax.grid(alpha=0.2)
|
||||
axes[0, 0].legend(fontsize=8)
|
||||
fig.suptitle(f"Alignment trajectories on held-out sample {example_id}")
|
||||
fig.savefig(output_dir / "typical_alignment_trajectories.png", dpi=170)
|
||||
plt.close(fig)
|
||||
|
||||
|
||||
def run(args: argparse.Namespace) -> dict[str, Any]:
|
||||
start_time = time.time()
|
||||
_seed_everything(args.seeds[0])
|
||||
if args.device == "auto":
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
else:
|
||||
device = torch.device(args.device)
|
||||
if device.type == "cuda" and not torch.cuda.is_available():
|
||||
raise RuntimeError("CUDA was requested but is not available in this WSL environment")
|
||||
|
||||
samples = load_feature_samples(args.feature_dir, args.manifest)
|
||||
if len(samples) != args.expected_samples:
|
||||
raise ValueError(f"expected {args.expected_samples} samples, found {len(samples)}")
|
||||
groups = [sample.group_id for sample in samples]
|
||||
if len(set(groups)) < args.folds:
|
||||
raise ValueError("fewer video_id groups than requested folds")
|
||||
split_iter = GroupKFold(n_splits=args.folds).split(np.zeros(len(samples)), groups=groups)
|
||||
splits = [(train.tolist(), val.tolist()) for train, val in split_iter]
|
||||
splits_json = [
|
||||
{
|
||||
"fold": fold + 1,
|
||||
"train_sample_ids": [samples[index].sample_id for index in train],
|
||||
"validation_sample_ids": [samples[index].sample_id for index in val],
|
||||
"train_video_ids": sorted({samples[index].group_id for index in train}),
|
||||
"validation_video_ids": sorted({samples[index].group_id for index in val}),
|
||||
}
|
||||
for fold, (train, val) in enumerate(splits)
|
||||
]
|
||||
args.output_dir.mkdir(parents=True, exist_ok=True)
|
||||
print(
|
||||
f"[start] samples={len(samples)} groups={len(set(groups))} folds={args.folds} "
|
||||
f"seeds={args.seeds} device={device}",
|
||||
flush=True,
|
||||
)
|
||||
(args.output_dir / "splits.json").write_text(
|
||||
json.dumps(splits_json, ensure_ascii=False, indent=2), encoding="utf-8"
|
||||
)
|
||||
|
||||
alignment_rows: list[dict[str, Any]] = []
|
||||
retrieval_rows: list[dict[str, Any]] = []
|
||||
reconstruction_rows: list[dict[str, Any]] = []
|
||||
emotion_rows: list[dict[str, Any]] = []
|
||||
training_rows: list[dict[str, Any]] = []
|
||||
examples: dict[str, dict[str, Mapping[str, np.ndarray]]] = defaultdict(dict)
|
||||
sample_by_id = {sample.sample_id: sample for sample in samples}
|
||||
preferred_example = args.example_id if args.example_id in sample_by_id else samples[0].sample_id
|
||||
preferred_seed = args.seeds[0]
|
||||
example_sample = sample_by_id[preferred_example]
|
||||
|
||||
for fold_index, (train_indices, val_indices) in enumerate(splits, start=1):
|
||||
train_samples = [samples[index] for index in train_indices]
|
||||
val_samples = [samples[index] for index in val_indices]
|
||||
stats = fit_feature_stats(train_samples)
|
||||
fold_output = args.output_dir / f"fold_{fold_index:02d}"
|
||||
print(
|
||||
f"[fold {fold_index}/{args.folds}] train={len(train_samples)} validation={len(val_samples)} "
|
||||
f"train_video_ids={len({sample.group_id for sample in train_samples})} "
|
||||
f"validation_video_ids={len({sample.group_id for sample in val_samples})}",
|
||||
flush=True,
|
||||
)
|
||||
|
||||
# M1 and M2 are deterministic methods with no learned alignment loss.
|
||||
for method in ("M1", "M2"):
|
||||
print(f"[fold {fold_index}] evaluate {method} and train identical probes", flush=True)
|
||||
train_aligned, _ = _collect_representations(
|
||||
method,
|
||||
train_samples,
|
||||
stats,
|
||||
device=device,
|
||||
grid_size=args.grid_size,
|
||||
batch_size=args.batch_size,
|
||||
)
|
||||
val_aligned, val_weights = _collect_representations(
|
||||
method,
|
||||
val_samples,
|
||||
stats,
|
||||
device=device,
|
||||
grid_size=args.grid_size,
|
||||
batch_size=args.batch_size,
|
||||
)
|
||||
combined = {**train_aligned, **val_aligned}
|
||||
rng = np.random.default_rng(args.seeds[0] + fold_index)
|
||||
for batch_samples in _batches(val_samples, args.batch_size, shuffle=False, rng=rng):
|
||||
sequences, durations, intervals = collate_feature_samples(batch_samples, stats, device)
|
||||
output = _baseline_output(method, sequences, durations, intervals, args.grid_size)
|
||||
alignment_rows.extend(
|
||||
_alignment_rows(method, "fixed", fold_index, batch_samples, output, device, args.mvr_epsilon)
|
||||
)
|
||||
for sample in val_samples:
|
||||
_save_alignment(args.output_dir, method, "fixed", sample, val_weights[sample.sample_id])
|
||||
if sample.sample_id == preferred_example:
|
||||
examples[sample.sample_id][method] = val_weights[sample.sample_id]
|
||||
probe_seed = args.seeds[0] + fold_index * 100
|
||||
retrieval_rows.extend(
|
||||
{
|
||||
"method": method,
|
||||
"seed": "fixed",
|
||||
"fold": fold_index,
|
||||
**row,
|
||||
}
|
||||
for row in run_retrieval_probe(
|
||||
[sample.sample_id for sample in train_samples],
|
||||
[sample.sample_id for sample in val_samples],
|
||||
combined,
|
||||
device=device,
|
||||
seed=probe_seed,
|
||||
epochs=args.retrieval_probe_epochs,
|
||||
batch_size=args.probe_batch_size,
|
||||
)
|
||||
)
|
||||
reconstruction_rows.extend(
|
||||
{"method": method, "seed": "fixed", "fold": fold_index, **row}
|
||||
for row in run_reconstruction_probe(
|
||||
[sample.sample_id for sample in train_samples],
|
||||
[sample.sample_id for sample in val_samples],
|
||||
combined,
|
||||
device=device,
|
||||
seed=probe_seed + 1,
|
||||
ratio=args.mask_ratio,
|
||||
epochs=args.reconstruction_probe_epochs,
|
||||
batch_size=args.probe_batch_size,
|
||||
)
|
||||
)
|
||||
emotion_rows.append(
|
||||
{
|
||||
"method": method,
|
||||
"seed": "fixed",
|
||||
"fold": fold_index,
|
||||
**run_frozen_emotion_probe(train_samples, val_samples, combined),
|
||||
}
|
||||
)
|
||||
|
||||
# M3/M4 share one unsupervised objective, split, and training budget.
|
||||
for seed in args.seeds:
|
||||
for method in ("M3", "M4"):
|
||||
print(f"[fold {fold_index}] train {method}, seed={seed}", flush=True)
|
||||
checkpoint = fold_output / f"seed_{seed}" / f"{method}.pt"
|
||||
model, training_info = _fit_learned_model(
|
||||
method,
|
||||
train_samples,
|
||||
val_samples,
|
||||
stats,
|
||||
device=device,
|
||||
seed=seed + fold_index * 1009,
|
||||
grid_size=args.grid_size,
|
||||
hidden_size=args.hidden_size,
|
||||
heads=args.heads,
|
||||
dropout=args.dropout,
|
||||
batch_size=args.batch_size,
|
||||
max_epochs=args.epochs,
|
||||
patience=args.patience,
|
||||
learning_rate=args.learning_rate,
|
||||
checkpoint_path=checkpoint,
|
||||
)
|
||||
training_rows.append(
|
||||
{
|
||||
"method": method,
|
||||
"seed": seed,
|
||||
"fold": fold_index,
|
||||
"best_epoch": training_info["best_epoch"],
|
||||
"best_validation_objective": training_info["best_validation_objective"],
|
||||
"checkpoint": training_info["checkpoint"],
|
||||
}
|
||||
)
|
||||
print(
|
||||
f"[fold {fold_index}] {method}, seed={seed} best_epoch={training_info['best_epoch']} "
|
||||
f"val_objective={training_info['best_validation_objective']:.5f}; running frozen probes",
|
||||
flush=True,
|
||||
)
|
||||
train_aligned, _ = _collect_representations(
|
||||
method,
|
||||
train_samples,
|
||||
stats,
|
||||
device=device,
|
||||
grid_size=args.grid_size,
|
||||
batch_size=args.batch_size,
|
||||
model=model,
|
||||
)
|
||||
val_aligned, val_weights = _collect_representations(
|
||||
method,
|
||||
val_samples,
|
||||
stats,
|
||||
device=device,
|
||||
grid_size=args.grid_size,
|
||||
batch_size=args.batch_size,
|
||||
model=model,
|
||||
)
|
||||
combined = {**train_aligned, **val_aligned}
|
||||
rng = np.random.default_rng(seed + fold_index)
|
||||
for batch_samples in _batches(val_samples, args.batch_size, shuffle=False, rng=rng):
|
||||
sequences, durations, _ = collate_feature_samples(batch_samples, stats, device)
|
||||
output = model(sequences)
|
||||
alignment_rows.extend(
|
||||
_alignment_rows(method, str(seed), fold_index, batch_samples, output, device, args.mvr_epsilon)
|
||||
)
|
||||
for sample in val_samples:
|
||||
_save_alignment(args.output_dir, method, str(seed), sample, val_weights[sample.sample_id])
|
||||
if sample.sample_id == preferred_example and seed == preferred_seed:
|
||||
examples[sample.sample_id][method] = val_weights[sample.sample_id]
|
||||
|
||||
probe_seed = seed + fold_index * 100 + (3 if method == "M3" else 7)
|
||||
retrieval_rows.extend(
|
||||
{
|
||||
"method": method,
|
||||
"seed": seed,
|
||||
"fold": fold_index,
|
||||
**row,
|
||||
}
|
||||
for row in run_retrieval_probe(
|
||||
[sample.sample_id for sample in train_samples],
|
||||
[sample.sample_id for sample in val_samples],
|
||||
combined,
|
||||
device=device,
|
||||
seed=probe_seed,
|
||||
epochs=args.retrieval_probe_epochs,
|
||||
batch_size=args.probe_batch_size,
|
||||
)
|
||||
)
|
||||
reconstruction_rows.extend(
|
||||
{"method": method, "seed": seed, "fold": fold_index, **row}
|
||||
for row in run_reconstruction_probe(
|
||||
[sample.sample_id for sample in train_samples],
|
||||
[sample.sample_id for sample in val_samples],
|
||||
combined,
|
||||
device=device,
|
||||
seed=probe_seed + 1,
|
||||
ratio=args.mask_ratio,
|
||||
epochs=args.reconstruction_probe_epochs,
|
||||
batch_size=args.probe_batch_size,
|
||||
)
|
||||
)
|
||||
emotion_rows.append(
|
||||
{
|
||||
"method": method,
|
||||
"seed": seed,
|
||||
"fold": fold_index,
|
||||
**run_frozen_emotion_probe(train_samples, val_samples, combined),
|
||||
}
|
||||
)
|
||||
del model
|
||||
if device.type == "cuda":
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
_write_csv(args.output_dir / "alignment_metrics.csv", alignment_rows)
|
||||
_write_csv(args.output_dir / "retrieval_probe_metrics.csv", retrieval_rows)
|
||||
_write_csv(args.output_dir / "reconstruction_probe_metrics.csv", reconstruction_rows)
|
||||
_write_csv(args.output_dir / "frozen_emotion_probe_metrics.csv", emotion_rows)
|
||||
_write_csv(args.output_dir / "training_summary.csv", training_rows)
|
||||
|
||||
summaries = {
|
||||
"alignment": _group_summary(
|
||||
alignment_rows,
|
||||
("method", "modality"),
|
||||
("mvr", "normalized_entropy", "width80_source_positions", "expected_time_start_s", "expected_time_end_s", "trajectory_span_fraction"),
|
||||
),
|
||||
"retrieval": _group_summary(
|
||||
retrieval_rows,
|
||||
("method", "direction"),
|
||||
("r_at_1", "r_at_5", "mrr"),
|
||||
),
|
||||
"reconstruction": _group_summary(
|
||||
reconstruction_rows,
|
||||
("method", "target_modality"),
|
||||
("mae_standardized", "smooth_l1_standardized"),
|
||||
),
|
||||
"emotion": _group_summary(
|
||||
emotion_rows,
|
||||
("method",),
|
||||
("accuracy", "macro_f1", "mae", "pearson"),
|
||||
),
|
||||
}
|
||||
(args.output_dir / "summary.json").write_text(
|
||||
json.dumps(summaries, ensure_ascii=False, indent=2, allow_nan=False), encoding="utf-8"
|
||||
)
|
||||
for name, rows in summaries.items():
|
||||
_write_csv(args.output_dir / f"{name}_summary.csv", rows)
|
||||
comparison_table = _comparison_table(summaries)
|
||||
_write_csv(args.output_dir / "comparison_summary.csv", comparison_table)
|
||||
|
||||
if preferred_example in examples and set(examples[preferred_example]) == {"M1", "M2", "M3", "M4"}:
|
||||
_make_figures(args.output_dir, preferred_example, example_sample, examples[preferred_example], args.grid_size)
|
||||
|
||||
manifest = {
|
||||
"created_utc": datetime.now(timezone.utc).isoformat(),
|
||||
"sample_count": len(samples),
|
||||
"group_count": len(set(groups)),
|
||||
"folds": args.folds,
|
||||
"seeds": args.seeds,
|
||||
"device": str(device),
|
||||
"gpu_name": torch.cuda.get_device_name(device) if device.type == "cuda" else None,
|
||||
"python": platform.python_version(),
|
||||
"torch": torch.__version__,
|
||||
"parameters": {
|
||||
"grid_size": args.grid_size,
|
||||
"hidden_size": args.hidden_size,
|
||||
"heads": args.heads,
|
||||
"dropout": args.dropout,
|
||||
"batch_size": args.batch_size,
|
||||
"epochs_max": args.epochs,
|
||||
"early_stopping_patience": args.patience,
|
||||
"learning_rate": args.learning_rate,
|
||||
"mask_ratio": args.mask_ratio,
|
||||
"retrieval_probe_epochs": args.retrieval_probe_epochs,
|
||||
"reconstruction_probe_epochs": args.reconstruction_probe_epochs,
|
||||
"mvr_epsilon": args.mvr_epsilon,
|
||||
},
|
||||
"objective": "masked reconstruction + cross-modal contrastive + temporal monotonicity; emotion labels unused",
|
||||
"split_rule": "GroupKFold by group_id/video_id",
|
||||
"example_sample_id": preferred_example,
|
||||
"elapsed_seconds": time.time() - start_time,
|
||||
"interpretation_limits": [
|
||||
"Grid-index retrieval is a representation-consistency probe, not independent temporal ground truth.",
|
||||
"Masked reconstruction uses a decoder trained on the training fold and reports standardized-feature errors.",
|
||||
"The emotion probe is a small-sample downstream utility check, not a claim of generalization to MOSEI.",
|
||||
"No human event timestamps are available, so human IoU/MATE is not reported.",
|
||||
],
|
||||
}
|
||||
(args.output_dir / "run_manifest.json").write_text(
|
||||
json.dumps(manifest, ensure_ascii=False, indent=2), encoding="utf-8"
|
||||
)
|
||||
(args.output_dir / "README.md").write_text(
|
||||
"# Q1 method comparison\n\n"
|
||||
"This folder contains grouped cross-validation results for M1–M4. M1/M2 are fixed rules; M3/M4 are trained without emotion labels. "
|
||||
"All learned models and probes use training-fold-only feature normalization, and folds are grouped by `video_id`.\n\n"
|
||||
"`alignment_metrics.csv` reports expected-time trajectories, monotonicity, attention entropy, and width diagnostics. "
|
||||
"`retrieval_probe_metrics.csv` uses a separately trained linear projection probe; its grid-index positives are not independent temporal ground truth. "
|
||||
"`reconstruction_probe_metrics.csv` reports held-out masked reconstruction error in training-fold standardized feature units. "
|
||||
"`frozen_emotion_probe_metrics.csv` is a small-sample downstream utility check.\n\n"
|
||||
"No human event-time annotation is present, so the results cannot establish direct human alignment accuracy. "
|
||||
"See `run_manifest.json` for parameters, seeds, device, and interpretation limits.\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
print(
|
||||
f"[done] elapsed_seconds={manifest['elapsed_seconds']:.1f} output={args.output_dir}",
|
||||
flush=True,
|
||||
)
|
||||
return manifest
|
||||
|
||||
|
||||
def build_parser() -> argparse.ArgumentParser:
|
||||
project_dir = Path(__file__).resolve().parents[1]
|
||||
parser = argparse.ArgumentParser(description="Compare Q1 M1-M4 alignment methods with grouped CV.")
|
||||
parser.add_argument("--feature-dir", type=Path, default=project_dir / "outputs/q1_features/features")
|
||||
parser.add_argument("--manifest", type=Path, default=project_dir / "outputs/audit/manifest.csv")
|
||||
parser.add_argument("--output-dir", type=Path, default=project_dir / "outputs/method_comparison")
|
||||
parser.add_argument("--device", default="auto", help="auto, cpu, or a torch device such as cuda:0")
|
||||
parser.add_argument("--expected-samples", type=int, default=100)
|
||||
parser.add_argument("--folds", type=int, default=5)
|
||||
parser.add_argument("--seeds", type=int, nargs="+", default=[42, 3407, 2026])
|
||||
parser.add_argument("--grid-size", type=int, default=50)
|
||||
parser.add_argument("--hidden-size", type=int, default=128)
|
||||
parser.add_argument("--heads", type=int, default=4)
|
||||
parser.add_argument("--dropout", type=float, default=0.1)
|
||||
parser.add_argument("--batch-size", type=int, default=8)
|
||||
parser.add_argument("--probe-batch-size", type=int, default=16)
|
||||
parser.add_argument("--epochs", type=int, default=50)
|
||||
parser.add_argument("--patience", type=int, default=8)
|
||||
parser.add_argument("--learning-rate", type=float, default=1e-4)
|
||||
parser.add_argument("--mask-ratio", type=float, default=0.2)
|
||||
parser.add_argument("--retrieval-probe-epochs", type=int, default=20)
|
||||
parser.add_argument("--reconstruction-probe-epochs", type=int, default=25)
|
||||
parser.add_argument("--mvr-epsilon", type=float, default=0.02)
|
||||
parser.add_argument("--example-id", default="-tPCytz4rww/12")
|
||||
return parser
|
||||
|
||||
|
||||
def main() -> int:
|
||||
args = build_parser().parse_args()
|
||||
manifest = run(args)
|
||||
print(json.dumps(manifest, ensure_ascii=False, indent=2))
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
Reference in New Issue
Block a user