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

226 lines
9.1 KiB
Python

"""군집화 기반 표절 판별 (계획서 2단계 표절 검출 기술 고도화, p.21).
목표: 기존 pairwise 삼중 유사도만으로는 "일부 등장인물만 바꾸거나 특정 요소만 바꾼
표절"의 표절률을 수치화하기 어렵다. 계획서는 고도화 전략으로 다음을 명시한다.
> 분류된 요소가 비슷한 텍스트들을 군집화하여 해당 군집 내의 요소 간 유사도를
> 비교하는 방식으로 표절 기술을 고도화. ... 일부 등장 인물만 바꾸거나 특정 요소만
> 바꾼 표절률도 표절여부의 수치화가 가능.
본 모듈은 데이터 없이(현 코퍼스만으로) 동작하는 군집화 1차 라우팅 계층을 제공한다.
설계:
1) 코퍼스 문서를 요소(인물/모티프/키워드) + lemma 시그니처로 묶어 군집 생성
(그래프 connected-components: 요소 자카드 ≥ link_threshold 이면 같은 군집).
2) query 를 가장 가까운 군집으로 라우팅 → 군집 내 문서끼리만 정밀 비교 (탐색량 감소).
3) 군집 내 "요소별 부분 표절 점수" 산출 — 어떤 요소(인물/모티프/키워드/lemma)가
얼마나 겹치는지를 분해해, 인물만 바꾼 표절을 별도 신호로 노출.
순수 함수 위주로 작성되어 단위테스트가 쉽다. detector 에 옵션으로 결합한다.
"""
from __future__ import annotations
from dataclasses import dataclass, field
from app.api.schemas import ExtractedElements
def _jaccard(a: set[str], b: set[str]) -> float:
if not a and not b:
return 0.0
return len(a & b) / max(1, len(a | b))
def _element_signature(elem: ExtractedElements, lemmas: list[str] | None = None) -> dict[str, set[str]]:
"""문서를 요소별 집합 시그니처로 변환 (소문자 정규화)."""
sig = {
"characters": {c.lower() for c in elem.characters},
"motifs": {m.lower() for m in elem.motifs},
"keywords": {k.lower() for k in elem.keywords},
"genre": {elem.genre.lower()} if elem.genre else set(),
}
if lemmas is not None:
sig["lemmas"] = set(lemmas)
return sig
# 군집 링크/요소 표절 점수에 쓰는 요소 가중치 (lemma·키워드 = 본문 차용, 인물·모티프 = 구조 차용)
_SIGNATURE_WEIGHTS = {
"lemmas": 0.40,
"keywords": 0.25,
"characters": 0.15,
"motifs": 0.15,
"genre": 0.05,
}
def signature_similarity(a: dict[str, set[str]], b: dict[str, set[str]]) -> float:
"""두 시그니처의 가중 자카드 결합 (군집 링크 판정용)."""
total_w = 0.0
acc = 0.0
for key, w in _SIGNATURE_WEIGHTS.items():
if key not in a or key not in b:
continue
# 양쪽 모두 비어있는 요소는 정보가 없으므로 제외 (중립)
if not a[key] and not b[key]:
continue
total_w += w
acc += w * _jaccard(a[key], b[key])
return acc / total_w if total_w else 0.0
@dataclass
class Cluster:
cluster_id: int
members: list[str] = field(default_factory=list) # doc_id 리스트
centroid: dict[str, set[str]] = field(default_factory=dict) # 요소별 합집합 시그니처
@dataclass
class PartialPlagiarismSignal:
"""요소별 부분 표절 분해 — '무엇을 그대로 두고 무엇만 바꿨는지'."""
cluster_id: int
per_element: dict[str, float] # characters/motifs/keywords/lemmas/genre 별 자카드
signature_score: float # 가중 결합
changed_elements: list[str] # 거의 안 겹치는(=바꾼) 요소
retained_elements: list[str] # 강하게 겹치는(=유지한) 요소
verdict: str # "element_swap_plagiarism" / "near_duplicate" / "weak" / "none"
class ClusterIndex:
"""코퍼스 요소 시그니처를 군집화하고 query 를 라우팅한다."""
def __init__(
self,
doc_ids: list[str],
doc_elements: list[ExtractedElements],
doc_lemmas: list[list[str]] | None = None,
link_threshold: float = 0.35,
):
if doc_lemmas is not None and len(doc_lemmas) != len(doc_ids):
raise ValueError("doc_lemmas length mismatch")
self.link_threshold = link_threshold
self._doc_ids = doc_ids
self._sigs: dict[str, dict[str, set[str]]] = {
did: _element_signature(elem, doc_lemmas[i] if doc_lemmas else None)
for i, (did, elem) in enumerate(zip(doc_ids, doc_elements))
}
self.clusters: list[Cluster] = self._build_clusters()
self._cluster_of: dict[str, int] = {
did: c.cluster_id for c in self.clusters for did in c.members
}
# ---------- 군집 구성 (그래프 연결요소) ----------
def _build_clusters(self) -> list[Cluster]:
ids = self._doc_ids
parent = {d: d for d in ids}
def find(x: str) -> str:
while parent[x] != x:
parent[x] = parent[parent[x]]
x = parent[x]
return x
def union(x: str, y: str) -> None:
parent[find(x)] = find(y)
for i in range(len(ids)):
for j in range(i + 1, len(ids)):
if signature_similarity(self._sigs[ids[i]], self._sigs[ids[j]]) >= self.link_threshold:
union(ids[i], ids[j])
groups: dict[str, list[str]] = {}
for d in ids:
groups.setdefault(find(d), []).append(d)
clusters: list[Cluster] = []
for cid, members in enumerate(groups.values()):
centroid: dict[str, set[str]] = {}
for m in members:
for key, s in self._sigs[m].items():
centroid.setdefault(key, set()).update(s)
clusters.append(Cluster(cluster_id=cid, members=members, centroid=centroid))
return clusters
def cluster_of(self, doc_id: str) -> int | None:
return self._cluster_of.get(doc_id)
# ---------- query 라우팅 + 부분 표절 분해 ----------
def route(self, query_elem: ExtractedElements, query_lemmas: list[str] | None = None) -> Cluster | None:
"""query 와 가장 유사한 군집 반환 (centroid 가중 자카드 최대)."""
if not self.clusters:
return None
qsig = _element_signature(query_elem, query_lemmas)
return max(self.clusters, key=lambda c: signature_similarity(qsig, c.centroid))
def candidate_ids(self, query_elem: ExtractedElements, query_lemmas: list[str] | None = None) -> set[str]:
"""라우팅된 군집의 멤버 doc_id (정밀 비교 후보 축소용)."""
c = self.route(query_elem, query_lemmas)
return set(c.members) if c else set()
def partial_signal(
self,
doc_id: str,
query_elem: ExtractedElements,
query_lemmas: list[str] | None = None,
retain_threshold: float = 0.6,
change_threshold: float = 0.2,
) -> PartialPlagiarismSignal | None:
"""특정 코퍼스 문서 대비 요소별 부분 표절 신호 분해.
인물만 바꾸고 본문(lemma)·키워드는 유지한 표절을 element_swap_plagiarism 으로 식별.
"""
if doc_id not in self._sigs:
return None
qsig = _element_signature(query_elem, query_lemmas)
dsig = self._sigs[doc_id]
per_element: dict[str, float] = {}
for key in _SIGNATURE_WEIGHTS:
if key in qsig and key in dsig:
# 양쪽 모두 비어있으면 신호 없음 → 제외 (거짓 '변경' 방지)
if not qsig[key] and not dsig[key]:
continue
per_element[key] = round(_jaccard(qsig[key], dsig[key]), 4)
sig_score = signature_similarity(qsig, dsig)
retained = [k for k, v in per_element.items() if v >= retain_threshold]
changed = [k for k, v in per_element.items() if v <= change_threshold]
content_retained = any(k in retained for k in ("lemmas", "keywords"))
structure_changed = any(k in changed for k in ("characters", "motifs"))
high_content = per_element.get("lemmas", 0.0) >= 0.85 or per_element.get("keywords", 0.0) >= 0.85
if content_retained and structure_changed:
verdict = "element_swap_plagiarism" # 본문 유지 + 인물/모티프만 교체 (구조 차용)
elif high_content:
verdict = "near_duplicate" # 본문·구조 모두 유지 = 사실상 복제
elif sig_score >= change_threshold:
verdict = "weak"
else:
verdict = "none"
cid = self._cluster_of.get(doc_id, -1)
return PartialPlagiarismSignal(
cluster_id=cid,
per_element=per_element,
signature_score=round(sig_score, 4),
changed_elements=changed,
retained_elements=retained,
verdict=verdict,
)
@property
def num_clusters(self) -> int:
return len(self.clusters)
def summary(self) -> list[dict]:
"""군집 구성 요약 (디버그/리포트용)."""
return [
{"cluster_id": c.cluster_id, "size": len(c.members), "members": c.members}
for c in sorted(self.clusters, key=lambda x: len(x.members), reverse=True)
]