CoolFace
Modelpublic

HebArabNlpProject/Semantic-Retrieval-2nd-place

sourceHugging Faceotherupdated 5mo agoView on Hugging Face
0likes
model.py698 linesDownload Raw Back to root
1import os, re, math, unicodedata, time, json, hashlib, importlib.util
2from collections import defaultdict, Counter
3from typing import List, Tuple, Dict, Optional
4import numpy as np
5import torch
6from transformers import AutoTokenizer, AutoModel, AutoModelForSequenceClassification
7from sentence_transformers import SentenceTransformer
8import sys, pathlib
9
10HERE = pathlib.Path(__file__).resolve().parent
11if str(HERE) not in sys.path:
12    sys.path.insert(0, str(HERE))
13
14# ======================= Tunables =======================
15BM25_K1   = 1.3
16BM25_B    = 0.7
17RRF_K     = 35    # RRF constant
18CE_MAXLEN = 640
19CE_BATCH  = 128
20
21TOP_BM25 = TOP_E5 = TOP_GEMMA = CE_POOL = 190
22
23# Weighted RRF stage-1 fusion (BM25 + E5 + Gemma)
24WRRF_BM25_W = 1.0
25WRRF_E5_W   = 1.2
26WRRF_GEMMA_W= 1.4
27
28# Weighting for the final (reranker) fusion
29FINAL_SCORE_BGE_WEIGHT = .07
30
31# Model & cache dirs
32USE_CACHE = True
33
34BGE_DIR              = r"models/bge-reranker-hsrc-pairwise-rrf-V1.4".strip()
35
36E5_DIR               = r"models/e5-large-ft_v6".strip()
37E5_EVAL_CACHE_DIR    = r"".strip() if USE_CACHE else r""
38
39GEMMA_DIR            = r"models/multilingual-e5-large".strip()
40GEMMA_EVAL_CACHE_DIR = r"".strip() if USE_CACHE else r""
41
42# Gemma dtype & max length (SentenceTransformers truncation)
43PREFER_BF16_GEMMA = True
44GEMMA_MAX_TOK     = 512
45# GEMMA_QUERY_TASK  = "search result"   # used only when we must fall back to prompt=...
46
47# ======================= Silence under eval_std ============================
48_EVAL_SILENT = os.environ.get("EVAL_STD_MODE","").strip() == "1"
49def _log(msg: str):
50    if not _EVAL_SILENT:
51        print(msg, flush=True)
52
53# ======================= Normalization / Tokenization =======================
54# Priority: 1) Relative import, 2) sys.path, 3) Dynamic import, 4) Fallback
55try:
56    from .text_utils import (  # type: ignore
57        tok_he, norm_bm25,
58        norm_e5_query, norm_e5_passage,
59        norm_gemma_query, norm_gemma_passage,
60        norm_bge_query, norm_bge_passage
61    )
62    _log("[Init] Loaded text_utils (relative import).")
63except (ImportError, ModuleNotFoundError):
64    try:
65        from text_utils import (
66            tok_he, norm_bm25,
67            norm_e5_query, norm_e5_passage,
68            norm_gemma_query, norm_gemma_passage,
69            norm_bge_query, norm_bge_passage
70        )
71        _log("[Init] Loaded text_utils (sys.path import).")
72    except (ImportError, ModuleNotFoundError):
73        try:
74            spec_path = HERE / "text_utils.py"
75            if not spec_path.is_file():
76                raise FileNotFoundError(f"{spec_path} not found.")
77            spec = importlib.util.spec_from_file_location("text_utils", spec_path)
78            text_utils_module = importlib.util.module_from_spec(spec)
79            spec.loader.exec_module(text_utils_module)
80            tok_he = text_utils_module.tok_he
81            norm_bm25 = text_utils_module.norm_bm25
82            norm_e5_query = text_utils_module.norm_e5_query
83            norm_e5_passage = text_utils_module.norm_e5_passage
84            # Gemma-specific normalizers if present; fallback to e5 norms
85            norm_gemma_query = getattr(text_utils_module, "norm_gemma_query", text_utils_module.norm_e5_query)
86            norm_gemma_passage = getattr(text_utils_module, "norm_gemma_passage", text_utils_module.norm_e5_passage)
87            norm_bge_query = text_utils_module.norm_bge_query
88            norm_bge_passage = text_utils_module.norm_bge_passage
89            _log("[Init] Loaded text_utils (dynamic import).")
90        except Exception:
91            _log("[Init] `text_utils` not found. Using generic Hebrew-friendly normalizers for all components.")
92            HEB_PREFIXES = ("ו","ה","ב","ל","כ","מ","ש")
93            STOPWORDS = {}
94            def _generic_norm_he(s:str)->str:
95                if not s: return ""
96                s=unicodedata.normalize("NFKC",s)
97                s=re.sub(r"[\u0591-\u05BD\u05BF-\u05C7]","",s)
98                s=(s.replace("״",'"').replace("׳","'")
99                     .replace("”",'"').replace("“",'"')
100                     .replace("–","-").replace("—","-"))
101                return re.sub(r"\s+"," ",s).strip()
102            norm_bm25 = norm_e5_query = norm_e5_passage = _generic_norm_he
103            norm_gemma_query = norm_gemma_passage = _generic_norm_he
104            norm_bge_query = norm_bge_passage = _generic_norm_he
105            def tok_he(text: str) -> List[str]:
106                s = norm_bm25(text)
107                toks = re.findall(r"[A-Za-z0-9\u0590-\u05FF]+", s)
108                out=[]
109                for t in toks:
110                    if len(t)>3 and t[0] in HEB_PREFIXES: out.append(t[1:])
111                    out.append(t)
112                return [t for t in out if t not in STOPWORDS]
113
114# =========================== BM25 Backends ================================
115get_bm25_backend = None
116_HAS_BM25_BACKENDS = False
117try:
118    from .bm25_backends import get_bm25_backend
119    _HAS_BM25_BACKENDS = True
120    _log("[Init] Loaded bm25_backends (relative import).")
121except (ImportError, ModuleNotFoundError):
122    try:
123        from bm25_backends import get_bm25_backend
124        _HAS_BM25_BACKENDS = True
125        _log("[Init] Loaded bm25_backends (sys.path import).")
126    except (ImportError, ModuleNotFoundError):
127        try:
128            spec_path = HERE / "bm25_backends.py"
129            if not spec_path.is_file():
130                 raise FileNotFoundError(f"{spec_path} not found.")
131            spec = importlib.util.spec_from_file_location("bm25_backends", spec_path)
132            bm25_module = importlib.util.module_from_spec(spec)
133            spec.loader.exec_module(bm25_module)
134            get_bm25_backend = bm25_module.get_bm25_backend
135            _HAS_BM25_BACKENDS = True
136            _log("[Init] Loaded bm25_backends (dynamic import).")
137        except Exception as e:
138            _log(f"[Init] Could not load bm25_backends.py ({e}). Will use built-in fallbacks.")
139            pass
140
141class _LocalBM25SBackend:
142    """Minimal local wrapper for bm25s when bm25_backends.py is missing."""
143    def __init__(self, tokenizer, k1: float = 1.3, b: float = 0.7, logger=_log):
144        import bm25s
145        self._BM25 = bm25s.BM25
146        self.tokenizer = tokenizer
147        self.k1, self.b = k1, b
148        self._bm25 = None
149        self.doc_ids: List[str] = []
150        self._logger = logger
151    @property
152    def name(self)->str:
153        return f"LocalBM25S(k1={self.k1}, b={self.b})"
154    def build(self, ids: List[str], texts: List[str]):
155        t0=time.time()
156        self.doc_ids = list(ids)
157        tokenized = [self.tokenizer(t) for t in texts]
158        self._bm25 = self._BM25(k1=self.k1, b=self.b)
159        self._bm25.index(tokenized)
160        if self._logger: self._logger(f"[{self.name}] Indexed {len(self.doc_ids):,} docs in {time.time()-t0:.2f}s")
161    def search(self, query: str, topk: int = 300) -> List[str]:
162        terms = self.tokenizer(query)
163        if not terms or self._bm25 is None: return []
164        k = min(topk, len(self.doc_ids))
165        idxs, scores = self._bm25.retrieve([terms], k=k)
166        idxs, scores = idxs[0], scores[0]
167        mask = np.isfinite(scores) & (scores > 0)
168        idxs, scores = idxs[mask], scores[mask]
169        if idxs.size == 0: return []
170        order = np.lexsort((idxs, -scores))
171        idxs = idxs[order]
172        return [self.doc_ids[int(i)] for i in idxs]
173
174class _DeterministicBM25Backend:
175    """Embedded pure-Python deterministic BM25. Guaranteed fallback."""
176    def __init__(self, tokenizer, k1: float = 1.3, b: float = 0.7, logger=_log):
177        self.tokenizer=tokenizer; self.k1=k1; self.b=b
178        self.doc_ids: List[str]=[]; self.N=0; self.avgdl=0.0
179        self.doc_lens=None; self.vocab: Dict[str,int]={}
180        self.postings: Dict[int,Tuple[np.ndarray,np.ndarray]]={}
181        self.idf=None; self._logger=logger
182    @property
183    def name(self)->str:
184        return f"DeterministicBM25(k1={self.k1}, b={self.b})"
185    def build(self, ids: List[str], texts: List[str]):
186        self.doc_ids=list(ids); self.N=len(ids)
187        lens=np.zeros(self.N,dtype=np.int32)
188        tmp=defaultdict(list)
189        t0=time.time()
190        for i, text in enumerate(texts):
191            terms=self.tokenizer(text); lens[i]=len(terms)
192            if not terms: continue
193            ctr=Counter(terms)
194            for t,tf in ctr.items():
195                tid=self.vocab.setdefault(t, len(self.vocab))
196                tmp[tid].append((i, tf))
197        self.doc_lens=lens; self.avgdl=float(np.maximum(1,lens).mean())
198        V=len(self.vocab); self.idf=np.zeros(V,dtype=np.float32)
199        self.postings={}
200        for tid, pairs in tmp.items():
201            docs=np.array([d for d,_ in pairs],dtype=np.int32)
202            tfs =np.array([tf for _,tf in pairs],dtype=np.float32)
203            df=float(len(docs))
204            idf=math.log((self.N-df+0.5)/(df+0.5)+1.0)
205            self.idf[tid]=idf
206            self.postings[tid]=(docs,tfs)
207        if self._logger: self._logger(f"[{self.name}] Indexed {self.N:,} docs in {time.time()-t0:.2f}s")
208    def search(self, query: str, topk: int = 300) -> List[str]:
209        terms=self.tokenizer(query)
210        if not terms: return []
211        seen: Dict[int,float] = {}
212        for t in terms:
213            tid=self.vocab.get(t)
214            if tid is None: continue
215            idf=float(self.idf[tid])
216            docs,tfs=self.postings[tid]
217            denom=tfs + self.k1*(1-self.b + self.b*(self.doc_lens[docs]/self.avgdl))
218            contrib = idf * (tfs*(self.k1+1)) / denom
219            for d, c in zip(docs, contrib):
220                seen[d]=seen.get(d,0.0)+float(c)
221        if not seen: return []
222        idx=np.fromiter(seen.keys(),dtype=np.int32)
223        scs=np.fromiter(seen.values(),dtype=np.float32)
224        k=min(topk,len(scs))
225        order = np.lexsort((idx, -scs))
226        order = order[:k]
227        idx = idx[order]
228        return [self.doc_ids[i] for i in idx]
229
230class BM25Index:
231    """Unified BM25 wrapper. Returns List[str] of doc IDs."""
232    def __init__(self, k1=1.3, b=0.70, logger=_log):
233        self.k1, self.b = k1, b
234        self.doc_ids: List[str] = []
235        self._be = None; self._backend_name = "unset"; self._logger = logger
236    def build(self, ids: List[str], texts_norm: List[str]):
237        if _HAS_BM25_BACKENDS and callable(get_bm25_backend):
238            try:
239                self._be = get_bm25_backend(use_bm25s=True, tokenizer=tok_he, k1=self.k1, b=self.b, logger=self._logger)
240                self._be.build(ids, texts_norm)
241                self.doc_ids = list(self._be.doc_ids)
242                self._backend_name = f"{self._be.name} (bm25_backends.py)"
243                if self._logger: self._logger(f"[BM25] Using backend: {self._backend_name}")
244                return
245            except Exception as e:
246                if self._logger: self._logger(f"[BM25] bm25_backends failed ({e}). Trying direct bm25s...)")
247        try:
248            self._be = _LocalBM25SBackend(tok_he, k1=self.k1, b=self.b, logger=self._logger)
249            self._be.build(ids, texts_norm)
250            self.doc_ids = list(self._be.doc_ids)
251            self._backend_name = f"{self._be.name} (direct)"
252            if self._logger: self._logger(f"[BM25] Using backend: {self._backend_name}")
253            return
254        except Exception as e:
255            if self._logger: self._logger(f"[BM25] bm25s unavailable ({e}). Falling back to pure-Python).")
256        self._be = _DeterministicBM25Backend(tok_he, k1=self.k1, b=self.b, logger=self._logger)
257        self._be.build(ids, texts_norm)
258        self.doc_ids = list(self._be.doc_ids)
259        self._backend_name = f"{self._be.name} (embedded)"
260        if self._logger: self._logger(f"[BM25] Using backend: {self._backend_name}")
261    def search(self, query: str, topk: int = 200) -> List[str]:
262        if self._be is None: return []
263        return self._be.search(query, topk=topk)
264
265# ======================= Model Path Resolution =======================
266def _resolve_model_path(primary_path: str, fallback_names: List[str]) -> str:
267    """
268    Resolves a model path: checks primary_path, then HERE/models, HERE, CWD, CWD/models.
269    Falls back to first fallback name (HF id/path).
270    """
271    if primary_path and pathlib.Path(primary_path).is_dir():
272        return primary_path
273    base_dirs = [HERE / "models", HERE, pathlib.Path.cwd(), pathlib.Path.cwd() / "models"]
274    for base in base_dirs:
275        for name in fallback_names:
276            candidate = base / name
277            if candidate.is_dir():
278                return str(candidate)
279    return fallback_names[0]
280
281def model_name_key(s: str) -> str:
282    if not s:
283        return ""
284    s = s.strip().rstrip("/\\")
285    last = re.split(r"[\\/]+", s)[-1] or s
286    return last.lower()
287
288# ======================= E5 embedder =============================
289class E5Embedder:
290    def __init__(self, device=None):
291        fallback_names = ["e5-large-ft_v4","multilingual-e5-large"]
292        all_fallbacks = [pathlib.Path(E5_DIR).name] + fallback_names if E5_DIR else fallback_names
293        self.model_path = _resolve_model_path(E5_DIR, all_fallbacks)
294        self.model_name = model_name_key(self.model_path)
295        self.device=device or ("cuda" if torch.cuda.is_available() else "cpu")
296        _log(f"[E5] Loading encoder from: {self.model_path}  (device={self.device})")
297        self.tok=AutoTokenizer.from_pretrained(self.model_path)
298        self.mdl=AutoModel.from_pretrained(self.model_path, torch_dtype=torch.bfloat16 if self.device=="cuda" else None).to(self.device) # changed dtype to bf16
299        self.mdl.eval()
300    @torch.inference_mode()
301    def encode(self, texts: List[str], is_query=False, batch=64, progress_desc="E5 encode"):
302        # Expects already-normalized texts
303        pref="query: " if is_query else "passage: "
304        # pref="" if is_query else ""
305        out=[]
306        n=len(texts)
307        if n==0: return np.zeros((0,768), dtype=np.float32)
308        total_batches = (n + batch - 1)//batch
309        t0=time.time()
310        for bi in range(total_batches):
311            i = bi*batch
312            chunk = texts[i:i+batch]
313            enc=self.tok([pref+t.strip() for t in chunk], padding=True, truncation=True, max_length=512, return_tensors="pt").to(self.device)
314            hs=self.mdl(**enc).last_hidden_state
315            mask=enc["attention_mask"].unsqueeze(-1).expand(hs.size()).float()
316            embs=(hs*mask).sum(1)/mask.sum(1).clamp(min=1e-9)
317            embs=torch.nn.functional.normalize(embs, p=2, dim=1)
318            out.append(embs.detach().cpu().to(dtype=torch.float32))
319            if not _EVAL_SILENT:
320                if (bi+1)%50==0 or bi==0 or (bi+1)==total_batches:
321                    pct = 100.0*(bi+1)/total_batches
322                    elapsed = time.time()-t0
323                    ips = (i+len(chunk))/max(elapsed,1e-6)
324                    print(f"[{progress_desc}] batch {bi+1}/{total_batches}  ({pct:.1f}%)  ~{ips:.0f} items/s")
325            del enc, hs, embs
326            if torch.cuda.is_available(): torch.cuda.empty_cache()
327        return torch.cat(out, dim=0).numpy()
328
329# ======================= EmbeddingGemma embedder =====================
330class GemmaEmbedder:
331    """
332    Uses SentenceTransformer('google/embeddinggemma-300m'), BF16 if available.
333    Returns L2-normalized 768-dim numpy arrays.
334    No manual prompt prefixing; let SentenceTransformers handle prompting.
335    """
336    def __init__(self, device=None):
337        fallback_names = ["google/embeddinggemma-300m","embeddinggemma-300m"]
338        all_fallbacks = [pathlib.Path(GEMMA_DIR).name] + fallback_names if GEMMA_DIR else fallback_names
339        self.model_path = _resolve_model_path(GEMMA_DIR, all_fallbacks)
340        self.model_name = model_name_key(self.model_path)
341        self.device = device or ("cuda" if torch.cuda.is_available() else "cpu")
342        # dtype selection
343        use_bf16 = bool(PREFER_BF16_GEMMA)
344        if self.device == "cuda":
345            try:
346                use_bf16 = use_bf16 and torch.cuda.is_bf16_supported()
347            except Exception:
348                major, _ = torch.cuda.get_device_capability()
349                use_bf16 = use_bf16 and (major >= 8)
350        dtype = torch.bfloat16 if use_bf16 else torch.float16
351        _log(f"[Gemma] Loading encoder from: {self.model_path}  (device={self.device}, dtype={'bf16' if use_bf16 else 'fp16'})")
352        self.mdl = SentenceTransformer(
353            self.model_path,
354            device=self.device,
355            model_kwargs={"torch_dtype": dtype},
356        )
357        # Tunable max tokens
358        try:
359            self.mdl.max_seq_length = int(GEMMA_MAX_TOK)
360        except Exception:
361            pass
362        self.dim = 768
363        self.mdl.eval()
364
365    @torch.inference_mode()
366    def encode(self, texts: List[str], is_query=False, batch=64, progress_desc="Gemma encode", max_length: Optional[int]=None):
367        if not texts:
368            return np.zeros((0, self.dim), dtype=np.float32)
369
370        # Per-call max length override
371        old_len = getattr(self.mdl, "max_seq_length", None)
372        if isinstance(max_length, int) and max_length > 0:
373            try:
374                self.mdl.max_seq_length = max_length
375            except Exception:
376                pass
377
378        show = not _EVAL_SILENT
379
380        # DO NOT manually add prompts. Prefer encode_query / encode_document when available.
381        try:
382            if is_query and hasattr(self.mdl, "encode_query"):
383                embs = self.mdl.encode_query(
384                    texts, batch_size=batch, convert_to_numpy=True,
385                    normalize_embeddings=True, show_progress_bar=show
386                )
387            elif (not is_query) and hasattr(self.mdl, "encode_document"):
388                embs = self.mdl.encode_document(
389                    texts, batch_size=batch, convert_to_numpy=True,
390                    normalize_embeddings=True, show_progress_bar=show
391                )
392            else:
393                # Fallback: use encode with prompt=... if supported (avoids manual concatenation)
394                prompt = (f"{'query: ' if is_query else 'passage: '}")
395                try:
396                    embs = self.mdl.encode(
397                        texts, batch_size=batch, convert_to_numpy=True,
398                        normalize_embeddings=True, show_progress_bar=show,
399                        prompt=prompt
400                    )
401                except TypeError:
402                    # Last resort: plain encode (no prompt)
403                    embs = self.mdl.encode(
404                        texts, batch_size=batch, convert_to_numpy=True,
405                        normalize_embeddings=True, show_progress_bar=show
406                    )
407        finally:
408            if old_len is not None:
409                try: self.mdl.max_seq_length = old_len
410                except Exception: pass
411
412        embs = np.asarray(embs)
413        if embs.ndim == 1:
414            embs = embs[None, :]
415        return embs.astype(np.float32)
416
417# ======================= BGE reranker ============================
418class BGEReranker:
419    def __init__(self, device=None):
420        fallback_names = ["bge-reranker-hsrc-pairwise-rrf-V1.4","bge-v2-m3","bge-m3"]
421        all_fallbacks = [pathlib.Path(BGE_DIR).name] + fallback_names if BGE_DIR else fallback_names
422        self.model_path = _resolve_model_path(BGE_DIR, all_fallbacks)
423        self.device=device or ("cuda" if torch.cuda.is_available() else "cpu")
424        _log(f"[BGE] Loading reranker from: {self.model_path}  (device={self.device})")
425        self.tok=AutoTokenizer.from_pretrained(self.model_path)
426        self.mdl=AutoModelForSequenceClassification.from_pretrained(
427            self.model_path, torch_dtype=torch.float16 if self.device=="cuda" else None, trust_remote_code=True
428        ).to(self.device)
429        self.mdl.eval()
430    @torch.inference_mode()
431    def score_pairs(self, q: str, passages: List[str], batch=32, max_len=512) -> List[float]:
432        out=[]
433        for i in range(0,len(passages), batch):
434            enc=self.tok([q]*len(passages[i:i+batch]), passages[i:i+batch],
435                         truncation="only_second", max_length=max_len, padding=True, return_tensors="pt").to(self.device)
436            logits=self.mdl(**enc).logits
437            if logits.ndim==1: s=logits
438            elif logits.shape[1]==1: s=logits.squeeze(-1)
439            else: s=logits[:,1]
440            out += s.detach().float().cpu().tolist()
441            del enc, logits
442        return [float(x) for x in out]
443
444# ======================== Hybrid Searcher ========================
445class HybridSearcher:
446    """
447    Stage-1 retrieval: WRRF(BM25, E5, Gemma) → candidate ids + WRRF scores.
448    Then stage-2 reranking is done outside in predict().
449    """
450    def __init__(self, bm25: BM25Index,
451                 e5: E5Embedder, e5_corpus: np.ndarray,
452                 gemma: GemmaEmbedder, gemma_corpus: np.ndarray,
453                 id2text: Dict[str,str], id2norm: Dict[str,str]):
454        self.bm25=bm25
455        self.e5=e5; self.e5_corpus=e5_corpus
456        self.gemma=gemma; self.gemma_corpus=gemma_corpus
457        self.id2text=id2text; self.id2norm=id2norm
458        self._last_q: Optional[str] = None
459        self._last_fused: List[Tuple[str, float]] = []
460
461    def _wrrf_fuse3(self, bm_ids: List[str], e5_ids: List[str], gm_ids: List[str], k=60,
462                    w_bm25=1.0, w_e5=1.0, w_gm=1.0) -> List[Tuple[str, float]]:
463        rankA={pid:i for i,pid in enumerate(bm_ids)}
464        rankB={pid:i for i,pid in enumerate(e5_ids)}
465        rankC={pid:i for i,pid in enumerate(gm_ids)}
466        scores=defaultdict(float)
467        for pid, r in rankA.items(): scores[pid]+=w_bm25*(1.0/(k+r+1))
468        for pid, r in rankB.items(): scores[pid]+=w_e5  *(1.0/(k+r+1))
469        for pid, r in rankC.items(): scores[pid]+=w_gm  *(1.0/(k+r+1))
470        return sorted(scores.items(), key=lambda x:-x[1])
471
472    def search(self, query: str, topk: int=200) -> List[Tuple[str, float]]:
473        if self._last_q == query and self._last_fused:
474            return self._last_fused[:topk]
475
476        # BM25 list
477        bm_ids = self.bm25.search(query, topk=TOP_BM25)
478
479        # E5 list
480        q_norm_e5 = norm_e5_query(query)          # per-query normalization
481        qe = self.e5.encode([q_norm_e5], is_query=True, batch=1, progress_desc="E5 query")[0]
482        sims_e5 = (self.e5_corpus @ qe)           # cosine (embeddings are L2-normalized)
483        k2 = min(TOP_E5, len(sims_e5))
484        top_idx_e5 = np.argpartition(-sims_e5, k2-1)[:k2]
485        top_idx_e5 = top_idx_e5[np.argsort(-sims_e5[top_idx_e5])]
486        e5_ids = [self.bm25.doc_ids[i] for i in top_idx_e5]
487
488        # Gemma list
489        q_norm_gm = norm_gemma_query(query)       # per-query normalization
490        qg = self.gemma.encode([q_norm_gm], is_query=True, batch=1, progress_desc="Gemma query", max_length=GEMMA_MAX_TOK)[0]
491        sims_gm = (self.gemma_corpus @ qg)        # cosine (normalized)
492        k3 = min(TOP_GEMMA, len(sims_gm))
493        top_idx_gm = np.argpartition(-sims_gm, k3-1)[:k3]
494        top_idx_gm = top_idx_gm[np.argsort(-sims_gm[top_idx_gm])]
495        gm_ids = [self.bm25.doc_ids[i] for i in top_idx_gm]
496
497        fused_with_scores = self._wrrf_fuse3(
498            bm_ids, e5_ids, gm_ids, k=RRF_K,
499            w_bm25=WRRF_BM25_W, w_e5=WRRF_E5_W, w_gm=WRRF_GEMMA_W
500        )
501        # seen=set(); out=[]
502        # for pid, score in fused_with_scores:
503        #     key=self.id2norm.get(pid,"")
504        #     if key in seen: continue
505        #     seen.add(key)
506        #     out.append((pid, score))
507        #     if len(out)>=topk: break
508        out = fused_with_scores[:topk]
509        
510        self._last_q = query
511        self._last_fused = out[:]
512        return out
513
514# =========================== Globals ===========================
515_STATE = {}
516
517# =========================== Helpers ===========================
518def _sha1_ids(ids: List[str]) -> str:
519    h = hashlib.sha1()
520    for pid in ids:
521        h.update(pid.encode("utf-8")); h.update(b"\n")
522    return h.hexdigest()
523
524def _normalize_min_max(scores: List[float]) -> List[float]:
525    """Scales a list of scores to the [0, 1] range."""
526    if not scores or len(scores) < 2:
527        return [0.5] * len(scores)
528    min_s, max_s = min(scores), max(scores)
529    delta = max_s - min_s
530    if delta < 1e-9:
531        return [0.5] * len(scores)
532    return [(s - min_s) / delta for s in scores]
533
534# =========================== API funcs =========================
535def preprocess(corpus_dict: Dict[str, Dict]) -> Dict:
536    ids, texts = [], []
537    bm25_norms = []
538
539    # -------- Per-paragraph normalization before indexing --------
540    e5_passage_norms = []
541    gm_passage_norms = []
542
543    for pid,obj in corpus_dict.items():
544        t = obj.get("passage") or obj.get("text") or ""
545        pid = str(pid)
546        ids.append(pid)
547        texts.append(t)
548        bm25_norms.append(norm_bm25(t))         # BM25 per paragraph
549        e5_passage_norms.append(norm_e5_passage(t))
550        gm_passage_norms.append(norm_gemma_passage(t))
551
552    _log("="*60)
553    _log(f"PREPROCESS: Building BM25 + E5 + Gemma embeddings + loading BGE")
554    _log("="*60)
555
556    # BM25
557    bm25 = BM25Index(k1=BM25_K1, b=BM25_B, logger=_log)
558    bm25.build(ids, bm25_norms)
559
560    # E5 encoder + caching
561    e5 = E5Embedder()
562    e5_mat = None
563    cache_note_e5 = None
564    if E5_EVAL_CACHE_DIR:
565        os.makedirs(E5_EVAL_CACHE_DIR, exist_ok=True)
566        meta_p = os.path.join(E5_EVAL_CACHE_DIR, "e5_meta.json")
567        npy_p  = os.path.join(E5_EVAL_CACHE_DIR, "e5_corpus.npy")
568        sha = _sha1_ids(ids)
569        if os.path.isfile(meta_p) and os.path.isfile(npy_p):
570            try:
571                with open(meta_p,"r",encoding="utf-8") as f: m=json.load(f)
572                if m.get("sha1_ids")==sha and model_name_key(m.get("model_path",""))==e5.model_name and m.get("num_docs")==len(ids):
573                    _log(f"[E5] Loading cached corpus embeddings from {npy_p}")
574                    e5_mat = np.load(npy_p, mmap_mode=None)
575                    cache_note_e5 = "loaded"
576            except Exception as e: _log(f"[E5] Cache read failed: {e} — recomputing.")
577    if e5_mat is None:
578        _log("[E5] Computing corpus embeddings...")
579        t0=time.time()
580        e5_mat = e5.encode(e5_passage_norms, is_query=False, batch=64, progress_desc="E5 corpus")
581        _log(f"[E5] Done in {time.time()-t0:.1f}s — shape={e5_mat.shape}")
582        if E5_EVAL_CACHE_DIR:
583            try:
584                np.save(os.path.join(E5_EVAL_CACHE_DIR,"e5_corpus.npy"), e5_mat)
585                meta = {"sha1_ids": _sha1_ids(ids), "num_docs": len(ids), "model_path": e5.model_path, "dim": int(e5_mat.shape[1]), "created": time.time()}
586                with open(os.path.join(E5_EVAL_CACHE_DIR,"e5_meta.json"),"w",encoding="utf-8") as f: json.dump(meta,f,ensure_ascii=False, indent=2)
587                cache_note_e5 = "saved"
588                _log(f"[E5] Saved cache to {E5_EVAL_CACHE_DIR}")
589            except Exception as e: _log(f"[E5] Cache save failed: {e}")
590
591    # Gemma encoder + caching
592    gemma = GemmaEmbedder()
593    gemma_mat = None
594    cache_note_gm = None
595    if GEMMA_EVAL_CACHE_DIR:
596        os.makedirs(GEMMA_EVAL_CACHE_DIR, exist_ok=True)
597        meta_p_gm = os.path.join(GEMMA_EVAL_CACHE_DIR, "gemma_meta.json")
598        npy_p_gm  = os.path.join(GEMMA_EVAL_CACHE_DIR, "gemma_corpus.npy")
599        sha = _sha1_ids(ids)
600        if os.path.isfile(meta_p_gm) and os.path.isfile(npy_p_gm):
601            try:
602                with open(meta_p_gm,"r",encoding="utf-8") as f: m=json.load(f)
603                if m.get("sha1_ids")==sha and model_name_key(m.get("model_path",""))==gemma.model_name and m.get("num_docs")==len(ids):
604                    _log(f"[Gemma] Loading cached corpus embeddings from {npy_p_gm}")
605                    gemma_mat = np.load(npy_p_gm, mmap_mode=None)
606                    cache_note_gm = "loaded"
607            except Exception as e: _log(f"[Gemma] Cache read failed: {e} — recomputing.")
608    if gemma_mat is None:
609        _log("[Gemma] Computing corpus embeddings...")
610        t0=time.time()
611        gemma_mat = gemma.encode(gm_passage_norms, is_query=False, batch=64, progress_desc="Gemma corpus", max_length=GEMMA_MAX_TOK)
612        _log(f"[Gemma] Done in {time.time()-t0:.1f}s — shape={gemma_mat.shape}")
613        if GEMMA_EVAL_CACHE_DIR:
614            try:
615                np.save(os.path.join(GEMMA_EVAL_CACHE_DIR,"gemma_corpus.npy"), gemma_mat)
616                meta_gm = {"sha1_ids": _sha1_ids(ids), "num_docs": len(ids), "model_path": gemma.model_path, "dim": int(gemma_mat.shape[1]), "created": time.time()}
617                with open(os.path.join(GEMMA_EVAL_CACHE_DIR,"gemma_meta.json"),"w",encoding="utf-8") as f: json.dump(meta_gm,f,ensure_ascii=False, indent=2)
618                cache_note_gm = "saved"
619                _log(f"[Gemma] Saved cache to {GEMMA_EVAL_CACHE_DIR}")
620            except Exception as e: _log(f"[Gemma] Cache save failed: {e}")
621
622    # Reranker
623    rr = BGEReranker()
624
625    id2text = dict(zip(ids,texts))
626    id2norm = dict(zip(ids,bm25_norms))
627
628    hybrid = HybridSearcher(bm25, e5, e5_mat, gemma, gemma_mat, id2text, id2norm)
629    _STATE.update({
630        "bm25": bm25, "id2text": id2text, "id2norm": id2norm,
631        "e5": e5, "e5_corpus": e5_mat,
632        "gemma": gemma, "gemma_corpus": gemma_mat,
633        "reranker": rr, "hybrid": hybrid
634    })
635
636    reranker_params = {
637        "CE_POOL": CE_POOL, "CE_MAXLEN": CE_MAXLEN, "CE_BATCH": CE_BATCH,
638        "FINAL_SCORE_BGE_WEIGHT": FINAL_SCORE_BGE_WEIGHT
639    }
640
641    meta = {
642        "stage1_name": "WRRF(BM25, E5, Gemma)",
643        "stage1_params": {
644            "TOP_BM25": TOP_BM25, "TOP_E5": TOP_E5, "TOP_GEMMA": TOP_GEMMA, "RRF_K": RRF_K,
645            "WRRF_WEIGHTS": {"bm25": WRRF_BM25_W, "e5": WRRF_E5_W, "gemma": WRRF_GEMMA_W}
646        },
647        "reranker_name": "BGE + Hybrid Fusion (Conditional Boost)",
648        "reranker_params": reranker_params,
649        "candidate_pool_cap": CE_POOL,
650        "stage1_search_key": "bm25",
651        "bm25_backend": getattr(bm25, "_backend_name", "unknown"),
652        "e5_model_path": e5.model_path,
653        "gemma_model_path": gemma.model_path,
654        "bge_model_path": rr.model_path,
655        "cache_dir_e5": E5_EVAL_CACHE_DIR or None,
656        "cache_dir_gemma": GEMMA_EVAL_CACHE_DIR or None,
657        "e5_cache": cache_note_e5 or ("unused" if not E5_EVAL_CACHE_DIR else "miss"),
658        "gemma_cache": cache_note_gm or ("unused" if not GEMMA_EVAL_CACHE_DIR else "miss"),
659    }
660
661    _log("✓ PREPROCESS complete.")
662    return {
663        "bm25": hybrid, "id2text": id2text, "id2norm": id2norm,
664        "reranker": rr, "num_documents": len(ids), "_eval": meta
665    }
666
667def predict(query: Dict, pre: Dict):
668    q = query.get("query","")
669    if not q: return []
670    hyb = _STATE.get("hybrid") or pre["bm25"]
671    rr  = _STATE.get("reranker") or pre["reranker"]
672    id2text = _STATE.get("id2text") or pre["id2text"]
673
674    # Stage-1: WRRF retrieval
675    cand_id_scores = hyb.search(q, topk=CE_POOL)
676    if not cand_id_scores: return []
677    cand_ids, rrf_scores = zip(*cand_id_scores)
678    passages = [id2text[pid] for pid in cand_ids]
679
680    # Stage-2: BGE reranker (with its own normalizers)
681    q_norm_bge = norm_bge_query(q)
682    passages_norm_bge = [norm_bge_passage(p) for p in passages]
683    bge_scores = rr.score_pairs(q_norm_bge, passages_norm_bge, batch=CE_BATCH, max_len=CE_MAXLEN)
684
685    # Stage-3: Normalize and combine (conditional boost)
686    norm_bge = _normalize_min_max(bge_scores)
687    norm_rrf = _normalize_min_max(list(rrf_scores))
688    final_scores = []
689    w_rrf = 1.0 - FINAL_SCORE_BGE_WEIGHT
690    for bge_score, rrf_score in zip(norm_bge, norm_rrf):
691        boost = w_rrf * rrf_score * (1.0 - bge_score)
692        final_scores.append(bge_score + boost)
693
694    # Final output
695    out = [{"paragraph_uuid": pid, "score": float(s)}
696           for pid, s in sorted(zip(cand_ids, final_scores), key=lambda x: -x[1])]
697    return out
698