from __future__ import annotations import argparse import json import sys import time from collections.abc import Sequence from pathlib import Path from typing import TextIO import torch from .benchmark_runner import BenchmarkConfig, run_benchmark, write_benchmark_outputs from .deberta import DebertaCandidateScorer from .domain import CandidateScorer def _parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser(prog="deberta-ime-benchmark") parser.add_argument("--data-dir", type=Path, default=Path("work/data/ud-japanese-gsd")) parser.add_argument("--output-dir", type=Path, default=Path("outputs")) parser.add_argument("--stem", default="benchmark_full") parser.add_argument("--limit", type=int, default=0, help="0 evaluates all rows") parser.add_argument("--pool-size", type=int, default=8) parser.add_argument("--bootstrap-samples", type=int, default=5000) parser.add_argument("--example-limit", type=int, default=20) parser.add_argument("--seed", type=int, default=20260810) parser.add_argument("--device", default="cpu") parser.add_argument("--model-cache-dir", type=Path) parser.add_argument("--offline", action="store_true") parser.add_argument("--threads", type=int, default=min(8, torch.get_num_threads())) parser.add_argument("--include-prototype-test-rows", action="store_true") return parser def run( argv: Sequence[str] | None = None, *, stdout: TextIO | None = None, stderr: TextIO | None = None, scorer: CandidateScorer | None = None, ) -> int: args = _parser().parse_args(argv) output_stream = stdout or sys.stdout error_stream = stderr or sys.stderr torch.set_num_threads(max(1, args.threads)) active_scorer = scorer or DebertaCandidateScorer( device=args.device, cache_dir=args.model_cache_dir, local_files_only=args.offline, ) model_load_seconds: float | None = None loader = getattr(active_scorer, "load", None) if callable(loader): started = time.perf_counter() loader() model_load_seconds = time.perf_counter() - started last_reported: dict[str, int] = {} def progress(stage: str, completed: int, total: int) -> None: previous = last_reported.get(stage, -1) if completed == total or completed == 0 or completed - previous >= 100: error_stream.write(f"[{stage}] {completed}/{total}\n") error_stream.flush() last_reported[stage] = completed benchmark = run_benchmark( active_scorer, data_dir=args.data_dir, config=BenchmarkConfig( pool_size=args.pool_size, limit_per_split=args.limit or None, seed=args.seed, bootstrap_samples=args.bootstrap_samples, example_limit=args.example_limit, prototype_test_exclusion_count=(0 if args.include_prototype_test_rows else 120), ), model_load_seconds=model_load_seconds, progress=progress, ) json_path, markdown_path = write_benchmark_outputs( benchmark, output_dir=args.output_dir, stem=args.stem, ) summary = { "ok": True, "json": str(json_path), "markdown": str(markdown_path), "test": benchmark.report["metrics"]["sealed_test_remainder"], "statistics": benchmark.report["paired_test_statistics"], } output_stream.write(json.dumps(summary, ensure_ascii=False, indent=2) + "\n") return 0 def main() -> None: raise SystemExit(run()) if __name__ == "__main__": main()