Files

1863 lines
92 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""TSFA Shared-Private Representation (SPR) experiment.
The five-fold M4_sourceTime temporal checkpoints remain frozen. Original
training-fold-standardized source features are pooled with their temporal
weights, then factorized without emotion supervision into shared and
modality-private streams. All outputs are written to a separate directory.
"""
from __future__ import annotations
import argparse
import csv
import hashlib
import json
import math
import platform
import random
import shutil
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
import torch.nn.functional as F
from sklearn.decomposition import PCA
from sklearn.linear_model import LogisticRegression, Ridge
from sklearn.metrics import accuracy_score, f1_score, roc_auc_score
from sklearn.pipeline import make_pipeline
from sklearn.preprocessing import StandardScaler
from torch import Tensor, nn
from .compare_emotion_probes import _scores
from .correspondence_eval import _write_csv
from .experiment_data import (
FeatureSample,
fit_feature_stats,
load_feature_samples,
standardized_features,
)
from .tsfa_emotion_probe import CLASS_NAMES, _class_from_sentiment, _pearson
from .tsfa_experiment import GRID_SIZE, _collect_fold_features
from .types import MODALITIES
MODS = tuple(MODALITIES)
PAIRS = (("text", "audio"), ("text", "vision"), ("audio", "vision"))
HARD_OFFSETS = (-5, -3, -2, 2, 3, 5)
VARIANTS = ("SP-noOrth-noRec", "SP+Orth", "SPR", "SPR-noSharedContrastive")
VARIANT_LOSSES = {
"SP-noOrth-noRec": (True, False, False),
"SP+Orth": (True, True, False),
"SPR": (True, True, True),
"SPR-noSharedContrastive": (False, True, True),
}
EMOTION_VIEWS = (
"shared_only_fused",
"shared_only_unfused",
"private_text",
"private_audio",
"private_vision",
"private_audio_vision",
"private_all",
"shared_fused+private_text",
"shared_fused+private_audio",
"shared_fused+private_vision",
"shared_fused+private_audio_vision",
"shared_fused+private_all",
"shared_unfused+private_all",
)
class SharedPrivateFactorizer(nn.Module):
"""Small MLP factorizer with shared weights and modality-private branches."""
def __init__(
self,
dimensions: Mapping[str, int],
*,
common_dim: int = 128,
shared_dim: int = 64,
private_dim: int = 64,
) -> None:
super().__init__()
self.dimensions = dict(dimensions)
self.shared_dim = shared_dim
self.private_dim = private_dim
self.adapters = nn.ModuleDict({
name: nn.Sequential(
nn.Linear(dimensions[name], common_dim),
nn.LayerNorm(common_dim),
nn.GELU(),
)
for name in MODS
})
# One encoder instance is deliberately reused for all three modalities.
self.shared_encoder = nn.Sequential(
nn.Linear(common_dim, common_dim), nn.GELU(), nn.Linear(common_dim, shared_dim)
)
self.private_encoders = nn.ModuleDict({
name: nn.Sequential(
nn.Linear(common_dim, common_dim), nn.GELU(), nn.Linear(common_dim, private_dim)
)
for name in MODS
})
self.decoders = nn.ModuleDict({
name: nn.Sequential(
nn.Linear(shared_dim + private_dim, common_dim),
nn.GELU(),
nn.Linear(common_dim, dimensions[name]),
)
for name in MODS
})
def forward(self, inputs: Mapping[str, Tensor]) -> tuple[dict[str, Tensor], dict[str, Tensor], dict[str, Tensor]]:
shared: dict[str, Tensor] = {}
private: dict[str, Tensor] = {}
reconstruction: dict[str, Tensor] = {}
for name in MODS:
adapted = self.adapters[name](inputs[name])
shared[name] = self.shared_encoder(adapted)
private[name] = self.private_encoders[name](adapted)
reconstruction[name] = self.decoders[name](torch.cat([shared[name], private[name]], dim=-1))
return shared, private, reconstruction
def _fit_shared_loss(shared: Mapping[str, Tensor], temperature: float) -> Tensor:
"""Same-slot InfoNCE with within-video +/-2,3,5 slot negatives."""
slot_ids = list(range(5, GRID_SIZE - 5))
offsets = (0, *HARD_OFFSETS)
losses: list[Tensor] = []
for left, right in PAIRS:
left_values = F.normalize(shared[left], dim=-1)
right_values = F.normalize(shared[right], dim=-1)
for query, target in ((left_values, right_values), (right_values, left_values)):
anchors = query[:, slot_ids]
candidate_rows = []
for offset in offsets:
indices = [slot + offset for slot in slot_ids]
candidate_rows.append(target[:, indices])
candidates = torch.stack(candidate_rows, dim=2)
logits = torch.einsum("bnd,bnkd->bnk", anchors, candidates) / temperature
labels = torch.zeros(logits.shape[:2], dtype=torch.long, device=logits.device)
losses.append(F.cross_entropy(logits.reshape(-1, len(offsets)), labels.reshape(-1)))
return torch.stack(losses).mean()
def _orthogonality_loss(shared: Mapping[str, Tensor], private: Mapping[str, Tensor]) -> Tensor:
terms: list[Tensor] = []
for name in MODS:
s = shared[name].reshape(-1, shared[name].shape[-1])
p = private[name].reshape(-1, private[name].shape[-1])
s = s - s.mean(dim=0, keepdim=True)
p = p - p.mean(dim=0, keepdim=True)
cross = s.transpose(0, 1) @ p
normalizer = torch.linalg.vector_norm(s) * torch.linalg.vector_norm(p)
terms.append((cross / (normalizer + 1e-8)).square().sum())
return torch.stack(terms).mean()
def _reconstruction_loss(inputs: Mapping[str, Tensor], reconstruction: Mapping[str, Tensor]) -> Tensor:
# Standardized source dimensions have unit scale; averaging modalities keeps
# the 768-dimensional text stream from dominating the smaller streams.
return torch.stack([F.mse_loss(reconstruction[name], inputs[name]) for name in MODS]).mean()
def _seed_everything(seed: int) -> None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)
def _pool_original_source(
samples: Sequence[FeatureSample],
stats: Any,
temporal_by_id: Mapping[str, Mapping[str, Any]],
) -> dict[str, dict[str, np.ndarray]]:
"""Pool standardized original source vectors with the frozen M4 weights."""
by_id: dict[str, dict[str, np.ndarray]] = {}
for sample in samples:
standardized = standardized_features(sample, stats)
pooled: dict[str, np.ndarray] = {}
record = temporal_by_id[sample.sample_id]
for name in MODS:
weights = np.asarray(record["weights"][name], dtype=np.float32)
source = np.asarray(standardized[name], dtype=np.float32)
if weights.shape != (GRID_SIZE, source.shape[0]):
raise ValueError(
f"M4 weights/source shape mismatch for {sample.sample_id} {name}: "
f"{weights.shape} vs {source.shape}"
)
pooled[name] = (weights @ source).astype(np.float32, copy=False)
by_id[sample.sample_id] = pooled
return by_id
def _tensor_batch(ids: Sequence[str], pooled: Mapping[str, Mapping[str, np.ndarray]], device: torch.device) -> dict[str, Tensor]:
return {
name: torch.from_numpy(np.stack([pooled[sample_id][name] for sample_id in ids])).to(device)
for name in MODS
}
def _train_factorizer(
*,
variant: str,
fold: int,
train_ids: Sequence[str],
pooled: Mapping[str, Mapping[str, np.ndarray]],
dimensions: Mapping[str, int],
device: torch.device,
seed: int,
epochs: int,
batch_size: int,
learning_rate: float,
temperature: float,
orth_weight: float,
reconstruction_weight: float,
) -> tuple[SharedPrivateFactorizer, list[dict[str, Any]]]:
use_shared, use_orth, use_rec = VARIANT_LOSSES[variant]
fold_seed = seed + fold * 101
_seed_everything(fold_seed)
model = SharedPrivateFactorizer(dimensions).to(device)
optimizer = torch.optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=1e-4)
rng = np.random.default_rng(fold_seed)
history: list[dict[str, Any]] = []
for epoch in range(1, epochs + 1):
model.train()
order = rng.permutation(len(train_ids))
epoch_components: list[list[float]] = []
for start in range(0, len(order), batch_size):
ids = [train_ids[int(index)] for index in order[start : start + batch_size]]
inputs = _tensor_batch(ids, pooled, device)
shared, private, reconstruction = model(inputs)
shared_loss = _fit_shared_loss(shared, temperature) if use_shared else shared["text"].new_zeros(())
orth_loss = _orthogonality_loss(shared, private) if use_orth else shared["text"].new_zeros(())
rec_loss = _reconstruction_loss(inputs, reconstruction) if use_rec else shared["text"].new_zeros(())
loss = shared_loss + orth_weight * orth_loss + reconstruction_weight * rec_loss
if not torch.isfinite(loss):
raise FloatingPointError(f"non-finite {variant} loss in fold {fold}, epoch {epoch}")
optimizer.zero_grad(set_to_none=True)
loss.backward()
nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
epoch_components.append([
float(loss.detach().item()),
float(shared_loss.detach().item()),
float(orth_loss.detach().item()),
float(rec_loss.detach().item()),
])
means = np.asarray(epoch_components).mean(axis=0)
history.append({
"variant": variant,
"fold": fold,
"seed": fold_seed,
"epoch": epoch,
"loss": means[0],
"shared_loss": means[1],
"orthogonality_loss": means[2],
"reconstruction_loss": means[3],
})
model.eval()
return model, history
def _encode_all(
model: SharedPrivateFactorizer,
ids: Sequence[str],
pooled: Mapping[str, Mapping[str, np.ndarray]],
device: torch.device,
batch_size: int,
) -> tuple[dict[str, dict[str, np.ndarray]], dict[str, dict[str, np.ndarray]]]:
encoded: dict[str, dict[str, np.ndarray]] = {}
decoded: dict[str, dict[str, np.ndarray]] = {}
with torch.no_grad():
for start in range(0, len(ids), batch_size):
batch_ids = ids[start : start + batch_size]
inputs = _tensor_batch(batch_ids, pooled, device)
shared, private, reconstruction = model(inputs)
zeros_s = {name: torch.zeros_like(shared[name]) for name in MODS}
zeros_p = {name: torch.zeros_like(private[name]) for name in MODS}
reconstruction_shared = {
name: model.decoders[name](torch.cat([shared[name], zeros_p[name]], dim=-1))
for name in MODS
}
reconstruction_private = {
name: model.decoders[name](torch.cat([zeros_s[name], private[name]], dim=-1))
for name in MODS
}
for index, sample_id in enumerate(batch_ids):
encoded[sample_id] = {
"shared": {name: shared[name][index].cpu().numpy().astype(np.float32) for name in MODS},
"private": {name: private[name][index].cpu().numpy().astype(np.float32) for name in MODS},
}
decoded[sample_id] = {
"reconstruction_both": {
name: reconstruction[name][index].cpu().numpy().astype(np.float32) for name in MODS
},
"reconstruction_shared_only": {
name: reconstruction_shared[name][index].cpu().numpy().astype(np.float32)
for name in MODS
},
"reconstruction_private_only": {
name: reconstruction_private[name][index].cpu().numpy().astype(np.float32)
for name in MODS
},
}
return encoded, decoded
def _pool_five(sequence: np.ndarray) -> np.ndarray:
values = np.asarray(sequence, dtype=np.float32)
if values.ndim != 2 or values.shape[0] != GRID_SIZE:
raise ValueError(f"expected [{GRID_SIZE}, D] slot sequence, got {values.shape}")
return values.reshape(5, GRID_SIZE // 5, values.shape[1]).mean(axis=1).reshape(-1)
def _raw_private_views(sample: Mapping[str, np.ndarray]) -> dict[str, np.ndarray]:
return {
"raw_private_text": sample["text"],
"raw_private_audio": sample["audio"],
"raw_private_vision": sample["vision"],
"raw_private_all": np.concatenate([sample[name] for name in MODS], axis=-1),
}
def _emotion_views(
representation: Mapping[str, Any],
raw_private_pca: np.ndarray,
) -> dict[str, np.ndarray]:
shared = representation["shared"]
private = representation["private"]
fused = (shared["text"] + shared["audio"] + shared["vision"]) / 3.0
private_av = np.concatenate([private["audio"], private["vision"]], axis=-1)
private_all = np.concatenate([private[name] for name in MODS], axis=-1)
per_slot: dict[str, np.ndarray] = {
"shared_only_fused": fused,
"shared_only_unfused": np.concatenate([shared[name] for name in MODS], axis=-1),
"private_text": private["text"],
"private_audio": private["audio"],
"private_vision": private["vision"],
"private_audio_vision": private_av,
"private_all": private_all,
"shared_fused+private_text": np.concatenate([fused, private["text"]], axis=-1),
"shared_fused+private_audio": np.concatenate([fused, private["audio"]], axis=-1),
"shared_fused+private_vision": np.concatenate([fused, private["vision"]], axis=-1),
"shared_fused+private_audio_vision": np.concatenate([fused, private_av], axis=-1),
"shared_fused+private_all": np.concatenate([fused, private_all], axis=-1),
"shared_unfused+private_all": np.concatenate([
np.concatenate([shared[name] for name in MODS], axis=-1), private_all
], axis=-1),
}
per_slot["RawPrivate-PCA"] = raw_private_pca
return {name: _pool_five(values) for name, values in per_slot.items()}
def _fit_probe_pair(
train_x: np.ndarray,
train_class: np.ndarray,
train_value: np.ndarray,
test_x: np.ndarray,
seed: int,
) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
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(test_x)
regressor = make_pipeline(StandardScaler(), Ridge(alpha=25.0))
regressor.fit(train_x, train_value)
predicted_unclipped = regressor.predict(test_x)
return predicted_class, np.clip(predicted_unclipped, -3.0, 3.0), predicted_unclipped
def _fold_metrics(rows: Sequence[Mapping[str, Any]]) -> dict[str, float]:
actual_class = np.asarray([row["true_class_id"] for row in rows], dtype=np.int64)
predicted_class = np.asarray([row["predicted_class_id"] for row in rows], dtype=np.int64)
actual = np.asarray([row["true_label"] for row in rows], dtype=np.float64)
predicted = np.asarray([row["predicted_label"] for row in rows], dtype=np.float64)
fold_f1 = []
for fold in sorted({int(row["fold"]) for row in rows}):
fold_rows = [row for row in rows if int(row["fold"]) == fold]
fold_f1.append(f1_score(
[row["true_class_id"] for row in fold_rows],
[row["predicted_class_id"] for row in fold_rows],
labels=[0, 1, 2], average="macro", zero_division=0,
))
return {
"sample_count": len(rows),
"feature_dimension": int(rows[0].get("feature_dimension", -1)) if rows else -1,
"accuracy": float(accuracy_score(actual_class, predicted_class)),
"macro_f1": float(f1_score(actual_class, predicted_class, labels=[0, 1, 2], average="macro", zero_division=0)),
"macro_f1_fold_mean": float(np.mean(fold_f1)),
"macro_f1_fold_sd": float(np.std(fold_f1, ddof=1)) if len(fold_f1) > 1 else 0.0,
"mae": float(np.mean(np.abs(actual - predicted))),
"pearson": _pearson(actual, predicted),
}
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 _video_bootstrap_ci(
values: Mapping[str, float], *, repeats: int = 2000, seed: int = 42
) -> tuple[float, float]:
groups = sorted(values)
if not groups:
return float("nan"), float("nan")
point = np.asarray([values[group] for group in groups], dtype=np.float64)
point = point[np.isfinite(point)]
if not point.size:
return float("nan"), float("nan")
rng = np.random.default_rng(seed)
estimates = np.empty(repeats, dtype=np.float64)
for index in range(repeats):
chosen = rng.integers(0, len(groups), size=len(groups))
estimates[index] = np.nanmean([values[groups[int(i)]] for i in chosen])
return float(np.quantile(estimates, 0.025)), float(np.quantile(estimates, 0.975))
def _representation_diagnostics(
*,
variant: str,
fold: int,
sample: FeatureSample,
encoded: Mapping[str, Any],
) -> dict[str, Any]:
result: dict[str, Any] = {"variant": variant, "fold": fold, "sample_id": sample.sample_id,
"video_id": sample.group_id}
for name in MODS:
s = np.asarray(encoded["shared"][name], dtype=np.float64)
p = np.asarray(encoded["private"][name], dtype=np.float64)
s_center = s - s.mean(axis=0, keepdims=True)
p_center = p - p.mean(axis=0, keepdims=True)
cross = s_center.T @ p_center
denom = np.linalg.norm(s_center) * np.linalg.norm(p_center)
result[f"{name}_cross_covariance_norm"] = float(np.linalg.norm(cross) / (denom + 1e-12))
s_norm = np.linalg.norm(s, axis=1).mean()
p_norm = np.linalg.norm(p, axis=1).mean()
result[f"{name}_shared_norm"] = float(s_norm)
result[f"{name}_private_norm"] = float(p_norm)
result[f"{name}_shared_norm_ratio"] = float(s_norm / (s_norm + p_norm + 1e-12))
return result
def _source_modality_probe(
*,
encoded: Mapping[str, Mapping[str, Any]],
train_samples: Sequence[FeatureSample],
heldout_samples: Sequence[FeatureSample],
branch: str,
seed: int,
) -> dict[str, float]:
def matrix(samples: Sequence[FeatureSample]) -> tuple[np.ndarray, np.ndarray]:
vectors = []
labels = []
for sample in samples:
for label, name in enumerate(MODS):
values = np.asarray(encoded[sample.sample_id][branch][name], dtype=np.float32)
vectors.append(values)
labels.extend([label] * len(values))
return np.concatenate(vectors, axis=0), np.asarray(labels, dtype=np.int64)
train_x, train_y = matrix(train_samples)
heldout_x, heldout_y = matrix(heldout_samples)
probe = make_pipeline(
StandardScaler(), LogisticRegression(C=1.0, max_iter=2000, random_state=seed)
)
probe.fit(train_x, train_y)
predicted = probe.predict(heldout_x)
return {
"accuracy": float(accuracy_score(heldout_y, predicted)),
"macro_f1": float(f1_score(heldout_y, predicted, labels=[0, 1, 2], average="macro", zero_division=0)),
"sample_count": len(heldout_y),
}
def _pair_correspondence(
left_values: np.ndarray,
right_values: np.ndarray,
*,
rng: np.random.Generator,
shuffle_repeats: int,
) -> dict[str, float]:
"""Within-clip slot correspondence diagnostics; not a human alignment GT."""
left = F.normalize(torch.from_numpy(left_values), dim=-1).numpy()
right = F.normalize(torch.from_numpy(right_values), dim=-1).numpy()
def score(a: np.ndarray, b: np.ndarray) -> dict[str, float]:
similarities = a @ b.T
slots = np.arange(len(a))
positive = similarities[slots, slots]
negative = []
for offset in HARD_OFFSETS:
q = slots[(slots + offset >= 0) & (slots + offset < len(b))]
negative.extend(similarities[q, q + offset].tolist())
auc = float(roc_auc_score(np.r_[np.ones(len(positive)), np.zeros(len(negative))],
np.r_[positive, negative]))
best = similarities.argmax(axis=1)
error = np.abs(best - slots)
return {
"auc": auc,
"exact_r1": float(np.mean(error == 0)),
"within_pm1_r1": float(np.mean(error <= 1)),
"mase": float(np.mean(error)),
}
forward = score(left, right)
backward = score(right, left)
actual = {key: (forward[key] + backward[key]) / 2 for key in forward}
shuffled = []
for _ in range(shuffle_repeats):
shuffled.append(score(left, right[rng.permutation(len(right))]))
actual.update({
f"slot_shuffle_{key}": float(np.mean([row[key] for row in shuffled]))
for key in ("auc", "exact_r1", "within_pm1_r1", "mase")
})
return actual
def _emotion_feature_sets(
encoded_by_id: Mapping[str, Mapping[str, Any]],
raw_pca_by_id: Mapping[str, np.ndarray],
) -> dict[str, dict[str, np.ndarray]]:
return {
sample_id: _emotion_views(encoded_by_id[sample_id], raw_pca_by_id[sample_id])
for sample_id in encoded_by_id
}
def _probe_fold(
*,
variant: str,
fold: int,
train_samples: Sequence[FeatureSample],
heldout_samples: Sequence[FeatureSample],
features_by_id: Mapping[str, Mapping[str, np.ndarray]],
seed: int,
view_names: Sequence[str] = EMOTION_VIEWS,
) -> list[dict[str, Any]]:
train_classes = np.asarray([_class_from_sentiment(sample.sentiment) for sample in train_samples])
heldout_classes = np.asarray([_class_from_sentiment(sample.sentiment) for sample in heldout_samples])
train_values = np.asarray([sample.sentiment for sample in train_samples], dtype=np.float64)
heldout_values = np.asarray([sample.sentiment for sample in heldout_samples], dtype=np.float64)
rows: list[dict[str, Any]] = []
for view in view_names:
train_x = np.stack([features_by_id[sample.sample_id][view] for sample in train_samples])
heldout_x = np.stack([features_by_id[sample.sample_id][view] for sample in heldout_samples])
predicted_class, predicted_value, predicted_unclipped = _fit_probe_pair(
train_x, train_classes, train_values, heldout_x, seed
)
for index, sample in enumerate(heldout_samples):
rows.append({
"method": variant,
"view": view,
"fold": fold,
"sample_id": sample.sample_id,
"video_id": sample.group_id,
"true_class_id": int(heldout_classes[index]),
"true_class": CLASS_NAMES[int(heldout_classes[index])],
"predicted_class_id": int(predicted_class[index]),
"predicted_class": CLASS_NAMES[int(predicted_class[index])],
"true_label": float(heldout_values[index]),
"predicted_label": float(predicted_value[index]),
"predicted_label_unclipped": float(predicted_unclipped[index]),
"feature_dimension": int(train_x.shape[1]),
})
return rows
def _legacy_rows(
*,
old_predictions: Path,
raw_private_predictions: Path,
samples_by_id: Mapping[str, FeatureSample],
fold_by_id: Mapping[str, int],
) -> list[dict[str, Any]]:
rows: list[dict[str, Any]] = []
source_specs = (
(old_predictions, "TSFA-main", "TSFA-old", 1920),
(raw_private_predictions, "TSFA-T+Private", "TSFA+RawPrivate", 6845),
)
for path, source_method, method, dimension in source_specs:
source_rows = _read_csv(path)
if method == "TSFA-old":
source_rows = [row for row in source_rows if row["method"] == source_method and row["view"] == "all_modalities"]
else:
source_rows = [row for row in source_rows if row["method"] == source_method]
by_id = {row["sample_id"]: row for row in source_rows}
if len(by_id) != len(source_rows) or set(by_id) != set(samples_by_id):
raise ValueError(f"legacy OOF predictions do not cover the same unique samples: {path}")
for sample_id, source in by_id.items():
sample = samples_by_id[sample_id]
fold = int(source["fold"])
if fold != fold_by_id[sample_id] or source["video_id"] != sample.group_id:
raise ValueError(f"legacy OOF fold/group mismatch for {sample_id} in {method}")
rows.append({
"method": method,
"view": "all_modalities",
"fold": fold,
"sample_id": sample_id,
"video_id": sample.group_id,
"true_class_id": _class_from_sentiment(sample.sentiment),
"true_class": CLASS_NAMES[_class_from_sentiment(sample.sentiment)],
"predicted_class_id": int(source["predicted_class_id"]),
"predicted_class": CLASS_NAMES[int(source["predicted_class_id"])],
"true_label": float(sample.sentiment),
"predicted_label": float(source["predicted_label"]),
"predicted_label_unclipped": float(source.get("predicted_label_unclipped", source["predicted_label"])),
"feature_dimension": dimension,
})
return rows
def _math_reference_rows(
*,
predictions_path: Path,
splits_path: Path,
samples_by_id: Mapping[str, FeatureSample],
) -> list[dict[str, Any]]:
source = _read_csv(predictions_path)
split_rows = _read_csv(splits_path)
split_by_id = {row["sample_id"]: int(row["fold"]) for row in split_rows}
math_by_id = {row["sample_id"]: row for row in source}
if set(math_by_id) != set(samples_by_id) or set(split_by_id) != set(samples_by_id):
raise ValueError("math OOF predictions/splits do not match the Q1 100 samples")
rows: list[dict[str, Any]] = []
for method in ("B0", "B4"):
for sample_id, sample in samples_by_id.items():
ref = math_by_id[sample_id]
if ref["video_id"] != sample.group_id or not np.isclose(float(ref["true_sentiment"]), sample.sentiment):
raise ValueError(f"math label or video_id mismatch for {sample_id}")
rows.append({
"method": method,
"view": "all_modalities",
"fold": split_by_id[sample_id],
"sample_id": sample_id,
"video_id": sample.group_id,
"true_class_id": int(ref["true_polarity"]),
"true_class": CLASS_NAMES[int(ref["true_polarity"])],
"predicted_class_id": int(ref[f"{method}_predicted_polarity"]),
"predicted_class": CLASS_NAMES[int(ref[f"{method}_predicted_polarity"])],
"true_label": float(ref["true_sentiment"]),
"predicted_label": float(np.clip(float(ref[f"{method}_predicted_sentiment"]), -3.0, 3.0)),
"predicted_label_unclipped": float(ref[f"{method}_predicted_sentiment"]),
"feature_dimension": -1,
})
return rows
def _summarize_emotion_predictions(rows: Sequence[Mapping[str, Any]]) -> list[dict[str, Any]]:
grouped: dict[tuple[str, str], list[Mapping[str, Any]]] = defaultdict(list)
for row in rows:
grouped[(str(row["method"]), str(row["view"]))].append(row)
summaries: list[dict[str, Any]] = []
for (method, view), group in sorted(grouped.items()):
summary = _fold_metrics(group)
summary.update({"method": method, "view": view})
summaries.append(summary)
return summaries
def _paired_contrasts(
predictions: Sequence[Mapping[str, Any]],
*,
candidate_method: str = "SPR",
candidate_view: str = "shared_unfused+private_all",
references: Sequence[tuple[str, str]] | None = None,
repeats: int = 2000,
seed: int = 42,
) -> list[dict[str, Any]]:
by_key: dict[tuple[str, str], dict[str, Mapping[str, Any]]] = defaultdict(dict)
for row in predictions:
by_key[(str(row["method"]), str(row["view"]))][str(row["sample_id"])] = row
candidate = by_key[(candidate_method, candidate_view)]
if not candidate:
raise ValueError(f"missing paired contrast candidate {candidate_method}/{candidate_view}")
if references is None:
references = (
("TSFA-old", "all_modalities"),
("TSFA+RawPrivate", "all_modalities"),
("B0", "all_modalities"),
("B4", "all_modalities"),
("SPR-dim-matched", "SPR-dim-matched"),
("RawPrivate-PCA", "RawPrivate-PCA"),
)
output: list[dict[str, Any]] = []
for reference_method, reference_view in references:
reference = by_key[(reference_method, reference_view)]
if not reference:
continue
sample_ids = sorted(candidate)
if set(sample_ids) != set(reference):
raise ValueError(f"paired OOF sample IDs differ: SPR and {reference_method}")
for sample_id in sample_ids:
a, b = candidate[sample_id], reference[sample_id]
for key in ("video_id", "fold", "true_class_id", "true_label"):
if str(a[key]) != str(b[key]) and not (
key == "true_label" and np.isclose(float(a[key]), float(b[key]))
):
raise ValueError(f"paired OOF {key} mismatch for {sample_id}: {reference_method}")
groups: dict[str, list[str]] = defaultdict(list)
for sample_id in sample_ids:
groups[str(candidate[sample_id]["video_id"])].append(sample_id)
group_names = sorted(groups)
rng = np.random.default_rng(seed)
deltas: dict[str, list[float]] = {name: [] for name in ("macro_f1", "mae", "pearson")}
for _ in range(repeats):
chosen_groups = rng.choice(group_names, size=len(group_names), replace=True)
chosen = [sample_id for group in chosen_groups for sample_id in groups[str(group)]]
aa = [candidate[sample_id] for sample_id in chosen]
bb = [reference[sample_id] for sample_id in chosen]
class_true = np.asarray([int(row["true_class_id"]) for row in aa])
value_true = np.asarray([float(row["true_label"]) for row in aa])
score_a = _scores(
class_true,
np.asarray([int(row["predicted_class_id"]) for row in aa]),
value_true,
np.asarray([float(row["predicted_label"]) for row in aa]),
)
score_b = _scores(
class_true,
np.asarray([int(row["predicted_class_id"]) for row in bb]),
value_true,
np.asarray([float(row["predicted_label"]) for row in bb]),
)
for metric in deltas:
deltas[metric].append(score_a[metric] - score_b[metric])
whole_a = [candidate[sample_id] for sample_id in sample_ids]
whole_b = [reference[sample_id] for sample_id in sample_ids]
class_true = np.asarray([int(row["true_class_id"]) for row in whole_a])
value_true = np.asarray([float(row["true_label"]) for row in whole_a])
point_a = _scores(class_true,
np.asarray([int(row["predicted_class_id"]) for row in whole_a]),
value_true,
np.asarray([float(row["predicted_label"]) for row in whole_a]))
point_b = _scores(class_true,
np.asarray([int(row["predicted_class_id"]) for row in whole_b]),
value_true,
np.asarray([float(row["predicted_label"]) for row in whole_b]))
for metric, values in deltas.items():
finite = np.asarray(values, dtype=np.float64)
finite = finite[np.isfinite(finite)]
output.append({
"comparison": f"{candidate_method}-{reference_method}",
"candidate_view": candidate_view,
"reference_view": reference_view,
"metric": metric,
"candidate_oof": point_a[metric],
"reference_oof": point_b[metric],
"delta_candidate_minus_reference": point_a[metric] - point_b[metric],
"video_cluster_bootstrap_ci95_low": float(np.quantile(finite, 0.025)),
"video_cluster_bootstrap_ci95_high": float(np.quantile(finite, 0.975)),
"bootstrap_repeats": repeats,
"video_group_count": len(group_names),
"multiple_comparison_correction": "none",
})
return output
def _paired_view_contrasts(
predictions: Sequence[Mapping[str, Any]],
*,
method: str = "SPR",
views: Sequence[str] = (
"private_text", "private_audio", "private_vision", "private_audio_vision", "private_all"
),
repeats: int = 2000,
seed: int = 42,
) -> list[dict[str, Any]]:
by_view: dict[str, dict[str, Mapping[str, Any]]] = defaultdict(dict)
for row in predictions:
if str(row["method"]) == method and str(row["view"]) in views:
by_view[str(row["view"])][str(row["sample_id"])] = row
output: list[dict[str, Any]] = []
for left_index, left_view in enumerate(views):
for right_index, right_view in enumerate(views[left_index + 1:], start=left_index + 1):
left, right = by_view[left_view], by_view[right_view]
sample_ids = sorted(set(left) & set(right))
if not sample_ids:
continue
groups: dict[str, list[str]] = defaultdict(list)
for sample_id in sample_ids:
if left[sample_id]["video_id"] != right[sample_id]["video_id"]:
raise ValueError(f"video_id mismatch in source ablation for {sample_id}")
groups[str(left[sample_id]["video_id"])].append(sample_id)
group_names = sorted(groups)
rng = np.random.default_rng(seed + left_index * 100 + right_index)
boot: dict[str, list[float]] = {metric: [] for metric in ("macro_f1", "mae", "pearson")}
for _ in range(repeats):
chosen_groups = rng.choice(group_names, size=len(group_names), replace=True)
chosen = [sample_id for group in chosen_groups for sample_id in groups[str(group)]]
a_rows = [left[sample_id] for sample_id in chosen]
b_rows = [right[sample_id] for sample_id in chosen]
actual_class = np.asarray([int(row["true_class_id"]) for row in a_rows])
actual_value = np.asarray([float(row["true_label"]) for row in a_rows])
score_a = _scores(actual_class, np.asarray([int(row["predicted_class_id"]) for row in a_rows]),
actual_value, np.asarray([float(row["predicted_label"]) for row in a_rows]))
score_b = _scores(actual_class, np.asarray([int(row["predicted_class_id"]) for row in b_rows]),
actual_value, np.asarray([float(row["predicted_label"]) for row in b_rows]))
for metric in boot:
boot[metric].append(score_a[metric] - score_b[metric])
a_rows = [left[sample_id] for sample_id in sample_ids]
b_rows = [right[sample_id] for sample_id in sample_ids]
actual_class = np.asarray([int(row["true_class_id"]) for row in a_rows])
actual_value = np.asarray([float(row["true_label"]) for row in a_rows])
score_a = _scores(actual_class, np.asarray([int(row["predicted_class_id"]) for row in a_rows]),
actual_value, np.asarray([float(row["predicted_label"]) for row in a_rows]))
score_b = _scores(actual_class, np.asarray([int(row["predicted_class_id"]) for row in b_rows]),
actual_value, np.asarray([float(row["predicted_label"]) for row in b_rows]))
for metric in boot:
finite = np.asarray(boot[metric], dtype=np.float64)
finite = finite[np.isfinite(finite)]
output.append({
"method": method,
"left_view_minus_right_view": f"{left_view}-{right_view}",
"left_view": left_view,
"right_view": right_view,
"metric": metric,
"left_oof": score_a[metric],
"right_oof": score_b[metric],
"delta_left_minus_right": score_a[metric] - score_b[metric],
"video_cluster_bootstrap_ci95_low": float(np.quantile(finite, 0.025)),
"video_cluster_bootstrap_ci95_high": float(np.quantile(finite, 0.975)),
"bootstrap_repeats": repeats,
"video_group_count": len(group_names),
"multiple_comparison_correction": "none",
})
return output
def _clip_metric_summary(
rows: Sequence[Mapping[str, Any]],
*,
method_key: str,
group_keys: Sequence[str],
metrics: Sequence[str],
repeats: int,
seed: int,
) -> list[dict[str, Any]]:
grouped: dict[tuple[Any, ...], list[Mapping[str, Any]]] = defaultdict(list)
for row in rows:
grouped[tuple(row[key] for key in group_keys)].append(row)
output: list[dict[str, Any]] = []
for keys, items in sorted(grouped.items(), key=lambda pair: tuple(map(str, pair[0]))):
base = dict(zip(group_keys, keys))
base.update({"sample_count": len(items), "video_count": len({str(x["video_id"]) for x in items})})
for metric in metrics:
by_video: dict[str, list[float]] = defaultdict(list)
for row in items:
by_video[str(row["video_id"])].append(float(row[metric]))
video_means = {video: float(np.mean(values)) for video, values in by_video.items()}
mean = float(np.mean(list(video_means.values()))) if video_means else float("nan")
low, high = _video_bootstrap_ci(video_means, repeats=repeats,
seed=seed + sum(ord(ch) for ch in metric))
base[f"{metric}_mean_video"] = mean
base[f"{metric}_video_bootstrap_ci95_low"] = low
base[f"{metric}_video_bootstrap_ci95_high"] = high
output.append(base)
return output
def _sha256(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as stream:
for block in iter(lambda: stream.read(1024 * 1024), b""):
digest.update(block)
return digest.hexdigest()
def _save_architecture_figure(path: Path) -> None:
fig, ax = plt.subplots(figsize=(12, 6.5))
ax.set_xlim(0, 12)
ax.set_ylim(0, 7)
ax.axis("off")
colors = {"input": "#dbeafe", "shared": "#dcfce7", "private": "#ffedd5", "decode": "#f3e8ff"}
def box(x: float, y: float, w: float, h: float, text: str, color: str, fontsize: int = 10) -> None:
patch = plt.Rectangle((x, y), w, h, linewidth=1.3, edgecolor="#334155", facecolor=colors[color],
joinstyle="round")
ax.add_patch(patch)
ax.text(x + w / 2, y + h / 2, text, ha="center", va="center", fontsize=fontsize)
def arrow(start: tuple[float, float], end: tuple[float, float], color: str = "#475569") -> None:
ax.annotate("", xy=end, xytext=start,
arrowprops={"arrowstyle": "->", "lw": 1.2, "color": color})
y_rows = [5.6, 3.7, 1.8]
modalities = [("Text", "BERT"), ("Audio", "eGeMAPS"), ("Vision", "DeiT")]
for y, (label, source) in zip(y_rows, modalities):
box(0.2, y, 1.5, 0.75, f"{label} source\n{source}", "input")
box(2.25, y, 1.5, 0.75, f"Frozen M4\nweighted pool", "input")
box(4.3, y, 1.45, 0.75, f"Adapter $P_{label[0]}$", "input")
box(6.3, y + 0.42, 1.55, 0.75, f"Shared\nencoder", "shared")
box(6.3, y - 0.42, 1.55, 0.75, f"Private {label}\nencoder", "private")
box(9.25, y - 0.12, 2.15, 0.75, f"{label} decoder\nreconstruct $h^{label}$", "decode")
arrow((1.7, y + 0.37), (2.25, y + 0.37))
arrow((3.75, y + 0.37), (4.3, y + 0.37))
arrow((5.75, y + 0.37), (6.3, y + 0.78))
arrow((5.75, y + 0.37), (6.3, y - 0.02))
arrow((7.85, y + 0.78), (9.25, y + 0.38))
arrow((7.85, y - 0.02), (9.25, y + 0.15))
ax.text(7.08, 6.85, "$L_{shared}$: same-slot T-A / T-V / A-V InfoNCE", ha="center", va="center", fontsize=11)
ax.text(7.08, 0.55, "$L_{orth}$ within each modality + $L_{rec}$ per modality", ha="center", fontsize=11)
ax.set_title("TSFA-SPR: frozen temporal pooling, shared semantics, private modality detail", fontsize=14, pad=12)
fig.tight_layout()
fig.savefig(path, dpi=180, bbox_inches="tight")
plt.close(fig)
def _plot_summaries(
output_dir: Path,
modality_rows: Sequence[Mapping[str, Any]],
correspondence_rows: Sequence[Mapping[str, Any]],
reconstruction_rows: Sequence[Mapping[str, Any]],
emotion_rows: Sequence[Mapping[str, Any]],
diagnostics_rows: Sequence[Mapping[str, Any]],
) -> list[str]:
figures: list[str] = []
_save_architecture_figure(output_dir / "architecture.png")
figures.append("architecture.png")
variants = list(VARIANTS)
x = np.arange(len(variants))
fig, ax = plt.subplots(figsize=(9, 5))
for offset, branch, color, label in ((-0.18, "shared", "#2563eb", "Shared"),
(0.18, "private", "#ea580c", "Private")):
rows = {str(row["variant"]): row for row in modality_rows if row["branch"] == branch}
ax.bar(x + offset, [float(rows[v]["accuracy_mean"]) for v in variants], width=0.34,
color=color, label=label)
ax.set_xticks(x, variants, rotation=15, ha="right")
ax.set_ylim(0, 1)
ax.set_ylabel("Modality label accuracy (held-out videos)")
ax.set_title("Modality information in shared vs private streams")
ax.legend()
fig.tight_layout()
fig.savefig(output_dir / "shared_private_modality_probe.png", dpi=180)
plt.close(fig)
figures.append("shared_private_modality_probe.png")
fig, ax = plt.subplots(figsize=(9, 5))
pairs = ["text-audio", "text-vision", "audio-vision"]
pair_x = np.arange(len(pairs))
width = 0.22
for index, variant in enumerate(variants):
rows = {str(row["pair"]): row for row in correspondence_rows if row["variant"] == variant}
ax.bar(pair_x + (index - 1.5) * width,
[float(rows[pair]["auc_mean_video"]) for pair in pairs], width=width, label=variant)
ax.axhline(0.5, color="#64748b", linestyle="--", linewidth=1)
ax.set_xticks(pair_x, pairs)
ax.set_ylim(0, 1)
ax.set_ylabel("Matched vs shifted AUC")
ax.set_title("Shared representation: cross-modal slot correspondence")
ax.legend(fontsize=8)
fig.tight_layout()
fig.savefig(output_dir / "shared_crossmodal_auc.png", dpi=180)
plt.close(fig)
figures.append("shared_crossmodal_auc.png")
trained = [row for row in reconstruction_rows
if row["reconstruction_trained"] is True or str(row["reconstruction_trained"]).lower() == "true"]
branches = ["shared_only", "private_only", "both"]
fig, ax = plt.subplots(figsize=(10, 5))
labels = []
values = []
for variant in VARIANTS:
for branch in branches:
subset = [float(row["mse_mean_video"]) for row in trained
if row["variant"] == variant and row["branch"] == branch]
if subset:
labels.append(f"{variant}\n{branch}")
values.append(float(np.mean(subset)))
ax.bar(np.arange(len(values)), values, color=["#2563eb", "#ea580c", "#7c3aed"] * 4)
ax.set_xticks(np.arange(len(labels)), labels, rotation=25, ha="right", fontsize=8)
ax.set_ylabel("Held-out normalized MSE")
ax.set_title("Source reconstruction from shared, private, and combined factors")
fig.tight_layout()
fig.savefig(output_dir / "reconstruction_branches.png", dpi=180)
plt.close(fig)
figures.append("reconstruction_branches.png")
selected = [
("TSFA-old", "all_modalities"), ("TSFA+RawPrivate", "all_modalities"),
("SPR", "shared_unfused+private_all"), ("SPR-dim-matched", "SPR-dim-matched"),
("RawPrivate-PCA", "RawPrivate-PCA"),
]
fig, axes = plt.subplots(1, 2, figsize=(12, 5))
plot_rows = []
for method, view in selected:
row = next((r for r in emotion_rows if r["method"] == method and r["view"] == view), None)
if row:
plot_rows.append(row)
labels = [str(row["method"]) for row in plot_rows]
axes[0].bar(np.arange(len(plot_rows)), [float(row["mae"]) for row in plot_rows], color="#0f766e")
axes[0].set_ylabel("OOF MAE (lower is better)")
axes[1].bar(np.arange(len(plot_rows)), [float(row["pearson"]) for row in plot_rows], color="#7c3aed")
axes[1].set_ylabel("OOF Pearson")
for axis in axes:
axis.set_xticks(np.arange(len(labels)), labels, rotation=22, ha="right")
fig.suptitle("Emotion probe: old baseline, private residual, SPR, and dimension controls")
fig.tight_layout()
fig.savefig(output_dir / "emotion_regression_ablation.png", dpi=180)
plt.close(fig)
figures.append("emotion_regression_ablation.png")
private_names = ["private_text", "private_audio", "private_vision", "private_audio_vision", "private_all"]
private_rows = [row for row in emotion_rows if row["method"] == "SPR" and row["view"] in private_names]
private_rows = sorted(private_rows, key=lambda row: private_names.index(str(row["view"])))
fig, ax = plt.subplots(figsize=(9, 5))
xpos = np.arange(len(private_rows))
ax.bar(xpos - 0.18, [float(row["macro_f1"]) for row in private_rows], width=0.36,
color="#2563eb", label="Macro-F1")
ax2 = ax.twinx()
ax2.bar(xpos + 0.18, [float(row["mae"]) for row in private_rows], width=0.36,
color="#ea580c", label="MAE")
ax.set_xticks(xpos, [str(row["view"]) for row in private_rows], rotation=15, ha="right")
ax.set_ylabel("Macro-F1")
ax2.set_ylabel("MAE")
ax.set_title("Which private sources contribute to emotion probes?")
ax.legend(loc="upper left")
ax2.legend(loc="upper right")
fig.tight_layout()
fig.savefig(output_dir / "private_source_contribution.png", dpi=180)
plt.close(fig)
figures.append("private_source_contribution.png")
dim_keys = {
("TSFA-old", "all_modalities"),
("TSFA+RawPrivate", "all_modalities"),
("SPR", "shared_unfused+private_all"),
("SPR-dim-matched", "SPR-dim-matched"),
("RawPrivate-PCA", "RawPrivate-PCA"),
}
dim_rows = [row for row in emotion_rows if (row["method"], row["view"]) in dim_keys]
fig, axes = plt.subplots(1, 2, figsize=(12, 5.8))
dim_labels = [f"{row['method']} ({row['feature_dimension']}D)" for row in dim_rows]
ypos = np.arange(len(dim_rows))
axes[0].barh(ypos, [float(row["macro_f1"]) for row in dim_rows], color="#2563eb")
axes[1].barh(ypos, [float(row["mae"]) for row in dim_rows], color="#ea580c")
axes[0].set_yticks(ypos, dim_labels)
axes[1].set_yticks(ypos, dim_labels)
axes[0].invert_yaxis()
axes[1].invert_yaxis()
axes[0].set_xlabel("OOF Macro-F1")
axes[1].set_xlabel("OOF MAE")
fig.suptitle("Emotion probes at native and dimension-matched representation sizes")
fig.tight_layout()
fig.savefig(output_dir / "dimension_control.png", dpi=180)
plt.close(fig)
figures.append("dimension_control.png")
fig, ax = plt.subplots(figsize=(9, 5))
for index, variant in enumerate(variants):
means = [float(np.mean([float(row[f"{name}_shared_norm_ratio"]) for row in diagnostics_rows
if row["variant"] == variant])) for name in MODS]
ax.plot(MODS, means, marker="o", label=variant)
ax.set_ylim(0, 1)
ax.set_ylabel(r"$\|s\|/(\|s\|+\|p\|)$")
ax.set_title("Shared/private norm ratio by modality")
ax.legend(fontsize=8)
fig.tight_layout()
fig.savefig(output_dir / "shared_private_norm_ratio.png", dpi=180)
plt.close(fig)
figures.append("shared_private_norm_ratio.png")
return figures
def run(args: argparse.Namespace) -> None:
started = time.time()
args.output_dir.mkdir(parents=True, exist_ok=True)
if args.device == "auto":
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
else:
device = torch.device(args.device)
if device.type == "cuda" and not torch.cuda.is_available():
raise RuntimeError("CUDA requested but unavailable")
if device.type == "cuda":
torch.backends.cudnn.benchmark = False
torch.backends.cuda.matmul.allow_tf32 = False
samples = load_feature_samples(args.feature_dir, args.manifest)
samples_by_id = {sample.sample_id: sample for sample in samples}
splits = json.loads(args.splits.read_text(encoding="utf-8"))
if len(splits) != 5:
raise ValueError(f"expected five grouped folds, got {len(splits)}")
fold_by_id: dict[str, int] = {}
for split in splits:
fold = int(split["fold"])
train_ids = set(split["train_sample_ids"])
valid_ids = set(split["validation_sample_ids"])
if train_ids & valid_ids or train_ids | valid_ids != set(samples_by_id):
raise ValueError(f"fold {fold} does not partition all samples")
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 fold {fold}: {sorted(train_groups & valid_groups)}")
for sample_id in valid_ids:
if sample_id in fold_by_id:
raise ValueError(f"sample appears in multiple validation folds: {sample_id}")
fold_by_id[sample_id] = fold
if set(fold_by_id) != set(samples_by_id):
raise ValueError("fold validation partitions do not cover all samples exactly once")
# Existing learned and math baselines are read-only controls. Their OOF IDs
# and folds are validated before including them in any paired comparison.
legacy_predictions = _legacy_rows(
old_predictions=args.old_tsfa_predictions,
raw_private_predictions=args.raw_private_predictions,
samples_by_id=samples_by_id,
fold_by_id=fold_by_id,
)
math_predictions = _math_reference_rows(
predictions_path=args.math_predictions,
splits_path=args.math_splits,
samples_by_id=samples_by_id,
)
for row in math_predictions:
if int(row["fold"]) != fold_by_id[str(row["sample_id"])]:
raise ValueError(f"math GroupKFold assignment differs from Q1 fixed folds for {row['sample_id']}")
prediction_rows: list[dict[str, Any]] = [*legacy_predictions, *math_predictions]
history_rows: list[dict[str, Any]] = []
representation_rows: list[dict[str, Any]] = []
correspondence_clip_rows: list[dict[str, Any]] = []
modality_fold_rows: list[dict[str, Any]] = []
reconstruction_clip_rows: list[dict[str, Any]] = []
fold_manifest: list[dict[str, Any]] = []
input_hashes: dict[str, str] = {}
for split in splits:
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_samples = [*train_samples, *heldout_samples]
all_ids = [sample.sample_id for sample in all_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.checkpoint_root,
device=device,
batch_size=args.batch_size,
)
pooled = _pool_original_source(all_samples, feature_stats, temporal_by_id)
dimensions = {name: int(pooled[all_ids[0]][name].shape[-1]) for name in MODS}
raw_slots = {
sample_id: np.concatenate([pooled[sample_id][name] for name in MODS], axis=-1)
for sample_id in all_ids
}
train_slot_matrix = np.concatenate([raw_slots[sample.sample_id] for sample in train_samples], axis=0)
pca_dim = min(args.pca_dim, train_slot_matrix.shape[0], train_slot_matrix.shape[1])
pca = PCA(n_components=pca_dim, svd_solver="randomized", random_state=args.seed + fold)
pca.fit(train_slot_matrix)
raw_pca_by_id = {
sample_id: pca.transform(raw_slots[sample_id]).astype(np.float32, copy=False)
for sample_id in all_ids
}
fold_checkpoint = args.output_dir / "checkpoints" / f"fold_{fold:02d}"
fold_checkpoint.mkdir(parents=True, exist_ok=True)
np.savez_compressed(fold_checkpoint / "raw_private_pca.npz", mean=pca.mean_, components=pca.components_,
explained_variance_ratio=pca.explained_variance_ratio_)
for variant in VARIANTS:
model, history = _train_factorizer(
variant=variant,
fold=fold,
train_ids=[sample.sample_id for sample in train_samples],
pooled=pooled,
dimensions=dimensions,
device=device,
seed=args.seed,
epochs=args.epochs,
batch_size=args.batch_size,
learning_rate=args.learning_rate,
temperature=args.temperature,
orth_weight=args.orth_weight,
reconstruction_weight=args.reconstruction_weight,
)
history_rows.extend(history)
torch.save({
"variant": variant,
"fold": fold,
"seed": args.seed + fold * 101,
"dimensions": dimensions,
"model_state_dict": {key: value.detach().cpu() for key, value in model.state_dict().items()},
"loss_flags": VARIANT_LOSSES[variant],
"hyperparameters": {
"epochs": args.epochs,
"common_dim": args.common_dim,
"shared_dim": args.shared_dim,
"private_dim": args.private_dim,
"orth_weight": args.orth_weight,
"reconstruction_weight": args.reconstruction_weight,
"temperature": args.temperature,
},
}, fold_checkpoint / f"{variant}.pt")
encoded, decoded = _encode_all(
model, all_ids, pooled, device, args.batch_size
)
train_ids = {sample.sample_id for sample in train_samples}
heldout_ids = {sample.sample_id for sample in heldout_samples}
use_rec = VARIANT_LOSSES[variant][2]
for sample in heldout_samples:
sample_id = sample.sample_id
representation_rows.append(_representation_diagnostics(
variant=variant, fold=fold, sample=sample, encoded=encoded[sample_id]
))
for modality in MODS:
target = pooled[sample_id][modality]
for branch, key in (
("shared_only", "reconstruction_shared_only"),
("private_only", "reconstruction_private_only"),
("both", "reconstruction_both"),
):
pred = decoded[sample_id][key][modality]
reconstruction_clip_rows.append({
"variant": variant,
"fold": fold,
"sample_id": sample_id,
"video_id": sample.group_id,
"modality": modality,
"branch": branch,
"mse": float(np.mean(np.square(target - pred))),
"reconstruction_trained": bool(use_rec),
})
for left, right in PAIRS:
rng = np.random.default_rng(args.seed + fold * 1009 + sum(map(ord, sample_id + variant + left + right)))
scores = _pair_correspondence(
encoded[sample_id]["shared"][left],
encoded[sample_id]["shared"][right],
rng=rng,
shuffle_repeats=args.shuffle_repeats,
)
correspondence_clip_rows.append({
"variant": variant,
"fold": fold,
"sample_id": sample_id,
"video_id": sample.group_id,
"pair": f"{left}-{right}",
**scores,
"delta_auc_vs_shuffle": scores["auc"] - scores["slot_shuffle_auc"],
"delta_exact_r1_vs_shuffle": scores["exact_r1"] - scores["slot_shuffle_exact_r1"],
})
for branch in ("shared", "private"):
score = _source_modality_probe(
encoded=encoded,
train_samples=train_samples,
heldout_samples=heldout_samples,
branch=branch,
seed=args.seed + fold,
)
modality_fold_rows.append({
"variant": variant,
"fold": fold,
"branch": branch,
**score,
})
feature_sets = _emotion_feature_sets(encoded, raw_pca_by_id)
prediction_rows.extend(_probe_fold(
variant=variant,
fold=fold,
train_samples=train_samples,
heldout_samples=heldout_samples,
features_by_id=feature_sets,
seed=args.seed,
))
raw_pca_features = {
sample_id: {"RawPrivate-PCA": _pool_five(raw_pca_by_id[sample_id])}
for sample_id in all_ids
}
if variant == "SPR":
prediction_rows.extend(_probe_fold(
variant="SPR-dim-matched",
fold=fold,
train_samples=train_samples,
heldout_samples=heldout_samples,
features_by_id={
sample_id: {"SPR-dim-matched": feature_sets[sample_id]["shared_fused+private_all"]}
for sample_id in all_ids
},
seed=args.seed,
view_names=("SPR-dim-matched",),
))
prediction_rows.extend(_probe_fold(
variant="RawPrivate-PCA",
fold=fold,
train_samples=train_samples,
heldout_samples=heldout_samples,
features_by_id=raw_pca_features,
seed=args.seed,
view_names=("RawPrivate-PCA",),
))
del model, encoded, decoded, feature_sets, raw_pca_features
if device.type == "cuda":
torch.cuda.empty_cache()
print(f"[TSFA-SPR fold {fold}] {variant} complete", flush=True)
m4_checkpoint = args.checkpoint_root / ("M4_sourceTime/checkpoint.pt" if fold == 1 else f"fold_{fold:02d}/M4_sourceTime/checkpoint.pt")
if m4_checkpoint.is_file():
input_hashes[f"fold_{fold:02d}/M4_sourceTime"] = _sha256(m4_checkpoint)
fold_manifest.append({
"fold": fold,
"train_count": len(train_samples),
"heldout_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}),
"video_id_overlap": [],
"pca_train_slot_rows": int(train_slot_matrix.shape[0]),
"raw_private_pca_components": int(pca_dim),
})
del pooled, temporal_by_id, raw_slots, raw_pca_by_id, pca
if device.type == "cuda":
torch.cuda.empty_cache()
print(f"[TSFA-SPR fold {fold}] completed: train={len(train_samples)} heldout={len(heldout_samples)}", flush=True)
# The RawPrivate-PCA and SPR dimension-matched views are separate methods;
# this alias allows their paired OOF rows to be verified and joined directly.
summaries = _summarize_emotion_predictions(prediction_rows)
by_method_view = {(row["method"], row["view"]): row for row in summaries}
private_source_rows = [
row for row in summaries
if row["method"] == "SPR" and row["view"] in {
"private_text", "private_audio", "private_vision", "private_audio_vision", "private_all",
"shared_only_fused", "shared_only_unfused", "shared_unfused+private_all",
"shared_fused+private_all",
}
]
dimension_keys = (
("TSFA-old", "all_modalities"),
("TSFA+RawPrivate", "all_modalities"),
("SPR", "shared_unfused+private_all"),
("SPR-dim-matched", "SPR-dim-matched"),
("RawPrivate-PCA", "RawPrivate-PCA"),
)
dimension_rows = []
for key in dimension_keys:
if key in by_method_view:
dimension_rows.append({**by_method_view[key],
"per_slot_dimension": int(by_method_view[key]["feature_dimension"] // 5),
"dimension_control_note": (
"matched learned fused-shared + all-private" if key[0] == "SPR-dim-matched"
else "training-fold PCA on raw private slots" if key[0] == "RawPrivate-PCA"
else "existing baseline; native dimension"
)})
paired_rows = _paired_contrasts(
prediction_rows,
repeats=args.bootstrap_repeats,
seed=args.seed,
)
paired_rows.extend(_paired_contrasts(
prediction_rows,
candidate_method="SPR-dim-matched",
candidate_view="SPR-dim-matched",
references=(("RawPrivate-PCA", "RawPrivate-PCA"),),
repeats=args.bootstrap_repeats,
seed=args.seed,
))
correspondence_summary = _clip_metric_summary(
correspondence_clip_rows,
method_key="variant",
group_keys=("variant", "pair"),
metrics=("auc", "exact_r1", "within_pm1_r1", "mase", "slot_shuffle_auc",
"slot_shuffle_exact_r1", "delta_auc_vs_shuffle", "delta_exact_r1_vs_shuffle"),
repeats=args.bootstrap_repeats,
seed=args.seed,
)
reconstruction_summary = _clip_metric_summary(
reconstruction_clip_rows,
method_key="variant",
group_keys=("variant", "modality", "branch", "reconstruction_trained"),
metrics=("mse",),
repeats=args.bootstrap_repeats,
seed=args.seed,
)
modality_summary: list[dict[str, Any]] = []
for variant in VARIANTS:
for branch in ("shared", "private"):
items = [row for row in modality_fold_rows if row["variant"] == variant and row["branch"] == branch]
modality_summary.append({
"variant": variant,
"branch": branch,
"fold_count": len(items),
"accuracy_mean": float(np.mean([row["accuracy"] for row in items])),
"accuracy_sd": float(np.std([row["accuracy"] for row in items], ddof=1)),
"macro_f1_mean": float(np.mean([row["macro_f1"] for row in items])),
"macro_f1_sd": float(np.std([row["macro_f1"] for row in items], ddof=1)),
"fold_accuracy": json.dumps([round(float(row["accuracy"]), 6) for row in items]),
})
seed_rows = []
for method, view in dimension_keys:
if (method, view) in by_method_view:
row = by_method_view[(method, view)]
seed_rows.append({"seed": args.seed, "method": method, "view": view,
"accuracy": row["accuracy"], "macro_f1": row["macro_f1"],
"mae": row["mae"], "pearson": row["pearson"],
"status": "first_seed_only; additional seeds conditional on separation diagnostics"})
figures = _plot_summaries(
args.output_dir,
modality_summary,
correspondence_summary,
reconstruction_summary,
summaries,
representation_rows,
)
args.output_dir.mkdir(parents=True, exist_ok=True)
_write_csv(args.output_dir / "training_history.csv", history_rows)
_write_csv(args.output_dir / "representation_diagnostics.csv", representation_rows)
_write_csv(args.output_dir / "shared_correspondence_by_clip.csv", correspondence_clip_rows)
_write_csv(args.output_dir / "shared_correspondence_summary.csv", correspondence_summary)
_write_csv(args.output_dir / "modality_probe_by_fold.csv", modality_fold_rows)
_write_csv(args.output_dir / "modality_probe_summary.csv", modality_summary)
_write_csv(args.output_dir / "reconstruction_by_clip.csv", reconstruction_clip_rows)
_write_csv(args.output_dir / "reconstruction_summary.csv", reconstruction_summary)
_write_csv(args.output_dir / "emotion_probe_predictions.csv", prediction_rows)
_write_csv(args.output_dir / "emotion_probe_metrics.csv", summaries)
_write_csv(args.output_dir / "paired_contrasts.csv", paired_rows)
_write_csv(args.output_dir / "private_source_ablation.csv", private_source_rows)
_write_csv(args.output_dir / "dimension_control_summary.csv", dimension_rows)
_write_csv(args.output_dir / "seed_summary.csv", seed_rows)
spr_modality = {row["branch"]: row for row in modality_summary if row["variant"] == "SPR"}
spr_correspondence = [row for row in correspondence_summary if row["variant"] == "SPR"]
spr_reconstruction = [row for row in reconstruction_summary
if row["variant"] == "SPR" and row["reconstruction_trained"]]
private_source_accuracy = float(spr_modality["private"]["accuracy_mean"])
shared_source_accuracy = float(spr_modality["shared"]["accuracy_mean"])
mean_shared_auc = float(np.mean([row["auc_mean_video"] for row in spr_correspondence]))
rec_by_branch = {
branch: float(np.mean([row["mse_mean_video"] for row in spr_reconstruction if row["branch"] == branch]))
for branch in ("shared_only", "private_only", "both")
}
split_supported = bool(
private_source_accuracy > shared_source_accuracy
and private_source_accuracy - shared_source_accuracy >= args.modality_gap_threshold
and mean_shared_auc > 0.5
and rec_by_branch["both"] < min(rec_by_branch["shared_only"], rec_by_branch["private_only"])
)
config = {
"experiment": "TSFA-SPR: Temporal-Shared-Private Representation",
"seed": args.seed,
"folds": 5,
"grouped_by": "video_id/group_id",
"grid_size": GRID_SIZE,
"feature_dimensions": {name: int(samples[0].features[name].shape[1]) for name in MODS},
"common_dim": args.common_dim,
"shared_dim": args.shared_dim,
"private_dim": args.private_dim,
"batch_size": args.batch_size,
"epochs": args.epochs,
"learning_rate": args.learning_rate,
"temperature": args.temperature,
"orthogonality_weight": args.orth_weight,
"reconstruction_weight": args.reconstruction_weight,
"contrastive_pairs": [list(pair) for pair in PAIRS],
"hard_negative_offsets": list(HARD_OFFSETS),
"loss_variants": VARIANT_LOSSES,
"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",
"all_scalers_fit_on_training_fold_only": True,
},
"dimension_controls": {
"SPR_main": "unfused shared 3x64 plus private 3x64 = 384 per slot; 1920 after five-bin pooling",
"SPR_dim_matched": "mean fused shared 64 plus all private 3x64 = 256 per slot; 1280 after pooling",
"RawPrivate_PCA": f"train-fold PCA from raw private 985-d slot vectors to {args.pca_dim} per slot; 1280 after pooling if 256 components",
},
"emotion_labels_used_in_factorizer_training": False,
"frozen_m4_temporal_branch": True,
"explicit_time_code_or_slot_index_in_shared_encoder": False,
"bootstrap_repeats": args.bootstrap_repeats,
}
(args.output_dir / "config.json").write_text(json.dumps(config, ensure_ascii=False, indent=2), encoding="utf-8")
manifest = {
"created_utc": datetime.now(timezone.utc).isoformat(),
"experiment": "TSFA-SPR: Temporal-Shared-Private Representation",
"sample_count": len(samples),
"video_id_count": len({sample.group_id for sample in samples}),
"fold_count": len(splits),
"folds": fold_manifest,
"variants": list(VARIANTS),
"baseline_controls": ["TSFA-old", "TSFA+RawPrivate", "B0", "B4"],
"alignment_models_retrained": False,
"temporal_branch": "Frozen M4_sourceTime fold checkpoint; its attention weights pool fold-standardized original BERT/eGeMAPS/DeiT features.",
"emotion_labels_used_for_factorizer_training": False,
"feature_extractors_changed": False,
"shared_encoder": "One common MLP reused across Text, Audio, Vision; no explicit time code/slot index input.",
"private_encoders": "Three modality-specific MLPs with independent parameters.",
"loss": "L_shared + lambda_o*L_orth + lambda_r*L_rec; L_orth only within modality; per-modality mean reconstruction loss.",
"shared_correspondence_interpretation": "Same-slot vs same-video shifted-slot proxy; not independent human alignment ground truth.",
"emotion_probe_protocol": config["emotion_probe"],
"pca_control": "PCA fit on training-fold raw-private vectors over training video slots only; no emotion labels used.",
"paired_bootstrap": "2,000 video_id-cluster resamples; percentile 95% intervals; no multiple-comparison correction.",
"decomposition_diagnostics": {
"SPR_shared_modality_probe_accuracy": shared_source_accuracy,
"SPR_private_modality_probe_accuracy": private_source_accuracy,
"SPR_private_minus_shared_modality_accuracy": private_source_accuracy - shared_source_accuracy,
"SPR_mean_shared_crossmodal_auc": mean_shared_auc,
"SPR_reconstruction_mse": rec_by_branch,
"shared_private_separation_supported_by_predeclared_diagnostics": split_supported,
},
"checkpoint_hashes": input_hashes,
"input_paths": {
"features": str(args.feature_dir.resolve()),
"feature_manifest": str(args.manifest.resolve()),
"grouped_splits": str(args.splits.resolve()),
"frozen_m4_checkpoint_root": str(args.checkpoint_root.resolve()),
"math_predictions_read_only": str(args.math_predictions.resolve()),
"math_splits_read_only": str(args.math_splits.resolve()),
"old_tsfa_predictions": str(args.old_tsfa_predictions.resolve()),
"raw_private_predictions": str(args.raw_private_predictions.resolve()),
},
"input_sha256": {
"feature_manifest": _sha256(args.manifest),
"grouped_splits": _sha256(args.splits),
"math_predictions": _sha256(args.math_predictions),
"math_splits": _sha256(args.math_splits),
"old_tsfa_predictions": _sha256(args.old_tsfa_predictions),
"raw_private_predictions": _sha256(args.raw_private_predictions),
**input_hashes,
},
"outputs": [
"training_history.csv", "representation_diagnostics.csv", "shared_correspondence_summary.csv",
"modality_probe_summary.csv", "reconstruction_summary.csv", "emotion_probe_metrics.csv",
"emotion_probe_predictions.csv", "paired_contrasts.csv", "private_source_ablation.csv",
"dimension_control_summary.csv", "seed_summary.csv", *figures,
],
"device": str(device),
"gpu_name": torch.cuda.get_device_name(device) if device.type == "cuda" else None,
"python": platform.python_version(),
"pytorch": torch.__version__,
"scikit_learn": sklearn.__version__,
"elapsed_seconds": time.time() - started,
"interpretation_limits": [
"Same-slot contrastive positives use M4 latent slots as a training convention, not human event-level ground truth.",
"Frozen M4 source-time weights can carry temporal prior information into the pooled features.",
"Emotion metrics are lightweight frozen-representation probes on only 100 clips from 37 source videos.",
"Only seed 42 is run in this first pass; bootstrap intervals do not capture representation seed uncertainty.",
"RawPrivate-PCA and SPR-dim-matched outputs have equal 256 dimensions per slot and 1280 pooled dimensions.",
],
}
(args.output_dir / "run_manifest.json").write_text(json.dumps(manifest, ensure_ascii=False, indent=2), encoding="utf-8")
final_factorization_supported = _finalize_existing(args)
print(f"[TSFA-SPR complete] samples={len(samples)} device={device} output={args.output_dir}", flush=True)
print(f" separation diagnostics support factorization: {final_factorization_supported}", flush=True)
for key in dimension_keys:
if key in by_method_view:
row = by_method_view[key]
print(f" {key[0]}/{key[1]}: F1={row['macro_f1']:.3f}, MAE={row['mae']:.3f}, Pearson={row['pearson']:.3f}, d={row['feature_dimension']}", flush=True)
def _write_readme_section(
output_dir: Path,
config: Mapping[str, Any],
manifest: Mapping[str, Any],
summaries: Sequence[Mapping[str, Any]],
paired_rows: Sequence[Mapping[str, Any]],
dimension_rows: Sequence[Mapping[str, Any]],
private_rows: Sequence[Mapping[str, Any]],
) -> None:
metrics_lookup = {(row["method"], row["view"]): row for row in summaries}
main = metrics_lookup.get(("SPR", "shared_unfused+private_all"), {})
dimensions = {row["method"]: row for row in dimension_rows}
private_best_mae = min(
(row for row in private_rows if row["method"] == "SPR" and row["view"].startswith("private_")),
key=lambda row: float(row["mae"]),
default={},
)
lines = [
"# TSFA-SPR run notes",
"",
"This directory contains the first seed (42) of the temporal shared/private factorization experiment.",
"The M4_sourceTime checkpoint remains frozen. Factorizer losses use only source features and M4 slot relationships; emotion labels are used only by the held-out probes.",
"",
"## Main probe",
"",
f"- SPR shared-unfused + all-private: Accuracy {main.get('accuracy', float('nan')):.3f}, Macro-F1 {main.get('macro_f1', float('nan')):.3f}, MAE {main.get('mae', float('nan')):.3f}, Pearson {main.get('pearson', float('nan')):.3f}.",
f"- Best private-only MAE view: {private_best_mae.get('view', 'n/a')} (MAE {private_best_mae.get('mae', float('nan')):.3f}).",
"",
"## Dimension controls",
"",
]
for name in ("SPR", "SPR-dim-matched", "RawPrivate-PCA", "TSFA-old", "TSFA+RawPrivate"):
row = dimensions.get(name)
if row:
lines.append(f"- {name}: dimension {row['feature_dimension']}; Macro-F1 {row['macro_f1']:.3f}; MAE {row['mae']:.3f}; Pearson {row['pearson']:.3f}.")
lines.extend([
"",
"## Output map",
"",
"- `training_history.csv`: per-fold, per-epoch loss terms.",
"- `representation_diagnostics.csv`: held-out shared/private norms and normalized cross-covariance by modality.",
"- `shared_correspondence_summary.csv`: same-slot vs shifted-slot correspondence and slot-shuffle control.",
"- `modality_probe_summary.csv`: modality-source classification from frozen shared/private vectors.",
"- `reconstruction_summary.csv`: held-out standardized-source reconstruction from shared, private, and combined streams.",
"- `emotion_probe_metrics.csv` and `emotion_probe_predictions.csv`: five-fold grouped sentiment probes.",
"- `paired_contrasts.csv`: video-cluster paired bootstrap against TSFA-old, TSFA+RawPrivate, math B0/B4, and dimension controls.",
"- `checkpoints/`: fold-specific factorizer weights and train-fold PCA components.",
"",
"## Caution",
"",
"Same-slot AUC/retrieval tests agreement with the M4 slot convention, not independent human event alignment. A claim of successful shared/private separation requires the modality-source, covariance, reconstruction, and slot-shuffle diagnostics to agree; emotion score gains alone are insufficient.",
])
(output_dir / "README.md").write_text("\n".join(lines) + "\n", encoding="utf-8")
def _finalize_existing(args: argparse.Namespace) -> bool:
"""Refresh derived summaries/plots without retraining any factorizer."""
output_dir = args.output_dir
manifest_path = output_dir / "run_manifest.json"
if not manifest_path.is_file():
raise FileNotFoundError(f"no existing TSFA-SPR run manifest: {manifest_path}")
manifest = json.loads(manifest_path.read_text(encoding="utf-8"))
modality_rows = _read_csv(output_dir / "modality_probe_summary.csv")
correspondence_rows = _read_csv(output_dir / "shared_correspondence_summary.csv")
reconstruction_rows = _read_csv(output_dir / "reconstruction_summary.csv")
diagnostic_rows = _read_csv(output_dir / "representation_diagnostics.csv")
emotion_rows = _read_csv(output_dir / "emotion_probe_metrics.csv")
prediction_rows = _read_csv(output_dir / "emotion_probe_predictions.csv")
paired_rows = _read_csv(output_dir / "paired_contrasts.csv")
dimension_rows = _read_csv(output_dir / "dimension_control_summary.csv")
private_rows = _read_csv(output_dir / "private_source_ablation.csv")
paired_rows = [row for row in paired_rows
if row["comparison"] != "SPR-dim-matched-RawPrivate-PCA"]
paired_rows.extend(_paired_contrasts(
prediction_rows,
candidate_method="SPR-dim-matched",
candidate_view="SPR-dim-matched",
references=(("RawPrivate-PCA", "RawPrivate-PCA"),),
repeats=args.bootstrap_repeats,
seed=args.seed,
))
_write_csv(output_dir / "paired_contrasts.csv", paired_rows)
spr_probe = {row["branch"]: row for row in modality_rows if row["variant"] == "SPR"}
untrained_private_probe = next(
(row for row in modality_rows
if row["variant"] == "SP-noOrth-noRec" and row["branch"] == "private"),
None,
)
spr_pairs = [row for row in correspondence_rows if row["variant"] == "SPR"]
pair_checks = []
for row in spr_pairs:
auc = float(row["auc_mean_video"])
shuffled_auc = float(row["slot_shuffle_auc_mean_video"])
pair_checks.append({
"pair": row["pair"],
"matched_vs_shifted_auc": auc,
"slot_shuffle_auc": shuffled_auc,
"auc_drop_after_shuffle": auc - shuffled_auc,
"shared_correspondence_supported": auc >= 0.60 and auc - shuffled_auc >= 0.05,
})
spr_reconstruction = [row for row in reconstruction_rows
if row["variant"] == "SPR" and str(row["reconstruction_trained"]).lower() == "true"]
rec_by_branch = {
branch: float(np.mean([float(row["mse_mean_video"]) for row in spr_reconstruction
if row["branch"] == branch]))
for branch in ("shared_only", "private_only", "both")
}
shared_accuracy = float(spr_probe["shared"]["accuracy_mean"])
private_accuracy = float(spr_probe["private"]["accuracy_mean"])
untrained_private_accuracy = (
float(untrained_private_probe["accuracy_mean"]) if untrained_private_probe else float("nan")
)
source_separation = (
private_accuracy - shared_accuracy >= 0.05
and private_accuracy - untrained_private_accuracy >= 0.05
)
shared_correspondence = bool(pair_checks) and all(row["shared_correspondence_supported"] for row in pair_checks)
combined_reconstruction = rec_by_branch["both"] < min(rec_by_branch["shared_only"], rec_by_branch["private_only"])
factorization_supported = source_separation and shared_correspondence and combined_reconstruction
cross_covariance = {
modality: float(np.mean([float(row[f"{modality}_cross_covariance_norm"]) for row in diagnostic_rows
if row["variant"] == "SPR"]))
for modality in MODS
}
norm_ratio = {
modality: float(np.mean([float(row[f"{modality}_shared_norm_ratio"]) for row in diagnostic_rows
if row["variant"] == "SPR"]))
for modality in MODS
}
assessment = {
"shared_private_factorization_supported": factorization_supported,
"criteria": {
"private_source_probe_exceeds_shared_and_untrained_private_control_by_at_least_0_05": source_separation,
"each_shared_pair_auc_at_least_0_60_and_shuffle_drop_at_least_0_05": shared_correspondence,
"both_stream_reconstruction_beats_each_single_stream": combined_reconstruction,
},
"spr_shared_source_probe_accuracy": shared_accuracy,
"spr_private_source_probe_accuracy": private_accuracy,
"untrained_private_control_source_probe_accuracy": untrained_private_accuracy,
"spr_private_probe_gain_over_untrained_control": private_accuracy - untrained_private_accuracy,
"private_minus_shared_source_probe_accuracy": private_accuracy - shared_accuracy,
"shared_pair_checks": pair_checks,
"reconstruction_mse_by_branch": rec_by_branch,
"normalized_cross_covariance_by_modality": cross_covariance,
"shared_norm_ratio_by_modality": norm_ratio,
"interpretation": (
"The shared/private branches are distinguishable by source classification, but the no-private-loss control already identifies modality perfectly, so this is not evidence that training learned private semantic content. The shared stream also does not establish robust same-slot cross-modal correspondence; do not claim successful shared/private semantic disentanglement."
if not factorization_supported else
"All operational diagnostics passed; the result supports further validation, not a claim of human event-level alignment."
),
"criteria_note": "These are operational descriptive checks for this small first pass, not inferential significance thresholds.",
}
(output_dir / "decomposition_assessment.json").write_text(
json.dumps(assessment, ensure_ascii=False, indent=2), encoding="utf-8"
)
manifest["decomposition_diagnostics"] = {
**assessment,
"shared_correspondence": pair_checks,
}
manifest["outputs"] = sorted(set(manifest.get("outputs", [])) | {
"private_source_paired_contrasts.csv", "decomposition_assessment.json", "README.md",
"shared_correspondence_by_clip.csv", "modality_probe_by_fold.csv", "reconstruction_by_clip.csv",
})
manifest_path.write_text(json.dumps(manifest, ensure_ascii=False, indent=2), encoding="utf-8")
view_contrasts = _paired_view_contrasts(
prediction_rows, repeats=args.bootstrap_repeats, seed=args.seed
)
_write_csv(output_dir / "private_source_paired_contrasts.csv", view_contrasts)
# Regenerate the dimension-control panel after adding all native-size controls.
_plot_summaries(output_dir, modality_rows, correspondence_rows, reconstruction_rows,
emotion_rows, diagnostic_rows)
by_metric = {(row["comparison"], row["metric"]): row for row in paired_rows}
best_f1 = max((row for row in private_rows if row["view"].startswith("private_")),
key=lambda row: float(row["macro_f1"]), default={})
best_mae = min((row for row in private_rows if row["view"].startswith("private_")),
key=lambda row: float(row["mae"]), default={})
emotion_lookup = {(row["method"], row["view"]): row for row in emotion_rows}
dim_lookup = {row["method"]: row for row in dimension_rows}
main_row = emotion_lookup[("SPR", "shared_unfused+private_all")]
auc_text = ", ".join(
f"{row['pair']}={float(row['matched_vs_shifted_auc']):.3f}" for row in pair_checks
)
lines = [
"# TSFA-SPR:时序—共享—私有分解",
"",
"## 首轮结果(seed 42)",
"",
"M4_sourceTime 的五折 temporal checkpoints 已冻结;共享/私有因子器只用折内标准化的原始 BERT、eGeMAPS、DeiT 特征及 M4 权重训练,没有使用情感标签。Probe 口径沿用现有五折 GroupKFold、五段顺序池化、LogisticRegression(C=0.05)、Ridge(alpha=25) 与 [-3,3] 限幅。",
"",
f"SPR 主表示(unfused shared + all private):Accuracy {float(main_row['accuracy']):.3f},Macro-F1 {float(main_row['macro_f1']):.3f},MAE {float(main_row['mae']):.3f},Pearson {float(main_row['pearson']):.3f}。",
f"来源模态 probe:shared accuracy {shared_accuracy:.3f},SPR private accuracy {private_accuracy:.3f};但不含 private 训练损失的 SP-noOrth-noRec 控制也达到 {untrained_private_accuracy:.3f},说明该识别率可能来自模态专属分支/特征分布,不能当成已学会私有语义。shared 三组匹配 AUC 为 {auc_text},slot shuffle 后几乎不下降。",
f"合并重构 MSE {rec_by_branch['both']:.3f},shared-only {rec_by_branch['shared_only']:.3f},private-only {rec_by_branch['private_only']:.3f}。",
"",
f"诊断结论:**{'支持' if factorization_supported else '不支持'}已成功实现 shared/private 语义解耦**。合并分支能改善源特征重构,但 shared 跨模态同槽检索约为机会水平,且 shuffle 控制没有明显下降;private 来源识别在随机未训练控制中也已饱和。因此不应把本轮写成成功的 shared/private 语义对齐。",
"",
"## 私有来源探针",
"",
f"- private-only Macro-F1 最高:{best_f1.get('view', 'n/a')}({float(best_f1.get('macro_f1', float('nan'))):.3f})。",
f"- private-only MAE 最低:{best_mae.get('view', 'n/a')}({float(best_mae.get('mae', float('nan'))):.3f})。",
"- Audio 的单独分类探针 Macro-F1 点估计领先,Vision 的单独回归 probe MAE 点估计最低;video-group paired 区间跨 0,不能据此确定模态贡献差异。",
"",
"## 维度控制和基线",
"",
]
for name in ("TSFA-old", "TSFA+RawPrivate", "SPR", "SPR-dim-matched", "RawPrivate-PCA"):
row = dim_lookup.get(name)
if row:
lines.append(f"- {name}: d={row['feature_dimension']} pooled,Macro-F1 {float(row['macro_f1']):.3f},MAE {float(row['mae']):.3f},Pearson {float(row['pearson']):.3f}。")
main_vs_raw = by_metric.get(("SPR-RawPrivate-PCA", "mae"))
matched_vs_raw = by_metric.get(("SPR-dim-matched-RawPrivate-PCA", "mae"))
if main_vs_raw:
lines.append(f"- SPR 主表示相对 RawPrivate-PCA 的 MAE 差为 {float(main_vs_raw['delta_candidate_minus_reference']):+.3f},视频组 bootstrap 95% CI [{float(main_vs_raw['video_cluster_bootstrap_ci95_low']):+.3f}, {float(main_vs_raw['video_cluster_bootstrap_ci95_high']):+.3f}]。"
" RawPrivate-PCA 有较低 MAE,维度匹配后没有证据支持 learned shared/private 分解带来强度回归增益。")
if matched_vs_raw:
lines.append(f"- 直接比较同为 1280 维 pooled 的 SPR-dim-matched 与 RawPrivate-PCA,MAE 差为 {float(matched_vs_raw['delta_candidate_minus_reference']):+.3f},95% CI [{float(matched_vs_raw['video_cluster_bootstrap_ci95_low']):+.3f}, {float(matched_vs_raw['video_cluster_bootstrap_ci95_high']):+.3f}]。"
" 这是控制输出维数后的成对对照。")
lines.extend([
"",
"RawPrivate-PCA 将每个 slot 的 985 维原始私有特征仅用训练折做 PCA 到 256 维;与 fused-shared + all-private 的 SPR-dim-matched 都是 256 维/slot、1280 维/clip。RawPrivate-PCA 的低 MAE 说明可压缩 raw residual 已足以保留强度信息,但分类和 Pearson 并未同步领先。",
"",
"## 文件",
"",
"- `emotion_probe_metrics.csv` / `emotion_probe_predictions.csv`:全部五折 OOF probe。",
"- `paired_contrasts.csv`:SPR 主表示对 TSFA-old、TSFA+RawPrivate、math B0/B4 和维度对照的 2,000 次 video_id paired bootstrap;未做多重比较校正。",
"- `private_source_ablation.csv` 与 `private_source_paired_contrasts.csv`:模态来源 probe 及来源间配对区间。",
"- `decomposition_assessment.json`:shared/private 诊断判定和逐项依据。",
"- `representation_diagnostics.csv`、`shared_correspondence_summary.csv`、`modality_probe_summary.csv`、`reconstruction_summary.csv`:结构诊断明细。",
"- `architecture.png` 及其余七张 `.png`:实验结构、对应、重构、probe 与维度控制图。",
"",
"## 解释范围",
"",
"当前只有一个表示学习 seed,样本为 100 条 clips / 37 个来源视频;bootstrap 区间反映来源视频抽样不确定性,不含 seed 不确定性。same-slot 检索依赖 M4 槽位约定,并非人工事件级对齐真值。分类与回归来自独立轻量 probe。",
"",
"## 下一步",
"",
"暂不升级现有 TSFA 为正式 TSFA-SPR。最值得先做的是: (1) 检查同槽对比目标是否能在独立视频内建立可打乱验证的 shared correspondence;(2) 若共享诊断改善,再运行 seed 3407 和 2026,确认模态来源与情感 probe 差异是否稳定。",
])
(output_dir / "README.md").write_text("\n".join(lines) + "\n", encoding="utf-8")
print(f"[TSFA-SPR finalize] factorization_supported={factorization_supported}; assessment={output_dir / 'decomposition_assessment.json'}")
return factorization_supported
def build_parser() -> argparse.ArgumentParser:
project = Path(__file__).resolve().parents[1]
repository = project.parents[1]
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto")
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--epochs", type=int, default=40)
parser.add_argument("--batch-size", type=int, default=8)
parser.add_argument("--learning-rate", type=float, default=1e-3)
parser.add_argument("--temperature", type=float, default=0.1)
parser.add_argument("--common-dim", type=int, default=128)
parser.add_argument("--shared-dim", type=int, default=64)
parser.add_argument("--private-dim", type=int, default=64)
parser.add_argument("--orth-weight", type=float, default=0.1)
parser.add_argument("--reconstruction-weight", type=float, default=1.0)
parser.add_argument("--pca-dim", type=int, default=256)
parser.add_argument("--shuffle-repeats", type=int, default=20)
parser.add_argument("--bootstrap-repeats", type=int, default=2000)
parser.add_argument("--modality-gap-threshold", type=float, default=0.05)
parser.add_argument("--feature-dir", type=Path, default=project / "outputs/q1_features/features")
parser.add_argument("--manifest", type=Path, default=project / "outputs/audit/manifest.csv")
parser.add_argument("--splits", type=Path, default=project / "outputs/method_comparison/splits.json")
parser.add_argument("--checkpoint-root", type=Path, default=project / "outputs/alignment_debug/heldout")
parser.add_argument("--output-dir", type=Path, default=project / "outputs/tsfa_shared_private")
parser.add_argument("--old-tsfa-predictions", type=Path,
default=project / "outputs/tsfa_emotion_probe/emotion_probe_predictions.csv")
parser.add_argument("--raw-private-predictions", type=Path,
default=project / "outputs/tsfa_av_private_ablation/predictions.csv")
parser.add_argument("--math-predictions", type=Path,
default=repository / "math/results/model_comparison/oof_predictions.csv")
parser.add_argument("--math-splits", type=Path,
default=repository / "math/results/model_comparison/split_assignments.csv")
parser.add_argument("--finalize-existing", action="store_true",
help="refresh derived diagnostics and report files without retraining")
return parser
def main() -> None:
args = build_parser().parse_args()
if args.finalize_existing:
_finalize_existing(args)
else:
run(args)
if __name__ == "__main__":
main()