memisislabs / scripts /tabdiff_evaluate.py
nnagesh101's picture
Add one-shot server setup script; lazy DataDesigner import; --label for external baselines
c95367f
Raw
History Blame Contribute Delete
3.93 kB
"""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()