from __future__ import annotations import argparse import json import sys import time from collections.abc import Callable, Sequence from pathlib import Path from typing import TextIO import torch from . import profiles as profile_config from .benchmark_v2_runner import ( FrozenProfile, V2BenchmarkConfig, run_dev_selection, run_external_evaluation, write_v2_outputs, ) from .deberta import DebertaCandidateScorer from .domain import CandidateScorer from .mozc import MozcDictionaryIndex from .ud_gsd import CorpusArtifact, load_pinned_split from .ud_pud import PUD_LICENSE, PUD_REPOSITORY, load_pinned_test ArtifactLoader = Callable[[str, Path], CorpusArtifact] def _parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser(prog="deberta-ime-benchmark-v2") subparsers = parser.add_subparsers(dest="stage", required=True) dev = subparsers.add_parser("dev", help="select a profile on pinned GSD dev only") external = subparsers.add_parser( "external", help="evaluate a frozen profile on pinned Japanese-PUD" ) for command in (dev, external): command.add_argument("--index", type=Path, required=True) command.add_argument("--data-dir", type=Path, default=Path("work/data")) command.add_argument("--output-dir", type=Path, default=Path("outputs")) command.add_argument("--stem", required=True) command.add_argument("--limit", type=int, default=0, help="0 evaluates all rows") command.add_argument("--pool-size", type=int, default=8) command.add_argument("--bootstrap-samples", type=int, default=5000) command.add_argument("--example-limit", type=int, default=20) command.add_argument("--seed", type=int, default=20260810) command.add_argument("--device", default="cpu") command.add_argument("--model-cache-dir", type=Path) command.add_argument("--offline", action="store_true") command.add_argument("--threads", type=int, default=min(8, torch.get_num_threads())) dev.add_argument("--context-mode", choices=("left_only", "bidirectional"), required=True) external.add_argument( "--profile", choices=tuple(profile_config.MOZC_PROFILES), required=True ) return parser def run( argv: Sequence[str] | None = None, *, stdout: TextIO | None = None, stderr: TextIO | None = None, scorer: CandidateScorer | None = None, artifact_loader: ArtifactLoader | None = None, ) -> int: args = _parser().parse_args(argv) output_stream = stdout or sys.stdout error_stream = stderr or sys.stderr if args.stage == "external" and profile_config.MOZC_PROFILE_STATE != "frozen": error_stream.write("error: Mozc profiles are not frozen from GSD dev\n") return 2 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 def progress(completed: int, total: int) -> None: if completed == total or completed % 100 == 0: error_stream.write(f"[{args.stage}] {completed}/{total}\n") error_stream.flush() active_loader = artifact_loader or _load_artifact config = V2BenchmarkConfig( pool_size=args.pool_size, limit=args.limit or None, seed=args.seed, bootstrap_samples=args.bootstrap_samples, example_limit=args.example_limit, ) artifact = active_loader(args.stage, args.data_dir) with MozcDictionaryIndex(args.index) as index: if args.stage == "dev": benchmark = run_dev_selection( active_scorer, index=index, artifact=artifact, context_mode=args.context_mode, config=config, model_load_seconds=model_load_seconds, progress=progress, ) else: active_profile = profile_config.MOZC_PROFILES[args.profile] benchmark = run_external_evaluation( active_scorer, index=index, artifact=artifact, dataset_repository=PUD_REPOSITORY, dataset_license=PUD_LICENSE, profile=FrozenProfile( name=args.profile, context_mode=( "left_only" if args.profile == "incremental" else "bidirectional" ), prior_weight=active_profile.prior_weight, min_margin=active_profile.min_margin, selected_on=( f"UD Japanese-GSD dev {profile_config.MOZC_PROFILE_SELECTION_REVISION}" ), ), config=config, model_load_seconds=model_load_seconds, progress=progress, ) json_path, markdown_path = write_v2_outputs( benchmark, output_dir=args.output_dir, stem=args.stem, ) summary = { "ok": True, "stage": benchmark.report["stage"], "json": str(json_path), "markdown": str(markdown_path), "coverage": benchmark.report["coverage"], } output_stream.write(json.dumps(summary, ensure_ascii=False, indent=2) + "\n") return 0 def _load_artifact(stage: str, data_dir: Path) -> CorpusArtifact: if stage == "dev": return load_pinned_split("dev", data_dir / "ud-japanese-gsd") return load_pinned_test(data_dir / "ud-japanese-pud") def main() -> None: raise SystemExit(run()) if __name__ == "__main__": main()