整理 Q1-Q3 实验代码与结果
This commit is contained in:
@@ -0,0 +1,232 @@
|
||||
"""Paired, video-clustered comparison of TSFA and the math B0-B4 OOF probes."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import csv
|
||||
import json
|
||||
import shutil
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
from sklearn.metrics import accuracy_score, f1_score, mean_absolute_error
|
||||
|
||||
|
||||
METHODS = ("B0", "B1", "B2", "B3", "B4")
|
||||
METRICS = ("accuracy", "macro_f1", "mae", "pearson")
|
||||
|
||||
|
||||
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 _pearson(actual: np.ndarray, predicted: np.ndarray) -> float:
|
||||
if np.std(actual) <= 1e-12 or np.std(predicted) <= 1e-12:
|
||||
return float("nan")
|
||||
return float(np.corrcoef(actual, predicted)[0, 1])
|
||||
|
||||
|
||||
def _scores(
|
||||
actual_class: np.ndarray,
|
||||
predicted_class: np.ndarray,
|
||||
actual_value: np.ndarray,
|
||||
predicted_value: np.ndarray,
|
||||
) -> dict[str, float]:
|
||||
return {
|
||||
"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,
|
||||
)),
|
||||
"mae": float(mean_absolute_error(actual_value, predicted_value)),
|
||||
"pearson": _pearson(actual_value, predicted_value),
|
||||
}
|
||||
|
||||
|
||||
def run(args: argparse.Namespace) -> None:
|
||||
math_rows = _read_csv(args.math_predictions)
|
||||
tsfa_rows = [
|
||||
row for row in _read_csv(args.tsfa_predictions)
|
||||
if row["method"] == args.candidate_method
|
||||
and row.get("view", args.candidate_view) == args.candidate_view
|
||||
]
|
||||
split_rows = _read_csv(args.math_splits)
|
||||
|
||||
math_by_id = {row["sample_id"]: row for row in math_rows}
|
||||
tsfa_by_id = {row["sample_id"]: row for row in tsfa_rows}
|
||||
folds_by_id = {row["sample_id"]: int(row["fold"]) for row in split_rows}
|
||||
if len(math_by_id) != len(math_rows) or len(tsfa_by_id) != len(tsfa_rows):
|
||||
raise ValueError("expected unique sample IDs in both OOF prediction files")
|
||||
sample_ids = sorted(math_by_id)
|
||||
if set(sample_ids) != set(tsfa_by_id) or set(sample_ids) != set(folds_by_id):
|
||||
raise ValueError("math, TSFA, and math split files do not cover identical samples")
|
||||
|
||||
for sample_id in sample_ids:
|
||||
math_row = math_by_id[sample_id]
|
||||
tsfa_row = tsfa_by_id[sample_id]
|
||||
if math_row["video_id"] != tsfa_row["video_id"]:
|
||||
raise ValueError(f"video_id mismatch for {sample_id}")
|
||||
if int(tsfa_row["fold"]) != folds_by_id[sample_id]:
|
||||
raise ValueError(f"fold mismatch for {sample_id}")
|
||||
if int(math_row["true_polarity"]) != int(tsfa_row["true_class_id"]):
|
||||
raise ValueError(f"class label mismatch for {sample_id}")
|
||||
if not np.isclose(float(math_row["true_sentiment"]), float(tsfa_row["true_label"])):
|
||||
raise ValueError(f"continuous label mismatch for {sample_id}")
|
||||
|
||||
video_ids = np.asarray([math_by_id[sample_id]["video_id"] for sample_id in sample_ids])
|
||||
actual_class = np.asarray([int(math_by_id[sample_id]["true_polarity"]) for sample_id in sample_ids])
|
||||
actual_value = np.asarray([float(math_by_id[sample_id]["true_sentiment"]) for sample_id in sample_ids])
|
||||
tsfa_class = np.asarray([int(tsfa_by_id[sample_id]["predicted_class_id"]) for sample_id in sample_ids])
|
||||
tsfa_value = np.clip(
|
||||
np.asarray([float(tsfa_by_id[sample_id]["predicted_label"]) for sample_id in sample_ids]),
|
||||
-3.0,
|
||||
3.0,
|
||||
)
|
||||
tsfa_scores = _scores(actual_class, tsfa_class, actual_value, tsfa_value)
|
||||
|
||||
group_indices = {
|
||||
video_id: np.flatnonzero(video_ids == video_id)
|
||||
for video_id in sorted(set(video_ids.tolist()))
|
||||
}
|
||||
group_names = np.asarray(list(group_indices))
|
||||
rng = np.random.default_rng(args.seed)
|
||||
boot_scores: dict[str, dict[str, list[float]]] = {
|
||||
method: {metric: [] for metric in METRICS} for method in METHODS
|
||||
}
|
||||
point_scores: dict[str, dict[str, float]] = {}
|
||||
for method in METHODS:
|
||||
predicted_class = np.asarray([
|
||||
int(math_by_id[sample_id][f"{method}_predicted_polarity"])
|
||||
for sample_id in sample_ids
|
||||
])
|
||||
predicted_value = np.clip(np.asarray([
|
||||
float(math_by_id[sample_id][f"{method}_predicted_sentiment"])
|
||||
for sample_id in sample_ids
|
||||
]), -3.0, 3.0)
|
||||
point_scores[method] = _scores(actual_class, predicted_class, actual_value, predicted_value)
|
||||
|
||||
for _ in range(args.bootstrap_repeats):
|
||||
sampled_groups = rng.choice(group_names, size=len(group_names), replace=True)
|
||||
indices = np.concatenate([group_indices[group] for group in sampled_groups])
|
||||
resampled_tsfa = _scores(
|
||||
actual_class[indices], tsfa_class[indices], actual_value[indices], tsfa_value[indices]
|
||||
)
|
||||
for method in METHODS:
|
||||
predicted_class = np.asarray([
|
||||
int(math_by_id[sample_ids[index]][f"{method}_predicted_polarity"])
|
||||
for index in indices
|
||||
])
|
||||
predicted_value = np.clip(np.asarray([
|
||||
float(math_by_id[sample_ids[index]][f"{method}_predicted_sentiment"])
|
||||
for index in indices
|
||||
]), -3.0, 3.0)
|
||||
resampled_baseline = _scores(
|
||||
actual_class[indices], predicted_class, actual_value[indices], predicted_value
|
||||
)
|
||||
for metric in METRICS:
|
||||
boot_scores[method][metric].append(resampled_tsfa[metric] - resampled_baseline[metric])
|
||||
|
||||
output_rows: list[dict[str, Any]] = []
|
||||
for method in METHODS:
|
||||
for metric in METRICS:
|
||||
deltas = np.asarray(boot_scores[method][metric], dtype=np.float64)
|
||||
finite = deltas[np.isfinite(deltas)]
|
||||
output_rows.append({
|
||||
"comparison": f"{args.candidate_method}-{method}",
|
||||
"metric": metric,
|
||||
"tsfa_oof": tsfa_scores[metric],
|
||||
"baseline_oof": point_scores[method][metric],
|
||||
"delta_tsfa_minus_baseline": tsfa_scores[metric] - point_scores[method][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": args.bootstrap_repeats,
|
||||
"video_group_count": len(group_indices),
|
||||
})
|
||||
|
||||
args.output_dir.mkdir(parents=True, exist_ok=True)
|
||||
output_path = args.output_dir / f"{args.output_stem}.csv"
|
||||
with output_path.open("w", encoding="utf-8-sig", newline="") as stream:
|
||||
writer = csv.DictWriter(stream, fieldnames=list(output_rows[0]))
|
||||
writer.writeheader()
|
||||
writer.writerows(output_rows)
|
||||
manifest = {
|
||||
"experiment": f"Paired OOF comparison of {args.candidate_method} against math B0-B4",
|
||||
"candidate_method": args.candidate_method,
|
||||
"candidate_view": args.candidate_view,
|
||||
"sample_count": len(sample_ids),
|
||||
"video_group_count": len(group_indices),
|
||||
"identical_sample_ids": True,
|
||||
"identical_video_ids": True,
|
||||
"identical_fold_assignments": True,
|
||||
"fold_assignment_validation": "TSFA prediction fold matched math split_assignments.csv for every sample_id.",
|
||||
"metrics": list(METRICS),
|
||||
"macro_f1_labels": [0, 1, 2],
|
||||
"regression_prediction_clipping": [-3.0, 3.0],
|
||||
"bootstrap": {
|
||||
"unit": "video_id cluster",
|
||||
"repeats": args.bootstrap_repeats,
|
||||
"seed": args.seed,
|
||||
"interval": "percentile 95% confidence interval for paired TSFA-minus-baseline metric differences",
|
||||
},
|
||||
"math_inputs_read_only": [str(args.math_predictions.resolve()), str(args.math_splits.resolve())],
|
||||
"tsfa_input": str(args.tsfa_predictions.resolve()),
|
||||
"output": str(output_path.resolve()),
|
||||
}
|
||||
manifest_path = args.output_dir / f"{args.output_stem}_manifest.json"
|
||||
manifest_path.write_text(
|
||||
json.dumps(manifest, ensure_ascii=False, indent=2), encoding="utf-8"
|
||||
)
|
||||
report_bundle = args.output_dir.parent / "tsfa" / "report_bundle"
|
||||
if report_bundle.is_dir() and args.output_stem == "paired_math_comparison":
|
||||
shutil.copy2(output_path, report_bundle / output_path.name)
|
||||
shutil.copy2(manifest_path, report_bundle / manifest_path.name)
|
||||
print(f"[paired comparison complete] samples={len(sample_ids)} groups={len(group_indices)} output={output_path}")
|
||||
for row in output_rows:
|
||||
if row["metric"] == "macro_f1":
|
||||
print(
|
||||
f" {row['comparison']}: delta={row['delta_tsfa_minus_baseline']:.3f} "
|
||||
f"CI=[{row['video_cluster_bootstrap_ci95_low']:.3f}, "
|
||||
f"{row['video_cluster_bootstrap_ci95_high']:.3f}]"
|
||||
)
|
||||
|
||||
|
||||
def build_parser() -> argparse.ArgumentParser:
|
||||
project = Path(__file__).resolve().parents[1]
|
||||
repository = project.parents[1]
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
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(
|
||||
"--tsfa-predictions",
|
||||
type=Path,
|
||||
default=project / "outputs/tsfa_emotion_probe/emotion_probe_predictions.csv",
|
||||
)
|
||||
parser.add_argument("--output-dir", type=Path, default=project / "outputs/tsfa_emotion_probe")
|
||||
parser.add_argument("--candidate-method", default="TSFA-main")
|
||||
parser.add_argument("--candidate-view", default="all_modalities")
|
||||
parser.add_argument("--output-stem", default="paired_math_comparison")
|
||||
parser.add_argument("--bootstrap-repeats", type=int, default=2000)
|
||||
parser.add_argument("--seed", type=int, default=42)
|
||||
return parser
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = build_parser().parse_args()
|
||||
run(args)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -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()
|
||||
@@ -0,0 +1,351 @@
|
||||
"""Grouped emotion probes for the frozen TSFA representations."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import platform
|
||||
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
|
||||
from sklearn.linear_model import LogisticRegression, Ridge
|
||||
from sklearn.metrics import accuracy_score, confusion_matrix, f1_score
|
||||
from sklearn.pipeline import make_pipeline
|
||||
from sklearn.preprocessing import StandardScaler
|
||||
|
||||
from .correspondence_eval import _write_csv
|
||||
from .experiment_data import FeatureSample, fit_feature_stats, load_feature_samples
|
||||
from .tsfa_experiment import (
|
||||
BASELINE_VARIANTS,
|
||||
GRID_SIZE,
|
||||
_collect_fold_features,
|
||||
_generate_tsfa_outputs,
|
||||
_load_semantic_checkpoint,
|
||||
)
|
||||
|
||||
|
||||
CLASS_NAMES = ("Negative", "Neutral", "Positive")
|
||||
METHODS = (*BASELINE_VARIANTS, "TSFA-main")
|
||||
VIEWS = {
|
||||
"text": ("text",),
|
||||
"audio": ("audio",),
|
||||
"vision": ("vision",),
|
||||
"all_modalities": ("text", "audio", "vision"),
|
||||
}
|
||||
|
||||
|
||||
def _class_from_sentiment(value: float) -> int:
|
||||
if value < 0:
|
||||
return 0
|
||||
if value == 0:
|
||||
return 1
|
||||
return 2
|
||||
|
||||
|
||||
def _pool_five_segments(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 an aligned [{GRID_SIZE}, D] representation, got {values.shape}")
|
||||
# Five contiguous equal-width bins over the shared 50-slot timeline.
|
||||
return values.reshape(5, GRID_SIZE // 5, values.shape[1]).mean(axis=1).reshape(-1)
|
||||
|
||||
|
||||
def _vector(content: Mapping[str, np.ndarray], view: str) -> np.ndarray:
|
||||
return np.concatenate([_pool_five_segments(content[name]) for name in VIEWS[view]])
|
||||
|
||||
|
||||
def _pearson(actual: np.ndarray, predicted: np.ndarray) -> float:
|
||||
if np.std(actual) <= 1e-12 or np.std(predicted) <= 1e-12:
|
||||
return float("nan")
|
||||
return float(np.corrcoef(actual, predicted)[0, 1])
|
||||
|
||||
|
||||
def _sample_vectors(
|
||||
samples: Sequence[FeatureSample],
|
||||
content_by_id: Mapping[str, Mapping[str, np.ndarray]],
|
||||
view: str,
|
||||
) -> np.ndarray:
|
||||
return np.stack([_vector(content_by_id[sample.sample_id], view) for sample in samples])
|
||||
|
||||
|
||||
def run(args: argparse.Namespace) -> None:
|
||||
started = time.time()
|
||||
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 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}
|
||||
labels_by_id = {sample.sample_id: _class_from_sentiment(sample.sentiment) for sample in samples}
|
||||
for sample in samples:
|
||||
if labels_by_id[sample.sample_id] != sample.polarity:
|
||||
raise ValueError(f"sign-derived class disagrees with annotation for {sample.sample_id}")
|
||||
class_counts = {
|
||||
name: sum(label == index for label in labels_by_id.values())
|
||||
for index, name in enumerate(CLASS_NAMES)
|
||||
}
|
||||
|
||||
splits = json.loads(args.splits.read_text(encoding="utf-8"))
|
||||
if len(splits) != 5:
|
||||
raise ValueError(f"expected five grouped folds, found {len(splits)}")
|
||||
validation_ids = [sample_id for split in splits for sample_id in split["validation_sample_ids"]]
|
||||
if len(validation_ids) != len(set(validation_ids)) or set(validation_ids) != set(samples_by_id):
|
||||
raise ValueError("five-fold validation partitions must cover all samples exactly once")
|
||||
|
||||
checkpoint_file = args.tsfa_output_dir / "probe_checkpoints.pt"
|
||||
checkpoint_store = torch.load(checkpoint_file, map_location="cpu", weights_only=False)
|
||||
prediction_rows: list[dict[str, Any]] = []
|
||||
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"]]
|
||||
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}: {sorted(train_groups & heldout_groups)}")
|
||||
feature_stats = fit_feature_stats(train_samples)
|
||||
baseline_content, _, 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,
|
||||
)
|
||||
all_samples = [*train_samples, *heldout_samples]
|
||||
all_ids = [sample.sample_id for sample in all_samples]
|
||||
main_branch = _load_semantic_checkpoint(checkpoint_store, fold, device)
|
||||
tsfa_content, _, _ = _generate_tsfa_outputs(
|
||||
method="TSFA-main",
|
||||
fold=fold,
|
||||
sample_ids=all_ids,
|
||||
samples_by_id=samples_by_id,
|
||||
temporal_by_id=temporal_by_id,
|
||||
branch=main_branch,
|
||||
device=device,
|
||||
delta=args.delta,
|
||||
seed=args.seed,
|
||||
draw=0,
|
||||
batch_size=args.batch_size,
|
||||
)
|
||||
content_by_method = dict(baseline_content)
|
||||
content_by_method["TSFA-main"] = tsfa_content
|
||||
|
||||
train_y_class = np.asarray([labels_by_id[s.sample_id] for s in train_samples], dtype=np.int64)
|
||||
heldout_y_class = np.asarray([labels_by_id[s.sample_id] for s in heldout_samples], dtype=np.int64)
|
||||
train_y_value = np.asarray([s.sentiment for s in train_samples], dtype=np.float64)
|
||||
heldout_y_value = np.asarray([s.sentiment for s in heldout_samples], dtype=np.float64)
|
||||
|
||||
for method in METHODS:
|
||||
for view in VIEWS:
|
||||
train_x = _sample_vectors(train_samples, content_by_method[method], view)
|
||||
heldout_x = _sample_vectors(heldout_samples, content_by_method[method], view)
|
||||
classifier = make_pipeline(
|
||||
StandardScaler(),
|
||||
LogisticRegression(C=0.05, max_iter=5000, solver="lbfgs", random_state=args.seed),
|
||||
)
|
||||
classifier.fit(train_x, train_y_class)
|
||||
predicted_class = classifier.predict(heldout_x)
|
||||
|
||||
regressor = make_pipeline(StandardScaler(), Ridge(alpha=25.0))
|
||||
regressor.fit(train_x, train_y_value)
|
||||
predicted_value_unclipped = regressor.predict(heldout_x)
|
||||
predicted_value = np.clip(predicted_value_unclipped, -3.0, 3.0)
|
||||
|
||||
for index, sample in enumerate(heldout_samples):
|
||||
prediction_rows.append({
|
||||
"method": method,
|
||||
"view": view,
|
||||
"fold": fold,
|
||||
"sample_id": sample.sample_id,
|
||||
"video_id": sample.group_id,
|
||||
"true_class_id": int(heldout_y_class[index]),
|
||||
"true_class": CLASS_NAMES[int(heldout_y_class[index])],
|
||||
"predicted_class_id": int(predicted_class[index]),
|
||||
"predicted_class": CLASS_NAMES[int(predicted_class[index])],
|
||||
"true_label": float(heldout_y_value[index]),
|
||||
"predicted_label": float(predicted_value[index]),
|
||||
"predicted_label_unclipped": float(predicted_value_unclipped[index]),
|
||||
})
|
||||
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": sorted(train_groups & heldout_groups),
|
||||
})
|
||||
print(f"[emotion probe fold {fold}] train={len(train_samples)} heldout={len(heldout_samples)}", flush=True)
|
||||
del main_branch, baseline_content, temporal_by_id, tsfa_content, content_by_method
|
||||
if device.type == "cuda":
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
summary_rows = []
|
||||
confusion_rows = []
|
||||
majority_id = max(range(3), key=lambda index: class_counts[CLASS_NAMES[index]])
|
||||
for method in METHODS:
|
||||
for view in VIEWS:
|
||||
rows = [row for row in prediction_rows if row["method"] == method and row["view"] == view]
|
||||
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_value = np.asarray([row["true_label"] for row in rows], dtype=np.float64)
|
||||
predicted_value = np.asarray([row["predicted_label"] for row in rows], dtype=np.float64)
|
||||
predicted_value_unclipped = np.asarray(
|
||||
[row["predicted_label_unclipped"] for row in rows], dtype=np.float64
|
||||
)
|
||||
majority_prediction = np.full_like(actual_class, majority_id)
|
||||
matrix = confusion_matrix(actual_class, predicted_class, labels=[0, 1, 2])
|
||||
fold_macro_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_macro_f1.append(float(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,
|
||||
)))
|
||||
summary_rows.append({
|
||||
"method": method,
|
||||
"view": view,
|
||||
"sample_count": len(rows),
|
||||
"negative_support": class_counts["Negative"],
|
||||
"neutral_support": class_counts["Neutral"],
|
||||
"positive_support": class_counts["Positive"],
|
||||
"majority_class_accuracy_baseline": max(class_counts.values()) / len(samples),
|
||||
"majority_class_macro_f1_baseline": float(f1_score(
|
||||
actual_class, majority_prediction, labels=[0, 1, 2], average="macro", zero_division=0
|
||||
)),
|
||||
"accuracy": float(accuracy_score(actual_class, predicted_class)),
|
||||
"macro_f1_fixed_three_classes": float(f1_score(
|
||||
actual_class, predicted_class, labels=[0, 1, 2], average="macro", zero_division=0
|
||||
)),
|
||||
"macro_f1_fold_mean": float(np.mean(fold_macro_f1)),
|
||||
"macro_f1_fold_sd": float(np.std(fold_macro_f1, ddof=1)),
|
||||
"mae": float(np.mean(np.abs(actual_value - predicted_value))),
|
||||
"pearson": _pearson(actual_value, predicted_value),
|
||||
"mae_unclipped": float(np.mean(np.abs(actual_value - predicted_value_unclipped))),
|
||||
"pearson_unclipped": _pearson(actual_value, predicted_value_unclipped),
|
||||
})
|
||||
for true_id, true_name in enumerate(CLASS_NAMES):
|
||||
for predicted_id, predicted_name in enumerate(CLASS_NAMES):
|
||||
confusion_rows.append({
|
||||
"method": method,
|
||||
"view": view,
|
||||
"true_class": true_name,
|
||||
"predicted_class": predicted_name,
|
||||
"count": int(matrix[true_id, predicted_id]),
|
||||
})
|
||||
|
||||
args.output_dir.mkdir(parents=True, exist_ok=True)
|
||||
_write_csv(args.output_dir / "emotion_probe_predictions.csv", prediction_rows)
|
||||
_write_csv(args.output_dir / "emotion_probe_metrics.csv", summary_rows)
|
||||
_write_csv(args.output_dir / "emotion_probe_confusion_matrix.csv", confusion_rows)
|
||||
run_manifest = {
|
||||
"created_utc": datetime.now(timezone.utc).isoformat(),
|
||||
"experiment": "Grouped five-fold emotion probes on frozen TSFA and M3/M4 representations",
|
||||
"sample_count": len(samples),
|
||||
"video_id_count": len({sample.group_id for sample in samples}),
|
||||
"fold_count": len(splits),
|
||||
"heldout_prediction_count_per_method_view": len(samples),
|
||||
"split_rule": "Fixed five-fold GroupKFold by group_id/video_id, loaded from the Q1 method-comparison split file.",
|
||||
"seed": args.seed,
|
||||
"batch_size_for_feature_inference": args.batch_size,
|
||||
"input_paths": {
|
||||
"feature_dir": str(args.feature_dir.resolve()),
|
||||
"feature_manifest": str(args.manifest.resolve()),
|
||||
"grouped_splits": str(args.splits.resolve()),
|
||||
"alignment_checkpoint_root": str(args.checkpoint_root.resolve()),
|
||||
"tsfa_output_dir": str(args.tsfa_output_dir.resolve()),
|
||||
"tsfa_probe_checkpoint": str(checkpoint_file.resolve()),
|
||||
},
|
||||
"class_mapping": {"label_lt_0": "Negative", "label_eq_0": "Neutral", "label_gt_0": "Positive"},
|
||||
"class_counts": class_counts,
|
||||
"classification": {
|
||||
"estimator": "LogisticRegression",
|
||||
"C": 0.05,
|
||||
"max_iter": 5000,
|
||||
"features": "StandardScaler fitted on the training fold, then 5-segment pooled aligned representations",
|
||||
"macro_f1": "fixed labels [Negative, Neutral, Positive]; zero_division=0",
|
||||
},
|
||||
"temporal_pooling": "Five contiguous equal-width bins over the 50 shared slots; mean each bin and concatenate.",
|
||||
"regression": {
|
||||
"estimator": "Ridge",
|
||||
"alpha": 25.0,
|
||||
"features": "StandardScaler fitted on the training fold, then the same 5-segment pooled representations",
|
||||
"prediction_clipping": [-3.0, 3.0],
|
||||
"reported_mae_pearson": "computed on clipped predictions; unclipped values are also retained for diagnosis",
|
||||
},
|
||||
"methods": list(METHODS),
|
||||
"views": {key: list(value) for key, value in VIEWS.items()},
|
||||
"folds": fold_manifest,
|
||||
"emotion_labels_used_to_train_alignment": False,
|
||||
"alignment_models_retrained": False,
|
||||
"feature_extractors_changed": False,
|
||||
"label_agreement": "Sign-derived classes were checked against the annotation class for all samples.",
|
||||
"device_for_feature_inference": str(device),
|
||||
"python": platform.python_version(),
|
||||
"scikit_learn": sklearn.__version__,
|
||||
"elapsed_seconds": time.time() - started,
|
||||
}
|
||||
(args.output_dir / "emotion_probe_manifest.json").write_text(
|
||||
json.dumps(run_manifest, ensure_ascii=False, indent=2), encoding="utf-8"
|
||||
)
|
||||
bundle = args.tsfa_output_dir / "report_bundle"
|
||||
bundle.mkdir(parents=True, exist_ok=True)
|
||||
for name in (
|
||||
"emotion_probe_metrics.csv",
|
||||
"emotion_probe_confusion_matrix.csv",
|
||||
"emotion_probe_manifest.json",
|
||||
):
|
||||
shutil.copy2(args.output_dir / name, bundle / name)
|
||||
bundle_readme = bundle / "README.md"
|
||||
note = (
|
||||
"\n`emotion_probe_metrics.csv` adds five-fold, video-group-held-out "
|
||||
"LogisticRegression/Ridge probes on five-segment pooled representations. "
|
||||
"These are small-sample downstream probes, not end-to-end emotion model scores.\n"
|
||||
)
|
||||
current = bundle_readme.read_text(encoding="utf-8")
|
||||
if "`emotion_probe_metrics.csv` adds" not in current:
|
||||
bundle_readme.write_text(current + note, encoding="utf-8")
|
||||
print(
|
||||
f"[emotion probe complete] samples={len(samples)} class_counts={class_counts} "
|
||||
f"output={args.output_dir}", 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("--batch-size", type=int, default=8)
|
||||
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-output-dir", type=Path, default=project / "outputs/tsfa")
|
||||
parser.add_argument("--output-dir", type=Path, default=project / "outputs/tsfa_emotion_probe")
|
||||
return parser
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = build_parser().parse_args()
|
||||
run(args)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user