60 lines
2.0 KiB
Python
60 lines
2.0 KiB
Python
"""HF 선호학습 데이터 파이프라인 단위테스트 (계획서 2단계 고도화)."""
|
|
from __future__ import annotations
|
|
|
|
from app.engine.preference import (
|
|
PreferenceCandidate,
|
|
build_candidate,
|
|
to_dpo_record,
|
|
to_dpo_dataset,
|
|
validate_labeled,
|
|
dataset_stats,
|
|
)
|
|
|
|
|
|
def test_build_candidate_is_pending():
|
|
c = build_candidate("p1", "원문", "글A", "글B", suggested_rejected="b")
|
|
assert c.label_status == "pending"
|
|
assert "원문" in c.prompt
|
|
assert to_dpo_record(c) is None # 미라벨은 학습 레코드 없음
|
|
|
|
|
|
def test_labeled_becomes_dpo_record():
|
|
c = build_candidate("p1", "원문", "정상글", "표절글", suggested_rejected="b")
|
|
c.label_status = "labeled"
|
|
c.chosen = "정상글"
|
|
c.rejected = "표절글"
|
|
rec = to_dpo_record(c)
|
|
assert rec is not None
|
|
assert rec["chosen"] == "정상글" and rec["rejected"] == "표절글"
|
|
|
|
|
|
def test_validate_catches_same_chosen_rejected():
|
|
c = PreferenceCandidate("p1", "p", "a", "b", label_status="labeled", chosen="x", rejected="x")
|
|
errors = validate_labeled(c)
|
|
assert any("chosen == rejected" in e for e in errors)
|
|
|
|
|
|
def test_validate_catches_missing():
|
|
c = PreferenceCandidate("p1", "p", "a", "b", label_status="labeled", chosen="x", rejected=None)
|
|
assert any("rejected 누락" in e for e in validate_labeled(c))
|
|
|
|
|
|
def test_dataset_stats():
|
|
pending = build_candidate("p1", "o", "a", "b")
|
|
labeled = build_candidate("p2", "o", "a", "b")
|
|
labeled.label_status = "labeled"
|
|
labeled.chosen, labeled.rejected = "a", "b"
|
|
stats = dataset_stats([pending, labeled])
|
|
assert stats["total"] == 2
|
|
assert stats["labeled"] == 1
|
|
assert stats["pending"] == 1
|
|
assert stats["trainable"] == 1
|
|
assert stats["errors"] == []
|
|
|
|
|
|
def test_to_dpo_dataset_filters_pending():
|
|
cs = [build_candidate(f"p{i}", "o", "a", "b") for i in range(3)]
|
|
cs[0].label_status = "labeled"
|
|
cs[0].chosen, cs[0].rejected = "a", "b"
|
|
assert len(to_dpo_dataset(cs)) == 1
|