352 lines
16 KiB
Python
352 lines
16 KiB
Python
"""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()
|