整理 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()
|
||||
Reference in New Issue
Block a user