"""Semantic search with keyword boosting over structured captions.""" from __future__ import annotations import json import re import numpy as np from sentence_transformers import SentenceTransformer from caption_store import all_entries _embed_model: SentenceTransformer | None = None _MODEL_NAME = "BAAI/bge-base-en-v1.5" # stronger than MiniLM MIN_RELEVANCE = 0.6 TOP_K = 20 # Weight given to keyword boost relative to semantic score (0–1 additive) KEYWORD_BOOST = 0.25 def _get_embed_model() -> SentenceTransformer: global _embed_model if _embed_model is None: _embed_model = SentenceTransformer(_MODEL_NAME) return _embed_model def _query_tokens(query: str) -> list[str]: """Lowercase words from query, 3+ chars.""" return [w for w in re.findall(r"\b\w+\b", query.lower()) if len(w) >= 3] def _keyword_score(query_tokens: list[str], search_text: str, raw_caption: str) -> float: """ Boost score if query tokens appear in high-signal fields (attire, tags, summary). Returns a value in [0, 1]. """ if not query_tokens: return 0.0 # Tokenize the target text strictly by word boundaries to avoid partial substring hits target_words = set(re.findall(r"\b\w+\b", search_text.lower())) # Calculate regular hits hits = sum(1 for token in query_tokens if token in target_words) # Extra weight for specific high-signal fields (attire, keywords, tags) high_signal_words = set() try: meta = json.loads(raw_caption) # Gather arrays that explicitly hold high-intent search terms for field in [meta.get("search_tags", []), meta.get("archive_keywords", []), meta.get("subjects", {}).get("attire", [])]: if isinstance(field, list): for item in field: high_signal_words.update(re.findall(r"\b\w+\b", str(item).lower())) elif isinstance(field, str): high_signal_words.update(re.findall(r"\b\w+\b", field.lower())) except (json.JSONDecodeError, TypeError): pass high_signal_hits = sum(1 for token in query_tokens if token in high_signal_words) # Instead of penalizing long queries, reward ANY valid keyword match heavily if hits == 0: return 0.0 base_score = hits / len(query_tokens) boost_bonus = 0.5 if high_signal_hits > 0 else 0.0 return min(base_score + boost_bonus, 1.0) def search(query: str, collection: str | None = None) -> list[dict]: """ Hybrid search: semantic similarity + keyword boost on structured fields. Returns up to TOP_K results with score >= MIN_RELEVANCE. """ all_indexed = all_entries() if not all_indexed: return [] # Filter by collection if specified if collection and collection != "All": entries = { p: data for p, data in all_indexed.items() if data.get("collection") == collection } else: entries = all_indexed if not entries: return [] model = _get_embed_model() query_tokens = _query_tokens(query) paths = list(entries.keys()) search_texts = [entries[p].get("search_text") or entries[p]["caption"] for p in paths] raw_captions = [entries[p]["caption"] for p in paths] # Query side gets an explicit instruction prefix query_text = f"Represent this sentence for searching relevant passages: {query}" query_vec = model.encode([query_text], normalize_embeddings=True) # Document side stays normal caption_vecs = model.encode(search_texts, normalize_embeddings=True) semantic_scores = (caption_vecs @ query_vec.T).flatten() final_scores = [] for i, (sem, st, rc) in enumerate(zip(semantic_scores, search_texts, raw_captions)): kw = _keyword_score(query_tokens, st, rc) final_scores.append(float(sem) + KEYWORD_BOOST * kw) ranked = sorted( zip(paths, search_texts, final_scores), key=lambda x: x[2], reverse=True, ) return [ {"path": p, "caption": c, "score": round(s, 4)} for p, c, s in ranked[:TOP_K] if s >= MIN_RELEVANCE ]