diff --git a/Dockerfile.gpu-embed b/Dockerfile.gpu-embed new file mode 100644 index 0000000..a1f69e0 --- /dev/null +++ b/Dockerfile.gpu-embed @@ -0,0 +1,8 @@ +# 운영 API 이미지와 분리된 2차 KoSimCSE 배치 전용 이미지. +FROM pytorch/pytorch:2.3.1-cuda12.1-cudnn8-runtime +WORKDIR /app +COPY requirements.txt . +RUN pip install --no-cache-dir -r requirements.txt sentence-transformers==3.0.1 +COPY app ./app +COPY scripts ./scripts +COPY data ./data diff --git a/scripts/run_embedding_infringement_batch.py b/scripts/run_embedding_infringement_batch.py new file mode 100644 index 0000000..55fc904 --- /dev/null +++ b/scripts/run_embedding_infringement_batch.py @@ -0,0 +1,79 @@ +#!/usr/bin/env python3 +"""GPU KoSimCSE 2차 침해 검토: 별도 임베딩 인덱스만 생성한다.""" +from __future__ import annotations +import argparse, json, logging, os, sys, time +from collections import Counter +from pathlib import Path +import numpy as np + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from app.api.schemas import DetectOptions +from app.core.config import get_settings +from app.engine.detector import PlagiarismDetector +from app.engine.persistent_index import _evidence_spans, _covered_length +from app.engine.provenance import CorpusStore +from app.engine.similarity import SimilarityHit, _element_similarities +from app.engine.structural import extract_lemmas, lemma_overlap_ratio +from scripts.run_infringement_batch import build_query_plan, build_xlsx, result_row + +LOG = logging.getLogger("embedding_batch") + +def build_or_load(store, index_dir: Path, model_name: str, device: str, batch: int): + """운영 hash 인덱스와 독립된 KoSimCSE 행렬을 생성한다.""" + from sentence_transformers import SentenceTransformer + index_dir.mkdir(parents=True, exist_ok=True) + meta_path, vec_path = index_dir / "index.json", index_dir / "embeddings.npz" + segments = list(store.iter_segments()) + ids = [s.segment_id for s in segments] + hashes = {s.segment_id: s.text_sha256 for s in segments} + if meta_path.exists() and vec_path.exists(): + meta = json.loads(meta_path.read_text()) + if meta.get("segment_ids") == ids and meta.get("text_hashes") == hashes and meta.get("model") == model_name: + return np.load(vec_path)["embeddings"], meta, SentenceTransformer(model_name, device=device) + model = SentenceTransformer(model_name, device=device) + LOG.info("gpu_index_encode segments=%d model=%s device=%s", len(segments), model_name, device) + vectors = model.encode([s.text[:2048] for s in segments], batch_size=batch, normalize_embeddings=True, show_progress_bar=False, convert_to_numpy=True) + vectors = np.asarray(vectors, dtype=np.float32) + np.savez_compressed(vec_path, embeddings=vectors) + meta = {"version": 1, "backend": "KoSimCSE", "model": model_name, "device": device, "segment_ids": ids, "segment_document_ids": [s.document_id for s in segments], "document_source_groups": store.document_source_groups(), "text_hashes": hashes} + meta_path.write_text(json.dumps(meta, ensure_ascii=False), encoding="utf-8") + return vectors, meta, model + +def main(): + p=argparse.ArgumentParser(); p.add_argument("--database",type=Path,required=True); p.add_argument("--index-dir",type=Path,required=True); p.add_argument("--out-jsonl",type=Path,required=True); p.add_argument("--out-xlsx",type=Path,required=True); p.add_argument("--first-jsonl",type=Path,required=True); p.add_argument("--model",default="BM-K/KoSimCSE-roberta-multitask"); p.add_argument("--batch-size",type=int,default=64); p.add_argument("--progress-every",type=int,default=25); a=p.parse_args() + logging.basicConfig(level=logging.INFO,format="%(asctime)s %(levelname)s %(message)s") + if os.getenv("USE_LLM_LEGAL_JUDGE", "").lower() not in ("", "0", "false", "no", "off"): raise RuntimeError("USE_LLM_LEGAL_JUDGE=false required") + import torch + if not torch.cuda.is_available(): raise RuntimeError("CUDA unavailable") + device="cuda:0"; LOG.info("gpu_ready device=%s name=%s",device,torch.cuda.get_device_name(0)) + store=CorpusStore(a.database); vectors,meta,model=build_or_load(store,a.index_dir,a.model,device,a.batch_size) + get_settings.cache_clear(); settings=get_settings().model_copy(update={"corpus_db_path":str(a.database),"persistent_index_dir":"/nonexistent","use_persistent_index":False,"use_llm_legal_judge":False}) + # 법령/태그/규칙 엔진만 재사용한다. CPU 운영 인덱스를 변경하지 않는다. + detector=PlagiarismDetector(settings); groups=store.document_source_groups(); records=store.get_segments(meta["segment_ids"]); items=build_query_plan(store) + a.out_jsonl.parent.mkdir(parents=True,exist_ok=True); started=time.monotonic(); rows=[] + for n,item in enumerate(items,1): + qv=np.asarray(model.encode([item.text[:2048]],normalize_embeddings=True,show_progress_bar=False,convert_to_numpy=True)[0],dtype=np.float32); scores=vectors@qv + allowed=np.array([d != item.document_id and (not item.source_group or groups.get(d,"") != item.source_group) for d in meta["segment_document_ids"]]) + ix=np.flatnonzero(allowed); take=ix[np.argsort(scores[ix])[-20:][::-1]] + qe=detector._extractor.extract(item.text); ql=extract_lemmas(item.text); hits=[]; intervals=[]; doc_intervals={} + for i in take: + seg=records[meta["segment_ids"][int(i)]]; ev,ints,longest=_evidence_spans(item.text,seg.text); intervals.extend(ints); doc_intervals.setdefault(seg.document_id,[]).extend(ints); relem, reels=(seg.lemmas or extract_lemmas(seg.text)), (seg.elements or detector._extractor.extract(seg.text).model_dump()); es=_element_similarities(qe, type(qe)(**reels)); combined=.30*float(scores[int(i)])+.45*lemma_overlap_ratio(ql,relem)+.15*es["characters"]+.10*es["motifs"] + hits.append((combined,seg,float(scores[int(i)]),lemma_overlap_ratio(ql,relem),es,ev,longest)) + hits.sort(key=lambda x:x[0],reverse=True); matches=[] + for combined,seg,textsim,lemmasim,es,ev,longest in hits[:5]: + coverage=_covered_length(doc_intervals.get(seg.document_id,[]))/max(1,len(item.text)); + if combined < .65 and longest < 80 and coverage < .30: continue + sh=SimilarityHit(seg.segment_id,str(seg.metadata.get("document_title",seg.document_id)),combined,textsim,lemmasim,es,[]) + match=detector._to_match(sh,True,None,None).model_copy(update={"source_document_id":seg.document_id,"source_segment_id":seg.segment_id,"matched_coverage":round(coverage,4),"longest_span":longest,"evidence_spans":ev,"match_reasons":["embedding_score" if combined>=.65 else "exact_span_or_coverage"]}) + matches.append(match) + coverage=_covered_length(intervals)/max(1,len(item.text)); legal=detector._legal_engine.assess(max_similarity=matches[0].similarity if matches else 0,coverage=coverage,longest_span=max((m.longest_span for m in matches),default=0),legal_tags=[t.tag for m in matches for t in m.tags],query_text=item.text,evidence=[]) + class R: pass + r=R(); r.matches=matches; r.is_infringement=bool(matches); r.confidence=round(matches[0].similarity if matches else (hits[0][0] if hits else 0),4); r.legal_risk=legal; r.score_semantics=type("S",(),{"union_coverage":round(coverage,4)})(); rows.append(result_row(item,r,retrieval_backend=f"KoSimCSE GPU ({a.model}, {device})")) + if n%a.progress_every==0: LOG.info("progress processed=%d/%d",n,len(items)) + a.out_jsonl.write_text("\n".join(json.dumps(r,ensure_ascii=False) for r in rows)+"\n",encoding="utf-8") + build_xlsx(rows,a.out_xlsx,backend=f"KoSimCSE GPU ({a.model}, {device})") + from openpyxl import load_workbook + first={json.loads(x)["query_key"]:json.loads(x) for x in a.first_jsonl.read_text(encoding="utf-8").splitlines() if x.strip()}; second={r["query_key"]:r for r in rows}; f={k for k,v in first.items() if v["침해 의심"]}; s={k for k,v in second.items() if v["침해 의심"]}; wb=load_workbook(a.out_xlsx); ws=wb.create_sheet("1차_vs_2차 비교"); ws.append(["구분","건수"]); ws.append(["1차에서만",len(f-s)]); ws.append(["2차 신규",len(s-f)]); ws.append(["양쪽 공통",len(f&s)]); ws.append([]); ws.append(["2차 신규 상위 20", "가명 제목", "결합유사도", "매칭 상대 doc"]) + for r in sorted((second[k] for k in s-f),key=lambda x:x["결합유사도"],reverse=True)[:20]: ws.append([r["doc_id"],r["가명 제목"],r["결합유사도"],r["매칭 상대 doc"]]) + wb.save(a.out_xlsx); LOG.info("completed total=%d suspected=%d first_only=%d second_only=%d common=%d",len(rows),len(s),len(f-s),len(s-f),len(f&s)) +if __name__=="__main__": main()