Files

1205 lines
53 KiB
Python
Raw Permalink 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.
"""Re-evaluate frozen M4 Shared Latent Timeline checkpoints structurally and functionally.
No alignment model is trained here. The five grouped D5 M4_sourceTime
checkpoints are evaluated with self-structure, induced pairwise maps, cycle and
triangle consistency, content-only probes, content shuffling, and shifted
reconstruction controls.
"""
from __future__ import annotations
import argparse
import csv
import json
import platform
import time
from collections import defaultdict
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Mapping, Sequence
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import torch
import torch.nn.functional as F
from sklearn.metrics import roc_auc_score
from torch import Tensor, nn
from .correspondence_eval import (
CorrespondenceProjection,
_cluster_bootstrap,
_fit_probe,
_sample_metrics,
_stack_ids,
_write_csv,
)
from .experiment_data import (
FeatureSample,
FeatureStats,
collate_feature_samples,
fit_feature_stats,
load_feature_samples,
)
from .models import SharedLatentTimeline
from .types import MODALITIES
GRID_SIZE = 50
HIDDEN_SIZE = 128
HEADS = 4
PAIRINGS = (("text", "audio"), ("text", "vision"), ("audio", "vision"))
STRUCTURE_RADIUS = 2
FAR_RADIUS = 10
SIGMA_SELF = 0.08
SIGMA_CYCLE = 0.05
SHIFTS = (1, 2, 5, 10)
DECODER_TARGETS = {
"audio": ("text", "vision", "vision"),
"vision": ("text", "audio", "audio"),
"text": ("audio", "vision", "vision"),
}
class ContentReconstructionProbe(nn.Module):
"""Small decoder that sees two content streams and no slot/position code."""
def __init__(self, dimension: int = HIDDEN_SIZE) -> None:
super().__init__()
self.decoders = nn.ModuleDict(
{
target: nn.Sequential(
nn.Linear(dimension * 2, dimension * 2),
nn.GELU(),
nn.Linear(dimension * 2, dimension),
)
for target in MODALITIES
}
)
def forward(self, target: str, left: Tensor, right: Tensor) -> Tensor:
return self.decoders[target](torch.cat((left, right), dim=-1))
def _bootstrap_summary(
rows: Sequence[Mapping[str, Any]],
group_keys: Sequence[str],
metric_names: Sequence[str],
*,
seed: int,
repetitions: int = 2000,
) -> list[dict[str, Any]]:
grouped: dict[tuple[Any, ...], list[Mapping[str, Any]]] = defaultdict(list)
for row in rows:
grouped[tuple(row[key] for key in group_keys)].append(row)
output: list[dict[str, Any]] = []
for key, values in sorted(grouped.items(), key=lambda item: tuple(str(x) for x in item[0])):
result: dict[str, Any] = dict(zip(group_keys, key))
result["clip_count"] = len({row.get("sample_id") for row in values})
result["video_id_count"] = len({str(row["video_id"]) for row in values})
for metric_index, metric in enumerate(metric_names):
selected = [row for row in values if row.get(metric) not in (None, "")]
if not selected:
continue
mean, low, high, group_count = _cluster_bootstrap(
selected,
metric,
seed=seed + metric_index + sum(ord(char) for char in str(key)),
repetitions=repetitions,
)
result[f"{metric}_video_macro_mean"] = mean
result[f"{metric}_ci95_low"] = low
result[f"{metric}_ci95_high"] = high
result["video_id_count"] = group_count
output.append(result)
return output
def _normalize_attention(
weights: np.ndarray, valid: np.ndarray
) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
"""Return source-position-to-slot probabilities and supported source indices."""
matrix = np.asarray(weights, dtype=np.float64)
valid_indices = np.flatnonzero(valid)
column_mass = matrix[:, valid_indices].sum(axis=0)
supported = column_mass > 1e-12
indices = valid_indices[supported]
if len(indices) == 0:
raise ValueError("attention has no source positions supported by any latent slot")
normalized = matrix[:, indices] / column_mass[supported][None, :]
return normalized, indices, column_mass
def _pair_map(
source_weights: np.ndarray,
source_valid: np.ndarray,
destination_weights: np.ndarray,
destination_valid: np.ndarray,
) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
"""Construct C^(source->destination) after column-normalizing source attention."""
source_to_slot, source_indices, _ = _normalize_attention(source_weights, source_valid)
destination_indices = np.flatnonzero(destination_valid)
slot_to_destination = np.asarray(destination_weights, dtype=np.float64)[:, destination_indices]
mapping = source_to_slot.T @ slot_to_destination
row_mass = mapping.sum(axis=1, keepdims=True)
mapping = mapping / np.maximum(row_mass, 1e-12)
return mapping.astype(np.float32), source_indices, destination_indices
def _self_structure(
weights: np.ndarray, *, sigma: float = SIGMA_SELF
) -> tuple[np.ndarray, dict[str, float]]:
normalized = weights / np.maximum(np.linalg.norm(weights, axis=1, keepdims=True), 1e-12)
gram = normalized @ normalized.T
positions = (np.arange(gram.shape[0], dtype=np.float64) + 0.5) / gram.shape[0]
distances = np.abs(positions[:, None] - positions[None, :])
near = distances <= STRUCTURE_RADIUS / GRID_SIZE
far = distances >= FAR_RADIUS / GRID_SIZE
far_leakage = distances > FAR_RADIUS / GRID_SIZE
near_off_diagonal = near & ~np.eye(len(gram), dtype=bool)
target = np.exp(-(distances**2) / (2 * sigma**2))
near_mean = float(gram[near].mean())
far_mean = float(gram[far].mean())
near_off_diagonal_mean = float(gram[near_off_diagonal].mean())
result = {
# Keep the literal <= r definition from the task, and report an
# off-diagonal version so the trivial unit diagonal cannot dominate.
"near_similarity": near_mean,
"near_similarity_offdiag": near_off_diagonal_mean,
"far_similarity": far_mean,
"d_self": float(near_mean - far_mean),
"d_self_offdiag": float(near_off_diagonal_mean - far_mean),
"gram_target_error": float(np.linalg.norm(gram - target) / max(np.linalg.norm(target), 1e-12)),
"far_slot_leakage": float(gram[far_leakage].sum() / max(gram.sum(), 1e-12)),
"c_row_offdiag": float(gram[~np.eye(len(gram), dtype=bool)].mean()),
}
return gram.astype(np.float32), result
def _pair_metrics(
mapping: np.ndarray,
source_times: np.ndarray,
destination_times: np.ndarray,
) -> dict[str, float]:
row_sums = mapping.sum(axis=1)
predicted = (mapping @ destination_times) / np.maximum(row_sums, 1e-12)
error = predicted - source_times
return {
"pairwise_time_mae": float(np.abs(error).mean()),
"pairwise_signed_lag": float(error.mean()),
"pairwise_time_corr": float(np.corrcoef(source_times, predicted)[0, 1])
if len(source_times) > 1 and np.std(predicted) > 1e-12 and np.std(source_times) > 1e-12
else 0.0,
"source_position_count": int(len(source_times)),
"destination_position_count": int(len(destination_times)),
}
def _band_target(times: np.ndarray, sigma: float) -> np.ndarray:
distances = times[:, None] - times[None, :]
target = np.exp(-(distances**2) / (2 * sigma**2))
return target / np.maximum(target.sum(axis=1, keepdims=True), 1e-12)
def _cycle_triangle_metrics(
maps: Mapping[str, tuple[np.ndarray, np.ndarray, np.ndarray]],
normalized_times: Mapping[str, np.ndarray],
) -> tuple[list[dict[str, Any]], dict[str, np.ndarray]]:
cycle_rows = []
cycles: dict[str, np.ndarray] = {}
for left, right in PAIRINGS:
forward_key = f"{left}_{right}"
reverse_key = f"{right}_{left}"
forward, source_idx, destination_idx = maps[forward_key]
reverse, reverse_source_idx, reverse_destination_idx = maps[reverse_key]
if not np.array_equal(destination_idx, reverse_source_idx) or not np.array_equal(
source_idx, reverse_destination_idx
):
raise ValueError(f"pair map source supports disagree for {forward_key}/{reverse_key}")
cycle = forward @ reverse
times = normalized_times[left][source_idx]
target = _band_target(times, SIGMA_CYCLE)
pred_times = cycle @ times
cycle_rows.append(
{
"kind": f"cycle_{left}_{right}_{left}",
"cycle_band_error": float(np.linalg.norm(cycle - target) / max(np.linalg.norm(target), 1e-12)),
"cycle_time_mae": float(np.abs(pred_times - times).mean()),
}
)
cycles[f"cycle_{left}_{right}_{left}"] = cycle.astype(np.float32)
c_ta = maps["text_audio"][0]
c_av = maps["audio_vision"][0]
c_tv = maps["text_vision"][0]
if c_ta.shape[1] != c_av.shape[0] or c_ta.shape[0] != c_tv.shape[0] or c_av.shape[1] != c_tv.shape[1]:
raise ValueError("T-A, A-V, and T-V supports disagree; cannot calculate triangle consistency")
path = c_ta @ c_av
triangle_residual = path - c_tv
cycles["triangle_TAV_residual"] = triangle_residual.astype(np.float32)
cycle_rows.append(
{
"kind": "triangle_TAV",
"triangle_relative_error": float(
np.linalg.norm(triangle_residual) / max(np.linalg.norm(c_tv), 1e-12)
),
"triangle_mean_absolute_residual": float(np.abs(triangle_residual).mean()),
}
)
return cycle_rows, cycles
def _checkpoint_path(root: Path, fold: int) -> Path:
if fold == 1:
return root / "M4_sourceTime" / "checkpoint.pt"
return root / f"fold_{fold:02d}" / "M4_sourceTime" / "checkpoint.pt"
def _collect_fold(
*,
fold: int,
train_samples: Sequence[FeatureSample],
validation_samples: Sequence[FeatureSample],
stats: FeatureStats,
checkpoint_root: Path,
device: torch.device,
batch_size: int,
) -> tuple[dict[str, dict[str, np.ndarray]], dict[str, dict[str, np.ndarray]]]:
checkpoint_path = _checkpoint_path(checkpoint_root, fold)
if not checkpoint_path.is_file():
raise FileNotFoundError(f"missing M4_sourceTime checkpoint for fold {fold}: {checkpoint_path}")
checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False)
if checkpoint.get("variant") != "M4_sourceTime" or not checkpoint.get("source_time_encoding"):
raise ValueError(f"checkpoint is not M4_sourceTime: {checkpoint_path}")
if set(checkpoint.get("train_sample_ids", [])) != {sample.sample_id for sample in train_samples}:
raise ValueError(f"M4 checkpoint training IDs do not match fold {fold}")
if set(checkpoint.get("validation_sample_ids", [])) != {sample.sample_id for sample in validation_samples}:
raise ValueError(f"M4 checkpoint validation IDs do not match fold {fold}")
dimensions = {name: train_samples[0].features[name].shape[1] for name in MODALITIES}
model = SharedLatentTimeline(
dimensions,
grid_size=GRID_SIZE,
hidden_size=HIDDEN_SIZE,
heads=HEADS,
dropout=0.0,
absolute_position_encoding=True,
source_time_encoding=True,
).to(device)
model.load_state_dict(checkpoint["model_state_dict"], strict=True)
model.eval()
weights_by_id: dict[str, dict[str, np.ndarray]] = {}
content_by_id: dict[str, dict[str, np.ndarray]] = {}
with torch.no_grad():
samples = [*train_samples, *validation_samples]
for start in range(0, len(samples), batch_size):
batch_samples = samples[start : start + batch_size]
sequences, durations, _ = collate_feature_samples(batch_samples, stats, device)
output = model(sequences, durations)
for index, sample in enumerate(batch_samples):
weights: dict[str, np.ndarray] = {}
content: dict[str, np.ndarray] = {}
for name in MODALITIES:
length = len(sample.features[name])
weights[name] = output.weights[name][index, :, :length].detach().cpu().numpy().astype(np.float32)
# M4's returned values are A^m V^m: content values pooled by attention.
# Positional/query vectors are not concatenated into this representation.
content[name] = output.aligned[name][index].detach().cpu().numpy().astype(np.float32)
weights_by_id[sample.sample_id] = weights
content_by_id[sample.sample_id] = content
del model
if device.type == "cuda":
torch.cuda.empty_cache()
return weights_by_id, content_by_id
def _pairwise_for_sample(
sample: FeatureSample,
weights: Mapping[str, np.ndarray],
) -> tuple[dict[str, tuple[np.ndarray, np.ndarray, np.ndarray]], dict[str, np.ndarray], list[dict[str, Any]], dict[str, np.ndarray]]:
normalized_times = {
name: np.asarray(sample.times[name], dtype=np.float64) / max(sample.duration_s, 1e-8)
for name in MODALITIES
}
pair_maps: dict[str, tuple[np.ndarray, np.ndarray, np.ndarray]] = {}
pair_rows: list[dict[str, Any]] = []
self_grams: dict[str, np.ndarray] = {}
for name in MODALITIES:
gram, metrics = _self_structure(weights[name])
self_grams[name] = gram
pair_rows.append(
{"kind": "self", "modality": name, **metrics}
)
directions = (*PAIRINGS, *((right, left) for left, right in PAIRINGS))
for left, right in directions:
mapping, source_idx, destination_idx = _pair_map(
weights[left], sample.valid[left], weights[right], sample.valid[right]
)
key = f"{left}_{right}"
pair_maps[key] = (mapping, source_idx, destination_idx)
source_times = normalized_times[left][source_idx]
destination_times = normalized_times[right][destination_idx]
pair_rows.append(
{
"kind": "pairwise",
"direction": f"{left}_to_{right}",
**_pair_metrics(mapping, source_times, destination_times),
}
)
cycle_rows, cycle_maps = _cycle_triangle_metrics(pair_maps, normalized_times)
for row in cycle_rows:
row["kind_group"] = "cycle_triangle"
pair_rows.append(row)
return pair_maps, normalized_times, pair_rows, {**self_grams, **cycle_maps}
def _fit_reconstruction_probe(
train_ids: Sequence[str],
content_by_id: Mapping[str, Mapping[str, np.ndarray]],
*,
device: torch.device,
seed: int,
epochs: int,
batch_size: int,
learning_rate: float,
) -> tuple[ContentReconstructionProbe, list[dict[str, Any]]]:
torch.manual_seed(seed)
if device.type == "cuda":
torch.cuda.manual_seed_all(seed)
train = _stack_ids(train_ids, content_by_id, device)
model = ContentReconstructionProbe(train["text"].shape[-1]).to(device)
optimizer = torch.optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=1e-4)
rng = np.random.default_rng(seed)
history = []
model.train()
for epoch in range(1, epochs + 1):
order = rng.permutation(len(train_ids))
losses = []
for start in range(0, len(order), batch_size):
indexes = torch.as_tensor(order[start : start + batch_size], device=device)
batch = {name: value.index_select(0, indexes) for name, value in train.items()}
targets = []
for target in MODALITIES:
left, right, _ = DECODER_TARGETS[target]
prediction = model(target, batch[left], batch[right])
targets.append(F.smooth_l1_loss(prediction, batch[target]))
loss = torch.stack(targets).mean()
if not torch.isfinite(loss):
raise FloatingPointError(f"non-finite content reconstruction loss at epoch {epoch}")
optimizer.zero_grad(set_to_none=True)
loss.backward()
nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
losses.append(float(loss.detach().item()))
history.append({"epoch": epoch, "train_loss": float(np.mean(losses))})
return model, history
def _content_pair_metrics(
scores: np.ndarray,
*,
method: str,
fold: int,
sample: FeatureSample,
pair: str,
control: str,
shuffle_id: int | None = None,
) -> dict[str, Any]:
k = scores.shape[0]
diag = np.diag(scores)
indexes = np.arange(k)
negative_mask = np.abs(indexes[:, None] - indexes[None, :]) > 2
auc = roc_auc_score(
np.r_[np.ones(k), np.zeros(int(negative_mask.sum()))],
np.r_[diag, scores[negative_mask]],
)
row_pred = scores.argmax(axis=1)
col_pred = scores.T.argmax(axis=1)
row_err = np.abs(row_pred - indexes)
col_err = np.abs(col_pred - indexes)
result: dict[str, Any] = {
"method": method,
"control": control,
"fold": fold,
"sample_id": sample.sample_id,
"video_id": sample.group_id,
"pair": pair,
"same_time_similarity": float(diag.mean()),
"shifted_far_similarity": float(
np.mean(
[scores[indexes[: k - d], indexes[d:]].mean() for d in range(3, 11)]
+ [scores[indexes[d:], indexes[: k - d]].mean() for d in range(3, 11)]
)
),
"same_minus_shifted_margin": float(
diag.mean()
- np.mean(
[scores[indexes[: k - d], indexes[d:]].mean() for d in range(3, 11)]
+ [scores[indexes[d:], indexes[: k - d]].mean() for d in range(3, 11)]
)
),
"matched_vs_shifted_auc": float(auc),
"exact_r1_left_to_right": float(np.mean(row_err == 0)),
"within_pm1_r1_left_to_right": float(np.mean(row_err <= 1)),
"mase_slots_left_to_right": float(row_err.mean()),
"exact_r1_right_to_left": float(np.mean(col_err == 0)),
"within_pm1_r1_right_to_left": float(np.mean(col_err <= 1)),
"mase_slots_right_to_left": float(col_err.mean()),
}
if shuffle_id is not None:
result["shuffle_id"] = shuffle_id
return result
def _fit_and_score_content(
*,
fold: int,
train_samples: Sequence[FeatureSample],
validation_samples: Sequence[FeatureSample],
content_by_id: Mapping[str, Mapping[str, np.ndarray]],
device: torch.device,
args: argparse.Namespace,
probe_seeds: dict[str, Any],
) -> tuple[list[dict[str, Any]], list[dict[str, Any]], list[dict[str, Any]], dict[str, Any]]:
train_ids = [sample.sample_id for sample in train_samples]
val_ids = [sample.sample_id for sample in validation_samples]
probe_seed = args.seed + fold * 101
projector, projection_history = _fit_probe(
train_ids,
content_by_id,
device=device,
seed=probe_seed,
epochs=args.probe_epochs,
batch_size=args.batch_size,
learning_rate=args.learning_rate,
temperature=args.temperature,
)
for row in projection_history:
probe_seeds.setdefault("history", []).append(
{"fold": fold, "probe": "content_projection", "seed": probe_seed, **row}
)
projector.eval()
with torch.no_grad():
validation = _stack_ids(val_ids, content_by_id, device)
projected = projector(validation)
normal_rows: list[dict[str, Any]] = []
curve_rows: list[dict[str, Any]] = []
for index, sample in enumerate(validation_samples):
one = {name: projected[name][index] for name in MODALITIES}
normal_rows.extend(
_sample_metrics(
method="M4_sourceTime_content",
fold=fold,
sample=sample,
projected=one,
curve_rows=curve_rows,
)
)
rng = np.random.default_rng(probe_seed + 17)
shuffle_rows = []
with torch.no_grad():
for shuffle_id in range(args.shuffle_repeats):
for index, sample in enumerate(validation_samples):
permuted = {}
for name in MODALITIES:
permutation = torch.as_tensor(rng.permutation(GRID_SIZE), device=device)
permuted[name] = projected[name][index].index_select(0, permutation)
for left, right in PAIRINGS:
scores = (permuted[left] @ permuted[right].T).detach().cpu().numpy()
shuffle_rows.append(
_content_pair_metrics(
scores,
method="M4_sourceTime_content",
fold=fold,
sample=sample,
pair=f"{left}_{right}",
control="independent_within_clip_permutation",
shuffle_id=shuffle_id,
)
)
probe_seeds[f"fold_{fold:02d}/content_projection"] = {
"seed": probe_seed,
"state_dict": {key: value.detach().cpu() for key, value in projector.state_dict().items()},
}
del projector
if device.type == "cuda":
torch.cuda.empty_cache()
return normal_rows, curve_rows, shuffle_rows, probe_seeds
def _fit_and_score_reconstruction(
*,
fold: int,
train_samples: Sequence[FeatureSample],
validation_samples: Sequence[FeatureSample],
content_by_id: Mapping[str, Mapping[str, np.ndarray]],
device: torch.device,
args: argparse.Namespace,
checkpoint_store: dict[str, Any],
history_store: list[dict[str, Any]],
) -> list[dict[str, Any]]:
seed = args.seed + fold * 211
decoder, history = _fit_reconstruction_probe(
[sample.sample_id for sample in train_samples],
content_by_id,
device=device,
seed=seed,
epochs=args.decoder_epochs,
batch_size=args.batch_size,
learning_rate=args.decoder_learning_rate,
)
history_store.extend(
{"fold": fold, "probe": "content_reconstruction", "seed": seed, **row}
for row in history
)
decoder.eval()
validation = _stack_ids([sample.sample_id for sample in validation_samples], content_by_id, device)
rows = []
with torch.no_grad():
for target, (left, right, shifted_modality) in DECODER_TARGETS.items():
for sample_index, sample in enumerate(validation_samples):
left_values = validation[left][sample_index]
right_values = validation[right][sample_index]
target_values = validation[target][sample_index]
for delta in (*(-value for value in SHIFTS), *SHIFTS):
source_indices = np.arange(max(0, -delta), min(GRID_SIZE, GRID_SIZE - delta))
shifted_indices = source_indices + delta
target_idx = torch.as_tensor(source_indices, device=device)
shifted_idx = torch.as_tensor(shifted_indices, device=device)
left_eval = left_values.index_select(0, target_idx)
right_aligned = right_values.index_select(0, target_idx)
right_shifted = (
right_values.index_select(0, shifted_idx)
if shifted_modality == right
else right_aligned
)
left_shifted = (
left_values.index_select(0, shifted_idx)
if shifted_modality == left
else left_eval
)
aligned_pred = decoder(target, left_eval, right_aligned)
shifted_pred = decoder(target, left_shifted, right_shifted)
target_eval = target_values.index_select(0, target_idx)
aligned_mae = float((aligned_pred - target_eval).abs().mean().item())
shifted_mae = float((shifted_pred - target_eval).abs().mean().item())
rows.append(
{
"method": "M4_sourceTime",
"fold": fold,
"sample_id": sample.sample_id,
"video_id": sample.group_id,
"target_modality": target,
"shifted_modality": shifted_modality,
"delta": delta,
"abs_delta": abs(delta),
"aligned_mae_same_support": aligned_mae,
"shifted_mae": shifted_mae,
"gain_shift_minus_aligned": shifted_mae - aligned_mae,
}
)
checkpoint_store[f"fold_{fold:02d}/content_reconstruction"] = {
"seed": seed,
"state_dict": {key: value.detach().cpu() for key, value in decoder.state_dict().items()},
}
del decoder
if device.type == "cuda":
torch.cuda.empty_cache()
return rows
def _plot_self_gram(grams: Mapping[str, np.ndarray], path: Path) -> None:
fig, axes = plt.subplots(1, 3, figsize=(13, 4.4), sharex=True, sharey=True, constrained_layout=True)
for axis, name in zip(axes, MODALITIES, strict=True):
image = axis.imshow(grams[name], origin="lower", aspect="equal", vmin=0, vmax=1, cmap="magma")
axis.set_title(name.title())
axis.set_xlabel("Latent slot k")
axis.set_xticks([0, 9, 19, 29, 39, 49])
axis.set_yticks([0, 9, 19, 29, 39, 49])
axes[0].set_ylabel("Latent slot i")
fig.colorbar(image, ax=axes, fraction=0.025, pad=0.02, label="Cosine similarity of attention rows")
fig.suptitle("M4 self-structure: pairwise similarity between latent slots")
fig.savefig(path, dpi=180, bbox_inches="tight")
plt.close(fig)
def _time_cell_edges(times: np.ndarray) -> np.ndarray:
"""Convert ordered sample centers to physical-time cell boundaries."""
times = np.asarray(times, dtype=np.float64)
if len(times) == 0:
raise ValueError("cannot plot an empty time axis")
if len(times) == 1:
return np.array([times[0] - 0.005, times[0] + 0.005])
if np.any(np.diff(times) < 0):
raise ValueError("feature times must be ordered for time-faithful heatmaps")
midpoints = (times[:-1] + times[1:]) / 2
first = times[0] - (midpoints[0] - times[0])
last = times[-1] + (times[-1] - midpoints[-1])
return np.concatenate(([first], midpoints, [last]))
def _plot_attention(
weights: Mapping[str, np.ndarray],
times: Mapping[str, np.ndarray],
valid: Mapping[str, np.ndarray],
path: Path,
) -> None:
fig, axes = plt.subplots(1, 3, figsize=(15, 4.8), sharey=True, constrained_layout=True)
slot_edges = np.arange(GRID_SIZE + 1, dtype=np.float64) - 0.5
for axis, name in zip(axes, MODALITIES, strict=True):
mask = np.asarray(valid[name], dtype=bool)
source_times = np.asarray(times[name], dtype=np.float64)[mask]
matrix = np.asarray(weights[name], dtype=np.float64)[:, mask]
image = axis.pcolormesh(
_time_cell_edges(source_times),
slot_edges,
matrix,
shading="flat",
cmap="magma",
)
axis.set_title(name.title())
axis.set_xlabel(f"{name.title()} normalized source time")
axis.set_yticks([0, 9, 19, 29, 39, 49])
axes[0].set_ylabel("Shared latent slot")
fig.colorbar(image, ax=axes, fraction=0.025, pad=0.02, label="Attention weight")
fig.suptitle("M4 latent slots attending to each modality's source timeline")
fig.savefig(path, dpi=180, bbox_inches="tight")
plt.close(fig)
def _plot_pairwise(
maps: Mapping[str, tuple[np.ndarray, np.ndarray, np.ndarray]],
times: Mapping[str, np.ndarray],
path: Path,
) -> None:
directions = (("text", "audio"), ("text", "vision"), ("audio", "vision"))
fig, axes = plt.subplots(1, 3, figsize=(16, 4.8), constrained_layout=True)
for axis, (left, right) in zip(axes, directions, strict=True):
matrix, source_idx, destination_idx = maps[f"{left}_{right}"]
source_times = times[left][source_idx]
destination_times = times[right][destination_idx]
image = axis.pcolormesh(
_time_cell_edges(destination_times),
_time_cell_edges(source_times),
matrix,
shading="flat",
cmap="magma",
)
axis.set_title(f"{left.title()} → {right.title()}")
axis.set_xlabel(f"{right.title()} normalized time")
axis.set_ylabel(f"{left.title()} normalized time")
fig.colorbar(image, ax=axes, fraction=0.025, pad=0.02, label="Pairwise transition probability")
fig.suptitle("M4 modality-to-modality maps induced through the shared latent timeline")
fig.savefig(path, dpi=180, bbox_inches="tight")
plt.close(fig)
def _plot_pair_trajectories(
maps: Mapping[str, tuple[np.ndarray, np.ndarray, np.ndarray]],
times: Mapping[str, np.ndarray],
path: Path,
) -> None:
fig, axes = plt.subplots(1, 3, figsize=(15, 4.6), sharex=True, sharey=True)
for axis, (left, right) in zip(axes, (("text", "audio"), ("text", "vision"), ("audio", "vision")), strict=True):
matrix, source_idx, destination_idx = maps[f"{left}_{right}"]
x = times[left][source_idx]
y = matrix @ times[right][destination_idx]
axis.plot(x, y, color="#2673b8", linewidth=1.2, marker=".", markersize=2)
axis.plot([0, 1], [0, 1], color="black", linestyle="--", linewidth=1)
axis.set_title(f"{left.title()} → {right.title()}")
axis.set_xlabel(f"{left.title()} source time")
axis.grid(alpha=0.2)
axes[0].set_ylabel("Expected destination time")
fig.suptitle("Pairwise temporal trajectories; diagonal indicates equal physical time")
fig.tight_layout()
fig.savefig(path, dpi=180, bbox_inches="tight")
plt.close(fig)
def _plot_cycles(cycles: Mapping[str, np.ndarray], path: Path) -> None:
names = ("cycle_text_audio_text", "cycle_text_vision_text", "cycle_audio_vision_audio")
fig, axes = plt.subplots(1, 4, figsize=(17, 4.4), constrained_layout=True)
for axis, name in zip(axes[:3], names, strict=True):
image = axis.imshow(cycles[name], origin="lower", aspect="auto", cmap="magma")
axis.set_title(name.replace("cycle_", "").replace("_", "→"))
axis.set_xlabel("Source position")
axis.set_ylabel("Source position")
residual = np.abs(cycles["triangle_TAV_residual"])
image = axes[3].imshow(residual, origin="lower", aspect="auto", cmap="viridis")
axes[3].set_title("|T→A→V − T→V|")
axes[3].set_xlabel("Vision position")
axes[3].set_ylabel("Text position")
fig.colorbar(image, ax=axes, fraction=0.024, pad=0.02, label="Return / path residual")
fig.suptitle("Cycle transition maps and triangle-path residual")
fig.savefig(path, dpi=180, bbox_inches="tight")
plt.close(fig)
def _plot_content_curve(rows: Sequence[Mapping[str, Any]], path: Path) -> None:
colors = {"text_audio": "#2878b5", "text_vision": "#e1812c", "audio_vision": "#55a868"}
fig, axes = plt.subplots(1, 3, figsize=(14, 4.3), sharey=True)
for axis, pair in zip(axes, colors, strict=True):
xs, means, lows, highs = [], [], [], []
for delta in range(-10, 11):
selected = [row for row in rows if row["pair"] == pair and int(row["delta"]) == delta]
if not selected:
continue
mean, low, high, _ = _cluster_bootstrap(
selected,
"similarity",
seed=7301 + delta + sum(ord(char) for char in pair),
repetitions=1000,
)
xs.append(delta)
means.append(mean)
lows.append(low)
highs.append(high)
axis.plot(xs, means, color=colors[pair], linewidth=1.1, marker="o", markersize=3)
axis.fill_between(xs, lows, highs, color=colors[pair], alpha=0.16)
axis.axvline(0, color="black", linestyle="--", linewidth=0.8)
axis.set_title(pair.replace("_", "–"))
axis.set_xlabel("Temporal shift Δ (slots)")
axis.grid(alpha=0.2)
axes[0].set_ylabel("Content-only cosine similarity")
fig.suptitle("M4 content-only same-slot vs shifted similarity (video_id bootstrap CI)")
fig.tight_layout()
fig.savefig(path, dpi=180, bbox_inches="tight")
plt.close(fig)
def _plot_shuffle(normal: Sequence[Mapping[str, Any]], shuffled: Sequence[Mapping[str, Any]], path: Path) -> None:
pairs = [f"{left}_{right}" for left, right in PAIRINGS]
normal_means = []
shuffle_means = []
for pair in pairs:
nr = [row for row in normal if row["pair"] == pair]
sr = [row for row in shuffled if row["pair"] == pair]
normal_means.append(_cluster_bootstrap(nr, "matched_vs_shifted_auc", seed=921)[0])
shuffle_means.append(_cluster_bootstrap(sr, "matched_vs_shifted_auc", seed=922)[0])
positions = np.arange(len(pairs))
fig, axis = plt.subplots(figsize=(8, 4.6))
width = 0.34
axis.bar(positions - width / 2, normal_means, width, label="Content as aligned", color="#2878b5")
axis.bar(positions + width / 2, shuffle_means, width, label="Independent slot shuffle", color="#c44e52")
axis.axhline(0.5, color="black", linestyle="--", linewidth=0.9, label="Chance AUC")
axis.set_xticks(positions, [pair.replace("_", "–") for pair in pairs])
axis.set_ylim(0.35, 0.8)
axis.set_ylabel("Matched-vs-shifted AUC")
axis.set_title("Content shuffle control")
axis.legend(frameon=False)
axis.grid(axis="y", alpha=0.2)
fig.tight_layout()
fig.savefig(path, dpi=180, bbox_inches="tight")
plt.close(fig)
def _plot_reconstruction_gain(rows: Sequence[Mapping[str, Any]], path: Path) -> None:
targets = ("audio", "vision", "text")
colors = {"audio": "#2878b5", "vision": "#e1812c", "text": "#55a868"}
fig, axis = plt.subplots(figsize=(8, 5))
for target in targets:
selected = [row for row in rows if row["target_modality"] == target]
by_abs: dict[int, list[dict[str, Any]]] = defaultdict(list)
for row in selected:
by_abs[int(row["abs_delta"])].append(dict(row))
xs = sorted(by_abs)
means, lows, highs = [], [], []
for distance in xs:
mean, low, high, _ = _cluster_bootstrap(
by_abs[distance],
"gain_shift_minus_aligned",
seed=991 + distance + sum(ord(char) for char in target),
repetitions=1000,
)
means.append(mean)
lows.append(low)
highs.append(high)
axis.plot(xs, means, marker="o", color=colors[target], label=f"Reconstruct {target}")
axis.fill_between(xs, lows, highs, color=colors[target], alpha=0.15)
axis.axhline(0, color="black", linestyle="--", linewidth=0.9)
axis.set_xlabel("Absolute temporal shift |Δ| (slots)")
axis.set_ylabel("MAE(shifted) − MAE(aligned)")
axis.set_title("Does the aligned partner help reconstruct the target content?")
axis.legend(frameon=False)
axis.grid(alpha=0.2)
fig.tight_layout()
fig.savefig(path, dpi=180, bbox_inches="tight")
plt.close(fig)
def run(args: argparse.Namespace) -> dict[str, Any]:
started = time.time()
if args.device == "auto":
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
else:
device = torch.device(args.device)
if device.type == "cuda" and not torch.cuda.is_available():
raise RuntimeError("CUDA was requested but is unavailable")
samples = load_feature_samples(args.feature_dir, args.manifest)
by_id = {sample.sample_id: sample for sample in samples}
splits = json.loads(args.splits.read_text(encoding="utf-8"))
output_dir = args.output_dir
output_dir.mkdir(parents=True, exist_ok=True)
self_rows: list[dict[str, Any]] = []
pair_rows: list[dict[str, Any]] = []
cycle_triangle_rows: list[dict[str, Any]] = []
content_rows: list[dict[str, Any]] = []
content_curve_rows: list[dict[str, Any]] = []
shuffle_rows: list[dict[str, Any]] = []
reconstruction_rows: list[dict[str, Any]] = []
probe_history: list[dict[str, Any]] = []
probe_checkpoints: dict[str, Any] = {}
fold_manifests = []
example_id = args.example_id
example_arrays: dict[str, np.ndarray] | None = None
example_times: dict[str, np.ndarray] | None = None
example_pair_maps: dict[str, tuple[np.ndarray, np.ndarray, np.ndarray]] | None = None
example_self_grams: dict[str, np.ndarray] | None = None
example_cycles: dict[str, np.ndarray] | None = None
example_weights: dict[str, np.ndarray] | None = None
example_valid: dict[str, np.ndarray] | None = None
for split in splits:
fold = int(split["fold"])
train_samples = [by_id[sample_id] for sample_id in split["train_sample_ids"]]
validation_samples = [by_id[sample_id] for sample_id in split["validation_sample_ids"]]
train_groups = {sample.group_id for sample in train_samples}
validation_groups = {sample.group_id for sample in validation_samples}
if train_groups & validation_groups:
raise ValueError(f"video_id leakage in fold {fold}")
if example_id is not None and example_id not in {sample.sample_id for sample in validation_samples} and fold == 1:
# The requested example can be left out; evaluation still covers all held-out clips.
example_id = None
stats = fit_feature_stats(train_samples)
weights_by_id, content_by_id = _collect_fold(
fold=fold,
train_samples=train_samples,
validation_samples=validation_samples,
stats=stats,
checkpoint_root=args.checkpoint_root,
device=device,
batch_size=args.batch_size,
)
fold_manifests.append(
{
"fold": fold,
"train_count": len(train_samples),
"heldout_count": len(validation_samples),
"train_video_ids": sorted(train_groups),
"heldout_video_ids": sorted(validation_groups),
"overlap": sorted(train_groups & validation_groups),
}
)
print(
f"[M4 evaluation fold {fold}] train={len(train_samples)} heldout={len(validation_samples)} "
f"video_ids={len(train_groups)}/{len(validation_groups)}",
flush=True,
)
for sample in validation_samples:
maps, normalized_times, sample_rows, arrays = _pairwise_for_sample(
sample, weights_by_id[sample.sample_id]
)
for row in sample_rows:
row.update({"method": "M4_sourceTime", "fold": fold, "sample_id": sample.sample_id, "video_id": sample.group_id})
if row["kind"] == "self":
self_rows.append(row)
elif row["kind"] == "pairwise":
pair_rows.append(row)
else:
cycle_triangle_rows.append(row)
if example_id == sample.sample_id or (example_id is None and sample.sample_id == "-3g5yACwYnA/13"):
example_id = sample.sample_id
example_pair_maps = maps
example_times = normalized_times
example_self_grams = {name: arrays[name] for name in MODALITIES}
example_cycles = {name: arrays[name] for name in arrays if name.startswith("cycle_") or name.startswith("triangle_")}
example_weights = weights_by_id[sample.sample_id]
example_valid = {name: np.asarray(sample.valid[name], dtype=bool) for name in MODALITIES}
example_arrays = {}
for left, right in (*PAIRINGS, *((right, left) for left, right in PAIRINGS)):
mapping, source_idx, destination_idx = maps[f"{left}_{right}"]
example_arrays[f"C_{left}_{right}"] = mapping
example_arrays[f"source_indices_{left}_{right}"] = source_idx
example_arrays[f"destination_indices_{left}_{right}"] = destination_idx
for name in MODALITIES:
example_arrays[f"G_{name}"] = arrays[name]
example_arrays[f"A_{name}"] = weights_by_id[sample.sample_id][name]
example_arrays[f"tau_{name}"] = normalized_times[name]
example_arrays[f"valid_{name}"] = np.asarray(sample.valid[name], dtype=bool)
example_arrays.update(example_cycles)
normal, curves, shuffled, probe_checkpoints = _fit_and_score_content(
fold=fold,
train_samples=train_samples,
validation_samples=validation_samples,
content_by_id=content_by_id,
device=device,
args=args,
probe_seeds=probe_checkpoints,
)
content_rows.extend(normal)
content_curve_rows.extend(curves)
shuffle_rows.extend(shuffled)
probe_history.extend(probe_checkpoints.pop("history", []))
reconstruction_rows.extend(
_fit_and_score_reconstruction(
fold=fold,
train_samples=train_samples,
validation_samples=validation_samples,
content_by_id=content_by_id,
device=device,
args=args,
checkpoint_store=probe_checkpoints,
history_store=probe_history,
)
)
print(
f"[M4 evaluation fold {fold}] content_probe={len(normal)} pair rows; "
f"shuffle repeats={args.shuffle_repeats}; reconstruction rows={len(reconstruction_rows)}",
flush=True,
)
del weights_by_id, content_by_id
if device.type == "cuda":
torch.cuda.empty_cache()
if example_arrays is not None:
np.savez_compressed(output_dir / "example_m4_shared_latent_maps.npz", **example_arrays)
assert example_pair_maps is not None and example_times is not None
assert example_self_grams is not None and example_cycles is not None
assert example_weights is not None and example_valid is not None
_plot_attention(example_weights, example_times, example_valid, output_dir / "latent_attention_heatmaps.png")
_plot_self_gram(example_self_grams, output_dir / "self_gram_heatmaps.png")
_plot_pairwise(example_pair_maps, example_times, output_dir / "pairwise_heatmaps.png")
_plot_pair_trajectories(example_pair_maps, example_times, output_dir / "pairwise_trajectories.png")
_plot_cycles(example_cycles, output_dir / "cycle_triangle_maps.png")
self_summary = _bootstrap_summary(
self_rows,
("method", "modality"),
("near_similarity", "near_similarity_offdiag", "far_similarity", "d_self", "d_self_offdiag", "gram_target_error", "far_slot_leakage", "c_row_offdiag"),
seed=args.seed,
)
pair_summary = _bootstrap_summary(
pair_rows,
("method", "direction"),
("pairwise_time_mae", "pairwise_signed_lag", "pairwise_time_corr"),
seed=args.seed + 1,
)
cycle_triangle_summary = _bootstrap_summary(
cycle_triangle_rows,
("method", "kind"),
("cycle_band_error", "cycle_time_mae", "triangle_relative_error", "triangle_mean_absolute_residual"),
seed=args.seed + 2,
)
content_summary, content_curve_summary = _content_summary(content_rows, content_curve_rows, args.seed + 3)
shuffle_summary = _bootstrap_summary(
shuffle_rows,
("method", "control", "pair"),
("same_time_similarity", "same_minus_shifted_margin", "matched_vs_shifted_auc",
"exact_r1_left_to_right", "within_pm1_r1_left_to_right", "mase_slots_left_to_right"),
seed=args.seed + 4,
)
reconstruction_summary = _bootstrap_summary(
reconstruction_rows,
("method", "target_modality", "shifted_modality", "abs_delta"),
("aligned_mae_same_support", "shifted_mae", "gain_shift_minus_aligned"),
seed=args.seed + 5,
)
_write_csv(output_dir / "self_structure_by_clip.csv", self_rows)
_write_csv(output_dir / "self_structure_summary.csv", self_summary)
_write_csv(output_dir / "pairwise_by_clip.csv", pair_rows)
_write_csv(output_dir / "pairwise_summary.csv", pair_summary)
_write_csv(output_dir / "cycle_triangle_by_clip.csv", cycle_triangle_rows)
_write_csv(output_dir / "cycle_triangle_summary.csv", cycle_triangle_summary)
_write_csv(output_dir / "content_only_metrics_by_clip.csv", content_rows)
_write_csv(output_dir / "content_only_metrics_summary.csv", content_summary)
_write_csv(output_dir / "content_shift_curve_by_clip.csv", content_curve_rows)
_write_csv(output_dir / "content_shift_curve_summary.csv", content_curve_summary)
_write_csv(output_dir / "content_shuffle_by_clip.csv", shuffle_rows)
_write_csv(output_dir / "content_shuffle_summary.csv", shuffle_summary)
_write_csv(output_dir / "shifted_reconstruction_by_clip.csv", reconstruction_rows)
_write_csv(output_dir / "shifted_reconstruction_summary.csv", reconstruction_summary)
_write_csv(output_dir / "probe_training_history.csv", probe_history)
torch.save(probe_checkpoints, output_dir / "content_probe_checkpoints.pt")
_plot_content_curve(content_curve_rows, output_dir / "content_only_shifted_similarity.png")
_plot_shuffle(content_rows, shuffle_rows, output_dir / "content_shuffle_control.png")
_plot_reconstruction_gain(reconstruction_rows, output_dir / "shifted_reconstruction_gain.png")
manifest = {
"created_utc": datetime.now(timezone.utc).isoformat(),
"experiment": "M4 Shared Latent Timeline re-evaluation",
"alignment_model_retrained": False,
"checkpoint_variant": "M4_sourceTime",
"sample_count": len(samples),
"fold_count": len(splits),
"folds": fold_manifests,
"example_sample_id": example_id,
"parameters": {
"seed": args.seed,
"batch_size": args.batch_size,
"grid_size": GRID_SIZE,
"hidden_size": HIDDEN_SIZE,
"heads": HEADS,
"sigma_self_normalized_time": SIGMA_SELF,
"near_radius_slots": STRUCTURE_RADIUS,
"far_radius_slots": FAR_RADIUS,
"sigma_cycle_normalized_time": SIGMA_CYCLE,
"content_projection_dimension": 64,
"content_probe_epochs": args.probe_epochs,
"content_probe_learning_rate": args.learning_rate,
"contrastive_temperature": args.temperature,
"content_shuffle_repeats": args.shuffle_repeats,
"reconstruction_probe_epochs": args.decoder_epochs,
"reconstruction_probe_learning_rate": args.decoder_learning_rate,
"reconstruction_shifts_slots": list(SHIFTS),
"probe_seed_rule": "seed + fold * 101",
"reconstruction_seed_rule": "seed + fold * 211",
},
"input_paths": {
"feature_dir": str(args.feature_dir.resolve()),
"feature_manifest": str(args.manifest.resolve()),
"grouped_splits": str(args.splits.resolve()),
"checkpoint_root": str(args.checkpoint_root.resolve()),
"alignment_checkpoints": [
str(_checkpoint_path(args.checkpoint_root, int(split["fold"])).resolve())
for split in splits
],
},
"content_only_definition": "M4 output.aligned = attention weights times projected source values; no latent query vector or positional embedding is concatenated into the probe input",
"functional_protocol": {
"probe_training": "train-fold only, symmetric within-clip same-slot InfoNCE; no emotion labels",
"shuffle": "independently permute each modality's 50 content rows within each held-out clip; keep row/slot index labels fixed; 20 permutations",
"reconstruction": "train-fold decoder predicts one M4 value stream from the other two; at evaluation shift one partner stream by signed offsets +/-1,2,5,10 and compare MAE on identical valid target slots",
"confidence_intervals": "95% cluster bootstrap over video_id groups, 2,000 repetitions for CSV summaries; 1,000 for figure ribbons",
},
"device": str(device),
"gpu_name": torch.cuda.get_device_name(device) if device.type == "cuda" else None,
"python": platform.python_version(),
"torch": torch.__version__,
"elapsed_seconds": time.time() - started,
"interpretation_limits": [
"Self Gram, pairwise maps, cycle, and triangle are structural consistency checks, not independent semantic ground truth.",
"The content projection and reconstruction decoder are learned on training video groups and evaluated on held-out groups; their results measure transferable probe utility.",
"The time-code M4 alignment checkpoint was trained with a timestamp-derived Gaussian prior, so its attention can still encode a position shortcut.",
"Human event IoU and center-error validation remain unavailable until event intervals are annotated.",
"The D5 alignment checkpoints use one alignment training seed; bootstrap intervals describe video-group sampling uncertainty, not seed uncertainty.",
],
}
(output_dir / "run_manifest.json").write_text(
json.dumps(manifest, ensure_ascii=False, indent=2, allow_nan=False), encoding="utf-8"
)
print(
f"[M4 evaluation complete] samples={len(samples)} folds={len(splits)} "
f"elapsed={manifest['elapsed_seconds']:.1f}s output={output_dir}",
flush=True,
)
return manifest
def _content_summary(
rows: Sequence[Mapping[str, Any]],
curves: Sequence[Mapping[str, Any]],
seed: int,
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
metrics = (
"same_time_similarity",
"shifted_far_similarity",
"same_minus_shifted_margin",
"matched_vs_shifted_auc",
"exact_r1_text_to_audio",
"within_pm1_r1_text_to_audio",
"mase_slots_text_to_audio",
"exact_r1_text_to_vision",
"within_pm1_r1_text_to_vision",
"mase_slots_text_to_vision",
"exact_r1_audio_to_vision",
"within_pm1_r1_audio_to_vision",
"mase_slots_audio_to_vision",
)
grouped: dict[str, list[Mapping[str, Any]]] = defaultdict(list)
for row in rows:
grouped[str(row["pair"])].append(row)
summary = []
for pair, values in grouped.items():
output: dict[str, Any] = {
"method": "M4_sourceTime_content",
"pair": pair,
"clip_count": len(values),
"video_id_count": len({str(row["video_id"]) for row in values}),
}
for i, metric in enumerate(metrics):
if metric not in values[0]:
continue
mean, low, high, count = _cluster_bootstrap(
values, metric, seed=seed + i + sum(ord(char) for char in pair)
)
output[f"{metric}_video_macro_mean"] = mean
output[f"{metric}_ci95_low"] = low
output[f"{metric}_ci95_high"] = high
output["video_id_count"] = count
summary.append(output)
curve_groups: dict[tuple[str, int], list[Mapping[str, Any]]] = defaultdict(list)
for row in curves:
curve_groups[(str(row["pair"]), int(row["delta"]))].append(row)
curve_summary = []
for (pair, delta), values in sorted(curve_groups.items()):
mean, low, high, count = _cluster_bootstrap(
values, "similarity", seed=seed + delta + sum(ord(char) for char in pair)
)
curve_summary.append(
{
"method": "M4_sourceTime_content",
"pair": pair,
"delta": delta,
"mean_similarity_video_macro": mean,
"ci95_low": low,
"ci95_high": high,
"video_id_count": count,
}
)
return summary, curve_summary
def build_parser() -> argparse.ArgumentParser:
project = Path(__file__).resolve().parents[1]
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto")
parser.add_argument("--batch-size", type=int, default=8)
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--probe-epochs", type=int, default=40)
parser.add_argument("--shuffle-repeats", type=int, default=20)
parser.add_argument("--decoder-epochs", type=int, default=40)
parser.add_argument("--learning-rate", type=float, default=1e-3)
parser.add_argument("--temperature", type=float, default=0.1)
parser.add_argument("--decoder-learning-rate", type=float, default=1e-3)
parser.add_argument("--example-id", type=str, default="-3g5yACwYnA/13")
parser.add_argument(
"--feature-dir", type=Path, default=project / "outputs/q1_features/features"
)
parser.add_argument("--manifest", type=Path, default=project / "outputs/audit/manifest.csv")
parser.add_argument("--splits", type=Path, default=project / "outputs/method_comparison/splits.json")
parser.add_argument(
"--checkpoint-root", type=Path, default=project / "outputs/alignment_debug/heldout"
)
parser.add_argument(
"--output-dir", type=Path, default=project / "outputs/m4_shared_latent_eval"
)
return parser
def main() -> None:
args = build_parser().parse_args()
run(args)
if __name__ == "__main__":
main()