feat: 로컬 GPU AI 표본 병렬 생성 지원
This commit is contained in:
parent
946e9e9474
commit
aabd7844d8
@ -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.")
|
||||
|
||||
@ -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):
|
||||
|
||||
Loading…
Reference in New Issue
Block a user