CoolFace
Modelpublic

WilliamCHN/Legal_Document_Segment_Model

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
infer_cli.py94 linesDownload Raw Back to root
1from __future__ import annotations2 3import argparse4import json5import sys6from pathlib import Path7 8from tqdm import tqdm9 10# Allow running without installation: `python infer_cli.py ...`11BUNDLE_ROOT = Path(__file__).resolve().parent12SRC_DIR = BUNDLE_ROOT / "src"13if str(SRC_DIR) not in sys.path:14    sys.path.insert(0, str(SRC_DIR))15 16from judgment_partition_infer.infer import Predictor, default_run_dir, write_run_meta  # noqa: E40217 18 19def main() -> int:20    parser = argparse.ArgumentParser(description="judgment_partition_infer (JSONL -> JSONL)")21    parser.add_argument("--input", type=str, required=True, help="Input jsonl")22    parser.add_argument(23        "--output-root",24        type=str,25        default=None,26        help="Root output dir. Default: ./output/<timestamp>/",27    )28    parser.add_argument(29        "--output",30        type=str,31        default=None,32        help="Explicit output jsonl path (overrides output-root/timestamp).",33    )34    parser.add_argument("--model", type=str, default=None, help="Model checkpoint (.pt)")35    parser.add_argument("--vocab", type=str, default=None, help="Vocab json")36    parser.add_argument("--device", type=str, default="cuda", help="cuda|cpu (cuda falls back to cpu)")37    parser.add_argument("--anchor", type=str, default="auto", choices=["auto", "off"])38    parser.add_argument("--max-samples", type=int, default=None, help="Process at most N samples")39    args = parser.parse_args()40 41    input_path = Path(args.input)42    if not input_path.exists():43        raise FileNotFoundError(f"Missing input: {input_path}")44 45    output_root = Path(args.output_root) if args.output_root else (BUNDLE_ROOT / "output")46    run_dir = default_run_dir(output_root)47    run_dir.mkdir(parents=True, exist_ok=True)48 49    output_path = Path(args.output) if args.output else (run_dir / "predictions.jsonl")50    output_path.parent.mkdir(parents=True, exist_ok=True)51 52    predictor = Predictor(53        model_path=Path(args.model) if args.model else None,54        vocab_path=Path(args.vocab) if args.vocab else None,55        device=args.device,56        anchor=args.anchor,57    )58 59    meta = {60        "input": str(input_path),61        "output": str(output_path),62        "run_dir": str(run_dir),63        "device_requested": args.device,64        "device_used": str(predictor.torch_device),65        "anchor": args.anchor,66        "model_path": str(predictor.model_path),67        "vocab_path": str(predictor.vocab_path),68    }69    write_run_meta(run_dir / "run_meta.json", meta)70 71    written = 072    with input_path.open("r", encoding="utf-8") as f_in, output_path.open("w", encoding="utf-8") as f_out:73        for line in tqdm(f_in, desc="Infer", unit="line"):74            if args.max_samples is not None and written >= args.max_samples:75                break76            line = line.strip()77            if not line:78                continue79            try:80                record = json.loads(line)81            except Exception:82                continue83            out = predictor.predict_record(record)84            f_out.write(json.dumps(out, ensure_ascii=False) + "\n")85            written += 186 87    print(f"[DONE] samples={written} -> {output_path}")88    return 089 90 91if __name__ == "__main__":92    raise SystemExit(main())93 94