o2o-plagiarism-ai/scripts/build_preference_dataset.py

137 lines
5.3 KiB
Python

#!/usr/bin/env python3
"""Human Feedback 선호학습(DPO) 데이터 파이프라인 — 계획서 2단계 표절검출 고도화.
데이터(사람 선호 라벨) 도착 전까지 다음을 미리 가능하게 한다.
template : 후보쌍 입력(JSONL) → 라벨링 대기 템플릿 생성
convert : 라벨 완료 파일 → DPO 학습 포맷(JSONL) 변환 + 검증/통계
후보쌍 입력 형식 (JSONL):
{"pair_id": "p1", "original": "원문 ...", "text_a": "글A ...",
"text_b": "글B ...", "suggested_rejected": "b", "source_doc": "ref-0003"}
라벨링 템플릿 출력에서 사람은 각 행에 chosen/rejected 와 label_status="labeled" 를 채운다.
사용:
# 1) 후보쌍 → 라벨링 템플릿
python -m scripts.build_preference_dataset template candidates.jsonl -o label_me.jsonl
# 2) 라벨 완료 파일 → DPO 학습셋 + 통계
python -m scripts.build_preference_dataset convert labeled.jsonl -o dpo_train.jsonl
# 인자 없이 실행하면 내장 샘플로 파이프라인 검증
"""
from __future__ import annotations
import argparse
import json
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from app.engine.preference import ( # noqa: E402
PreferenceCandidate,
build_candidate,
dataset_stats,
to_dpo_dataset,
)
_SAMPLE_CANDIDATES = [
{
"pair_id": "demo-1",
"original": "홍길동은 활빈당을 만들어 탐관오리의 재물을 빼앗아 백성에게 나눠주었다.",
"text_a": "임꺽정은 도적단을 조직해 부패한 양반의 곡식을 빼앗아 농민에게 베풀었다.", # 정상 변형
"text_b": "홍길동은 활빈당을 만들어 탐관오리의 재물을 빼앗아 백성에게 나눠주었다.", # 거의 복제
"suggested_rejected": "b",
"source_doc": "ref-0003",
},
]
# convert 데모용: 사람이 라벨을 완료했다고 가정한 샘플
_SAMPLE_LABELED = [
{
"pair_id": "demo-1",
"prompt": "다음 원문을 참고한 두 글 중 ...\n\n[원문]\n홍길동은 ...",
"candidate_a": "임꺽정은 도적단을 조직해 ...",
"candidate_b": "홍길동은 활빈당을 만들어 ...",
"label_status": "labeled",
"chosen": "임꺽정은 도적단을 조직해 부패한 양반의 곡식을 빼앗아 농민에게 베풀었다.",
"rejected": "홍길동은 활빈당을 만들어 탐관오리의 재물을 빼앗아 백성에게 나눠주었다.",
},
]
def _read_jsonl(path: str, sample: list[dict]) -> list[dict]:
p = Path(path)
if not p.exists():
print(f"[warn] 파일 없음: {path} → 내장 샘플 사용", file=sys.stderr)
return sample
return [json.loads(l) for l in p.read_text(encoding="utf-8").splitlines() if l.strip()]
def _write_jsonl(rows: list[dict], path: str | None) -> None:
out = "\n".join(json.dumps(r, ensure_ascii=False) for r in rows)
if path:
Path(path).write_text(out + "\n", encoding="utf-8")
print(f"[ok] {len(rows)}건 저장 → {path}")
else:
print(out)
def cmd_template(args) -> None:
rows = _read_jsonl(args.input, _SAMPLE_CANDIDATES)
candidates = [
build_candidate(
pair_id=r.get("pair_id", f"p{i}"),
original=r["original"],
text_a=r["text_a"],
text_b=r["text_b"],
suggested_rejected=r.get("suggested_rejected"),
source_doc=r.get("source_doc"),
)
for i, r in enumerate(rows)
]
_write_jsonl([c.to_dict() for c in candidates], args.output)
print(f"\n라벨링 안내: 각 행의 chosen/rejected 를 candidate_a/b 텍스트로 채우고 "
f"label_status 를 'labeled' 로 변경하세요.", file=sys.stderr)
def cmd_convert(args) -> None:
rows = _read_jsonl(args.input, _SAMPLE_LABELED)
candidates = [PreferenceCandidate(**{k: v for k, v in r.items() if k in PreferenceCandidate.__dataclass_fields__})
for r in rows]
stats = dataset_stats(candidates)
print(f"\n=== 선호 데이터 통계 ===", file=sys.stderr)
print(f"전체 {stats['total']} / 라벨완료 {stats['labeled']} / 대기 {stats['pending']} / "
f"학습가능 {stats['trainable']}", file=sys.stderr)
if stats["errors"]:
print("[검증 오류]", file=sys.stderr)
for e in stats["errors"]:
print(f" - {e}", file=sys.stderr)
dpo = to_dpo_dataset(candidates)
_write_jsonl(dpo, args.output)
def main() -> None:
ap = argparse.ArgumentParser(description="HF 선호학습 데이터 파이프라인")
sub = ap.add_subparsers(dest="cmd")
t = sub.add_parser("template", help="후보쌍 → 라벨링 템플릿")
t.add_argument("input", nargs="?", default="__sample__")
t.add_argument("-o", "--output", default=None)
t.set_defaults(func=cmd_template)
c = sub.add_parser("convert", help="라벨 완료 → DPO 학습셋")
c.add_argument("input", nargs="?", default="__sample__")
c.add_argument("-o", "--output", default=None)
c.set_defaults(func=cmd_convert)
args = ap.parse_args()
if not args.cmd:
print("[info] 서브커맨드 없음 → template 데모 실행\n", file=sys.stderr)
cmd_template(argparse.Namespace(input="__sample__", output=None))
return
args.func(args)
if __name__ == "__main__":
main()