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)