| 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() |
|
|