lnp_ml/tests/paired_test.py

65 lines
2.0 KiB
Python

import json
import glob
import os
import numpy as np
from scipy.stats import wilcoxon, ttest_rel
BASE = "models/abl"
# (task, metric, higher_better)
ROWS = [
("delivery", "r2", True),
("delivery", "rmse", False),
("biodist", "kl_divergence", False),
("size", "r2", True),
("size", "rmse", False),
("pdi", "f1", True),
("ee", "f1", True),
("toxic", "f1", True),
]
def latest(v: str) -> str:
dirs = sorted(glob.glob(f"{BASE}/{v}/*/"), key=os.path.getmtime)
if not dirs:
raise FileNotFoundError(f"variant '{v}'{BASE} 下没有 run 目录")
return dirs[-1]
def collect(run: str, task: str, metric: str) -> dict:
"""按相对路径(fold/repeat)收集某任务某指标,便于跨变体配对。"""
out = {}
for f in glob.glob(os.path.join(run, "**", "test_metrics.json"), recursive=True):
with open(f) as fh:
d = json.load(fh)
if task in d and metric in d[task]:
out[os.path.relpath(f, run)] = d[task][metric]
return out
def paired(vA: str, vB: str, task: str, metric: str, higher_better: bool = True) -> None:
a, b = collect(latest(vA), task, metric), collect(latest(vB), task, metric)
keys = sorted(set(a) & set(b))
if not keys:
print(f"{task}.{metric}: 无可配对数据")
return
da = np.array([a[k] for k in keys])
db = np.array([b[k] for k in keys])
diff = (db - da) if higher_better else (da - db) # 正 = B 更好
tp = ttest_rel(db, da).pvalue
try:
wp = wilcoxon(db, da).pvalue
except ValueError:
wp = float("nan")
print(
f"{task + '.' + metric:18s} {vA}={da.mean():.4f} {vB}={db.mean():.4f} "
f"Δ(B优)={diff.mean():+.4f} n={len(keys)} Wilcoxon p={wp:.3f} t p={tp:.3f}"
)
if __name__ == "__main__":
for vB in ["moe", "llm", "both"]:
print(f"\n===== baseline vs {vB} =====")
for task, metric, hb in ROWS:
paired("baseline", vB, task, metric, hb)