1617 lines
77 KiB
Python
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()
|