leaderboard / scripts /sync_results.py
Paul Hager
Replace Gradio runtime with static leaderboard
3fc79a4
Raw
History Blame Contribute Delete
4.42 kB
#!/usr/bin/env python3
"""Build the static leaderboard data bundle from MIMIC-CDM result repositories."""
from __future__ import annotations
import argparse
import json
import re
import subprocess
from collections import defaultdict
from datetime import datetime, timezone
from decimal import Decimal, ROUND_HALF_EVEN
from pathlib import Path
from typing import Any
TASKS = ("Appendicitis", "Cholecystitis", "Diverticulitis", "Pancreatitis")
VARIANT_SUFFIX = re.compile(r"\s+(?:4|8|16)bit$", re.IGNORECASE)
def read_git_metadata(path: Path) -> dict[str, str]:
try:
output = subprocess.run(
["git", "-C", str(path), "log", "-1", "--format=%H%n%cI"],
check=True,
capture_output=True,
text=True,
).stdout.splitlines()
return {"revision": output[0], "updated_at": output[1]}
except (OSError, subprocess.CalledProcessError, IndexError):
return {
"revision": "local",
"updated_at": datetime.now(timezone.utc).isoformat(),
}
def model_hub_id(model_name: str) -> str:
return VARIANT_SUFFIX.sub("", model_name).strip()
def load_result_files(results_path: Path) -> list[dict[str, Any]]:
merged: dict[tuple[str, int], dict[str, Any]] = defaultdict(
lambda: {"scores": {}}
)
for result_file in sorted(results_path.rglob("*.json")):
try:
payload = json.loads(result_file.read_text())
config = payload["config"]
model_name = str(config["model_name"]).strip()
bits = int(config.get("model_quantization_bits") or 0)
key = (model_name, bits)
run = merged[key]
run.update(
{
"model": model_name,
"hub_model": model_hub_id(model_name),
"params_b": config.get("params"),
"context_length": config.get("max_sequence_length"),
"quantization_bits": bits or None,
}
)
for task in TASKS:
score = payload.get("results", {}).get(task, {}).get("acc")
if score is not None:
run["scores"][task] = float(
(Decimal(str(score)) * 100).quantize(
Decimal("0.01"), rounding=ROUND_HALF_EVEN
)
)
except (KeyError, TypeError, ValueError, json.JSONDecodeError) as error:
raise ValueError(f"Invalid result file {result_file}: {error}") from error
rows = []
for run in merged.values():
if not all(task in run["scores"] for task in TASKS):
continue
run["average"] = float(
(
sum(Decimal(str(run["scores"][task])) for task in TASKS)
/ len(TASKS)
).quantize(Decimal("0.01"), rounding=ROUND_HALF_EVEN)
)
rows.append(run)
rows.sort(key=lambda row: (-row["average"], row["model"].lower()))
for rank, row in enumerate(rows, start=1):
row["rank"] = rank
return rows
def build_bundle(cdm_path: Path, cdm_fi_path: Path) -> dict[str, Any]:
sources = {
"cdm": {
"repo": "MIMIC-CDM/results-CDM",
**read_git_metadata(cdm_path),
},
"cdm_fi": {
"repo": "MIMIC-CDM/results-CDM-FI",
**read_git_metadata(cdm_fi_path),
},
}
return {
"schema_version": 1,
"tasks": list(TASKS),
"sources": sources,
"leaderboards": {
"cdm": load_result_files(cdm_path),
"cdm_fi": load_result_files(cdm_fi_path),
},
}
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--cdm-dir", required=True, type=Path)
parser.add_argument("--cdm-fi-dir", required=True, type=Path)
parser.add_argument(
"--output",
type=Path,
default=Path("data/leaderboards.json"),
)
args = parser.parse_args()
bundle = build_bundle(args.cdm_dir, args.cdm_fi_dir)
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(json.dumps(bundle, indent=2) + "\n")
print(
f"Wrote {len(bundle['leaderboards']['cdm'])} CDM and "
f"{len(bundle['leaderboards']['cdm_fi'])} CDM-FI entries to {args.output}"
)
if __name__ == "__main__":
main()