"""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()