feat: support sharded infringement batches

This commit is contained in:
hbyang 2026-09-02 09:28:38 +09:00
parent c35c5aabfc
commit cbeb1e1084

View File

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