Files
modeling_zhaocui/deep_learning/Q1/q1/shared_definition_comparison.py

1617 lines
77 KiB
Python

"""Compare similarity, correlation, and predictability as shared information.
All representation methods use the same fold-standardized original BERT,
eGeMAPS, and DeiT streams pooled with frozen M4_sourceTime attention.
Emotion labels are used only by the downstream, outer-fold probes.
"""
from __future__ import annotations
import argparse
import csv
import hashlib
import json
import platform
import random
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 sklearn
import torch
from sklearn.decomposition import PCA
from sklearn.linear_model import LogisticRegression, Ridge
from sklearn.metrics import accuracy_score, f1_score, mean_squared_error, r2_score
from sklearn.model_selection import GroupKFold
from sklearn.pipeline import make_pipeline
from sklearn.preprocessing import StandardScaler
from .experiment_data import FeatureSample, fit_feature_stats, load_feature_samples
from .tsfa_emotion_probe import _class_from_sentiment
from .tsfa_experiment import (
GRID_SIZE,
_collect_fold_features,
)
from .tsfa_shared_private import (
MODS,
SharedPrivateFactorizer,
_encode_all,
_pool_five,
_pool_original_source,
)
PREDICTIVE_CONTEXTS = {"same_slot": 0, "local_pm1": 1}
PREDICTIVE_TARGETS = MODS
PAIRS = (("text", "audio"), ("text", "vision"), ("audio", "vision"))
ALPHAS = (1.0, 10.0, 100.0)
PCA_COMPONENTS = {"text": 64, "audio": 24, "vision": 64}
CORE_SLOTS = np.arange(6, GRID_SIZE - 6)
MAIN_DIM_PER_SLOT = 256
MAIN_CLIP_DIM = 5 * MAIN_DIM_PER_SLOT
EMOTION_CLASSES = (0, 1, 2)
def _read_csv(path: Path) -> list[dict[str, str]]:
with path.open("r", encoding="utf-8-sig", newline="") as stream:
return list(csv.DictReader(stream))
def _r2(y_true: np.ndarray, y_pred: np.ndarray) -> float:
return float(r2_score(y_true, y_pred, multioutput="uniform_average", force_finite=True))
def _clip_r2_mse(y_true: np.ndarray, y_pred: np.ndarray, slots: Sequence[int] | None = None) -> tuple[float, float]:
if slots is not None:
y_true = y_true[slots]
y_pred = y_pred[slots]
return _r2(y_true, y_pred), float(mean_squared_error(y_true, y_pred))
def _ridge(alpha: float) -> Any:
# Cholesky solves the same L2-regularized least-squares objective exactly
# and handles the dense, multi-output source features much faster here.
return make_pipeline(StandardScaler(), Ridge(alpha=float(alpha), solver="cholesky"))
def _predictor_input(
sample_id: str,
target: str,
slot: int,
radius: int,
pooled: Mapping[str, Mapping[str, np.ndarray]],
*,
shift: int = 0,
permutations: Mapping[str, np.ndarray] | None = None,
) -> np.ndarray:
parts: list[np.ndarray] = []
for source_modality in MODS:
if source_modality == target:
continue
source = pooled[sample_id][source_modality]
for offset in range(-radius, radius + 1):
index = int(np.clip(slot + shift + offset, 0, GRID_SIZE - 1))
if permutations is not None:
index = int(permutations[source_modality][index])
parts.append(source[index])
return np.concatenate(parts, axis=0)
def _matrix_for_samples(
samples: Sequence[FeatureSample],
target: str,
radius: int,
pooled: Mapping[str, Mapping[str, np.ndarray]],
*,
slots: Sequence[int] | None = None,
) -> tuple[np.ndarray, np.ndarray, np.ndarray, list[tuple[str, int]]]:
use_slots = list(range(GRID_SIZE)) if slots is None else [int(slot) for slot in slots]
xs: list[np.ndarray] = []
ys: list[np.ndarray] = []
groups: list[str] = []
keys: list[tuple[str, int]] = []
for sample in samples:
sample_id = sample.sample_id
for slot in use_slots:
xs.append(_predictor_input(sample_id, target, slot, radius, pooled))
ys.append(pooled[sample_id][target][slot])
groups.append(sample.group_id)
keys.append((sample_id, slot))
return np.stack(xs), np.stack(ys), np.asarray(groups), keys
def _select_ridge_alpha(
x: np.ndarray,
y: np.ndarray,
groups: np.ndarray,
*,
alphas: Sequence[float] = ALPHAS,
) -> tuple[float, dict[float, float]]:
unique_groups = np.unique(groups)
folds = min(3, len(unique_groups))
if folds < 2:
raise ValueError("inner grouped CV requires at least two training video groups")
cv = GroupKFold(n_splits=folds)
scores: dict[float, float] = {}
for alpha in alphas:
fold_scores = []
for train_index, valid_index in cv.split(x, y, groups):
model = _ridge(alpha)
model.fit(x[train_index], y[train_index])
fold_scores.append(_r2(y[valid_index], model.predict(x[valid_index])))
scores[float(alpha)] = float(np.mean(fold_scores))
best = max(scores, key=lambda alpha: (scores[alpha], -alpha))
return best, scores
def _crossfit_predictions(x: np.ndarray, y: np.ndarray, groups: np.ndarray, alpha: float) -> np.ndarray:
unique_groups = np.unique(groups)
cv = GroupKFold(n_splits=min(3, len(unique_groups)))
output = np.empty_like(y, dtype=np.float32)
seen = np.zeros(len(y), dtype=bool)
for train_index, valid_index in cv.split(x, y, groups):
model = _ridge(alpha)
model.fit(x[train_index], y[train_index])
output[valid_index] = model.predict(x[valid_index]).astype(np.float32, copy=False)
seen[valid_index] = True
if not seen.all():
raise RuntimeError("cross-fit residual predictions do not cover all training slots")
return output
def _predict_sequences(
model: Any,
sample_ids: Sequence[str],
target: str,
radius: int,
pooled: Mapping[str, Mapping[str, np.ndarray]],
*,
shift: int = 0,
permutations_by_id: Mapping[str, Mapping[str, np.ndarray]] | None = None,
slots: Sequence[int] | None = None,
) -> dict[str, np.ndarray]:
use_slots = list(range(GRID_SIZE)) if slots is None else [int(slot) for slot in slots]
output: dict[str, np.ndarray] = {}
for sample_id in sample_ids:
perms = None if permutations_by_id is None else permutations_by_id[sample_id]
x = np.stack([
_predictor_input(sample_id, target, slot, radius, pooled,
shift=shift, permutations=perms)
for slot in use_slots
])
predictions = model.predict(x).astype(np.float32, copy=False)
sequence = np.full_like(pooled[sample_id][target], np.nan, dtype=np.float32)
sequence[use_slots] = predictions
output[sample_id] = sequence
return output
def _predictive_fold(
*,
fold: int,
context_name: str,
train_samples: Sequence[FeatureSample],
heldout_samples: Sequence[FeatureSample],
pooled: Mapping[str, Mapping[str, np.ndarray]],
seed: int,
shuffle_repeats: int,
) -> tuple[dict[str, dict[str, dict[str, np.ndarray]]], list[dict[str, Any]],
list[dict[str, Any]], list[dict[str, Any]], list[dict[str, Any]]]:
radius = PREDICTIVE_CONTEXTS[context_name]
train_ids = [sample.sample_id for sample in train_samples]
heldout_ids = [sample.sample_id for sample in heldout_samples]
train_representations = {sample_id: {"shared": {}, "private": {}} for sample_id in train_ids}
heldout_representations = {sample_id: {"shared": {}, "private": {}} for sample_id in heldout_ids}
summary_rows: list[dict[str, Any]] = []
control_rows: list[dict[str, Any]] = []
residual_rows: list[dict[str, Any]] = []
for target in PREDICTIVE_TARGETS:
x_train, y_train, groups_train, keys_train = _matrix_for_samples(
train_samples, target, radius, pooled
)
alpha, alpha_cv = _select_ridge_alpha(x_train, y_train, groups_train)
first_stage = _ridge(alpha)
first_stage.fit(x_train, y_train)
yhat_train_cf = _crossfit_predictions(x_train, y_train, groups_train, alpha)
yhat_valid = _predict_sequences(
first_stage, heldout_ids, target, radius, pooled
)
yhat_train: dict[str, np.ndarray] = {
sample_id: np.full_like(pooled[sample_id][target], np.nan, dtype=np.float32)
for sample_id in train_ids
}
for index, (sample_id, slot) in enumerate(keys_train):
yhat_train[sample_id][slot] = yhat_train_cf[index]
residual_y_train = y_train - yhat_train_cf
# Reuse the first-stage alpha so original and residual prediction use the
# same regularization budget; this avoids selecting a second hyperparameter
# on a noisier target residual.
residual_alpha, residual_alpha_cv = alpha, {}
residual_model = _ridge(residual_alpha)
residual_model.fit(x_train, residual_y_train)
p_hat_valid = _predict_sequences(
residual_model, heldout_ids, target, radius, pooled
)
for sample_id in train_ids:
train_representations[sample_id]["shared"][target] = yhat_train[sample_id]
train_representations[sample_id]["private"][target] = (
pooled[sample_id][target] - yhat_train[sample_id]
).astype(np.float32)
for sample_id in heldout_ids:
heldout_representations[sample_id]["shared"][target] = yhat_valid[sample_id]
heldout_representations[sample_id]["private"][target] = (
pooled[sample_id][target] - yhat_valid[sample_id]
).astype(np.float32)
original_true = []
original_pred = []
private_true = []
private_predictions = []
clip_metric_rows: list[dict[str, Any]] = []
for sample in heldout_samples:
sample_id = sample.sample_id
true = pooled[sample_id][target][CORE_SLOTS]
pred = yhat_valid[sample_id][CORE_SLOTS]
private_true_seq = pooled[sample_id][target] - yhat_valid[sample_id]
pred_private_seq = p_hat_valid[sample_id]
private_target = private_true_seq[CORE_SLOTS]
pred_private = pred_private_seq[CORE_SLOTS]
original_true.append(true)
original_pred.append(pred)
private_true.append(private_target)
private_predictions.append(pred_private)
clip_metric_rows.append({
"sample_id": sample_id,
"video_id": sample.group_id,
"original_r2": _r2(true, pred),
"original_mse": float(mean_squared_error(true, pred)),
"private_r2": _r2(private_target, pred_private),
"private_mse": float(mean_squared_error(private_target, pred_private)),
})
original_true_array = np.concatenate(original_true)
original_pred_array = np.concatenate(original_pred)
private_true_array = np.concatenate(private_true)
private_pred_array = np.concatenate(private_predictions)
summary_rows.append({
"fold": fold,
"context": context_name,
"target_modality": target,
"selected_alpha_original": alpha,
"inner_cv_r2_by_alpha": json.dumps({str(k): v for k, v in alpha_cv.items()}),
"selected_alpha_private": residual_alpha,
"private_inner_cv_r2_by_alpha": json.dumps({str(k): v for k, v in residual_alpha_cv.items()}),
"private_alpha_selection": "reused original-target grouped inner-CV choice",
"heldout_r2_original_from_other_modalities": _r2(original_true_array, original_pred_array),
"heldout_mse_original_from_other_modalities": float(mean_squared_error(original_true_array, original_pred_array)),
"heldout_r2_private_residual_from_other_modalities": _r2(private_true_array, private_pred_array),
"heldout_mse_private_residual_from_other_modalities": float(mean_squared_error(private_true_array, private_pred_array)),
"delta_r2_original_minus_private": _r2(original_true_array, original_pred_array) - _r2(private_true_array, private_pred_array),
})
for row in clip_metric_rows:
residual_rows.append({"fold": fold, "context": context_name,
"target_modality": target, **row,
"delta_r2_original_minus_private": row["original_r2"] - row["private_r2"]})
# The exact same trained predictor is evaluated at every source shift.
for delta in range(-5, 6):
shifted_predictions = _predict_sequences(
first_stage, heldout_ids, target, radius, pooled,
shift=delta, slots=CORE_SLOTS,
)
for sample in heldout_samples:
sample_id = sample.sample_id
true = pooled[sample_id][target][CORE_SLOTS]
pred = shifted_predictions[sample_id][CORE_SLOTS]
control_rows.append({
"fold": fold,
"context": context_name,
"target_modality": target,
"control_type": "shift",
"delta": delta,
"repeat": 0,
"sample_id": sample_id,
"video_id": sample.group_id,
"r2": _r2(true, pred),
"mse": float(mean_squared_error(true, pred)),
})
for repeat in range(shuffle_repeats):
perm_by_id = {
sample_id: {
name: np.random.default_rng(seed + fold * 1009 + repeat * 17 + sum(map(ord, sample_id + name))).permutation(GRID_SIZE)
for name in MODS if name != target
}
for sample_id in heldout_ids
}
shuffled_predictions = _predict_sequences(
first_stage, heldout_ids, target, radius, pooled,
shift=0, permutations_by_id=perm_by_id, slots=CORE_SLOTS,
)
for sample in heldout_samples:
sample_id = sample.sample_id
true = pooled[sample_id][target][CORE_SLOTS]
pred = shuffled_predictions[sample_id][CORE_SLOTS]
control_rows.append({
"fold": fold,
"context": context_name,
"target_modality": target,
"control_type": "within_video_shuffle",
"delta": 0,
"repeat": repeat,
"sample_id": sample_id,
"video_id": sample.group_id,
"r2": _r2(true, pred),
"mse": float(mean_squared_error(true, pred)),
})
return train_representations, heldout_representations, summary_rows, control_rows, residual_rows
def _flatten_sequences(
sample_ids: Sequence[str],
pooled: Mapping[str, Mapping[str, np.ndarray]],
modality: str,
) -> np.ndarray:
return np.concatenate([pooled[sample_id][modality] for sample_id in sample_ids], axis=0)
def _inverse_sqrt(matrix: np.ndarray, ridge: float) -> np.ndarray:
eigenvalues, eigenvectors = np.linalg.eigh(matrix.astype(np.float64, copy=False))
eigenvalues = np.maximum(eigenvalues, 0.0) + ridge
return (eigenvectors * (1.0 / np.sqrt(eigenvalues))[None, :]) @ eigenvectors.T
def _gcca_fold(
*,
fold: int,
train_samples: Sequence[FeatureSample],
heldout_samples: Sequence[FeatureSample],
pooled: Mapping[str, Mapping[str, np.ndarray]],
seed: int,
common_dim: int = 32,
ridge: float = 1e-3,
) -> tuple[dict[str, dict[str, dict[str, np.ndarray]]], list[dict[str, Any]], list[dict[str, Any]]]:
"""Fit fold-only regularized MAXVAR/GCCA and return common/private streams."""
train_ids = [sample.sample_id for sample in train_samples]
heldout_ids = [sample.sample_id for sample in heldout_samples]
all_ids = [*train_ids, *heldout_ids]
train_views: dict[str, np.ndarray] = {}
all_views: dict[str, np.ndarray] = {}
pca_rows: list[dict[str, Any]] = []
for modality in MODS:
x_train = _flatten_sequences(train_ids, pooled, modality)
x_all = _flatten_sequences(all_ids, pooled, modality)
n_components = min(PCA_COMPONENTS[modality], x_train.shape[0] - 1, x_train.shape[1])
pca = PCA(n_components=n_components, svd_solver="randomized", random_state=seed + fold)
scores_train_raw = pca.fit_transform(x_train)
score_scaler = StandardScaler().fit(scores_train_raw)
train_views[modality] = score_scaler.transform(scores_train_raw).astype(np.float64)
scores_all_raw = pca.transform(x_all)
all_views[modality] = score_scaler.transform(scores_all_raw).astype(np.float64)
pca_rows.append({
"fold": fold,
"modality": modality,
"pca_components": n_components,
"pca_explained_variance_ratio_sum": float(pca.explained_variance_ratio_.sum()),
})
n_train = train_views[MODS[0]].shape[0]
whitened_train: list[np.ndarray] = []
covariances: dict[str, np.ndarray] = {}
for modality in MODS:
x = train_views[modality]
covariance = (x.T @ x) / n_train
covariances[modality] = covariance
whitened_train.append((x @ _inverse_sqrt(covariance, ridge)) / np.sqrt(n_train))
stacked = np.concatenate(whitened_train, axis=1)
left_vectors, singular_values, _ = np.linalg.svd(stacked, full_matrices=False)
dimension = min(common_dim, left_vectors.shape[1])
z_train_target = left_vectors[:, :dimension] * np.sqrt(n_train)
energy = singular_values**2
energy_ratio = energy / max(float(energy.sum()), 1e-12)
cumulative = np.cumsum(energy_ratio)
components_90 = int(np.searchsorted(cumulative, 0.90) + 1)
effective_rank = float(np.exp(-np.sum(energy_ratio * np.log(np.maximum(energy_ratio, 1e-12)))))
shared_train: dict[str, np.ndarray] = {}
shared_all: dict[str, np.ndarray] = {}
for modality in MODS:
width = train_views[modality].shape[1]
x_train = train_views[modality]
# Ridge-regularized least squares maps each view into one common latent basis.
weight = np.linalg.solve(
covariances[modality] + ridge * np.eye(width),
(x_train.T @ z_train_target) / n_train,
)
all_flat = all_views[modality]
shared_all[modality] = (all_flat @ weight).astype(np.float32)
shared_train[modality] = shared_all[modality][:n_train]
sequences: dict[str, dict[str, dict[str, np.ndarray]]] = {
sample_id: {"shared": {}, "private": {}} for sample_id in all_ids
}
explained_rows: list[dict[str, Any]] = []
for sample_index, sample_id in enumerate(all_ids):
slot_start = sample_index * GRID_SIZE
slot_stop = slot_start + GRID_SIZE
for modality in MODS:
sequences[sample_id]["shared"][modality] = shared_all[modality][slot_start:slot_stop]
z_common_train = np.mean([shared_train[name] for name in MODS], axis=0)
z_common_all = np.mean([shared_all[name] for name in MODS], axis=0)
for sample_index, sample_id in enumerate(all_ids):
slot_start = sample_index * GRID_SIZE
slot_stop = slot_start + GRID_SIZE
sequences[sample_id]["common"] = z_common_all[slot_start:slot_stop].astype(np.float32)
for modality in MODS:
target_train = _flatten_sequences(train_ids, pooled, modality)
target_all = _flatten_sequences(all_ids, pooled, modality)
decoder = make_pipeline(StandardScaler(), Ridge(alpha=10.0, solver="lsqr"))
decoder.fit(z_common_train, target_train)
prediction_all = decoder.predict(z_common_all).astype(np.float32)
residual_all = target_all - prediction_all
for sample_index, sample_id in enumerate(all_ids):
slot_start = sample_index * GRID_SIZE
slot_stop = slot_start + GRID_SIZE
sequences[sample_id]["private"][modality] = residual_all[slot_start:slot_stop]
heldout_start = len(train_ids) * GRID_SIZE
heldout_target = target_all[heldout_start:]
heldout_prediction = prediction_all[heldout_start:]
explained_rows.append({
"fold": fold,
"modality": modality,
"heldout_shared_reconstruction_r2": _r2(heldout_target, heldout_prediction),
"heldout_shared_reconstruction_mse": float(mean_squared_error(heldout_target, heldout_prediction)),
"heldout_private_residual_variance": float(np.var(heldout_target - heldout_prediction)),
})
pair_rows: list[dict[str, Any]] = []
val_start = len(train_ids) * GRID_SIZE
for left, right in PAIRS:
left_matrix = shared_all[left][val_start:]
right_matrix = shared_all[right][val_start:]
correlations = []
for component in range(dimension):
x = left_matrix[:, component]
y = right_matrix[:, component]
if np.std(x) > 1e-12 and np.std(y) > 1e-12:
correlations.append(float(np.corrcoef(x, y)[0, 1]))
pair_rows.append({
"fold": fold,
"pair": f"{left}-{right}",
"shared_dimension": dimension,
"heldout_mean_component_correlation": float(np.mean(correlations)) if correlations else float("nan"),
"heldout_median_component_correlation": float(np.median(correlations)) if correlations else float("nan"),
"valid_component_count": len(correlations),
})
stability_row = {
"fold": fold,
"metric": "spectrum_stability",
"shared_dimension": dimension,
"components_for_90pct_energy": components_90,
"effective_rank": effective_rank,
"gcca_ridge": ridge,
"singular_energy_top10": json.dumps([float(v) for v in energy_ratio[:10]]),
"pca_components_text": PCA_COMPONENTS["text"],
"pca_components_audio": PCA_COMPONENTS["audio"],
"pca_components_vision": PCA_COMPONENTS["vision"],
}
return sequences, [stability_row, *pair_rows, *pca_rows], explained_rows
def _spr_fold(
*,
fold: int,
sample_ids: Sequence[str],
pooled: Mapping[str, Mapping[str, np.ndarray]],
device: torch.device,
checkpoint_root: Path,
batch_size: int,
) -> dict[str, dict[str, dict[str, np.ndarray]]]:
checkpoint_path = checkpoint_root / f"fold_{fold:02d}" / "SPR.pt"
if not checkpoint_path.is_file():
raise FileNotFoundError(f"frozen SPR checkpoint missing: {checkpoint_path}")
dimensions = {name: pooled[sample_ids[0]][name].shape[-1] for name in MODS}
model = SharedPrivateFactorizer(dimensions).to(device)
checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False)
if checkpoint.get("variant") != "SPR":
raise ValueError(f"expected SPR checkpoint at {checkpoint_path}, got {checkpoint.get('variant')}")
model.load_state_dict(checkpoint.get("model_state_dict", checkpoint.get("state_dict", checkpoint)), strict=True)
model.eval()
encoded, _ = _encode_all(model, sample_ids, pooled, device, batch_size)
del model, checkpoint
if device.type == "cuda":
torch.cuda.empty_cache()
return encoded
def _representation_views(
representation: Mapping[str, Mapping[str, np.ndarray]],
) -> dict[str, np.ndarray]:
shared = representation["shared"]
private = representation["private"]
shared_unfused = np.concatenate([shared[name] for name in MODS], axis=-1)
private_all = np.concatenate([private[name] for name in MODS], axis=-1)
private_av = np.concatenate([private["audio"], private["vision"]], axis=-1)
views = {
"shared_only": shared_unfused,
"private_only": private_all,
"private_text": private["text"],
"private_audio": private["audio"],
"private_vision": private["vision"],
"private_audio_vision": private_av,
"private_all": private_all,
"shared+private": np.concatenate([shared_unfused, private_all], axis=-1),
"shared+private_text": np.concatenate([shared_unfused, private["text"]], axis=-1),
"shared+private_audio": np.concatenate([shared_unfused, private["audio"]], axis=-1),
"shared+private_vision": np.concatenate([shared_unfused, private["vision"]], axis=-1),
"shared+private_audio_vision": np.concatenate([shared_unfused, private_av], axis=-1),
"shared+private_all": np.concatenate([shared_unfused, private_all], axis=-1),
}
dims = [shared[name].shape[-1] for name in MODS]
if len(set(dims)) == 1:
views["shared_fused"] = np.mean([shared[name] for name in MODS], axis=0)
if "common" in representation:
views["shared_only"] = representation["common"]
views["shared_fused"] = representation["common"]
views["shared+private"] = np.concatenate([representation["common"], private_all], axis=-1)
views["shared+private_all"] = views["shared+private"]
views["main"] = views["shared+private"]
return views
def _dimension_matched_sequences(
train_samples: Sequence[FeatureSample],
heldout_samples: Sequence[FeatureSample],
sequences_by_id: Mapping[str, np.ndarray],
*,
seed: int,
components: int = MAIN_DIM_PER_SLOT,
) -> tuple[dict[str, np.ndarray], dict[str, np.ndarray], PCA]:
train_ids = [sample.sample_id for sample in train_samples]
heldout_ids = [sample.sample_id for sample in heldout_samples]
x_train = np.concatenate([sequences_by_id[sample_id] for sample_id in train_ids], axis=0)
x_valid = np.concatenate([sequences_by_id[sample_id] for sample_id in heldout_ids], axis=0)
n_components = min(components, x_train.shape[0] - 1, x_train.shape[1])
pca = PCA(n_components=n_components, svd_solver="randomized", random_state=seed)
pca.fit(x_train)
train_projected = pca.transform(x_train).astype(np.float32)
valid_projected = pca.transform(x_valid).astype(np.float32)
train_by_id = {
sample_id: train_projected[index * GRID_SIZE : (index + 1) * GRID_SIZE]
for index, sample_id in enumerate(train_ids)
}
valid_by_id = {
sample_id: valid_projected[index * GRID_SIZE : (index + 1) * GRID_SIZE]
for index, sample_id in enumerate(heldout_ids)
}
return train_by_id, valid_by_id, pca
def _raw_private_pca_sequences(
sample_ids: Sequence[str],
pooled: Mapping[str, Mapping[str, np.ndarray]],
checkpoint_path: Path,
) -> dict[str, np.ndarray]:
with np.load(checkpoint_path, allow_pickle=False) as archive:
mean = np.asarray(archive["mean"], dtype=np.float32)
components = np.asarray(archive["components"], dtype=np.float32)
output = {}
for sample_id in sample_ids:
raw = np.concatenate([pooled[sample_id][name] for name in MODS], axis=-1)
if raw.shape[1] != len(mean):
raise ValueError(f"raw-private PCA width mismatch for {sample_id}: {raw.shape[1]} vs {len(mean)}")
output[sample_id] = ((raw - mean) @ components.T).astype(np.float32)
return output
def _fit_probe_predictions(
*,
method: str,
view: str,
fold: int,
train_samples: Sequence[FeatureSample],
heldout_samples: Sequence[FeatureSample],
vectors_by_id: Mapping[str, np.ndarray],
seed: int,
) -> list[dict[str, Any]]:
train_x = np.stack([_pool_five(vectors_by_id[sample.sample_id]) for sample in train_samples])
heldout_x = np.stack([_pool_five(vectors_by_id[sample.sample_id]) for sample in heldout_samples])
train_class = np.asarray([_class_from_sentiment(sample.sentiment) for sample in train_samples], dtype=np.int64)
heldout_class = np.asarray([_class_from_sentiment(sample.sentiment) for sample in heldout_samples], dtype=np.int64)
train_value = np.asarray([sample.sentiment for sample in train_samples], dtype=np.float64)
heldout_value = np.asarray([sample.sentiment for sample in heldout_samples], dtype=np.float64)
classifier = make_pipeline(
StandardScaler(),
LogisticRegression(C=0.05, max_iter=5000, solver="lbfgs", random_state=seed),
)
classifier.fit(train_x, train_class)
predicted_class = classifier.predict(heldout_x)
regressor = make_pipeline(StandardScaler(), Ridge(alpha=25.0))
regressor.fit(train_x, train_value)
predicted_unclipped = regressor.predict(heldout_x)
predicted_value = np.clip(predicted_unclipped, -3.0, 3.0)
dimension = int(train_x.shape[1])
return [
{
"method": method,
"view": view,
"fold": fold,
"sample_id": sample.sample_id,
"video_id": sample.group_id,
"true_class_id": int(heldout_class[index]),
"true_class": ("Negative", "Neutral", "Positive")[int(heldout_class[index])],
"predicted_class_id": int(predicted_class[index]),
"predicted_class": ("Negative", "Neutral", "Positive")[int(predicted_class[index])],
"true_label": float(heldout_value[index]),
"predicted_label": float(predicted_value[index]),
"predicted_label_unclipped": float(predicted_unclipped[index]),
"feature_dimension": dimension,
}
for index, sample in enumerate(heldout_samples)
]
def _summarize_probe_rows(prediction_rows: Sequence[Mapping[str, Any]]) -> list[dict[str, Any]]:
keys = sorted({(str(row["method"]), str(row["view"])) for row in prediction_rows})
summaries: list[dict[str, Any]] = []
for method, view in keys:
rows = [row for row in prediction_rows if row["method"] == method and row["view"] == view]
y_class = np.asarray([int(row["true_class_id"]) for row in rows], dtype=np.int64)
p_class = np.asarray([int(row["predicted_class_id"]) for row in rows], dtype=np.int64)
y_value = np.asarray([float(row["true_label"]) for row in rows], dtype=np.float64)
p_value = np.asarray([float(row["predicted_label"]) for row in rows], dtype=np.float64)
per_fold_f1 = []
for fold in sorted({int(row["fold"]) for row in rows if row.get("fold") not in (None, "")}):
selected = [row for row in rows if int(row["fold"]) == fold]
per_fold_f1.append(float(f1_score(
[int(row["true_class_id"]) for row in selected],
[int(row["predicted_class_id"]) for row in selected],
labels=list(EMOTION_CLASSES), average="macro", zero_division=0,
)))
pearson = float(np.corrcoef(y_value, p_value)[0, 1]) if np.std(y_value) > 1e-12 and np.std(p_value) > 1e-12 else float("nan")
summaries.append({
"method": method,
"view": view,
"sample_count": len(rows),
"feature_dimension": int(rows[0].get("feature_dimension", -1) or -1),
"accuracy": float(accuracy_score(y_class, p_class)),
"macro_f1": float(f1_score(y_class, p_class, labels=list(EMOTION_CLASSES), average="macro", zero_division=0)),
"macro_f1_fold_mean": float(np.mean(per_fold_f1)) if per_fold_f1 else float("nan"),
"macro_f1_fold_sd": float(np.std(per_fold_f1, ddof=1)) if len(per_fold_f1) > 1 else 0.0,
"mae": float(np.mean(np.abs(y_value - p_value))),
"pearson": pearson,
})
return summaries
def _paired_video_bootstrap(
candidate_rows: Sequence[Mapping[str, Any]],
reference_rows: Sequence[Mapping[str, Any]],
*,
candidate_name: str,
reference_name: str,
repeats: int,
seed: int,
) -> list[dict[str, Any]]:
candidate = {str(row["sample_id"]): row for row in candidate_rows}
reference = {str(row["sample_id"]): row for row in reference_rows}
common = sorted(set(candidate) & set(reference))
if len(common) < 2:
raise ValueError(f"paired contrast has only {len(common)} shared sample IDs: {candidate_name} vs {reference_name}")
groups = sorted({str(candidate[sample_id]["video_id"]) for sample_id in common})
ids_by_group = {
group: [sample_id for sample_id in common if str(candidate[sample_id]["video_id"]) == group]
for group in groups
}
metrics: dict[str, Any] = {
"macro_f1": lambda rows: f1_score(
[int(row["true_class_id"]) for row in rows],
[int(row["predicted_class_id"]) for row in rows],
labels=list(EMOTION_CLASSES), average="macro", zero_division=0,
),
"mae": lambda rows: float(np.mean(np.abs(
np.asarray([float(row["true_label"]) for row in rows])
- np.asarray([float(row["predicted_label"]) for row in rows])
))),
"pearson": lambda rows: float(np.corrcoef(
[float(row["true_label"]) for row in rows],
[float(row["predicted_label"]) for row in rows],
)[0, 1]),
}
candidate_aligned = [candidate[sample_id] for sample_id in common]
reference_aligned = [reference[sample_id] for sample_id in common]
rng = np.random.default_rng(seed)
boot: dict[str, np.ndarray] = {name: np.empty(repeats, dtype=np.float64) for name in metrics}
for repeat in range(repeats):
chosen_groups = rng.choice(groups, size=len(groups), replace=True)
sampled_ids = [sample_id for group in chosen_groups for sample_id in ids_by_group[str(group)]]
c_rows = [candidate[sample_id] for sample_id in sampled_ids]
r_rows = [reference[sample_id] for sample_id in sampled_ids]
for name, metric in metrics.items():
try:
boot[name][repeat] = float(metric(c_rows) - metric(r_rows))
except (FloatingPointError, ValueError, ZeroDivisionError):
boot[name][repeat] = np.nan
observed = {name: float(metric(candidate_aligned) - metric(reference_aligned)) for name, metric in metrics.items()}
rows = []
for name, values in boot.items():
finite = values[np.isfinite(values)]
rows.append({
"candidate": candidate_name,
"reference": reference_name,
"metric": name,
"delta_candidate_minus_reference": observed[name],
"ci95_low": float(np.quantile(finite, 0.025)) if len(finite) else float("nan"),
"ci95_high": float(np.quantile(finite, 0.975)) if len(finite) else float("nan"),
"bootstrap_repeats": repeats,
"video_group_count": len(groups),
"paired_sample_count": len(common),
})
return rows
def _summarize_shift_controls(rows: Sequence[Mapping[str, Any]]) -> list[dict[str, Any]]:
grouped: dict[tuple[str, str, str, int], dict[str, list[float]]] = defaultdict(lambda: defaultdict(list))
for row in rows:
key = (str(row["context"]), str(row["target_modality"]), str(row["control_type"]), int(row["delta"]))
value = float(row["mse"])
if np.isfinite(value):
grouped[key][str(row["video_id"])].append(value)
per_video: dict[tuple[str, str, str, int], dict[str, float]] = {
key: {video: float(np.mean(values)) for video, values in by_video.items()}
for key, by_video in grouped.items()
}
baseline_by_key: dict[tuple[str, str, str], dict[str, float]] = {}
for (context, target, control, delta), values in per_video.items():
if control == "shift" and delta == 0:
baseline_by_key[(context, target, control)] = values
output = []
for (context, target, control, delta), video_mse in sorted(per_video.items()):
baseline = baseline_by_key.get((context, target, "shift"), {})
common_videos = sorted(set(video_mse) & set(baseline))
paired_delta = [video_mse[video] - baseline[video] for video in common_videos]
paired_ratio = [video_mse[video] / baseline[video] for video in common_videos if baseline[video] > 1e-12]
output.append({
"context": context,
"target_modality": target,
"control_type": control,
"delta": delta,
"video_macro_mean_mse": float(np.mean(list(video_mse.values()))),
"video_macro_sd_mse": float(np.std(list(video_mse.values()), ddof=1)) if len(video_mse) > 1 else 0.0,
"delta_mse_vs_same_slot": float(np.mean(paired_delta)) if paired_delta else float("nan"),
"mse_ratio_vs_same_slot": float(np.mean(paired_ratio)) if paired_ratio else float("nan"),
"video_count": len(video_mse),
"sample_count": len({str(row["sample_id"]) for row in rows if row["context"] == context and row["target_modality"] == target and row["control_type"] == control and int(row["delta"]) == delta}),
"control_row_count": sum(row["context"] == context and row["target_modality"] == target and row["control_type"] == control and int(row["delta"]) == delta for row in rows),
})
return output
def _load_reference_predictions(
*,
tsfa_predictions_path: Path,
math_predictions_path: Path,
math_splits_path: Path,
expected_fold_by_id: Mapping[str, int],
samples_by_id: Mapping[str, FeatureSample],
) -> list[dict[str, Any]]:
rows: list[dict[str, Any]] = []
tsfa_rows = _read_csv(tsfa_predictions_path)
selected_tsfa = [
row for row in tsfa_rows
if row.get("method") == "TSFA+RawPrivate" and row.get("view") == "all_modalities"
]
if len(selected_tsfa) != len(samples_by_id):
raise ValueError(
f"expected 100 TSFA+RawPrivate OOF rows in {tsfa_predictions_path}, found {len(selected_tsfa)}"
)
for row in selected_tsfa:
sample_id = row["sample_id"]
if sample_id not in expected_fold_by_id:
raise ValueError(f"TSFA+RawPrivate has unexpected sample ID: {sample_id}")
fold = int(row["fold"])
if fold != expected_fold_by_id[sample_id]:
raise ValueError(f"TSFA+RawPrivate fold mismatch for {sample_id}: {fold} vs {expected_fold_by_id[sample_id]}")
rows.append({
"method": "TSFA+RawPrivate",
"view": "main_external_nonmatched",
"fold": fold,
"sample_id": sample_id,
"video_id": samples_by_id[sample_id].group_id,
"true_class_id": int(row["true_class_id"]),
"true_class": row["true_class"],
"predicted_class_id": int(row["predicted_class_id"]),
"predicted_class": row["predicted_class"],
"true_label": float(row["true_label"]),
"predicted_label": float(row["predicted_label"]),
"predicted_label_unclipped": float(row.get("predicted_label_unclipped", row["predicted_label"])),
"feature_dimension": int(row.get("feature_dimension", 6845)),
})
math_split_rows = _read_csv(math_splits_path)
math_fold_by_id = {
row["sample_id"]: int(row["fold"])
for row in math_split_rows
if row.get("split", "valid_oof").lower() in {"valid_oof", "validation", "valid"}
}
if set(math_fold_by_id) != set(expected_fold_by_id):
raise ValueError("math OOF split assignment does not cover the same 100 sample IDs")
mismatched = [sample_id for sample_id in expected_fold_by_id if math_fold_by_id[sample_id] != expected_fold_by_id[sample_id]]
if mismatched:
raise ValueError(f"math OOF folds differ from Q1 fixed folds, examples: {mismatched[:5]}")
math_rows = _read_csv(math_predictions_path)
if len(math_rows) != len(samples_by_id):
raise ValueError(f"expected one math OOF row per sample, found {len(math_rows)}")
for source in math_rows:
sample_id = source["sample_id"]
if sample_id not in expected_fold_by_id:
raise ValueError(f"math OOF predictions have unexpected sample ID: {sample_id}")
sample = samples_by_id[sample_id]
if int(source["true_polarity"]) != _class_from_sentiment(sample.sentiment):
raise ValueError(f"math true class disagrees with sign-derived label for {sample_id}")
for baseline in ("B0", "B4"):
rows.append({
"method": f"Math-{baseline}",
"view": "math_all_modalities",
"fold": expected_fold_by_id[sample_id],
"sample_id": sample_id,
"video_id": sample.group_id,
"true_class_id": int(source["true_polarity"]),
"true_class": source["true_polarity_name"],
"predicted_class_id": int(source[f"{baseline}_predicted_polarity"]),
"predicted_class": source[f"{baseline}_predicted_polarity"],
"true_label": float(source["true_sentiment"]),
"predicted_label": float(source[f"{baseline}_predicted_sentiment"]),
"predicted_label_unclipped": float(source[f"{baseline}_predicted_sentiment"]),
"feature_dimension": -1,
})
return rows
def _write_rows(path: Path, rows: Sequence[Mapping[str, Any]]) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
fieldnames: list[str] = []
for row in rows:
for key in row:
if key not in fieldnames:
fieldnames.append(key)
if not fieldnames:
path.write_text("", encoding="utf-8")
return
with path.open("w", encoding="utf-8-sig", newline="") as stream:
writer = csv.DictWriter(stream, fieldnames=fieldnames, extrasaction="ignore")
writer.writeheader()
for row in rows:
cooked = {
key: json.dumps(value, ensure_ascii=False) if isinstance(value, (dict, list, tuple)) else value
for key, value in row.items()
}
writer.writerow(cooked)
def _write_json(path: Path, value: Any) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(json.dumps(value, ensure_ascii=False, indent=2, allow_nan=True) + "\n", encoding="utf-8")
def _make_private_ablation_rows(metrics: Sequence[Mapping[str, Any]]) -> list[dict[str, Any]]:
metric_map = {(row["method"], row["view"]): row for row in metrics}
rows = []
for method in sorted({str(row["method"]) for row in metrics if row["method"] in {
"Similarity-SPR", "GCCA", "Predictive-Same", "Predictive-Local1"
}}):
base = metric_map.get((method, "shared_only"))
for view in (
"private_text", "private_audio", "private_vision", "private_audio_vision", "private_all",
"shared+private_text", "shared+private_audio", "shared+private_vision",
"shared+private_audio_vision", "shared+private", "main_1280d",
):
row = metric_map.get((method, view))
if row is None:
continue
rows.append({
"method": method,
"view": view,
"accuracy": row["accuracy"],
"macro_f1": row["macro_f1"],
"mae": row["mae"],
"pearson": row["pearson"],
"macro_f1_delta_vs_shared_only": float(row["macro_f1"] - base["macro_f1"]) if base else float("nan"),
"mae_delta_vs_shared_only": float(row["mae"] - base["mae"]) if base else float("nan"),
"feature_dimension": row["feature_dimension"],
})
return rows
def _save_figures(
output_dir: Path,
*,
predictive_rows: Sequence[Mapping[str, Any]],
shift_rows: Sequence[Mapping[str, Any]],
gcca_rows: Sequence[Mapping[str, Any]],
explained_rows: Sequence[Mapping[str, Any]],
emotion_metrics: Sequence[Mapping[str, Any]],
dimension_metrics: Sequence[Mapping[str, Any]],
ablation_rows: Sequence[Mapping[str, Any]],
) -> list[str]:
figure_dir = output_dir / "figures"
figure_dir.mkdir(parents=True, exist_ok=True)
created: list[str] = []
fig, axes = plt.subplots(1, 3, figsize=(13, 3.8))
titles = ("Similarity", "Correlation", "Predictability")
captions = (
"Map each modality into a shared space\nthen compare same-slot vectors",
"Find projections whose paired values\nco-vary across modalities",
"Predict one modality from the others;\nresidual is source-specific content",
)
for axis, title, caption in zip(axes, titles, captions):
axis.set_title(title, fontsize=13, weight="bold")
axis.set_xlim(0, 1)
axis.set_ylim(0, 1)
axis.axis("off")
axis.text(0.5, 0.76, caption, ha="center", va="center", fontsize=9)
for y, label, color in ((0.48, "Text", "#4C78A8"), (0.34, "Audio", "#F58518"), (0.20, "Vision", "#54A24B")):
axes[0].text(0.08, y, label, color=color, weight="bold")
axes[0].scatter([0.40, 0.70], [y, y], s=100, color=color)
axes[0].annotate("", xy=(0.66, y), xytext=(0.44, y), arrowprops={"arrowstyle": "<->", "color": color})
axes[1].text(0.08, y, label, color=color, weight="bold")
axes[1].plot([0.36, 0.58, 0.80], [y, y + 0.035, y - 0.035], marker="o", color=color)
axes[2].text(0.04, y, label, color=color, weight="bold")
axes[2].annotate("", xy=(0.76, y), xytext=(0.32, y), arrowprops={"arrowstyle": "->", "color": color, "lw": 2})
axes[2].text(0.51, y + 0.045, "predict", ha="center", fontsize=8, color=color)
fig.suptitle("Three testable meanings of shared information", fontsize=14)
fig.tight_layout()
path = figure_dir / "01_shared_definitions.png"
fig.savefig(path, dpi=170, bbox_inches="tight")
plt.close(fig)
created.append(path.name)
shift_rows_only = [row for row in shift_rows if row.get("control_type") == "shift"]
fig, axes = plt.subplots(1, 2, figsize=(12, 4), sharey=True)
colors = {"text": "#4C78A8", "audio": "#F58518", "vision": "#54A24B"}
for axis, context in zip(axes, ("same_slot", "local_pm1")):
for modality in MODS:
selected = [row for row in shift_rows_only if row["context"] == context and row["target_modality"] == modality]
if not selected:
continue
deltas = sorted({int(row["delta"]) for row in selected})
means = []
for delta in deltas:
by_video: dict[str, list[float]] = defaultdict(list)
for row in selected:
if int(row["delta"]) == delta:
by_video[str(row["video_id"])].append(float(row["mse"]))
means.append(float(np.mean([np.mean(values) for values in by_video.values()])))
axis.plot(deltas, means, marker="o", ms=3, label=modality.title(), color=colors[modality])
axis.axvline(0, color="black", lw=0.8, ls="--")
axis.set_title(context.replace("_", " "))
axis.set_xlabel("source shift in slots")
axis.grid(alpha=0.25)
axis.legend(frameon=False)
axes[0].set_ylabel("held-out MSE (lower is better)")
fig.suptitle("Predictive correspondence across source-time shifts")
fig.tight_layout()
path = figure_dir / "02_predictive_shift_curves.png"
fig.savefig(path, dpi=170, bbox_inches="tight")
plt.close(fig)
created.append(path.name)
contexts = sorted({str(row["context"]) for row in predictive_rows})
targets = list(MODS)
fig, axes = plt.subplots(1, len(contexts), figsize=(5 * len(contexts), 4), squeeze=False)
for axis, context in zip(axes[0], contexts):
selected = [row for row in predictive_rows if row["context"] == context]
for index, modality in enumerate(targets):
modality_rows = [item for item in selected if item["target_modality"] == modality]
if modality_rows:
axis.scatter([index - 0.12], [np.mean([float(row["heldout_r2_original_from_other_modalities"]) for row in modality_rows])], s=75, label="original", color="#4C78A8")
axis.scatter([index + 0.12], [np.mean([float(row["heldout_r2_private_residual_from_other_modalities"]) for row in modality_rows])], s=75, label="private residual", color="#E45756")
axis.axhline(0, color="black", lw=0.8)
axis.set_xticks(range(len(targets)), [v.title() for v in targets])
axis.set_title(context.replace("_", " "))
axis.grid(axis="y", alpha=0.25)
axis.set_ylabel("held-out R²")
handles, labels = axes[0, 0].get_legend_handles_labels()
if handles:
fig.legend(handles, labels, loc="upper center", ncol=2, frameon=False)
fig.suptitle("Can other modalities predict the private residual?", y=1.04)
fig.tight_layout()
path = figure_dir / "03_original_vs_private_predictability.png"
fig.savefig(path, dpi=170, bbox_inches="tight")
plt.close(fig)
created.append(path.name)
pair_rows = [row for row in gcca_rows if str(row.get("pair", "")) in {"text-audio", "text-vision", "audio-vision"}]
pairs = ["text-audio", "text-vision", "audio-vision"]
fig, axis = plt.subplots(figsize=(7, 4))
means = [np.mean([float(row["heldout_mean_component_correlation"]) for row in pair_rows if row["pair"] == pair]) for pair in pairs]
errors = [np.std([float(row["heldout_mean_component_correlation"]) for row in pair_rows if row["pair"] == pair], ddof=1) if sum(row["pair"] == pair for row in pair_rows) > 1 else 0 for pair in pairs]
axis.bar(pairs, means, yerr=errors, color=["#4C78A8", "#54A24B", "#F58518"], capsize=4)
axis.axhline(0, color="black", lw=0.8)
axis.set_ylabel("mean held-out canonical component correlation")
axis.set_title("Regularized GCCA: shared-view correlations by pair")
axis.grid(axis="y", alpha=0.25)
fig.tight_layout()
path = figure_dir / "04_gcca_pair_correlations.png"
fig.savefig(path, dpi=170, bbox_inches="tight")
plt.close(fig)
created.append(path.name)
branch_views = ("shared_only", "private_only", "shared+private")
methods = ("Similarity-SPR", "GCCA", "Predictive-Same", "Predictive-Local1")
fig, axis = plt.subplots(figsize=(10, 4.5))
width = 0.24
x = np.arange(len(methods))
for offset, view in enumerate(branch_views):
scores = [
next((float(row["macro_f1"]) for row in emotion_metrics if row["method"] == method and row["view"] == view), np.nan)
for method in methods
]
axis.bar(x + (offset - 1) * width, scores, width, label=view)
axis.set_xticks(x, methods, rotation=15, ha="right")
axis.set_ylabel("OOF Macro-F1")
axis.set_title("Emotion probe: shared-only, private-only, and both")
axis.legend(frameon=False)
axis.grid(axis="y", alpha=0.25)
fig.tight_layout()
path = figure_dir / "05_shared_private_emotion_probe.png"
fig.savefig(path, dpi=170, bbox_inches="tight")
plt.close(fig)
created.append(path.name)
main_rows = [row for row in dimension_metrics if row["view"] in {"main_1280d", "main_external_nonmatched", "math_all_modalities"}]
fig, axis = plt.subplots(figsize=(8, 5))
for row in main_rows:
axis.scatter(float(row["mae"]), float(row["pearson"]), s=75)
axis.annotate(str(row["method"]), (float(row["mae"]), float(row["pearson"])), xytext=(4, 4), textcoords="offset points", fontsize=8)
axis.set_xlabel("OOF MAE (lower is better)")
axis.set_ylabel("OOF Pearson (higher is better)")
axis.set_title("Emotion strength probe across frozen representations")
axis.grid(alpha=0.25)
fig.tight_layout()
path = figure_dir / "06_emotion_mae_pearson.png"
fig.savefig(path, dpi=170, bbox_inches="tight")
plt.close(fig)
created.append(path.name)
fixed_methods = ("RawPrivate-PCA", "Similarity-SPR", "GCCA", "Predictive-Same", "Predictive-Local1")
fig, axis = plt.subplots(figsize=(9, 4.5))
fixed_rows = [row for row in dimension_metrics if row["view"] == "main_1280d" and row["method"] in fixed_methods]
scores = [next((float(row["macro_f1"]) for row in fixed_rows if row["method"] == method), np.nan) for method in fixed_methods]
axis.bar(fixed_methods, scores, color=["#9D755D", "#4C78A8", "#72B7B2", "#F58518", "#E45756"])
axis.set_ylabel("OOF Macro-F1")
axis.set_title("Dimension-matched 1,280D emotion representations")
axis.tick_params(axis="x", labelrotation=18)
axis.grid(axis="y", alpha=0.25)
fig.tight_layout()
path = figure_dir / "07_dimension_matched_macro_f1.png"
fig.savefig(path, dpi=170, bbox_inches="tight")
plt.close(fig)
created.append(path.name)
ablation_methods = ("Similarity-SPR", "GCCA", "Predictive-Same", "Predictive-Local1")
sources = ("private_text", "private_audio", "private_vision", "private_audio_vision", "private_all")
fig, axes = plt.subplots(1, 2, figsize=(13, 4.5))
width = 0.16
x = np.arange(len(sources))
for index, method in enumerate(ablation_methods):
f1_values = [next((float(row["macro_f1"]) for row in ablation_rows if row["method"] == method and row["view"] == source), np.nan) for source in sources]
mae_values = [next((float(row["mae"]) for row in ablation_rows if row["method"] == method and row["view"] == source), np.nan) for source in sources]
axes[0].bar(x + (index - 1.5) * width, f1_values, width, label=method)
axes[1].bar(x + (index - 1.5) * width, mae_values, width, label=method)
for axis in axes:
axis.set_xticks(x, sources, rotation=25, ha="right")
axis.grid(axis="y", alpha=0.25)
axes[0].set_ylabel("OOF Macro-F1")
axes[1].set_ylabel("OOF MAE")
axes[0].set_title("Private-source classification")
axes[1].set_title("Private-source strength prediction")
axes[1].legend(fontsize=8, frameon=False)
fig.suptitle("Private modality ablation")
fig.tight_layout()
path = figure_dir / "08_private_source_ablation.png"
fig.savefig(path, dpi=170, bbox_inches="tight")
plt.close(fig)
created.append(path.name)
fig, axis = plt.subplots(figsize=(7.5, 4.5))
for context in ("same_slot", "local_pm1"):
values = []
deltas = list(range(-5, 6))
for delta in deltas:
selected = [row for row in shift_rows if row.get("control_type") == "shift" and row["context"] == context and int(row["delta"]) == delta]
by_video: dict[str, list[float]] = defaultdict(list)
for row in selected:
by_video[str(row["video_id"])].append(float(row["mse"]))
values.append(float(np.mean([np.mean(video_values) for video_values in by_video.values()])) if by_video else np.nan)
axis.plot(deltas, values, marker="o", label=context.replace("_", " "))
shuffle_by_video: dict[str, list[float]] = defaultdict(list)
for row in shift_rows:
if row.get("control_type") == "within_video_shuffle":
shuffle_by_video[str(row["video_id"])].append(float(row["mse"]))
shuffle_values = [float(np.mean(values)) for values in shuffle_by_video.values()]
if shuffle_values:
axis.axhline(float(np.mean(shuffle_values)), color="#E45756", ls="--", label="within-video shuffle")
axis.axvline(0, color="black", lw=0.8, ls=":")
axis.set_xlabel("source shift in slots")
axis.set_ylabel("held-out MSE, averaged over targets (lower is better)")
axis.set_title("Same-slot, shifted, and shuffled controls")
axis.legend(frameon=False)
axis.grid(alpha=0.25)
fig.tight_layout()
path = figure_dir / "09_shift_shuffle_controls.png"
fig.savefig(path, dpi=170, bbox_inches="tight")
plt.close(fig)
created.append(path.name)
stability_rows = [row for row in gcca_rows if row.get("metric") == "spectrum_stability"]
modality_rows = list(explained_rows)
fig, axes = plt.subplots(1, 2, figsize=(10, 4))
axes[0].plot([int(row["fold"]) for row in stability_rows], [float(row["effective_rank"]) for row in stability_rows], marker="o", label="effective rank")
axes[0].plot([int(row["fold"]) for row in stability_rows], [int(row["components_for_90pct_energy"]) for row in stability_rows], marker="s", label="components for 90%")
axes[0].set_xlabel("outer fold")
axes[0].set_title("GCCA shared-spectrum stability")
axes[0].legend(frameon=False, fontsize=8)
axes[0].grid(alpha=0.25)
for modality in MODS:
selected = [row for row in modality_rows if row["modality"] == modality]
axes[1].plot([int(row["fold"]) for row in selected], [float(row["heldout_shared_reconstruction_r2"]) for row in selected], marker="o", label=modality.title())
axes[1].set_xlabel("outer fold")
axes[1].set_title("Held-out variance explained from common GCCA")
axes[1].legend(frameon=False)
axes[1].grid(alpha=0.25)
fig.tight_layout()
path = figure_dir / "10_gcca_stability_and_variance.png"
fig.savefig(path, dpi=170, bbox_inches="tight")
plt.close(fig)
created.append(path.name)
return created
def run(args: argparse.Namespace) -> None:
started = time.time()
device = torch.device(args.device if args.device != "auto" else ("cuda" if torch.cuda.is_available() else "cpu"))
if device.type == "cuda" and not torch.cuda.is_available():
raise RuntimeError("CUDA was requested but is unavailable in this uv environment")
random.seed(args.seed)
np.random.seed(args.seed)
torch.manual_seed(args.seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(args.seed)
output_dir: Path = args.output_dir
output_dir.mkdir(parents=True, exist_ok=True)
samples = load_feature_samples(args.feature_dir, args.manifest)
samples_by_id = {sample.sample_id: sample for sample in samples}
if len(samples) != 100:
raise ValueError(f"expected the complete 100-sample Q1 feature set, found {len(samples)}")
splits = json.loads(args.splits.read_text(encoding="utf-8"))
if len(splits) != 5:
raise ValueError(f"expected five fixed video-group folds, found {len(splits)}")
all_valid_ids = [sample_id for split in splits for sample_id in split["validation_sample_ids"]]
if len(all_valid_ids) != len(set(all_valid_ids)) or set(all_valid_ids) != set(samples_by_id):
raise ValueError("fixed held-out folds do not cover the 100 feature samples exactly once")
expected_fold_by_id: dict[str, int] = {}
fold_rows: list[dict[str, Any]] = []
for split in splits:
fold = int(split["fold"])
train_ids = list(split["train_sample_ids"])
valid_ids = list(split["validation_sample_ids"])
train_groups = {samples_by_id[sample_id].group_id for sample_id in train_ids}
valid_groups = {samples_by_id[sample_id].group_id for sample_id in valid_ids}
if train_groups & valid_groups:
raise ValueError(f"video_id leakage in fixed fold {fold}: {sorted(train_groups & valid_groups)}")
if set(split.get("train_video_ids", train_groups)) != train_groups:
raise ValueError(f"train video_id list disagrees with sample manifest in fold {fold}")
if set(split.get("validation_video_ids", valid_groups)) != valid_groups:
raise ValueError(f"validation video_id list disagrees with sample manifest in fold {fold}")
for sample_id in valid_ids:
if sample_id in expected_fold_by_id:
raise ValueError(f"sample appears in multiple held-out folds: {sample_id}")
expected_fold_by_id[sample_id] = fold
fold_rows.extend({
"fold": fold,
"sample_id": sample_id,
"video_id": samples_by_id[sample_id].group_id,
"split": "train" if sample_id in set(train_ids) else "valid_oof",
} for sample_id in [*train_ids, *valid_ids])
_write_rows(output_dir / "fold_assignments.csv", fold_rows)
config = {
"experiment": "Similarity vs Correlation vs Predictability shared information definitions",
"seed": args.seed,
"device": str(device),
"sample_count": len(samples),
"video_id_count": len({sample.group_id for sample in samples}),
"fold_count": len(splits),
"grouped_by": "video_id/group_id",
"grid_size": GRID_SIZE,
"frozen_temporal_model": "M4_sourceTime fold checkpoints",
"source_features": {name: int(samples[0].features[name].shape[1]) for name in MODS},
"gcca_pca_components": PCA_COMPONENTS,
"gcca_shared_dimension": 32,
"predictive_contexts": PREDICTIVE_CONTEXTS,
"predictive_ridge_alphas": ALPHAS,
"predictive_ridge_solver": "cholesky",
"private_residual_alpha_selection": "reuse the selected original-target alpha; no separate residual alpha search",
"predictive_shift_slots": list(range(-5, 6)),
"predictive_eval_core_slots": CORE_SLOTS.tolist(),
"within_video_shuffle_repeats": args.shuffle_repeats,
"dimension_matched_per_slot": MAIN_DIM_PER_SLOT,
"dimension_matched_clip": MAIN_CLIP_DIM,
"emotion_probe": {
"classification": "StandardScaler + LogisticRegression(C=0.05, max_iter=5000)",
"regression": "StandardScaler + Ridge(alpha=25), clipped to [-3, 3]",
"temporal_pooling": "50 slots to five consecutive 10-slot mean bins",
"emotion_labels_used_for_representation_learning": False,
},
"paired_video_bootstrap_repeats": args.bootstrap_repeats,
"multiple_comparison_correction": False,
}
_write_json(output_dir / "config.json", config)
predictive_summary_rows: list[dict[str, Any]] = []
predictive_control_rows: list[dict[str, Any]] = []
residual_predictability_rows: list[dict[str, Any]] = []
gcca_summary_rows: list[dict[str, Any]] = []
shared_explained_rows: list[dict[str, Any]] = []
emotion_prediction_rows: list[dict[str, Any]] = []
fold_runtime_rows: list[dict[str, Any]] = []
for split in splits:
fold_start = time.time()
fold = int(split["fold"])
train_samples = [samples_by_id[sample_id] for sample_id in split["train_sample_ids"]]
heldout_samples = [samples_by_id[sample_id] for sample_id in split["validation_sample_ids"]]
all_fold_samples = [*train_samples, *heldout_samples]
all_ids = [sample.sample_id for sample in all_fold_samples]
feature_stats = fit_feature_stats(train_samples)
_, _, temporal_by_id = _collect_fold_features(
fold=fold,
train_samples=train_samples,
validation_samples=heldout_samples,
feature_stats=feature_stats,
checkpoint_root=args.m4_checkpoint_root,
device=device,
batch_size=args.batch_size,
)
pooled = _pool_original_source(all_fold_samples, feature_stats, temporal_by_id)
spr_representations = _spr_fold(
fold=fold,
sample_ids=all_ids,
pooled=pooled,
device=device,
checkpoint_root=args.spr_checkpoint_root,
batch_size=args.batch_size,
)
gcca_representations, fold_gcca_rows, fold_explained_rows = _gcca_fold(
fold=fold,
train_samples=train_samples,
heldout_samples=heldout_samples,
pooled=pooled,
seed=args.seed,
)
gcca_summary_rows.extend(fold_gcca_rows)
shared_explained_rows.extend(fold_explained_rows)
same_train, same_valid, same_summary, same_controls, same_residuals = _predictive_fold(
fold=fold,
context_name="same_slot",
train_samples=train_samples,
heldout_samples=heldout_samples,
pooled=pooled,
seed=args.seed,
shuffle_repeats=args.shuffle_repeats,
)
local_train, local_valid, local_summary, local_controls, local_residuals = _predictive_fold(
fold=fold,
context_name="local_pm1",
train_samples=train_samples,
heldout_samples=heldout_samples,
pooled=pooled,
seed=args.seed,
shuffle_repeats=args.shuffle_repeats,
)
predictive_summary_rows.extend(same_summary)
predictive_summary_rows.extend(local_summary)
predictive_control_rows.extend(same_controls)
predictive_control_rows.extend(local_controls)
residual_predictability_rows.extend(same_residuals)
residual_predictability_rows.extend(local_residuals)
method_representations: dict[str, Mapping[str, Mapping[str, Any]]] = {
"Similarity-SPR": spr_representations,
"GCCA": gcca_representations,
"Predictive-Same": {**same_train, **same_valid},
"Predictive-Local1": {**local_train, **local_valid},
}
main_vectors_by_method: dict[str, dict[str, np.ndarray]] = {}
for method, representations in method_representations.items():
views_by_id = {sample_id: _representation_views(representations[sample_id]) for sample_id in all_ids}
all_view_names = sorted(set.intersection(*(set(views_by_id[sample_id]) for sample_id in all_ids)))
for view in all_view_names:
vectors = {sample_id: views_by_id[sample_id][view] for sample_id in all_ids}
emotion_prediction_rows.extend(_fit_probe_predictions(
method=method,
view=view,
fold=fold,
train_samples=train_samples,
heldout_samples=heldout_samples,
vectors_by_id=vectors,
seed=args.seed,
))
main_sequences = {
sample_id: views_by_id[sample_id]["shared+private"]
for sample_id in all_ids
}
train_matched, valid_matched, pca = _dimension_matched_sequences(
train_samples, heldout_samples, main_sequences, seed=args.seed + fold
)
main_vectors = {**train_matched, **valid_matched}
main_vectors_by_method[method] = main_vectors
emotion_prediction_rows.extend(_fit_probe_predictions(
method=method,
view="main_1280d",
fold=fold,
train_samples=train_samples,
heldout_samples=heldout_samples,
vectors_by_id=main_vectors,
seed=args.seed,
))
raw_private_path = args.spr_checkpoint_root / f"fold_{fold:02d}" / "raw_private_pca.npz"
raw_private_sequences = _raw_private_pca_sequences(all_ids, pooled, raw_private_path)
emotion_prediction_rows.extend(_fit_probe_predictions(
method="RawPrivate-PCA",
view="main_1280d",
fold=fold,
train_samples=train_samples,
heldout_samples=heldout_samples,
vectors_by_id=raw_private_sequences,
seed=args.seed,
))
fold_runtime_rows.append({
"fold": fold,
"train_sample_count": len(train_samples),
"heldout_sample_count": len(heldout_samples),
"train_video_count": len({sample.group_id for sample in train_samples}),
"heldout_video_count": len({sample.group_id for sample in heldout_samples}),
"gpu": torch.cuda.get_device_name(device) if device.type == "cuda" else "cpu",
"runtime_seconds": time.time() - fold_start,
})
_write_json(output_dir / "progress.json", {
"completed_folds": [int(row["fold"]) for row in fold_runtime_rows],
"total_folds": len(splits),
"last_fold_runtime_seconds": fold_runtime_rows[-1]["runtime_seconds"],
"updated_utc": datetime.now(timezone.utc).isoformat(),
})
print(
f"[shared definitions fold {fold}/5] train={len(train_samples)} heldout={len(heldout_samples)} "
f"runtime={fold_runtime_rows[-1]['runtime_seconds']:.1f}s",
flush=True,
)
del temporal_by_id, pooled, spr_representations, gcca_representations
del same_train, same_valid, local_train, local_valid, method_representations
if device.type == "cuda":
torch.cuda.empty_cache()
reference_rows = _load_reference_predictions(
tsfa_predictions_path=args.tsfa_predictions,
math_predictions_path=args.math_predictions,
math_splits_path=args.math_splits,
expected_fold_by_id=expected_fold_by_id,
samples_by_id=samples_by_id,
)
emotion_prediction_rows.extend(reference_rows)
emotion_metrics = _summarize_probe_rows(emotion_prediction_rows)
dimension_metrics = [
row for row in emotion_metrics
if row["view"] in {"main_1280d", "main_external_nonmatched", "math_all_modalities"}
]
for row in dimension_metrics:
row["dimension_matched"] = row["view"] == "main_1280d"
row["comparison_note"] = (
"1280D, five-bin pooled" if row["dimension_matched"] else
("external 6845D comparator" if row["view"] == "main_external_nonmatched" else "math OOF comparator; feature dimension not exported")
)
paired_specs = (
("Predictive-Same", "RawPrivate-PCA", "main_1280d", "main_1280d"),
("GCCA", "RawPrivate-PCA", "main_1280d", "main_1280d"),
("Predictive-Same", "Similarity-SPR", "main_1280d", "main_1280d"),
("Predictive-Same", "TSFA+RawPrivate", "main_1280d", "main_external_nonmatched"),
("Predictive-Same", "Math-B0", "main_1280d", "math_all_modalities"),
("Predictive-Same", "Math-B4", "main_1280d", "math_all_modalities"),
("Predictive-Same", "Predictive-Local1", "main_1280d", "main_1280d"),
)
paired_rows: list[dict[str, Any]] = []
for contrast_index, (candidate_name, reference_name, candidate_view, reference_view) in enumerate(paired_specs):
candidate_rows = [row for row in emotion_prediction_rows if row["method"] == candidate_name and row["view"] == candidate_view]
reference_rows_for_pair = [row for row in emotion_prediction_rows if row["method"] == reference_name and row["view"] == reference_view]
paired_rows.extend(_paired_video_bootstrap(
candidate_rows,
reference_rows_for_pair,
candidate_name=f"{candidate_name}/{candidate_view}",
reference_name=f"{reference_name}/{reference_view}",
repeats=args.bootstrap_repeats,
seed=args.seed + 5000 + contrast_index,
))
private_ablation_rows = _make_private_ablation_rows(emotion_metrics)
predictive_shift_summary = _summarize_shift_controls(predictive_control_rows)
seed_summary = [
{"seed": args.seed, **row}
for row in emotion_metrics
]
figure_files = _save_figures(
output_dir,
predictive_rows=predictive_summary_rows,
shift_rows=predictive_control_rows,
gcca_rows=gcca_summary_rows,
explained_rows=shared_explained_rows,
emotion_metrics=emotion_metrics,
dimension_metrics=dimension_metrics,
ablation_rows=private_ablation_rows,
)
for filename, rows in (
("predictive_shared_summary.csv", predictive_summary_rows),
("predictive_shift_controls.csv", predictive_control_rows),
("predictive_shift_summary.csv", predictive_shift_summary),
("private_residual_predictability.csv", residual_predictability_rows),
("gcca_summary.csv", gcca_summary_rows),
("shared_explained_variance.csv", shared_explained_rows),
("emotion_probe_metrics.csv", emotion_metrics),
("emotion_probe_predictions.csv", emotion_prediction_rows),
("paired_contrasts.csv", paired_rows),
("dimension_matched_summary.csv", dimension_metrics),
("private_source_ablation.csv", private_ablation_rows),
("seed_summary.csv", seed_summary),
("fold_runtime.csv", fold_runtime_rows),
):
_write_rows(output_dir / filename, rows)
checkpoint_hashes: dict[str, str] = {}
for fold in range(1, 6):
for path in (
args.m4_checkpoint_root / ("" if fold == 1 else f"fold_{fold:02d}/") / "M4_sourceTime" / "checkpoint.pt",
args.spr_checkpoint_root / f"fold_{fold:02d}" / "SPR.pt",
):
if path.is_file():
checkpoint_hashes[str(path.relative_to(args.q1_root))] = hashlib.sha256(path.read_bytes()).hexdigest()
manifest = {
"created_utc": datetime.now(timezone.utc).isoformat(),
"runtime_seconds": time.time() - started,
"python": platform.python_version(),
"torch": torch.__version__,
"sklearn": sklearn.__version__,
"cuda_available": torch.cuda.is_available(),
"gpu_name": torch.cuda.get_device_name(0) if torch.cuda.is_available() else None,
"device_used": str(device),
"folds": fold_runtime_rows,
"checkpoint_hashes": checkpoint_hashes,
"emotion_labels_used_for_representation_learning": False,
"neural_spr_retrained": False,
"frozen_m4_temporal_branch": True,
"fixed_video_group_folds": str(args.splits),
"math_read_only_inputs": [str(args.math_predictions), str(args.math_splits)],
"bootstrap_note": "Paired percentile bootstrap resamples video_id clusters; no multiple-comparison correction.",
"predictive_shift_summary_metric": "video_id-macro mean MSE; per-clip R2 is retained as a diagnostic only because short clips can have near-zero target variance.",
"figures": figure_files,
"output_csvs": [
"fold_assignments.csv", "gcca_summary.csv", "predictive_shared_summary.csv",
"predictive_shift_controls.csv", "predictive_shift_summary.csv",
"private_residual_predictability.csv", "shared_explained_variance.csv",
"emotion_probe_metrics.csv", "emotion_probe_predictions.csv", "paired_contrasts.csv",
"dimension_matched_summary.csv", "private_source_ablation.csv", "seed_summary.csv",
],
}
_write_json(output_dir / "run_manifest.json", manifest)
(output_dir / "progress.json").unlink(missing_ok=True)
print(f"[shared definitions complete] outputs={output_dir} runtime={manifest['runtime_seconds']:.1f}s", flush=True)
def finalize_existing(args: argparse.Namespace) -> None:
"""Rebuild diagnostics and figures from completed fold outputs only."""
output_dir: Path = args.output_dir
required = (
"predictive_shift_controls.csv", "predictive_shared_summary.csv", "gcca_summary.csv",
"shared_explained_variance.csv", "emotion_probe_metrics.csv",
"dimension_matched_summary.csv", "private_source_ablation.csv",
)
missing = [name for name in required if not (output_dir / name).is_file()]
if missing:
raise FileNotFoundError(f"cannot finalize incomplete experiment outputs: {missing}")
controls = _read_csv(output_dir / "predictive_shift_controls.csv")
predictive = _read_csv(output_dir / "predictive_shared_summary.csv")
for row in predictive:
if row.get("private_alpha_selection") == "reused original-target grouped inner-CV choice":
row["private_inner_cv_r2_by_alpha"] = "{}"
_write_rows(output_dir / "predictive_shared_summary.csv", predictive)
config_path = output_dir / "config.json"
if config_path.is_file():
config = json.loads(config_path.read_text(encoding="utf-8"))
config["predictive_eval_core_slots"] = CORE_SLOTS.tolist()
config["predictive_shift_metric"] = "video_id-macro mean MSE; core slots 6-43"
config["private_residual_alpha_selection"] = "reuse the selected original-target alpha; no separate residual alpha search"
_write_json(config_path, config)
gcca = _read_csv(output_dir / "gcca_summary.csv")
explained = _read_csv(output_dir / "shared_explained_variance.csv")
metrics = _read_csv(output_dir / "emotion_probe_metrics.csv")
dimensions = _read_csv(output_dir / "dimension_matched_summary.csv")
ablation = _read_csv(output_dir / "private_source_ablation.csv")
shift_summary = _summarize_shift_controls(controls)
_write_rows(output_dir / "predictive_shift_summary.csv", shift_summary)
figures = _save_figures(
output_dir,
predictive_rows=predictive,
shift_rows=controls,
gcca_rows=gcca,
explained_rows=explained,
emotion_metrics=metrics,
dimension_metrics=dimensions,
ablation_rows=ablation,
)
manifest_path = output_dir / "run_manifest.json"
manifest = json.loads(manifest_path.read_text(encoding="utf-8")) if manifest_path.is_file() else {}
manifest["figures"] = figures
manifest["predictive_shift_summary_metric"] = (
"video_id-macro mean MSE; per-clip R2 is diagnostic only because individual clips can have near-zero target variance."
)
manifest["finalized_utc"] = datetime.now(timezone.utc).isoformat()
_write_json(manifest_path, manifest)
print(f"[shared definitions finalized] outputs={output_dir}", flush=True)
def build_arg_parser() -> argparse.ArgumentParser:
q1_root = Path(__file__).resolve().parents[1]
repo_root = q1_root.parents[1]
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--device", default="auto", choices=("auto", "cpu", "cuda"))
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--batch-size", type=int, default=8)
parser.add_argument("--shuffle-repeats", type=int, default=20)
parser.add_argument("--bootstrap-repeats", type=int, default=2000)
parser.add_argument("--finalize-existing", action="store_true", help="rebuild summaries/figures from completed outputs without rerunning folds")
parser.add_argument("--q1-root", type=Path, default=q1_root)
parser.add_argument("--feature-dir", type=Path, default=q1_root / "outputs/q1_features/features")
parser.add_argument("--manifest", type=Path, default=q1_root / "outputs/audit/manifest.csv")
parser.add_argument("--splits", type=Path, default=q1_root / "outputs/method_comparison/splits.json")
parser.add_argument("--m4-checkpoint-root", dest="m4_checkpoint_root", type=Path,
default=q1_root / "outputs/alignment_debug/heldout")
parser.add_argument("--spr-checkpoint-root", dest="spr_checkpoint_root", type=Path,
default=q1_root / "outputs/tsfa_shared_private/checkpoints")
parser.add_argument("--tsfa-predictions", type=Path,
default=q1_root / "outputs/tsfa_shared_private/emotion_probe_predictions.csv")
parser.add_argument("--math-predictions", type=Path,
default=repo_root / "math/results/model_comparison/oof_predictions.csv")
parser.add_argument("--math-splits", type=Path,
default=repo_root / "math/results/model_comparison/split_assignments.csv")
parser.add_argument("--output-dir", type=Path,
default=q1_root / "outputs/shared_definition_comparison")
return parser
def main() -> None:
args = build_arg_parser().parse_args()
if args.finalize_existing:
finalize_existing(args)
else:
run(args)
if __name__ == "__main__":
main()