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