整理 Q1-Q3 实验代码与结果
This commit is contained in:
@@ -0,0 +1,719 @@
|
||||
"""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()
|
||||
Reference in New Issue
Block a user