feat: add isolated gpu embedding review batch
This commit is contained in:
parent
8762fce51e
commit
8717fa4425
8
Dockerfile.gpu-embed
Normal file
8
Dockerfile.gpu-embed
Normal file
@ -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
|
||||||
79
scripts/run_embedding_infringement_batch.py
Normal file
79
scripts/run_embedding_infringement_batch.py
Normal file
@ -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()
|
||||||
Loading…
Reference in New Issue
Block a user