555 lines
21 KiB
Python
555 lines
21 KiB
Python
"""AI 생성 판별 학습셋 빌더 — human(xlsx) + AI(JSONL/CSV) 결합.
|
||
|
||
목적:
|
||
컴북스에서 받은 xlsx 의 '에피소드' 본문(= human 라벨)과, 별도로 준비한 AI
|
||
생성 텍스트(JSONL/CSV)를 결합해 group-aware 분할이 가능한 학습셋을 만든다.
|
||
|
||
⚠️ 외부 API 를 호출하지 않는다. AI 샘플은 **이미 생성되어 파일로 존재하는 것만**
|
||
읽는다. 원고를 외부 서비스로 내보내는 경로를 이 스크립트에 만들지 말 것.
|
||
|
||
누출(leakage) 방지:
|
||
· 분할 단위는 개별 텍스트가 아니라 **source_group**(기본: 책/도서 식별자)이다.
|
||
같은 책의 에피소드가 train 과 test 에 동시에 들어가면, 모델이 문체가 아니라
|
||
'그 책'을 외워 성능이 부풀려진다.
|
||
· AI 샘플도 생성기·프롬프트 출처 단위로 묶는다(`ai:<generator>` 기본).
|
||
· 한 group 에 human/AI 라벨이 섞이면 경고한다. 원본을 AI로 재작성한 페어라면
|
||
**같은 group 으로 묶여야** 원본-생성물이 분할을 가로지르지 않는다.
|
||
· 정규화 후 완전 중복 텍스트는 제거한다(9,633행 중복 전례).
|
||
|
||
사용:
|
||
# 1) 컬럼 확인만 (아무것도 쓰지 않음)
|
||
python scripts/build_ai_training_dataset.py --xlsx episodes.xlsx --inspect
|
||
|
||
# 2) 실제 빌드
|
||
python scripts/build_ai_training_dataset.py \
|
||
--xlsx episodes.xlsx --text-column 본문 --book-column 도서명 \
|
||
--ai-jsonl data/training/ai_samples.jsonl \
|
||
--out data/training/ai_dataset.jsonl
|
||
|
||
출력 JSONL 1행:
|
||
{"text": "...", "label": 0, "origin": "human", "source_group": "book:홍길동전",
|
||
"book": "홍길동전", "split": "train", "char_count": 812, "meta": {...}}
|
||
label: 0=human, 1=ai
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import argparse
|
||
import csv
|
||
import hashlib
|
||
import json
|
||
import logging
|
||
import os
|
||
import re
|
||
import sys
|
||
import unicodedata
|
||
from collections import Counter, defaultdict
|
||
from dataclasses import dataclass, field
|
||
from pathlib import Path
|
||
|
||
ROOT = Path(__file__).resolve().parent.parent
|
||
sys.path.insert(0, str(ROOT))
|
||
|
||
from app.engine.training_data import pseudonymous_id, redact_direct_identifiers # noqa: E402
|
||
|
||
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s: %(message)s")
|
||
logger = logging.getLogger("build-ai-dataset")
|
||
|
||
LABEL_HUMAN = 0
|
||
LABEL_AI = 1
|
||
|
||
#: 헤더 자동탐지 힌트 (부분일치, 소문자 비교). 실제 파일을 못 본 상태라
|
||
#: 오탐 가능성이 있으니 --inspect 로 먼저 확인하고 --text-column 으로 고정할 것.
|
||
TEXT_HINTS = ("본문", "에피소드", "내용", "원고", "텍스트", "story", "text", "body", "content")
|
||
BOOK_HINTS = ("도서", "책", "서명", "제목", "book", "title", "작품")
|
||
ID_HINTS = ("id", "번호", "no", "식별")
|
||
AUTHOR_HINTS = ("저자", "작가", "author", "writer")
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 정규화 / 중복 제거
|
||
# ---------------------------------------------------------------------------
|
||
|
||
def normalize(text: str) -> str:
|
||
if not text:
|
||
return ""
|
||
t = unicodedata.normalize("NFKC", str(text))
|
||
t = t.replace("\r\n", "\n").replace("\r", "\n")
|
||
t = re.sub(r"[ \t ]+", " ", t)
|
||
t = re.sub(r"\n{3,}", "\n\n", t)
|
||
return t.strip()
|
||
|
||
|
||
def dedup_key(text: str) -> str:
|
||
"""공백까지 제거한 형태의 해시 — 서식만 다른 중복을 잡는다."""
|
||
compact = re.sub(r"\s+", "", text)
|
||
return hashlib.sha1(compact.encode("utf-8")).hexdigest()
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 레코드
|
||
# ---------------------------------------------------------------------------
|
||
|
||
@dataclass
|
||
class Record:
|
||
text: str
|
||
label: int
|
||
origin: str # "human" | "ai"
|
||
source_group: str
|
||
book: str = ""
|
||
meta: dict = field(default_factory=dict)
|
||
split: str = ""
|
||
|
||
@property
|
||
def char_count(self) -> int:
|
||
return len(self.text)
|
||
|
||
def to_json(self) -> dict:
|
||
return {
|
||
"text": self.text,
|
||
"label": self.label,
|
||
"origin": self.origin,
|
||
"source_group": self.source_group,
|
||
"book": self.book,
|
||
"split": self.split,
|
||
"char_count": self.char_count,
|
||
"meta": self.meta,
|
||
}
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# xlsx 로드
|
||
# ---------------------------------------------------------------------------
|
||
|
||
def _load_workbook(path: Path):
|
||
try:
|
||
from openpyxl import load_workbook
|
||
except ImportError:
|
||
raise SystemExit(
|
||
"openpyxl 이 필요합니다. `pip install openpyxl` 또는 "
|
||
"requirements.txt 설치 후 다시 실행하세요."
|
||
)
|
||
return load_workbook(filename=str(path), read_only=True, data_only=True)
|
||
|
||
|
||
def _match_column(headers: list[str], hints: tuple[str, ...]) -> str | None:
|
||
"""헤더 목록에서 힌트에 부분일치하는 첫 컬럼명."""
|
||
lowered = [(h, (h or "").strip().lower()) for h in headers]
|
||
for hint in hints:
|
||
for original, low in lowered:
|
||
if hint in low:
|
||
return original
|
||
return None
|
||
|
||
|
||
def inspect_xlsx(path: Path, sheet: str | None = None) -> None:
|
||
"""헤더와 표본을 출력만 한다. 컬럼 지정 전에 반드시 한 번 실행할 것."""
|
||
wb = _load_workbook(path)
|
||
sheets = [sheet] if sheet else wb.sheetnames
|
||
for name in sheets:
|
||
ws = wb[name]
|
||
rows = ws.iter_rows(values_only=True)
|
||
try:
|
||
header = [str(c) if c is not None else "" for c in next(rows)]
|
||
except StopIteration:
|
||
logger.info("[%s] 빈 시트", name)
|
||
continue
|
||
logger.info("[시트 %s] 컬럼 %d개", name, len(header))
|
||
for i, h in enumerate(header):
|
||
logger.info(" [%2d] %s", i, h)
|
||
logger.info(
|
||
" 자동탐지 → text=%s / book=%s / id=%s",
|
||
_match_column(header, TEXT_HINTS),
|
||
_match_column(header, BOOK_HINTS),
|
||
_match_column(header, ID_HINTS),
|
||
)
|
||
for n, row in enumerate(rows):
|
||
if n >= 2:
|
||
break
|
||
preview = {
|
||
h: (str(v)[:60] + "…" if v is not None and len(str(v)) > 60 else v)
|
||
for h, v in zip(header, row)
|
||
}
|
||
logger.info(" 샘플%d: %s", n + 1, preview)
|
||
wb.close()
|
||
|
||
|
||
def load_human_from_xlsx(
|
||
path: Path,
|
||
text_column: str | None,
|
||
book_column: str | None,
|
||
id_column: str | None,
|
||
sheet: str | None,
|
||
min_chars: int,
|
||
group_column: str | None = None,
|
||
anonymization_salt: str | None = None,
|
||
provenance: str = "unknown",
|
||
) -> list[Record]:
|
||
wb = _load_workbook(path)
|
||
sheets = [sheet] if sheet else wb.sheetnames
|
||
records: list[Record] = []
|
||
|
||
for name in sheets:
|
||
ws = wb[name]
|
||
rows = ws.iter_rows(values_only=True)
|
||
try:
|
||
header = [str(c) if c is not None else "" for c in next(rows)]
|
||
except StopIteration:
|
||
continue
|
||
|
||
tcol = text_column or _match_column(header, TEXT_HINTS)
|
||
if tcol is None or tcol not in header:
|
||
logger.warning(
|
||
"[%s] 본문 컬럼을 찾지 못했습니다(지정=%s). --inspect 로 확인 후 "
|
||
"--text-column 으로 지정하세요. 이 시트는 건너뜁니다.",
|
||
name, text_column,
|
||
)
|
||
continue
|
||
bcol = book_column or _match_column(header, BOOK_HINTS)
|
||
icol = id_column or _match_column(header, ID_HINTS)
|
||
acol = _match_column(header, AUTHOR_HINTS)
|
||
gcol = group_column or acol or icol or bcol
|
||
|
||
ti = header.index(tcol)
|
||
bi = header.index(bcol) if bcol in header else None
|
||
ii = header.index(icol) if icol in header else None
|
||
ai = header.index(acol) if acol in header else None
|
||
gi = header.index(gcol) if gcol in header else None
|
||
|
||
logger.info(
|
||
"[%s] text=%r book=%r id=%r author=%r", name, tcol, bcol, icol, acol
|
||
)
|
||
|
||
for rownum, row in enumerate(rows, start=2):
|
||
if ti >= len(row):
|
||
continue
|
||
text = normalize(row[ti] if row[ti] is not None else "")
|
||
if anonymization_salt:
|
||
text = redact_direct_identifiers(text)
|
||
if len(text) < min_chars:
|
||
continue
|
||
book = ""
|
||
if bi is not None and bi < len(row) and row[bi] is not None:
|
||
book = str(row[bi]).strip()
|
||
raw_group = ""
|
||
if gi is not None and gi < len(row) and row[gi] is not None:
|
||
raw_group = str(row[gi]).strip()
|
||
if anonymization_salt and raw_group:
|
||
group = pseudonymous_id(raw_group, anonymization_salt)
|
||
elif raw_group:
|
||
group = f"source:{raw_group}"
|
||
else:
|
||
group = f"sheet:{name}"
|
||
meta = {
|
||
"sheet": name,
|
||
"row": rownum,
|
||
"source_file": path.name,
|
||
"provenance": provenance,
|
||
"human_verified": True,
|
||
"ai_assistance": False,
|
||
}
|
||
if ii is not None and ii < len(row) and row[ii] is not None:
|
||
raw_id = str(row[ii]).strip()
|
||
if anonymization_salt:
|
||
meta["row_id_hash"] = pseudonymous_id(raw_id, anonymization_salt)
|
||
else:
|
||
meta["row_id"] = raw_id
|
||
if ai is not None and ai < len(row) and row[ai] is not None:
|
||
raw_author = str(row[ai]).strip()
|
||
if anonymization_salt:
|
||
meta["author_hash"] = pseudonymous_id(raw_author, anonymization_salt)
|
||
else:
|
||
meta["author"] = raw_author
|
||
records.append(
|
||
Record(
|
||
text=text, label=LABEL_HUMAN, origin="human",
|
||
source_group=group,
|
||
book="" if anonymization_salt else book,
|
||
meta=meta,
|
||
)
|
||
)
|
||
wb.close()
|
||
logger.info("human 레코드 %d건 (xlsx=%s)", len(records), path.name)
|
||
return records
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# AI 샘플 (JSONL / CSV)
|
||
# ---------------------------------------------------------------------------
|
||
|
||
def load_ai_records(
|
||
jsonl_paths: list[Path],
|
||
csv_paths: list[Path],
|
||
text_field: str,
|
||
group_field: str | None,
|
||
generator_field: str,
|
||
min_chars: int,
|
||
) -> list[Record]:
|
||
records: list[Record] = []
|
||
|
||
def _push(obj: dict, src: str, idx: int) -> None:
|
||
raw = obj.get(text_field)
|
||
if raw is None:
|
||
return
|
||
text = normalize(raw)
|
||
if len(text) < min_chars:
|
||
return
|
||
generator = str(obj.get(generator_field) or "unknown").strip()
|
||
if group_field and obj.get(group_field):
|
||
group = str(obj[group_field]).strip()
|
||
else:
|
||
# 생성기 단위 묶음. 같은 모델이 만든 글끼리는 문체가 닮아
|
||
# 분할을 가로지르면 성능이 부풀려진다.
|
||
group = f"ai:{generator}"
|
||
meta = {k: v for k, v in obj.items() if k != text_field}
|
||
meta["source_file"] = src
|
||
meta.setdefault("row", idx)
|
||
records.append(
|
||
Record(
|
||
text=text, label=LABEL_AI, origin="ai",
|
||
source_group=group, book=str(obj.get("book") or ""), meta=meta,
|
||
)
|
||
)
|
||
|
||
for p in jsonl_paths:
|
||
with p.open(encoding="utf-8") as fh:
|
||
for i, line in enumerate(fh, start=1):
|
||
line = line.strip()
|
||
if not line:
|
||
continue
|
||
try:
|
||
_push(json.loads(line), p.name, i)
|
||
except json.JSONDecodeError:
|
||
logger.warning("%s:%d JSON 파싱 실패 — 건너뜀", p.name, i)
|
||
|
||
for p in csv_paths:
|
||
with p.open(encoding="utf-8", newline="") as fh:
|
||
for i, row in enumerate(csv.DictReader(fh), start=2):
|
||
_push(row, p.name, i)
|
||
|
||
logger.info("ai 레코드 %d건", len(records))
|
||
return records
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 중복 제거 + 그룹 분할
|
||
# ---------------------------------------------------------------------------
|
||
|
||
def deduplicate(records: list[Record]) -> tuple[list[Record], int]:
|
||
seen: dict[str, Record] = {}
|
||
dropped = 0
|
||
for rec in records:
|
||
key = dedup_key(rec.text)
|
||
if key in seen:
|
||
dropped += 1
|
||
continue
|
||
seen[key] = rec
|
||
return list(seen.values()), dropped
|
||
|
||
|
||
def assign_splits(
|
||
records: list[Record],
|
||
ratios: tuple[float, float, float],
|
||
seed: int,
|
||
) -> dict[str, str]:
|
||
"""source_group 단위 결정적 분할. 같은 group 은 반드시 같은 split.
|
||
|
||
sklearn 없이도 재현 가능하도록 해시 기반으로 자른다. 다만 라벨 비율이
|
||
한쪽으로 쏠리지 않도록, 라벨별로 group 을 나눠 각각 비율을 맞춘다.
|
||
"""
|
||
train_r, val_r, _ = ratios
|
||
by_label: dict[int, list[str]] = defaultdict(list)
|
||
group_labels: dict[str, set[int]] = defaultdict(set)
|
||
|
||
for rec in records:
|
||
group_labels[rec.source_group].add(rec.label)
|
||
|
||
for group, labels in group_labels.items():
|
||
if len(labels) > 1:
|
||
logger.warning(
|
||
"group %r 에 human/AI 라벨이 함께 있습니다. 원본-생성물 페어라면 "
|
||
"의도된 것이며 분할 누출은 방지됩니다.", group
|
||
)
|
||
# 대표 라벨(다수)로 비율 배분만 결정
|
||
by_label[sorted(labels)[0]].append(group)
|
||
|
||
assignment: dict[str, str] = {}
|
||
for label, groups in by_label.items():
|
||
# 해시로 결정적 순서 부여 (입력 순서에 의존하지 않음)
|
||
ordered = sorted(
|
||
groups,
|
||
key=lambda g: hashlib.sha1(f"{seed}:{g}".encode("utf-8")).hexdigest(),
|
||
)
|
||
n = len(ordered)
|
||
n_train = int(round(n * train_r))
|
||
n_val = int(round(n * val_r))
|
||
# 그룹이 적을 때 test 가 0이 되지 않도록 보정
|
||
if n >= 3:
|
||
n_train = min(n_train, n - 2)
|
||
n_val = min(n_val, n - n_train - 1)
|
||
for i, g in enumerate(ordered):
|
||
if i < n_train:
|
||
assignment[g] = "train"
|
||
elif i < n_train + n_val:
|
||
assignment[g] = "val"
|
||
else:
|
||
assignment[g] = "test"
|
||
logger.info(
|
||
"label=%d groups=%d → train=%d val=%d test=%d",
|
||
label, n, n_train, n_val, n - n_train - n_val,
|
||
)
|
||
return assignment
|
||
|
||
|
||
def summarize(records: list[Record]) -> dict:
|
||
by_split = Counter(r.split for r in records)
|
||
by_label = Counter(r.label for r in records)
|
||
per_split_label = Counter((r.split, r.label) for r in records)
|
||
groups = {r.source_group for r in records}
|
||
split_groups: dict[str, set[str]] = defaultdict(set)
|
||
for r in records:
|
||
split_groups[r.split].add(r.source_group)
|
||
|
||
overlap: list[str] = []
|
||
for a in ("train", "val", "test"):
|
||
for b in ("train", "val", "test"):
|
||
if a >= b:
|
||
continue
|
||
shared = split_groups[a] & split_groups[b]
|
||
if shared:
|
||
overlap.append(f"{a}∩{b}={len(shared)}")
|
||
|
||
return {
|
||
"total": len(records),
|
||
"by_label": {str(k): v for k, v in sorted(by_label.items())},
|
||
"by_split": dict(by_split),
|
||
"by_split_label": {f"{s}/{l}": c for (s, l), c in sorted(per_split_label.items())},
|
||
"group_count": len(groups),
|
||
"mean_chars": round(
|
||
sum(r.char_count for r in records) / len(records), 1
|
||
) if records else 0,
|
||
"group_overlap_between_splits": overlap,
|
||
}
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# CLI
|
||
# ---------------------------------------------------------------------------
|
||
|
||
def main() -> int:
|
||
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||
ap.add_argument("--xlsx", type=Path, help="human 에피소드 xlsx")
|
||
ap.add_argument("--sheet", default=None, help="시트명 (생략 시 전체)")
|
||
ap.add_argument("--inspect", action="store_true", help="헤더/표본만 출력하고 종료")
|
||
ap.add_argument("--text-column", default=None, help="본문 컬럼명 (미지정 시 자동탐지)")
|
||
ap.add_argument("--book-column", default=None, help="도서 컬럼명 (분할 그룹 기준)")
|
||
ap.add_argument("--id-column", default=None, help="행 식별자 컬럼명")
|
||
ap.add_argument("--group-column", default=None,
|
||
help="분할 그룹 컬럼. 자서전.net은 id를 지정해 작성자 누출 방지")
|
||
ap.add_argument("--anonymize", action="store_true",
|
||
help="그룹/행 ID를 HMAC 가명화하고 본문의 직접 식별자를 제거")
|
||
ap.add_argument("--anonymization-salt-env", default="DATA_ANONYMIZATION_SALT")
|
||
ap.add_argument("--human-provenance", default="combooks_confirmed_human")
|
||
ap.add_argument("--ai-jsonl", type=Path, action="append", default=[], help="AI 샘플 JSONL (반복 가능)")
|
||
ap.add_argument("--ai-csv", type=Path, action="append", default=[], help="AI 샘플 CSV (반복 가능)")
|
||
ap.add_argument("--ai-text-field", default="text", help="AI 파일의 본문 필드명")
|
||
ap.add_argument("--ai-group-field", default=None, help="AI 파일의 그룹 필드명")
|
||
ap.add_argument("--ai-generator-field", default="generator", help="생성 모델 필드명")
|
||
ap.add_argument("--min-chars", type=int, default=200, help="이 미만 길이는 제외")
|
||
ap.add_argument("--train-ratio", type=float, default=0.7)
|
||
ap.add_argument("--val-ratio", type=float, default=0.15)
|
||
ap.add_argument("--seed", type=int, default=20260810)
|
||
ap.add_argument("--out", type=Path, default=Path("data/training/ai_dataset.jsonl"))
|
||
ap.add_argument("--report", type=Path, default=None, help="요약 JSON 경로 (기본: <out>.summary.json)")
|
||
args = ap.parse_args()
|
||
|
||
if args.inspect:
|
||
if not args.xlsx:
|
||
ap.error("--inspect 에는 --xlsx 가 필요합니다")
|
||
inspect_xlsx(args.xlsx, args.sheet)
|
||
return 0
|
||
|
||
records: list[Record] = []
|
||
anonymization_salt = None
|
||
if args.anonymize:
|
||
anonymization_salt = os.environ.get(args.anonymization_salt_env, "")
|
||
if not anonymization_salt:
|
||
logger.error(
|
||
"--anonymize 사용 시 %s 환경변수가 필요합니다.",
|
||
args.anonymization_salt_env,
|
||
)
|
||
return 2
|
||
if args.xlsx:
|
||
if not args.xlsx.exists():
|
||
logger.error("xlsx 없음: %s", args.xlsx)
|
||
return 2
|
||
records += load_human_from_xlsx(
|
||
args.xlsx, args.text_column, args.book_column,
|
||
args.id_column, args.sheet, args.min_chars,
|
||
args.group_column, anonymization_salt, args.human_provenance,
|
||
)
|
||
|
||
jsonl = [p for p in args.ai_jsonl if p.exists()]
|
||
csvs = [p for p in args.ai_csv if p.exists()]
|
||
for p in list(args.ai_jsonl) + list(args.ai_csv):
|
||
if not p.exists():
|
||
logger.error("AI 샘플 파일 없음: %s", p)
|
||
return 2
|
||
if jsonl or csvs:
|
||
records += load_ai_records(
|
||
jsonl, csvs, args.ai_text_field, args.ai_group_field,
|
||
args.ai_generator_field, args.min_chars,
|
||
)
|
||
|
||
if not records:
|
||
logger.error("레코드가 0건입니다. --xlsx 또는 --ai-jsonl/--ai-csv 를 확인하세요.")
|
||
return 2
|
||
|
||
records, dropped = deduplicate(records)
|
||
logger.info("중복 제거: %d건 제외 → %d건", dropped, len(records))
|
||
|
||
labels = {r.label for r in records}
|
||
if len(labels) < 2:
|
||
logger.warning(
|
||
"라벨이 한 종류(%s)뿐입니다. 학습은 불가하며, 평가/특징분석 용도로만 "
|
||
"쓸 수 있습니다. train_ai_detector.py 는 이 데이터로 실패합니다.",
|
||
labels,
|
||
)
|
||
|
||
ratios = (args.train_ratio, args.val_ratio, 1 - args.train_ratio - args.val_ratio)
|
||
if ratios[2] <= 0:
|
||
logger.error("train+val 비율이 1 이상입니다: %s", ratios)
|
||
return 2
|
||
|
||
assignment = assign_splits(records, ratios, args.seed)
|
||
for rec in records:
|
||
rec.split = assignment.get(rec.source_group, "train")
|
||
|
||
args.out.parent.mkdir(parents=True, exist_ok=True)
|
||
with args.out.open("w", encoding="utf-8") as fh:
|
||
for rec in records:
|
||
fh.write(json.dumps(rec.to_json(), ensure_ascii=False) + "\n")
|
||
|
||
summary = summarize(records)
|
||
summary["dropped_duplicates"] = dropped
|
||
summary["seed"] = args.seed
|
||
summary["min_chars"] = args.min_chars
|
||
report_path = args.report or args.out.with_suffix(".summary.json")
|
||
report_path.write_text(
|
||
json.dumps(summary, ensure_ascii=False, indent=2), encoding="utf-8"
|
||
)
|
||
|
||
logger.info("저장: %s (%d건)", args.out, len(records))
|
||
logger.info("요약: %s", json.dumps(summary, ensure_ascii=False))
|
||
if summary["group_overlap_between_splits"]:
|
||
logger.error(
|
||
"분할 간 그룹 중복이 발견되었습니다: %s — 누출 위험!",
|
||
summary["group_overlap_between_splits"],
|
||
)
|
||
return 1
|
||
return 0
|
||
|
||
|
||
if __name__ == "__main__":
|
||
raise SystemExit(main())
|