File size: 3,605 Bytes
54c3e65 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 | 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()
|