diff --git a/scripts/run_infringement_batch.py b/scripts/run_infringement_batch.py index d88726b..4b15604 100644 --- a/scripts/run_infringement_batch.py +++ b/scripts/run_infringement_batch.py @@ -216,12 +216,20 @@ def main() -> int: parser.add_argument("--out-xlsx", type=Path, required=True) parser.add_argument("--progress-every", type=int, default=50) parser.add_argument("--resume", action="store_true") + parser.add_argument("--shard-count", type=int, default=1) + parser.add_argument("--shard-index", type=int, default=0) + parser.add_argument( + "--no-xlsx", action="store_true", + help="분할 작업용: JSONL만 쓰고 XLSX 요약은 병합 단계에서 생성", + ) args = parser.parse_args() logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s") if not _env_is_false(os.environ.get("USE_LLM_LEGAL_JUDGE")): raise RuntimeError("USE_LLM_LEGAL_JUDGE=false 이어야 원고 외부 전송 없이 실행할 수 있습니다") if not args.database.exists(): parser.error(f"database does not exist: {args.database}") + if args.shard_count < 1 or not 0 <= args.shard_index < args.shard_count: + parser.error("shard-index는 0 이상이고 shard-count보다 작아야 합니다") get_settings.cache_clear() settings = get_settings().model_copy(update={ @@ -234,7 +242,11 @@ def main() -> int: if not detector.uses_persistent_index: raise RuntimeError("영속 인덱스가 준비되지 않아 배치를 중단합니다") store = CorpusStore(args.database) - items = build_query_plan(store) + all_items = build_query_plan(store) + items = [ + item for position, item in enumerate(all_items) + if position % args.shard_count == args.shard_index + ] completed = load_completed(args.out_jsonl) if args.resume else set() if args.out_jsonl.exists() and not args.resume: args.out_jsonl.unlink() @@ -286,6 +298,9 @@ def main() -> int: rows = [json.loads(line) for line in args.out_jsonl.read_text(encoding="utf-8").splitlines() if line.strip()] if len(rows) != len(items): raise RuntimeError(f"결과 행 수 불일치: {len(rows)} != {len(items)}") + if args.no_xlsx: + LOGGER.info("batch_shard_completed shard=%d/%d total=%d jsonl=%s", args.shard_index, args.shard_count, len(rows), args.out_jsonl) + return 0 build_xlsx(rows, args.out_xlsx, backend=backend) suspected = sum(bool(row["침해 의심"]) for row in rows) LOGGER.info("batch_completed total=%d suspected=%d jsonl=%s xlsx=%s", len(rows), suspected, args.out_jsonl, args.out_xlsx)