from __future__ import annotations import numpy as np import onnxruntime as ort from tokenizers import Tokenizer from pageparse.config import settings from pageparse.store import Store class SemanticSearch: def __init__(self) -> None: self.store = Store() self._session: ort.InferenceSession | None = None self._tokenizer: Tokenizer | None = None self._init_semantic() def _init_semantic(self) -> None: model_path = settings.model_path(settings.embedding_model) tokenizer_path = settings.model_path("tokenizer.json") if model_path.exists() and tokenizer_path.exists(): try: self._session = ort.InferenceSession( str(model_path), providers=["CPUExecutionProvider"], ) self._tokenizer = Tokenizer.from_file(str(tokenizer_path)) except Exception as e: print(f"Failed to load embedding model: {e}") def _embed(self, text: str) -> np.ndarray: if self._session is None or self._tokenizer is None: return np.zeros(384, dtype=np.float32) try: encoded = self._tokenizer.encode(text) input_ids = np.array([encoded.ids], dtype=np.int64) attention_mask = np.array([encoded.attention_mask], dtype=np.int64) if hasattr(encoded, "attention_mask") else np.ones_like(input_ids) outputs = self._session.run( None, { "input_ids": input_ids, "attention_mask": attention_mask, }, ) embedding = outputs[0].squeeze() norm = np.linalg.norm(embedding) return embedding / norm if norm > 0 else embedding except Exception as e: print(f"Embedding failed: {e}") return np.zeros(384, dtype=np.float32) def search(self, query: str, top_k: int = 5) -> list[dict]: records = self.store.get_records() query_embedding = self._embed(query) use_semantic = not np.all(query_embedding == 0) scored = [] query_lower = query.lower() for rec in records: if use_semantic: content = rec.get("content", "") rec_embedding = self._embed(content) norm = np.linalg.norm(rec_embedding) if norm > 0: similarity = float(np.dot(query_embedding, rec_embedding) / norm) else: similarity = 0.0 keyword_score = content.lower().count(query_lower) * 0.1 score = similarity + keyword_score else: content_lower = rec.get("content", "").lower() score = content_lower.count(query_lower) scored.append((score, rec)) scored.sort(key=lambda x: x[0], reverse=True) return [r for s, r in scored[:top_k]]