o2o-plagiarism-ai/tests/test_ai_training_split.py

168 lines
6.4 KiB
Python

"""AI 탐지기 학습 데이터 분리 검증 (#F).
핵심 불변식 두 가지:
1) fit=train / threshold=val / report=test 로 역할이 섞이지 않는다.
2) 어떤 source_group 도 두 split 에 동시에 나타나지 않는다(누출 금지).
"""
from __future__ import annotations
import json
import subprocess
import sys
from pathlib import Path
import pytest
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from scripts.build_ai_training_dataset import Record, assign_splits, summarize # noqa: E402
sklearn = pytest.importorskip("sklearn", reason="scikit-learn 미설치")
pytest.importorskip("joblib", reason="joblib 미설치")
# ---------------------------------------------------------------------------
# 분할 자체의 불변식 (sklearn 불필요 부분)
# ---------------------------------------------------------------------------
def _records(n_groups_per_label: int = 6, per_group: int = 6) -> list[Record]:
out: list[Record] = []
for label in (0, 1):
origin = "human" if label == 0 else "ai"
for g in range(n_groups_per_label):
group = f"{origin}:group-{g}"
for i in range(per_group):
out.append(Record(
text=f"{origin} 문단 {g}-{i}", label=label, origin=origin,
source_group=group, book=group,
))
return out
def test_assign_splits_never_shares_a_group():
records = _records()
assignment = assign_splits(records, (0.7, 0.15, 0.15), seed=7)
for rec in records:
rec.split = assignment[rec.source_group]
by_split: dict[str, set[str]] = {}
for rec in records:
by_split.setdefault(rec.split, set()).add(rec.source_group)
splits = list(by_split)
for i, a in enumerate(splits):
for b in splits[i + 1:]:
assert not (by_split[a] & by_split[b]), f"{a}/{b} 그룹 중복 = 누출"
assert summarize(records)["group_overlap_between_splits"] == []
def test_assign_splits_is_deterministic_for_same_seed():
a = assign_splits(_records(), (0.7, 0.15, 0.15), seed=11)
b = assign_splits(_records(), (0.7, 0.15, 0.15), seed=11)
assert a == b
def test_assign_splits_covers_all_three_splits():
assignment = assign_splits(_records(), (0.7, 0.15, 0.15), seed=3)
assert set(assignment.values()) == {"train", "val", "test"}
# ---------------------------------------------------------------------------
# 학습 CLI end-to-end
# ---------------------------------------------------------------------------
def _write_dataset(path: Path) -> None:
human = "비가 왔다. 나는 그날 학교에 가지 않았고 대신 뒷산에 올라가 온종일 앉아 있었다. 춥지는 않았다. 형이 왔다."
ai = "그날의 기억은 오래도록 남아, 지금까지도 선명하게 떠오르는 장면이 되었다. 아침의 공기는 서늘했고, 발걸음은 조용히 이어졌다."
rows = []
for label, body, origin in ((0, human, "human"), (1, ai, "ai")):
for g in range(6):
split = "train" if g < 4 else ("val" if g == 4 else "test")
for i in range(8):
rows.append({
"text": (body + f" 변형 {g}-{i}. ") * 3,
"label": label, "origin": origin,
"source_group": f"{origin}:g{g}", "book": f"{origin}-{g}",
"split": split,
})
path.write_text(
"\n".join(json.dumps(r, ensure_ascii=False) for r in rows), encoding="utf-8"
)
def _run_trainer(data: Path, out: Path, *extra: str) -> subprocess.CompletedProcess:
return subprocess.run(
[sys.executable, "scripts/train_ai_detector.py", "--data", str(data),
"--out", str(out), *extra],
cwd=ROOT, capture_output=True, text=True,
)
def test_trainer_uses_train_val_test_roles(tmp_path):
data = tmp_path / "ds.jsonl"
out = tmp_path / "model.joblib"
_write_dataset(data)
proc = _run_trainer(data, out)
assert proc.returncode == 0, proc.stderr[-2000:]
metrics = json.loads((tmp_path / "model.metrics.json").read_text(encoding="utf-8"))
# 세 역할이 모두 기록되어야 한다
assert metrics["n_train"] > 0 and metrics["n_val"] > 0 and metrics["n_test"] > 0
assert {"train", "validation", "test"} <= set(metrics)
# 임계값은 val 에서 뽑혔음이 지표에 남아야 한다
assert "validation_low_point" in metrics["cuts"]
assert "validation_high_point" in metrics["cuts"]
# 학습에 쓰인 표본 수와 보고 표본 수가 서로 다른 집합이어야 한다
assert metrics["n_train"] != metrics["n_test"] or metrics["n_val"] != metrics["n_test"]
def test_trainer_rejects_group_overlap_between_splits(tmp_path):
data = tmp_path / "leaky.jsonl"
out = tmp_path / "model.joblib"
_write_dataset(data)
rows = [json.loads(line) for line in data.read_text(encoding="utf-8").splitlines()]
for row in rows:
# train 그룹 하나를 test 에도 등장시켜 누출을 주입
if row["source_group"] == "human:g0" and row["split"] == "train":
row["split"] = "test"
break
data.write_text(
"\n".join(json.dumps(r, ensure_ascii=False) for r in rows), encoding="utf-8"
)
proc = _run_trainer(data, out)
assert proc.returncode == 2, "그룹 누출은 학습을 중단시켜야 한다"
assert "누출" in proc.stderr or "중복" in proc.stderr
def test_trainer_fails_on_single_class(tmp_path):
data = tmp_path / "one.jsonl"
_write_dataset(data)
rows = [json.loads(line) for line in data.read_text(encoding="utf-8").splitlines()]
kept = [r for r in rows if r["label"] == 0]
data.write_text(
"\n".join(json.dumps(r, ensure_ascii=False) for r in kept), encoding="utf-8"
)
proc = _run_trainer(data, tmp_path / "m.joblib")
assert proc.returncode == 2
assert "단일 클래스" in proc.stderr
def test_trainer_fails_when_a_split_is_empty(tmp_path):
data = tmp_path / "noval.jsonl"
_write_dataset(data)
rows = [json.loads(line) for line in data.read_text(encoding="utf-8").splitlines()]
for row in rows:
if row["split"] == "val":
row["split"] = "train"
data.write_text(
"\n".join(json.dumps(r, ensure_ascii=False) for r in rows), encoding="utf-8"
)
proc = _run_trainer(data, tmp_path / "m.joblib")
assert proc.returncode == 2
assert "빈 split" in proc.stderr