From aabd7844d84305cf0f609673956cb2478def6710 Mon Sep 17 00:00:00 2001 From: hbyang Date: Thu, 20 Aug 2026 17:13:19 +0900 Subject: [PATCH] =?UTF-8?q?feat:=20=EB=A1=9C=EC=BB=AC=20GPU=20AI=20?= =?UTF-8?q?=ED=91=9C=EB=B3=B8=20=EB=B3=91=EB=A0=AC=20=EC=83=9D=EC=84=B1=20?= =?UTF-8?q?=EC=A7=80=EC=9B=90?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- scripts/generate_ai_samples.py | 183 ++++++++++++++++++------------ tests/test_generate_ai_samples.py | 4 + 2 files changed, 114 insertions(+), 73 deletions(-) diff --git a/scripts/generate_ai_samples.py b/scripts/generate_ai_samples.py index c5191b0..7165036 100644 --- a/scripts/generate_ai_samples.py +++ b/scripts/generate_ai_samples.py @@ -52,6 +52,7 @@ import logging import os import random import sys +from concurrent.futures import ThreadPoolExecutor from pathlib import Path ROOT = Path(__file__).resolve().parent.parent @@ -156,6 +157,14 @@ def load_done_ids(out_path: Path) -> set[str]: return done +def _hangul_ratio(text: str) -> float: + """영문 오류·깨진 출력을 거르기 위한 경량 품질 게이트.""" + letters = [ch for ch in text if ch.isalpha()] + if not letters: + return 0.0 + return sum("가" <= ch <= "힣" for ch in letters) / len(letters) + + def main() -> int: ap = argparse.ArgumentParser( description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter @@ -178,11 +187,19 @@ def main() -> int: help="OpenAI 호환 엔드포인트 (사내 vLLM 등)") ap.add_argument("--limit", type=int, default=100, help="생성할 표본 수") ap.add_argument("--temperature", type=float, default=1.0) + ap.add_argument("--concurrency", type=int, default=1, + help="동시 생성 요청 수(로컬 서버 슬롯 수 이하)") + ap.add_argument("--max-tokens", type=int, default=2200, + help="요청당 최대 생성 토큰") + ap.add_argument("--min-hangul-ratio", type=float, default=0.55, + help="알파벳 문자 중 최소 한글 비율") ap.add_argument("--seed", type=int, default=20260819) ap.add_argument("--out", type=Path, default=Path("data/training/ai_samples.jsonl")) ap.add_argument("--dry-run", action="store_true", help="프롬프트만 출력하고 API 를 호출하지 않는다 (비용 0)") args = ap.parse_args() + if args.concurrency < 1: + ap.error("--concurrency는 1 이상이어야 합니다") import pandas as pd @@ -237,7 +254,8 @@ def main() -> int: raise SystemExit("OPENAI_API_KEY 가 없습니다. --dry-run 으로 먼저 확인하세요.") client = OpenAI(api_key=api_key or "local", base_url=args.base_url, max_retries=2) - written = skipped_len = failed = 0 + target_pool = [n for n in lengths if low_chars <= n <= high_chars] + written = skipped_len = skipped_quality = failed = 0 attempted = 0 # dry-run 은 파일시스템도 건드리지 않는다. 빈 출력 파일이 남으면 다음 실행이 @@ -248,84 +266,103 @@ def main() -> int: args.out.parent.mkdir(parents=True, exist_ok=True) sink = args.out.open("a", encoding="utf-8") + def prepare(idx: int): + nonlocal attempted + prompt_id = hashlib.sha1(f"{args.xlsx.name}:{idx}".encode()).hexdigest()[:16] + if prompt_id in done: + return None + meta = { + c: sanitize_prompt_metadata(str(frame[c].iloc[idx])) + for c in usable_meta + if isinstance(frame[c].iloc[idx], str) and frame[c].iloc[idx].strip() + } + if not meta: + return None + source_group = f"ai-topic:{prompt_id}" + if args.group_column: + raw_group = str(frame[args.group_column].iloc[idx] or "").strip() + if raw_group and raw_group.lower() != "nan": + source_group = pseudonymous_id(raw_group, group_salt) + target = rng.choice(target_pool) + style = STYLE_VARIANTS[_stable_int(prompt_id) % len(STYLE_VARIANTS)] + model = models[attempted % len(models)] + prompt = _build_prompt(meta, target, style) + attempted += 1 + return prompt_id, meta, source_group, target, style, model, prompt + + def generate(task): + prompt_id, _, _, _, _, model, prompt = task + try: + response = client.chat.completions.create( + model=model, + temperature=args.temperature, + max_tokens=args.max_tokens, + messages=[ + {"role": "system", "content": SYSTEM_PROMPT}, + {"role": "user", "content": prompt}, + ], + ) + return task, (response.choices[0].message.content or "").strip(), None + except Exception as exc: # noqa: BLE001 — 개별 실패는 기록 후 계속 + logger.warning("생성 실패 (%s, %s): %s", model, prompt_id, exc) + return task, "", exc + + tasks = filter(None, (prepare(idx) for idx in order)) with sink as sink: - for idx in order: - if written >= args.limit: - break - prompt_id = hashlib.sha1(f"{args.xlsx.name}:{idx}".encode()).hexdigest()[:16] - if prompt_id in done: - continue - - meta = { - c: sanitize_prompt_metadata(str(frame[c].iloc[idx])) - for c in usable_meta - if isinstance(frame[c].iloc[idx], str) and frame[c].iloc[idx].strip() - } - if not meta: - continue - - source_group = f"ai-topic:{prompt_id}" - if args.group_column: - raw_group = str(frame[args.group_column].iloc[idx] or "").strip() - if raw_group and raw_group.lower() != "nan": - source_group = pseudonymous_id(raw_group, group_salt) - - target = rng.choice(lengths) - style = STYLE_VARIANTS[_stable_int(prompt_id) % len(STYLE_VARIANTS)] - model = models[attempted % len(models)] - prompt = _build_prompt(meta, target, style) - attempted += 1 - - if args.dry_run: + if args.dry_run: + for task in tasks: + if written >= args.limit: + break + _, _, _, target, _, model, prompt = task print(f"\n--- [{written + 1}] model={model} target={target}자 ---") print(prompt) written += 1 - continue - - try: - response = client.chat.completions.create( - model=model, - temperature=args.temperature, - messages=[ - {"role": "system", "content": SYSTEM_PROMPT}, - {"role": "user", "content": prompt}, - ], - ) - text = (response.choices[0].message.content or "").strip() - except Exception as exc: # noqa: BLE001 — 어떤 실패든 기록하고 계속 - failed += 1 - logger.warning("생성 실패 (%s, %s): %s", model, prompt_id, exc) - continue - - # human 과 같은 정규화를 거쳐야 양쪽 표기가 같은 형태가 된다. - text = normalize_ocr(text) - if not (low_chars <= len(text) <= high_chars): - # 길이가 어긋난 표본은 버린다. 남겨두면 판별기가 길이를 배운다. - skipped_len += 1 - continue - - sink.write(json.dumps({ - "text": text, - "generator": model, - "book": "", - "prompt_id": prompt_id, - "target_chars": target, - "char_count": len(text), - "style": style, - "source_group": source_group, - "generation_type": "pure_ai", - "provenance": "ai_generated_controlled", - "meta": meta, - }, ensure_ascii=False) + "\n") - sink.flush() # 중간에 죽어도 여기까지는 남는다 - written += 1 - if written % 50 == 0: - logger.info("생성 %d/%d (길이미달 폐기 %d, 실패 %d)", - written, args.limit, skipped_len, failed) + else: + with ThreadPoolExecutor(max_workers=args.concurrency) as executor: + while written < args.limit: + batch = [] + for _ in range(min(args.concurrency, args.limit - written)): + task = next(tasks, None) + if task is not None: + batch.append(task) + if not batch: + break + for task, text, error in executor.map(generate, batch): + prompt_id, meta, source_group, target, style, model, _ = task + if error is not None: + failed += 1 + continue + text = normalize_ocr(text) + if not (low_chars <= len(text) <= high_chars): + skipped_len += 1 + continue + if _hangul_ratio(text) < args.min_hangul_ratio: + skipped_quality += 1 + continue + sink.write(json.dumps({ + "text": text, + "generator": model, + "book": "", + "prompt_id": prompt_id, + "target_chars": target, + "char_count": len(text), + "style": style, + "source_group": source_group, + "generation_type": "pure_ai", + "provenance": "ai_generated_controlled", + "meta": meta, + }, ensure_ascii=False) + "\n") + sink.flush() + written += 1 + if written % 50 == 0: + logger.info( + "생성 %d/%d (길이 폐기 %d, 품질 폐기 %d, 실패 %d)", + written, args.limit, skipped_len, skipped_quality, failed, + ) logger.info( - "완료: %d건 생성 / 길이 벗어나 폐기 %d건 / 호출 실패 %d건 → %s", - written, skipped_len, failed, args.out, + "완료: %d건 생성 / 길이 폐기 %d건 / 품질 폐기 %d건 / 호출 실패 %d건 → %s", + written, skipped_len, skipped_quality, failed, args.out, ) if args.dry_run: logger.info("dry-run 이라 API 를 호출하지 않았습니다. 비용 0.") diff --git a/tests/test_generate_ai_samples.py b/tests/test_generate_ai_samples.py index 74fd8c8..5351e0a 100644 --- a/tests/test_generate_ai_samples.py +++ b/tests/test_generate_ai_samples.py @@ -74,6 +74,10 @@ class TestLengthBand: with pytest.raises(SystemExit): gen.load_human_length_band(frame, "에피소드") + def test_hangul_quality_gate(self): + assert gen._hangul_ratio("오늘은 시장에 갔다.") > 0.9 + assert gen._hangul_ratio("This is an English model failure.") == 0.0 + class TestResume: def test_reads_done_ids(self, tmp_path):