o2o-plagiarism-ai/app/engine/preference.py

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,
}