pageparse.ai / src /pageparse /search.py
Varun2007's picture
initial clean deployment commit with compilers
8c3e275
Raw
History Blame Contribute Delete
2.98 kB
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]]