Files

233 lines
9.9 KiB
Python

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