File size: 2,976 Bytes
8c3e275
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
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]]