rishik1111's picture
fix: Add safe response handling and error alerts for text and voice queries in web UI
bf340fa
Raw
History Blame Contribute Delete
10.5 kB
"""
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