"""Evaluate TabDiff's synthetic output with the full MemisisLabs metric suite. TabDiff reports its own metrics (Density, MLE, C2ST, alpha-precision/beta-recall, DCR). This runs OUR arena evaluators on the same synthetic table so TabDiff gets directly-comparable numbers — fidelity (SDMetrics), privacy (NewRowSynthesis), ML utility (TSTR), fairness (DPR/EOR) — and lands on the shared leaderboard alongside every other generator. Usage: python scripts/tabdiff_evaluate.py --dataset openml_45040 \ --synthetic /path/to/TabDiff/eval/report_runs/learnable_schedule/schizophrenia/sample_0.csv \ [--record] """ from __future__ import annotations import argparse import sys from pathlib import Path import pandas as pd sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) from pipeline import datasets, leaderboard # noqa: E402 from pipeline.evaluate import evaluate_all # noqa: E402 from pipeline.metadata import build_metadata # noqa: E402 def main() -> None: ap = argparse.ArgumentParser(description=__doc__) ap.add_argument("--dataset", default="openml_45040") ap.add_argument("--synthetic", required=True, help="sampled CSV (or a dir of samples)") ap.add_argument("--label", default="tabdiff · diffusion", help="leaderboard model name (e.g. 'tabsyn · latent-diffusion')") ap.add_argument("--record", action="store_true", help="append to the shared leaderboard") args = ap.parse_args() real, target, protected, task = datasets.load(args.dataset) metadata = build_metadata(real) # Accept a single CSV or a directory of TabDiff samples (evaluate each, average). path = Path(args.synthetic) files = sorted(path.glob("*.csv")) if path.is_dir() else [path] if not files: sys.exit(f"no CSVs found at {path}") results = [] for f in files: syn = pd.read_csv(f) syn = syn[[c for c in real.columns if c in syn.columns]] # align columns r = evaluate_all(real, syn, metadata, target=target, protected=protected, task=task) r["synthetic"] = syn results.append(r) fair = r.get("fairness") or {} print(f"{f.name}: fidelity={r['overall_score']:.3f} privacy={r['new_row_synthesis']} " f"utility={r['ml_efficacy']} fairness={fair.get('fairness_score')}") # Average across samples for the leaderboard row. def avg(key): vals = [x.get(key) for x in results if isinstance(x.get(key), (int, float))] return sum(vals) / len(vals) if vals else None fair_scores = [(x.get("fairness") or {}).get("fairness_score") for x in results] fair_scores = [v for v in fair_scores if isinstance(v, (int, float))] mean_metrics = { "overall_score": avg("overall_score"), "new_row_synthesis": avg("new_row_synthesis"), "ml_efficacy": avg("ml_efficacy"), "fairness": {"fairness_score": (sum(fair_scores) / len(fair_scores)) if fair_scores else None}, } print("\n=== TabDiff · diffusion — averaged over " f"{len(files)} sample(s) on {datasets.display_name(args.dataset)} ===") print(f" fidelity : {mean_metrics['overall_score']}") print(f" privacy : {mean_metrics['new_row_synthesis']}") print(f" utility : {mean_metrics['ml_efficacy']}") print(f" fairness : {mean_metrics['fairness']['fairness_score']}") if args.record: backend = args.label.split(" · ")[0].strip() or "external" entry = {"backend": backend, "method": args.label.split(" · ")[-1].strip(), "label": args.label, "metrics": mean_metrics} leaderboard.record([entry], dataset=args.dataset, run_id=args.label.replace(" ", "_").replace("·", ""), num_records=len(results[0]["synthetic"])) print(f"\nrecorded -> leaderboard ({args.label})") if __name__ == "__main__": main()