File size: 10,546 Bytes
9f6ffb8
 
 
 
 
 
 
 
 
 
 
 
 
 
bf340fa
9f6ffb8
 
 
 
bf340fa
 
 
 
 
 
9f6ffb8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
"""
Embedding Model Wrapper for multilingual-e5-small with ONNX Runtime CPU Acceleration.

CRITICAL REQUIREMENT:
`intfloat/multilingual-e5-small` is a retrieval-trained model.
All query encodings MUST use the 'query: ' prefix.
All passage/document encodings MUST use the 'passage: ' prefix.
"""

import logging
import os
from pathlib import Path
from typing import List, Union
import numpy as np
import torch
import config

logger = logging.getLogger(__name__)

# Optimize PyTorch CPU parallelism
try:
    torch.set_num_threads(max(1, torch.get_num_threads()))
except Exception:
    pass

_EMBEDDER_INSTANCE = None


class ONNXMultilingualE5Embedder:
    """
    High-performance ONNX Runtime CPU Embedder for multilingual-e5-small.
    Uses INT8 dynamic quantization and static graph execution for sub-10ms query vectorization.
    """
    def __init__(self, model_name: str = config.EMBEDDING_MODEL_NAME):
        self.model_name = model_name
        self.dim = config.EMBEDDING_DIM
        self.onnx_dir = Path(getattr(config, "ONNX_MODELS_DIR", config.DATA_DIR / "onnx_models"))
        self.onnx_dir.mkdir(parents=True, exist_ok=True)
        
        self.onnx_int8_path = self.onnx_dir / "e5_small_int8.onnx"
        self.onnx_fp32_path = self.onnx_dir / "e5_small.onnx"
        
        from transformers import AutoTokenizer
        self.tokenizer = AutoTokenizer.from_pretrained(model_name)
        
        # Ensure ONNX model exists
        self._ensure_onnx_model()
        
        import onnxruntime as ort
        opts = ort.SessionOptions()
        num_threads = getattr(config, "ONNX_NUM_THREADS", 2)
        opts.intra_op_num_threads = num_threads
        opts.inter_op_num_threads = 1
        opts.execution_mode = ort.ExecutionMode.ORT_SEQUENTIAL
        opts.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
        
        load_path = self.onnx_int8_path if self.onnx_int8_path.exists() else self.onnx_fp32_path
        logger.info(f"Loading ONNX embedding model from: {load_path} (threads={num_threads})")
        self.session = ort.InferenceSession(str(load_path), opts, providers=["CPUExecutionProvider"])
        
        # Warmup ONNX inference graph to avoid cold-start JIT latency
        try:
            dummy_in = self.tokenizer(["query: warmup"], padding=True, return_tensors="np")
            self.session.run(None, {
                "input_ids": dummy_in["input_ids"].astype(np.int64),
                "attention_mask": dummy_in["attention_mask"].astype(np.int64),
            })
        except Exception:
            pass
        logger.info("ONNX embedding session initialized and warmed up successfully.")

    def _ensure_onnx_model(self):
        """Auto-export and quantize PyTorch model if ONNX files do not exist."""
        if self.onnx_int8_path.exists() or self.onnx_fp32_path.exists():
            return
            
        logger.info("Exporting multilingual-e5-small to ONNX format...")
        import torch.nn as nn
        from transformers import AutoModel
        from onnxruntime.quantization import quantize_dynamic, QuantType
        
        class E5Wrapper(nn.Module):
            def __init__(self, m):
                super().__init__()
                self.m = m
            def forward(self, input_ids, attention_mask):
                out = self.m(input_ids=input_ids, attention_mask=attention_mask, return_dict=False)
                return out[0]

        base_model = AutoModel.from_pretrained(self.model_name)
        base_model.eval()
        wrapper = E5Wrapper(base_model)
        wrapper.eval()
        
        dummy = self.tokenizer(["query 1", "query 2"], padding=True, return_tensors="pt")
        try:
            torch.onnx.export(
                wrapper,
                (dummy["input_ids"], dummy["attention_mask"]),
                str(self.onnx_fp32_path),
                input_names=["input_ids", "attention_mask"],
                output_names=["last_hidden_state"],
                dynamic_axes={
                    "input_ids": {0: "batch", 1: "seq"},
                    "attention_mask": {0: "batch", 1: "seq"},
                    "last_hidden_state": {0: "batch", 1: "seq"},
                },
                opset_version=14,
                do_constant_folding=True,
            )
            logger.info("Exported ONNX embedding model with dynamic shapes.")
        except Exception as e:
            logger.warning(f"ONNX export failed: {e}. PyTorch fallback will be used.")

    def _mean_pool_and_normalize(self, token_embeddings: np.ndarray, attention_mask: np.ndarray, normalize: bool = True) -> np.ndarray:
        """Vectorized mean pooling over active attention mask tokens with L2 normalization."""
        input_mask_expanded = np.expand_dims(attention_mask, -1)
        sum_embeddings = np.sum(token_embeddings * input_mask_expanded, axis=1)
        sum_mask = np.clip(input_mask_expanded.sum(axis=1), a_min=1e-9, a_max=None)
        pooled = sum_embeddings / sum_mask
        if normalize:
            norm = np.linalg.norm(pooled, axis=1, keepdims=True)
            pooled = pooled / np.clip(norm, a_min=1e-9, a_max=None)
        return np.ascontiguousarray(pooled, dtype=np.float32)

    def encode_queries(
        self, queries: Union[str, List[str]], normalize: bool = True
    ) -> np.ndarray:
        """
        Encodes one or more queries with mandatory 'query: ' prefix using ONNX Runtime.
        """
        if isinstance(queries, str):
            queries = [queries]
        prefixed = [f"{config.QUERY_PREFIX}{q.strip()}" for q in queries]
        
        inputs = self.tokenizer(
            prefixed,
            padding=True,
            truncation=True,
            max_length=getattr(config, "CONTEXT_BOUNDING_MAX_TOKENS", 64),
            return_tensors="np",
        )
        ort_inputs = {
            "input_ids": inputs["input_ids"].astype(np.int64),
            "attention_mask": inputs["attention_mask"].astype(np.int64),
        }
        outputs = self.session.run(None, ort_inputs)
        token_embeddings = outputs[0]
        return self._mean_pool_and_normalize(token_embeddings, inputs["attention_mask"], normalize=normalize)

    def encode_passages(
        self, passages: Union[str, List[str]], batch_size: int = 64, normalize: bool = True
    ) -> np.ndarray:
        """
        Encodes passages with mandatory 'passage: ' prefix in batches.
        """
        if isinstance(passages, str):
            passages = [passages]
        prefixed = [f"{config.PASSAGE_PREFIX}{p.strip()}" for p in passages]
        
        all_embeddings = []
        for i in range(0, len(prefixed), batch_size):
            batch = prefixed[i : i + batch_size]
            inputs = self.tokenizer(
                batch,
                padding=True,
                truncation=True,
                max_length=getattr(config, "CONTEXT_BOUNDING_MAX_TOKENS", 64),
                return_tensors="np",
            )
            ort_inputs = {
                "input_ids": inputs["input_ids"].astype(np.int64),
                "attention_mask": inputs["attention_mask"].astype(np.int64),
            }
            outputs = self.session.run(None, ort_inputs)
            pooled = self._mean_pool_and_normalize(outputs[0], inputs["attention_mask"], normalize=normalize)
            all_embeddings.append(pooled)
            
        if not all_embeddings:
            return np.empty((0, self.dim), dtype=np.float32)
        return np.vstack(all_embeddings)

    def encode_sentences(self, sentences: List[str]) -> np.ndarray:
        """Encodes consecutive sentences for semantic distance analysis."""
        return self.encode_passages(sentences, normalize=True)


class PyTorchMultilingualE5Embedder:
    """
    PyTorch fallback wrapper for sentence-transformers multilingual-e5-small.
    """
    def __init__(self, model_name: str = config.EMBEDDING_MODEL_NAME):
        from sentence_transformers import SentenceTransformer
        logger.info(f"Loading PyTorch fallback embedding model: '{model_name}'...")
        self.model_name = model_name
        try:
            self.model = SentenceTransformer(model_name, local_files_only=True)
        except Exception:
            self.model = SentenceTransformer(model_name)
        self.dim = config.EMBEDDING_DIM
        logger.info(f"PyTorch embedding model loaded (dim={self.dim}).")

    def encode_queries(
        self, queries: Union[str, List[str]], normalize: bool = True
    ) -> np.ndarray:
        if isinstance(queries, str):
            queries = [queries]
        prefixed = [f"{config.QUERY_PREFIX}{q.strip()}" for q in queries]
        vectors = self.model.encode(
            prefixed,
            normalize_embeddings=normalize,
            show_progress_bar=False,
            convert_to_numpy=True,
        )
        return np.ascontiguousarray(vectors, dtype=np.float32)

    def encode_passages(
        self, passages: Union[str, List[str]], batch_size: int = 64, normalize: bool = True
    ) -> np.ndarray:
        if isinstance(passages, str):
            passages = [passages]
        prefixed = [f"{config.PASSAGE_PREFIX}{p.strip()}" for p in passages]
        vectors = self.model.encode(
            prefixed,
            batch_size=batch_size,
            normalize_embeddings=normalize,
            show_progress_bar=(len(passages) > 200),
            convert_to_numpy=True,
        )
        return np.ascontiguousarray(vectors, dtype=np.float32)

    def encode_sentences(self, sentences: List[str]) -> np.ndarray:
        return self.encode_passages(sentences, normalize=True)


def get_embedder():
    """
    Get or initialize the global singleton embedder instance with ONNX-first policy.
    """
    global _EMBEDDER_INSTANCE
    if _EMBEDDER_INSTANCE is None:
        onnx_int8 = config.ONNX_MODELS_DIR / "e5_small_int8.onnx"
        onnx_fp32 = config.ONNX_MODELS_DIR / "e5_small.onnx"
        if getattr(config, "ENABLE_ONNX_EMBEDDING", True) and (onnx_int8.exists() or onnx_fp32.exists()):
            try:
                _EMBEDDER_INSTANCE = ONNXMultilingualE5Embedder()
            except Exception as e:
                logger.warning(f"Failed to initialize ONNX Embedder: {e}. Falling back to PyTorch.")
                _EMBEDDER_INSTANCE = PyTorchMultilingualE5Embedder()
        else:
            logger.info("ONNX embedding model file not cached. Using PyTorch SentenceTransformer embedder.")
            _EMBEDDER_INSTANCE = PyTorchMultilingualE5Embedder()
    return _EMBEDDER_INSTANCE