720 lines
32 KiB
Python
720 lines
32 KiB
Python
"""Five-fold TSFA ablation: local Audio-Vision edge x raw modality residuals."""
|
||
|
||
from __future__ import annotations
|
||
|
||
import argparse
|
||
import csv
|
||
import json
|
||
import math
|
||
import platform
|
||
import random
|
||
import shutil
|
||
import time
|
||
from datetime import datetime, timezone
|
||
from pathlib import Path
|
||
from typing import Any, Mapping, Sequence
|
||
|
||
import numpy as np
|
||
import sklearn
|
||
import torch
|
||
import matplotlib
|
||
|
||
matplotlib.use("Agg")
|
||
import matplotlib.pyplot as plt
|
||
from sklearn.linear_model import LogisticRegression, Ridge
|
||
from sklearn.metrics import confusion_matrix, f1_score
|
||
from sklearn.pipeline import make_pipeline
|
||
from sklearn.preprocessing import StandardScaler
|
||
from torch import Tensor, nn
|
||
|
||
from .compare_emotion_probes import METRICS, _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, _pool_five_segments
|
||
from .tsfa_experiment import (
|
||
GRID_SIZE,
|
||
HIDDEN_SIZE,
|
||
TSFASemanticBranch,
|
||
_candidate_mask,
|
||
_collate_temporal,
|
||
_collect_fold_features,
|
||
_generate_tsfa_outputs,
|
||
_load_semantic_checkpoint,
|
||
_local_contrastive_loss,
|
||
)
|
||
from .types import MODALITIES
|
||
|
||
|
||
VARIANTS = {
|
||
(False, False): "TSFA-T",
|
||
(True, False): "TSFA-AV",
|
||
(False, True): "TSFA-T+Private",
|
||
(True, True): "TSFA-AV+Private",
|
||
}
|
||
CONTRASTS = {
|
||
"AV_without_private": {"TSFA-AV": 1, "TSFA-T": -1},
|
||
"Private_without_AV": {"TSFA-T+Private": 1, "TSFA-T": -1},
|
||
"AV_with_private": {"TSFA-AV+Private": 1, "TSFA-T+Private": -1},
|
||
"Private_with_AV": {"TSFA-AV+Private": 1, "TSFA-AV": -1},
|
||
"Both_vs_original": {"TSFA-AV+Private": 1, "TSFA-T": -1},
|
||
"AV_x_Private_interaction": {
|
||
"TSFA-AV+Private": 1,
|
||
"TSFA-AV": -1,
|
||
"TSFA-T+Private": -1,
|
||
"TSFA-T": 1,
|
||
},
|
||
}
|
||
|
||
|
||
class TSFAAVBranch(TSFASemanticBranch):
|
||
"""Add reciprocal A-V messages within the unchanged frozen-M4 candidate masks."""
|
||
|
||
def __init__(self, dimension: int = HIDDEN_SIZE) -> None:
|
||
super().__init__(dimension)
|
||
self.av_queries = nn.ModuleDict({
|
||
name: nn.Linear(dimension, dimension, bias=False) for name in ("audio", "vision")
|
||
})
|
||
self.av_keys = nn.ModuleDict({
|
||
name: nn.Linear(dimension, dimension, bias=False) for name in ("audio", "vision")
|
||
})
|
||
self.av_values = nn.ModuleDict({
|
||
name: nn.Linear(dimension, dimension, bias=False) for name in ("audio", "vision")
|
||
})
|
||
|
||
def av_attend(
|
||
self,
|
||
query_content: Tensor,
|
||
source_values: Tensor,
|
||
*,
|
||
query_modality: str,
|
||
source_modality: str,
|
||
candidate_mask: Tensor,
|
||
) -> tuple[Tensor, Tensor]:
|
||
query = self.av_queries[query_modality](query_content)
|
||
keys = self.av_keys[source_modality](source_values)
|
||
scores = torch.bmm(query, keys.transpose(1, 2)) / math.sqrt(query.shape[-1])
|
||
scores = scores.masked_fill(~candidate_mask, torch.finfo(scores.dtype).min)
|
||
weights = torch.softmax(scores, dim=-1)
|
||
message = torch.bmm(weights, self.av_values[source_modality](source_values))
|
||
return weights, message
|
||
|
||
|
||
def _av_forward(
|
||
branch: TSFAAVBranch,
|
||
batch: Mapping[str, Any],
|
||
m4_audio: Tensor,
|
||
m4_vision: Tensor,
|
||
*,
|
||
delta: float,
|
||
) -> tuple[dict[str, Tensor], dict[str, Tensor], dict[str, Tensor]]:
|
||
masks = {}
|
||
for modality in ("audio", "vision"):
|
||
masks[modality], _ = _candidate_mask(
|
||
batch["times"][modality],
|
||
batch["valid"][modality],
|
||
batch["centers"][modality],
|
||
delta=delta,
|
||
mode="local",
|
||
)
|
||
_, text_to_audio = branch.attend(
|
||
batch["text_content"], batch["values"]["audio"], "audio", masks["audio"]
|
||
)
|
||
_, text_to_vision = branch.attend(
|
||
batch["text_content"], batch["values"]["vision"], "vision", masks["vision"]
|
||
)
|
||
a_to_v_weights, vision_to_audio = branch.av_attend(
|
||
m4_audio,
|
||
batch["values"]["vision"],
|
||
query_modality="audio",
|
||
source_modality="vision",
|
||
candidate_mask=masks["vision"],
|
||
)
|
||
v_to_a_weights, audio_to_vision = branch.av_attend(
|
||
m4_vision,
|
||
batch["values"]["audio"],
|
||
query_modality="vision",
|
||
source_modality="audio",
|
||
candidate_mask=masks["audio"],
|
||
)
|
||
shared = {
|
||
"text": batch["text_content"],
|
||
"audio": text_to_audio + vision_to_audio,
|
||
"vision": text_to_vision + audio_to_vision,
|
||
}
|
||
weights = {"audio_to_vision": a_to_v_weights, "vision_to_audio": v_to_a_weights}
|
||
return shared, weights, masks
|
||
|
||
|
||
def _m4_content(
|
||
ids: Sequence[str], temporal_by_id: Mapping[str, Mapping[str, Any]], modality: str,
|
||
device: torch.device,
|
||
) -> Tensor:
|
||
return torch.from_numpy(np.stack([
|
||
temporal_by_id[sample_id]["content"][modality] for sample_id in ids
|
||
])).to(device)
|
||
|
||
|
||
def _fit_av_branch(
|
||
*,
|
||
fold: int,
|
||
train_samples: Sequence[FeatureSample],
|
||
temporal_by_id: Mapping[str, Mapping[str, Any]],
|
||
device: torch.device,
|
||
args: argparse.Namespace,
|
||
) -> tuple[TSFAAVBranch, list[dict[str, Any]]]:
|
||
fold_seed = args.seed + fold * 101
|
||
random.seed(fold_seed)
|
||
np.random.seed(fold_seed)
|
||
torch.manual_seed(fold_seed)
|
||
if device.type == "cuda":
|
||
torch.cuda.manual_seed_all(fold_seed)
|
||
branch = TSFAAVBranch().to(device)
|
||
optimizer = torch.optim.AdamW(branch.parameters(), lr=args.semantic_learning_rate, weight_decay=1e-4)
|
||
rng = np.random.default_rng(fold_seed)
|
||
train_ids = [sample.sample_id for sample in train_samples]
|
||
history = []
|
||
branch.train()
|
||
for epoch in range(1, args.semantic_epochs + 1):
|
||
order = rng.permutation(len(train_ids))
|
||
losses = []
|
||
for start in range(0, len(order), args.batch_size):
|
||
ids = [train_ids[int(index)] for index in order[start : start + args.batch_size]]
|
||
batch = _collate_temporal(ids, temporal_by_id, device)
|
||
shared, _, _ = _av_forward(
|
||
branch,
|
||
batch,
|
||
_m4_content(ids, temporal_by_id, "audio", device),
|
||
_m4_content(ids, temporal_by_id, "vision", device),
|
||
delta=args.delta,
|
||
)
|
||
loss = _local_contrastive_loss(
|
||
branch,
|
||
shared["text"],
|
||
shared["audio"],
|
||
shared["vision"],
|
||
temperature=args.local_temperature,
|
||
)
|
||
if not torch.isfinite(loss):
|
||
raise FloatingPointError(f"non-finite AV semantic loss in fold {fold}, epoch {epoch}")
|
||
optimizer.zero_grad(set_to_none=True)
|
||
loss.backward()
|
||
nn.utils.clip_grad_norm_(branch.parameters(), 1.0)
|
||
optimizer.step()
|
||
losses.append(float(loss.detach().item()))
|
||
history.append({
|
||
"fold": fold,
|
||
"epoch": epoch,
|
||
"seed": fold_seed,
|
||
"train_loss": float(np.mean(losses)),
|
||
})
|
||
branch.eval()
|
||
return branch, history
|
||
|
||
|
||
def _generate_av_outputs(
|
||
*,
|
||
branch: TSFAAVBranch,
|
||
sample_ids: Sequence[str],
|
||
temporal_by_id: Mapping[str, Mapping[str, Any]],
|
||
samples_by_id: Mapping[str, FeatureSample],
|
||
device: torch.device,
|
||
args: argparse.Namespace,
|
||
) -> tuple[dict[str, dict[str, np.ndarray]], list[dict[str, Any]]]:
|
||
content_by_id = {}
|
||
diagnostics = []
|
||
branch.eval()
|
||
with torch.no_grad():
|
||
for start in range(0, len(sample_ids), args.batch_size):
|
||
ids = list(sample_ids[start : start + args.batch_size])
|
||
batch = _collate_temporal(ids, temporal_by_id, device)
|
||
shared, weights, masks = _av_forward(
|
||
branch,
|
||
batch,
|
||
_m4_content(ids, temporal_by_id, "audio", device),
|
||
_m4_content(ids, temporal_by_id, "vision", device),
|
||
delta=args.delta,
|
||
)
|
||
for index, sample_id in enumerate(ids):
|
||
content_by_id[sample_id] = {
|
||
modality: shared[modality][index].cpu().numpy().astype(np.float32, copy=False)
|
||
for modality in MODALITIES
|
||
}
|
||
for direction, source in (("audio_to_vision", "vision"), ("vision_to_audio", "audio")):
|
||
length = len(samples_by_id[sample_id].features[source])
|
||
one_weights = weights[direction][index, :, :length]
|
||
one_times = batch["times"][source][index, :length]
|
||
one_center = batch["centers"][source][index]
|
||
candidate_count = masks[source][index, :, :length].sum(dim=-1)
|
||
diagnostics.append({
|
||
"sample_id": sample_id,
|
||
"video_id": samples_by_id[sample_id].group_id,
|
||
"direction": direction,
|
||
"source_modality": source,
|
||
"candidate_count_mean": float(candidate_count.float().mean().item()),
|
||
"attention_center_error_normalized_time": float(
|
||
(one_weights @ one_times - one_center).abs().mean().item()
|
||
),
|
||
"attention_row_sum_max_error": float(
|
||
(one_weights.sum(dim=-1) - 1).abs().max().item()
|
||
),
|
||
})
|
||
return content_by_id, diagnostics
|
||
|
||
|
||
def _private_content(
|
||
samples: Sequence[FeatureSample],
|
||
feature_stats: Any,
|
||
temporal_by_id: Mapping[str, Mapping[str, Any]],
|
||
) -> dict[str, dict[str, np.ndarray]]:
|
||
result = {}
|
||
for sample in samples:
|
||
standardized = standardized_features(sample, feature_stats)
|
||
record = temporal_by_id[sample.sample_id]
|
||
result[sample.sample_id] = {
|
||
modality: (
|
||
np.asarray(record["weights"][modality], dtype=np.float32)
|
||
@ standardized[modality]
|
||
).astype(np.float32, copy=False)
|
||
for modality in MODALITIES
|
||
}
|
||
return result
|
||
|
||
|
||
def _vector(
|
||
sample_id: str,
|
||
shared_by_id: Mapping[str, Mapping[str, np.ndarray]],
|
||
private_by_id: Mapping[str, Mapping[str, np.ndarray]],
|
||
use_private: bool,
|
||
) -> np.ndarray:
|
||
parts = [_pool_five_segments(shared_by_id[sample_id][modality]) for modality in MODALITIES]
|
||
if use_private:
|
||
parts.extend(_pool_five_segments(private_by_id[sample_id][modality]) for modality in MODALITIES)
|
||
return np.concatenate(parts)
|
||
|
||
|
||
def _fit_fold_probe(
|
||
*,
|
||
fold: int,
|
||
train_samples: Sequence[FeatureSample],
|
||
heldout_samples: Sequence[FeatureSample],
|
||
shared_by_id: Mapping[str, Mapping[str, np.ndarray]],
|
||
private_by_id: Mapping[str, Mapping[str, np.ndarray]],
|
||
method: str,
|
||
use_private: bool,
|
||
seed: int,
|
||
) -> list[dict[str, Any]]:
|
||
train_x = np.stack([
|
||
_vector(sample.sample_id, shared_by_id, private_by_id, use_private)
|
||
for sample in train_samples
|
||
])
|
||
heldout_x = np.stack([
|
||
_vector(sample.sample_id, shared_by_id, private_by_id, use_private)
|
||
for sample in heldout_samples
|
||
])
|
||
train_class = np.asarray([_class_from_sentiment(sample.sentiment) for sample in train_samples])
|
||
train_value = np.asarray([sample.sentiment for sample in train_samples], dtype=np.float64)
|
||
classifier = make_pipeline(
|
||
StandardScaler(),
|
||
LogisticRegression(C=0.05, max_iter=5000, solver="lbfgs", random_state=seed),
|
||
)
|
||
regressor = make_pipeline(StandardScaler(), Ridge(alpha=25.0))
|
||
classifier.fit(train_x, train_class)
|
||
regressor.fit(train_x, train_value)
|
||
predicted_class = classifier.predict(heldout_x)
|
||
predicted_unclipped = regressor.predict(heldout_x)
|
||
predicted_value = np.clip(predicted_unclipped, -3.0, 3.0)
|
||
return [{
|
||
"method": method,
|
||
"av_edge": int(method.startswith("TSFA-AV")),
|
||
"private_residual": int(use_private),
|
||
"fold": fold,
|
||
"sample_id": sample.sample_id,
|
||
"video_id": sample.group_id,
|
||
"true_class_id": _class_from_sentiment(sample.sentiment),
|
||
"predicted_class_id": int(predicted_class[index]),
|
||
"true_label": float(sample.sentiment),
|
||
"predicted_label": float(predicted_value[index]),
|
||
"predicted_label_unclipped": float(predicted_unclipped[index]),
|
||
"feature_dimension": int(train_x.shape[1]),
|
||
} for index, sample in enumerate(heldout_samples)]
|
||
|
||
|
||
def _score_rows(rows: Sequence[Mapping[str, Any]]) -> dict[str, float]:
|
||
return _scores(
|
||
np.asarray([int(row["true_class_id"]) for row in rows]),
|
||
np.asarray([int(row["predicted_class_id"]) for row in rows]),
|
||
np.asarray([float(row["true_label"]) for row in rows]),
|
||
np.asarray([float(row["predicted_label"]) for row in rows]),
|
||
)
|
||
|
||
|
||
def _summarize(
|
||
prediction_rows: Sequence[Mapping[str, Any]],
|
||
*,
|
||
bootstrap_repeats: int,
|
||
seed: int,
|
||
) -> tuple[list[dict[str, Any]], list[dict[str, Any]], list[dict[str, Any]]]:
|
||
by_method = {
|
||
method: {str(row["sample_id"]): row for row in prediction_rows if row["method"] == method}
|
||
for method in VARIANTS.values()
|
||
}
|
||
sample_ids = sorted(next(iter(by_method.values())))
|
||
if any(set(rows) != set(sample_ids) for rows in by_method.values()):
|
||
raise ValueError("four ablation variants do not cover identical samples")
|
||
groups = {sample_id: str(by_method["TSFA-T"][sample_id]["video_id"]) for sample_id in sample_ids}
|
||
if any(
|
||
str(row["video_id"]) != groups[sample_id]
|
||
for rows in by_method.values() for sample_id, row in rows.items()
|
||
):
|
||
raise ValueError("video_id mismatch across ablation cells")
|
||
|
||
metrics_rows = []
|
||
confusion_rows = []
|
||
point_scores = {}
|
||
for method, mapping in by_method.items():
|
||
rows = [mapping[sample_id] for sample_id in sample_ids]
|
||
scores = _score_rows(rows)
|
||
point_scores[method] = scores
|
||
fold_f1 = [float(f1_score(
|
||
[int(row["true_class_id"]) for row in rows if int(row["fold"]) == fold],
|
||
[int(row["predicted_class_id"]) for row in rows if int(row["fold"]) == fold],
|
||
labels=[0, 1, 2], average="macro", zero_division=0,
|
||
)) for fold in range(1, 6)]
|
||
metrics_rows.append({
|
||
"method": method,
|
||
"av_edge": rows[0]["av_edge"],
|
||
"private_residual": rows[0]["private_residual"],
|
||
"sample_count": len(rows),
|
||
"feature_dimension": rows[0]["feature_dimension"],
|
||
**scores,
|
||
"macro_f1_fold_mean": float(np.mean(fold_f1)),
|
||
"macro_f1_fold_sd": float(np.std(fold_f1, ddof=1)),
|
||
})
|
||
matrix = confusion_matrix(
|
||
[int(row["true_class_id"]) for row in rows],
|
||
[int(row["predicted_class_id"]) for row in rows],
|
||
labels=[0, 1, 2],
|
||
)
|
||
confusion_rows.extend({
|
||
"method": method,
|
||
"true_class": CLASS_NAMES[true_id],
|
||
"predicted_class": CLASS_NAMES[predicted_id],
|
||
"count": int(matrix[true_id, predicted_id]),
|
||
} for true_id in range(3) for predicted_id in range(3))
|
||
|
||
group_names = sorted(set(groups.values()))
|
||
group_ids = {group: [sample_id for sample_id in sample_ids if groups[sample_id] == group]
|
||
for group in group_names}
|
||
rng = np.random.default_rng(seed)
|
||
bootstrap_values = {
|
||
contrast: {metric: [] for metric in METRICS} for contrast in CONTRASTS
|
||
}
|
||
for _ in range(bootstrap_repeats):
|
||
drawn = rng.choice(group_names, size=len(group_names), replace=True)
|
||
drawn_ids = [sample_id for group in drawn for sample_id in group_ids[str(group)]]
|
||
scores = {
|
||
method: _score_rows([mapping[sample_id] for sample_id in drawn_ids])
|
||
for method, mapping in by_method.items()
|
||
}
|
||
for contrast, weights in CONTRASTS.items():
|
||
for metric in METRICS:
|
||
bootstrap_values[contrast][metric].append(
|
||
sum(weight * scores[method][metric] for method, weight in weights.items())
|
||
)
|
||
contrast_rows = []
|
||
for contrast, weights in CONTRASTS.items():
|
||
for metric in METRICS:
|
||
values = np.asarray(bootstrap_values[contrast][metric])
|
||
values = values[np.isfinite(values)]
|
||
contrast_rows.append({
|
||
"contrast": contrast,
|
||
"metric": metric,
|
||
"point_estimate": sum(
|
||
weight * point_scores[method][metric] for method, weight in weights.items()
|
||
),
|
||
"ci95_low": float(np.quantile(values, 0.025)),
|
||
"ci95_high": float(np.quantile(values, 0.975)),
|
||
"video_group_count": len(group_names),
|
||
"bootstrap_repeats": bootstrap_repeats,
|
||
})
|
||
return metrics_rows, contrast_rows, confusion_rows
|
||
|
||
|
||
def _verify_original_control(
|
||
prediction_rows: Sequence[Mapping[str, Any]], path: Path,
|
||
) -> None:
|
||
with path.open("r", encoding="utf-8-sig", newline="") as stream:
|
||
previous = {
|
||
row["sample_id"]: row for row in csv.DictReader(stream)
|
||
if row["method"] == "TSFA-main" and row["view"] == "all_modalities"
|
||
}
|
||
original = {str(row["sample_id"]): row for row in prediction_rows if row["method"] == "TSFA-T"}
|
||
if set(original) != set(previous):
|
||
raise ValueError("original TSFA control sample IDs do not match previous probe")
|
||
for sample_id, row in original.items():
|
||
old = previous[sample_id]
|
||
if int(row["fold"]) != int(old["fold"]):
|
||
raise ValueError(f"original TSFA fold mismatch for {sample_id}")
|
||
if int(row["predicted_class_id"]) != int(old["predicted_class_id"]):
|
||
raise ValueError(f"original TSFA class prediction changed for {sample_id}")
|
||
if not np.isclose(float(row["predicted_label"]), float(old["predicted_label"]), atol=1e-5):
|
||
raise ValueError(f"original TSFA regression prediction changed for {sample_id}")
|
||
|
||
|
||
def _plot_factorial(metrics_path: Path, output_path: Path) -> None:
|
||
with metrics_path.open("r", encoding="utf-8-sig", newline="") as stream:
|
||
rows = list(csv.DictReader(stream))
|
||
if len(rows) != 4:
|
||
raise ValueError("factorial plot requires exactly four ablation cells")
|
||
fig, axes = plt.subplots(1, 2, figsize=(10, 4.2), constrained_layout=True)
|
||
panels = (("mae", "Emotion intensity MAE ↓"), ("macro_f1", "Polarity Macro-F1 ↑"))
|
||
for axis, (metric, title) in zip(axes, panels, strict=True):
|
||
for private, label, color in ((0, "No private residual", "#4c78a8"),
|
||
(1, "With private residual", "#e45756")):
|
||
selected = sorted(
|
||
(row for row in rows if int(row["private_residual"]) == private),
|
||
key=lambda row: int(row["av_edge"]),
|
||
)
|
||
values = [float(row[metric]) for row in selected]
|
||
axis.plot([0, 1], values, marker="o", markersize=7, linewidth=2,
|
||
color=color, label=label)
|
||
for x, value in enumerate(values):
|
||
axis.annotate(f"{value:.3f}", (x, value), xytext=(0, 7),
|
||
textcoords="offset points", ha="center", fontsize=9)
|
||
axis.set_xticks([0, 1], ["No A–V edge", "Local A–V edge"])
|
||
axis.set_xlim(-0.18, 1.18)
|
||
axis.set_title(title)
|
||
axis.grid(axis="y", alpha=0.25)
|
||
axes[0].set_ylim(0.50, 0.88)
|
||
axes[1].set_ylim(0.27, 0.47)
|
||
axes[0].set_ylabel("OOF metric on 100 clips")
|
||
axes[0].legend(frameon=False, loc="upper right", fontsize=8)
|
||
fig.suptitle("TSFA local A–V edge × preserved modality features")
|
||
fig.savefig(output_path, dpi=190, bbox_inches="tight")
|
||
plt.close(fig)
|
||
|
||
|
||
def _refresh_report_bundle(output_dir: Path) -> None:
|
||
bundle = output_dir / "report_bundle"
|
||
bundle.mkdir(parents=True, exist_ok=True)
|
||
for name in (
|
||
"metrics.csv", "predictions.csv", "paired_contrasts.csv", "confusion_matrix.csv",
|
||
"av_attention_diagnostics.csv", "av_training_history.csv", "factorial_effects.png",
|
||
"private_vs_math.csv", "private_vs_math_manifest.json", "run_manifest.json",
|
||
):
|
||
source = output_dir / name
|
||
if source.is_file():
|
||
shutil.copy2(source, bundle / name)
|
||
(bundle / "README.md").write_text(
|
||
"# TSFA A–V edge × private residual ablation\n\n"
|
||
"Four OOF cells on the same 100 clips and five video-group folds. "
|
||
"`metrics.csv` gives Accuracy, fixed-three-class Macro-F1, clipped MAE, "
|
||
"and Pearson. `paired_contrasts.csv` reports video-cluster bootstrap intervals "
|
||
"for the two factors and their interaction. `private_vs_math.csv` compares "
|
||
"the strongest regression cell with math B0–B4 when available. The "
|
||
"original math files were read only. See `run_manifest.json` for the "
|
||
"protocol and limitations. Full AV branch checkpoints remain one level up.\n",
|
||
encoding="utf-8",
|
||
)
|
||
|
||
|
||
def run(args: argparse.Namespace) -> None:
|
||
started = time.time()
|
||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu") if args.device == "auto" else torch.device(args.device)
|
||
if device.type == "cuda" and not torch.cuda.is_available():
|
||
raise RuntimeError("CUDA was requested but is unavailable")
|
||
samples = load_feature_samples(args.feature_dir, args.manifest)
|
||
samples_by_id = {sample.sample_id: sample for sample in samples}
|
||
splits = json.loads(args.splits.read_text(encoding="utf-8"))
|
||
heldout_ids = [sample_id for split in splits for sample_id in split["validation_sample_ids"]]
|
||
if len(splits) != 5 or len(heldout_ids) != len(set(heldout_ids)) or set(heldout_ids) != set(samples_by_id):
|
||
raise ValueError("expected exactly five grouped folds covering all samples once")
|
||
checkpoint_store = torch.load(args.tsfa_checkpoints, map_location="cpu", weights_only=False)
|
||
prediction_rows = []
|
||
diagnostic_rows = []
|
||
history_rows = []
|
||
av_checkpoints = {}
|
||
fold_manifest = []
|
||
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]
|
||
train_groups = {sample.group_id for sample in train_samples}
|
||
heldout_groups = {sample.group_id for sample in heldout_samples}
|
||
if train_groups & heldout_groups:
|
||
raise ValueError(f"video_id leakage in fold {fold}")
|
||
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,
|
||
)
|
||
original_branch = _load_semantic_checkpoint(checkpoint_store, fold, device)
|
||
original_content, _, _ = _generate_tsfa_outputs(
|
||
method="TSFA-main",
|
||
fold=fold,
|
||
sample_ids=[sample.sample_id for sample in all_samples],
|
||
samples_by_id=samples_by_id,
|
||
temporal_by_id=temporal_by_id,
|
||
branch=original_branch,
|
||
device=device,
|
||
delta=args.delta,
|
||
seed=args.seed,
|
||
batch_size=args.batch_size,
|
||
)
|
||
av_branch, history = _fit_av_branch(
|
||
fold=fold,
|
||
train_samples=train_samples,
|
||
temporal_by_id=temporal_by_id,
|
||
device=device,
|
||
args=args,
|
||
)
|
||
history_rows.extend(history)
|
||
av_checkpoints[f"fold_{fold:02d}/TSFA-AV"] = {
|
||
"seed": args.seed + fold * 101,
|
||
"train_sample_ids": [sample.sample_id for sample in train_samples],
|
||
"validation_sample_ids": [sample.sample_id for sample in heldout_samples],
|
||
"state_dict": {key: value.detach().cpu() for key, value in av_branch.state_dict().items()},
|
||
}
|
||
av_content, diagnostics = _generate_av_outputs(
|
||
branch=av_branch,
|
||
sample_ids=[sample.sample_id for sample in all_samples],
|
||
temporal_by_id=temporal_by_id,
|
||
samples_by_id=samples_by_id,
|
||
device=device,
|
||
args=args,
|
||
)
|
||
heldout_set = {sample.sample_id for sample in heldout_samples}
|
||
diagnostic_rows.extend({**row, "fold": fold} for row in diagnostics if row["sample_id"] in heldout_set)
|
||
private_by_id = _private_content(all_samples, feature_stats, temporal_by_id)
|
||
for (av_edge, use_private), method in VARIANTS.items():
|
||
shared = av_content if av_edge else original_content
|
||
prediction_rows.extend(_fit_fold_probe(
|
||
fold=fold,
|
||
train_samples=train_samples,
|
||
heldout_samples=heldout_samples,
|
||
shared_by_id=shared,
|
||
private_by_id=private_by_id,
|
||
method=method,
|
||
use_private=use_private,
|
||
seed=args.seed,
|
||
))
|
||
fold_manifest.append({
|
||
"fold": fold,
|
||
"train_count": len(train_samples),
|
||
"heldout_count": len(heldout_samples),
|
||
"train_video_id_count": len(train_groups),
|
||
"heldout_video_id_count": len(heldout_groups),
|
||
"video_id_overlap": [],
|
||
})
|
||
print(f"[AV/private fold {fold}] train={len(train_samples)} heldout={len(heldout_samples)}", flush=True)
|
||
del original_branch, av_branch, original_content, av_content, private_by_id, temporal_by_id
|
||
if device.type == "cuda":
|
||
torch.cuda.empty_cache()
|
||
|
||
_verify_original_control(prediction_rows, args.baseline_predictions)
|
||
metrics_rows, contrast_rows, confusion_rows = _summarize(
|
||
prediction_rows, bootstrap_repeats=args.bootstrap_repeats, seed=args.seed
|
||
)
|
||
args.output_dir.mkdir(parents=True, exist_ok=True)
|
||
_write_csv(args.output_dir / "predictions.csv", prediction_rows)
|
||
_write_csv(args.output_dir / "metrics.csv", metrics_rows)
|
||
_write_csv(args.output_dir / "paired_contrasts.csv", contrast_rows)
|
||
_write_csv(args.output_dir / "confusion_matrix.csv", confusion_rows)
|
||
_write_csv(args.output_dir / "av_attention_diagnostics.csv", diagnostic_rows)
|
||
_write_csv(args.output_dir / "av_training_history.csv", history_rows)
|
||
_plot_factorial(args.output_dir / "metrics.csv", args.output_dir / "factorial_effects.png")
|
||
torch.save(av_checkpoints, args.output_dir / "av_branch_checkpoints.pt")
|
||
manifest = {
|
||
"created_utc": datetime.now(timezone.utc).isoformat(),
|
||
"experiment": "2x2 TSFA ablation: local reciprocal Audio-Vision edge x private modality residual",
|
||
"sample_count": len(samples),
|
||
"video_id_count": len({sample.group_id for sample in samples}),
|
||
"fold_count": len(splits),
|
||
"folds": fold_manifest,
|
||
"original_tsfa_control_predictions_identical": True,
|
||
"av_edge": "M4 aligned Audio queries raw projected Vision, and M4 aligned Vision queries raw projected Audio; both directions use the existing M4 source-center +/-delta masks. Messages add to the original Text-to-Audio/Vision content.",
|
||
"private_residual": "For each modality, frozen M4 temporal attention pools its training-fold-standardized original BERT/eGeMAPS/DeiT source features onto 50 slots; the same five-segment probe concatenates all three private streams after shared streams.",
|
||
"semantic_training": "same 40-epoch Text-Audio/Text-Vision local contrastive loss, optimizer, temperature, batch size and seed as original TSFA; no emotion labels",
|
||
"probe": "StandardScaler + LogisticRegression(C=0.05) and StandardScaler + Ridge(alpha=25); trained per fold; regression predictions clipped to [-3,3]",
|
||
"bootstrap": "2,000 paired video_id-cluster samples; percentile 95% intervals; no multiple-comparison correction",
|
||
"parameters": {
|
||
"seed": args.seed,
|
||
"delta": args.delta,
|
||
"semantic_epochs": args.semantic_epochs,
|
||
"semantic_learning_rate": args.semantic_learning_rate,
|
||
"local_temperature": args.local_temperature,
|
||
"batch_size": args.batch_size,
|
||
"bootstrap_repeats": args.bootstrap_repeats,
|
||
},
|
||
"input_paths": {
|
||
"feature_dir": str(args.feature_dir.resolve()),
|
||
"feature_manifest": str(args.manifest.resolve()),
|
||
"splits": str(args.splits.resolve()),
|
||
"frozen_alignment_checkpoint_root": str(args.checkpoint_root.resolve()),
|
||
"original_tsfa_checkpoints": str(args.tsfa_checkpoints.resolve()),
|
||
"original_tsfa_predictions": str(args.baseline_predictions.resolve()),
|
||
},
|
||
"device": str(device),
|
||
"python": platform.python_version(),
|
||
"pytorch": torch.__version__,
|
||
"scikit_learn": sklearn.__version__,
|
||
"elapsed_seconds": time.time() - started,
|
||
"interpretation_limits": [
|
||
"An added edge also adds trainable parameters; an edge effect alone does not prove correct semantic alignment.",
|
||
"Private residuals increase probe input dimension; any gain cannot be assigned to a particular raw feature without further controlled study.",
|
||
"Only one semantic-branch seed was trained on 100 clips and 37 source videos.",
|
||
"Clip-level emotion labels cannot directly validate event-level correspondence.",
|
||
],
|
||
}
|
||
(args.output_dir / "run_manifest.json").write_text(
|
||
json.dumps(manifest, ensure_ascii=False, indent=2), encoding="utf-8"
|
||
)
|
||
_refresh_report_bundle(args.output_dir)
|
||
print(f"[AV/private complete] output={args.output_dir}", flush=True)
|
||
for row in metrics_rows:
|
||
print(
|
||
f" {row['method']}: Acc={row['accuracy']:.3f} F1={row['macro_f1']:.3f} "
|
||
f"MAE={row['mae']:.3f} Pearson={row['pearson']:.3f}", flush=True
|
||
)
|
||
|
||
|
||
def build_parser() -> argparse.ArgumentParser:
|
||
project = Path(__file__).resolve().parents[1]
|
||
parser = argparse.ArgumentParser(description=__doc__)
|
||
parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto")
|
||
parser.add_argument("--seed", type=int, default=42)
|
||
parser.add_argument("--delta", type=float, default=0.10)
|
||
parser.add_argument("--semantic-epochs", type=int, default=40)
|
||
parser.add_argument("--semantic-learning-rate", type=float, default=1e-3)
|
||
parser.add_argument("--local-temperature", type=float, default=0.1)
|
||
parser.add_argument("--batch-size", type=int, default=8)
|
||
parser.add_argument("--bootstrap-repeats", type=int, default=2000)
|
||
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("--tsfa-checkpoints", type=Path, default=project / "outputs/tsfa/probe_checkpoints.pt")
|
||
parser.add_argument("--baseline-predictions", type=Path, default=project / "outputs/tsfa_emotion_probe/emotion_probe_predictions.csv")
|
||
parser.add_argument("--output-dir", type=Path, default=project / "outputs/tsfa_av_private_ablation")
|
||
parser.add_argument("--plot-existing", action="store_true", help="refresh the figure from existing metrics.csv")
|
||
return parser
|
||
|
||
|
||
def main() -> None:
|
||
args = build_parser().parse_args()
|
||
if args.plot_existing:
|
||
_plot_factorial(args.output_dir / "metrics.csv", args.output_dir / "factorial_effects.png")
|
||
_refresh_report_bundle(args.output_dir)
|
||
else:
|
||
run(args)
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|