107 lines
4.1 KiB
Python
107 lines
4.1 KiB
Python
"""Human Feedback 기반 Preference Optimization 데이터 파이프라인 (계획서 p.22 고도화).
|
|
|
|
계획서 2단계 표절 검출 고도화:
|
|
> Human Feedback 형태의 Preference Optimization 을 통해서 표절 문서를
|
|
> 비선호하게끔 학습하는 형태로 sLLM 을 고도화.
|
|
|
|
선호 학습(DPO/ORPO)은 (prompt, chosen, rejected) 삼중쌍을 입력으로 한다.
|
|
표절 도메인에서:
|
|
chosen = 정당한 글(원본 또는 정상 2차 창작) ← 선호
|
|
rejected = 표절 글(무단 복제/요소 교체) ← 비선호
|
|
|
|
본 모듈은 데이터(사람 선호 라벨) 도착 전까지:
|
|
1) 라벨링 '후보쌍 템플릿' 생성 — 사람이 chosen/rejected 를 확정할 수 있는 골격
|
|
2) 라벨 완료 파일 → DPO 학습 포맷(JSONL) 변환 + 검증/통계
|
|
까지를 제공한다. 라벨이 들어오면 즉시 sLLM 선호학습에 투입 가능.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
from dataclasses import dataclass, field, asdict
|
|
|
|
# 선호 학습 기본 지시문 (표절 비선호 방향 고정)
|
|
DEFAULT_PROMPT = (
|
|
"다음 원문을 참고한 두 글 중, 저작권을 침해하지 않은 정당한 글을 선호하라."
|
|
)
|
|
|
|
|
|
@dataclass
|
|
class PreferenceCandidate:
|
|
"""라벨링 전 후보쌍. label_status 가 'pending' 이면 사람 확정 대기."""
|
|
pair_id: str
|
|
prompt: str
|
|
candidate_a: str
|
|
candidate_b: str
|
|
source_doc: str | None = None # 비교 기준 원본 doc_id
|
|
suggested_rejected: str | None = None # 엔진이 표절로 추정한 쪽 ("a"/"b") — 약한 신호
|
|
label_status: str = "pending" # pending | labeled
|
|
chosen: str | None = None # 라벨 결과: candidate_a/b 중 선호 텍스트
|
|
rejected: str | None = None # 라벨 결과: 비선호 텍스트
|
|
meta: dict = field(default_factory=dict)
|
|
|
|
def to_dict(self) -> dict:
|
|
return asdict(self)
|
|
|
|
|
|
def build_candidate(
|
|
pair_id: str,
|
|
original: str,
|
|
text_a: str,
|
|
text_b: str,
|
|
suggested_rejected: str | None = None,
|
|
source_doc: str | None = None,
|
|
prompt: str = DEFAULT_PROMPT,
|
|
) -> PreferenceCandidate:
|
|
"""원본 + 두 후보 글 → 라벨링 대기 후보쌍.
|
|
|
|
suggested_rejected: 표절 탐지 결과로 추정한 비선호 쪽('a'/'b'). 사람이 검토·확정.
|
|
"""
|
|
return PreferenceCandidate(
|
|
pair_id=pair_id,
|
|
prompt=f"{prompt}\n\n[원문]\n{original}",
|
|
candidate_a=text_a,
|
|
candidate_b=text_b,
|
|
source_doc=source_doc,
|
|
suggested_rejected=suggested_rejected,
|
|
meta={"original_len": len(original)},
|
|
)
|
|
|
|
|
|
def to_dpo_record(c: PreferenceCandidate) -> dict | None:
|
|
"""라벨 완료 후보쌍 → DPO/ORPO 학습 레코드. 미라벨이면 None."""
|
|
if c.label_status != "labeled" or not c.chosen or not c.rejected:
|
|
return None
|
|
return {"prompt": c.prompt, "chosen": c.chosen, "rejected": c.rejected, "pair_id": c.pair_id}
|
|
|
|
|
|
def to_dpo_dataset(candidates: list[PreferenceCandidate]) -> list[dict]:
|
|
return [r for c in candidates if (r := to_dpo_record(c)) is not None]
|
|
|
|
|
|
def validate_labeled(c: PreferenceCandidate) -> list[str]:
|
|
"""라벨 완료 후보쌍 검증 — 학습 투입 전 무결성 체크."""
|
|
errors: list[str] = []
|
|
if c.label_status == "labeled":
|
|
if not c.chosen:
|
|
errors.append(f"{c.pair_id}: chosen 누락")
|
|
if not c.rejected:
|
|
errors.append(f"{c.pair_id}: rejected 누락")
|
|
if c.chosen and c.rejected and c.chosen.strip() == c.rejected.strip():
|
|
errors.append(f"{c.pair_id}: chosen == rejected (구분 불가)")
|
|
return errors
|
|
|
|
|
|
def dataset_stats(candidates: list[PreferenceCandidate]) -> dict:
|
|
labeled = [c for c in candidates if c.label_status == "labeled"]
|
|
pending = [c for c in candidates if c.label_status != "labeled"]
|
|
errors: list[str] = []
|
|
for c in labeled:
|
|
errors.extend(validate_labeled(c))
|
|
return {
|
|
"total": len(candidates),
|
|
"labeled": len(labeled),
|
|
"pending": len(pending),
|
|
"trainable": len(to_dpo_dataset(candidates)),
|
|
"errors": errors,
|
|
}
|