1863 lines
92 KiB
Python
1863 lines
92 KiB
Python
"""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()
|