codekingpro/portable-devtools
114k
1import os2from multiprocessing import get_all_start_methods3from pathlib import Path4from typing import Any, Iterable, Sequence, Type5 6import numpy as np7from tokenizers import Encoding8 9from fastembed.common.onnx_model import (10 EmbeddingWorker,11 OnnxModel,12 OnnxOutputContext,13 OnnxProvider,14)15from fastembed.common.types import NumpyArray, Device16from fastembed.common.preprocessor_utils import load_tokenizer17from fastembed.common.utils import iter_batch18from fastembed.parallel_processor import ParallelWorkerPool19 20 21class OnnxCrossEncoderModel(OnnxModel[float]):22 ONNX_OUTPUT_NAMES: list[str] | None = None23 24 @classmethod25 def _get_worker_class(cls) -> Type["TextRerankerWorker"]:26 raise NotImplementedError("Subclasses must implement this method")27 28 def _load_onnx_model(29 self,30 model_dir: Path,31 model_file: str,32 threads: int | None,33 providers: Sequence[OnnxProvider] | None = None,34 cuda: bool | Device = Device.AUTO,35 device_id: int | None = None,36 extra_session_options: dict[str, Any] | None = None,37 ) -> None:38 super()._load_onnx_model(39 model_dir=model_dir,40 model_file=model_file,41 threads=threads,42 providers=providers,43 cuda=cuda,44 device_id=device_id,45 extra_session_options=extra_session_options,46 )47 self.tokenizer, _ = load_tokenizer(model_dir=model_dir)48 assert self.tokenizer is not None49 50 def tokenize(self, pairs: list[tuple[str, str]], **_: Any) -> list[Encoding]:51 return self.tokenizer.encode_batch(pairs) # type: ignore[union-attr]52 53 def _build_onnx_input(self, tokenized_input: list[Encoding]) -> dict[str, NumpyArray]:54 input_names: set[str] = {node.name for node in self.model.get_inputs()} # type: ignore[union-attr]55 inputs: dict[str, NumpyArray] = {56 "input_ids": np.array([enc.ids for enc in tokenized_input], dtype=np.int64),57 }58 if "token_type_ids" in input_names:59 inputs["token_type_ids"] = np.array(60 [enc.type_ids for enc in tokenized_input], dtype=np.int6461 )62 if "attention_mask" in input_names:63 inputs["attention_mask"] = np.array(64 [enc.attention_mask for enc in tokenized_input], dtype=np.int6465 )66 return inputs67 68 def onnx_embed(self, query: str, documents: list[str], **kwargs: Any) -> OnnxOutputContext:69 pairs = [(query, doc) for doc in documents]70 return self.onnx_embed_pairs(pairs, **kwargs)71 72 def onnx_embed_pairs(self, pairs: list[tuple[str, str]], **kwargs: Any) -> OnnxOutputContext:73 tokenized_input = self.tokenize(pairs, **kwargs)74 inputs = self._build_onnx_input(tokenized_input)75 onnx_input = self._preprocess_onnx_input(inputs, **kwargs)76 outputs = self.model.run(self.ONNX_OUTPUT_NAMES, onnx_input) # type: ignore[union-attr]77 relevant_output = outputs[0]78 scores: NumpyArray = relevant_output[:, 0]79 return OnnxOutputContext(model_output=scores)80 81 def _rerank_documents(82 self, query: str, documents: Iterable[str], batch_size: int, **kwargs: Any83 ) -> Iterable[float]:84 if not hasattr(self, "model") or self.model is None:85 self.load_onnx_model()86 for batch in iter_batch(documents, batch_size):87 yield from self._post_process_onnx_output(self.onnx_embed(query, batch, **kwargs))88 89 def _rerank_pairs(90 self,91 model_name: str,92 cache_dir: str,93 pairs: Iterable[tuple[str, str]],94 batch_size: int,95 parallel: int | None = None,96 providers: Sequence[OnnxProvider] | None = None,97 cuda: bool | Device = Device.AUTO,98 device_ids: list[int] | None = None,99 local_files_only: bool = False,100 specific_model_path: str | None = None,101 extra_session_options: dict[str, Any] | None = None,102 **kwargs: Any,103 ) -> Iterable[float]:104 is_small = False105 106 if isinstance(pairs, tuple):107 pairs = [pairs]108 is_small = True109 110 if isinstance(pairs, list):111 if len(pairs) < batch_size:112 is_small = True113 114 if parallel is None or is_small:115 if not hasattr(self, "model") or self.model is None:116 self.load_onnx_model()117 for batch in iter_batch(pairs, batch_size):118 yield from self._post_process_onnx_output(self.onnx_embed_pairs(batch, **kwargs))119 else:120 if parallel == 0:121 parallel = os.cpu_count()122 123 start_method = "forkserver" if "forkserver" in get_all_start_methods() else "spawn"124 params = {125 "model_name": model_name,126 "cache_dir": cache_dir,127 "providers": providers,128 "local_files_only": local_files_only,129 "specific_model_path": specific_model_path,130 **kwargs,131 }132 133 if extra_session_options is not None:134 params.update(extra_session_options)135 136 pool = ParallelWorkerPool(137 num_workers=parallel or 1,138 worker=self._get_worker_class(),139 cuda=cuda,140 device_ids=device_ids,141 start_method=start_method,142 )143 for batch in pool.ordered_map(iter_batch(pairs, batch_size), **params):144 yield from self._post_process_onnx_output(batch) # type: ignore145 146 def _post_process_onnx_output(147 self, output: OnnxOutputContext, **kwargs: Any148 ) -> Iterable[float]:149 """Post-process the ONNX model output to convert it into a usable format.150 151 Args:152 output (OnnxOutputContext): The raw output from the ONNX model.153 **kwargs: Additional keyword arguments that may be needed by specific implementations.154 155 Returns:156 Iterable[float]: Post-processed output as an iterable of float values.157 """158 raise NotImplementedError("Subclasses must implement this method")159 160 def _preprocess_onnx_input(161 self, onnx_input: dict[str, NumpyArray], **kwargs: Any162 ) -> dict[str, NumpyArray]:163 """164 Preprocess the onnx input.165 """166 return onnx_input167 168 def _token_count(169 self, pairs: Iterable[tuple[str, str]], batch_size: int = 1024, **_: Any170 ) -> int:171 if not hasattr(self, "model") or self.model is None:172 self.load_onnx_model() # loads the tokenizer as well173 174 token_num = 0175 assert self.tokenizer is not None176 for batch in iter_batch(pairs, batch_size):177 for tokens in self.tokenizer.encode_batch(batch):178 token_num += sum(tokens.attention_mask)179 180 return token_num181 182 183class TextRerankerWorker(EmbeddingWorker[float]):184 def __init__(185 self,186 model_name: str,187 cache_dir: str,188 **kwargs: Any,189 ):190 self.model: OnnxCrossEncoderModel191 super().__init__(model_name, cache_dir, **kwargs)192 193 def init_embedding(194 self,195 model_name: str,196 cache_dir: str,197 **kwargs: Any,198 ) -> OnnxCrossEncoderModel:199 raise NotImplementedError()200 201 def process(self, items: Iterable[tuple[int, Any]]) -> Iterable[tuple[int, Any]]:202 for idx, batch in items:203 onnx_output = self.model.onnx_embed_pairs(batch)204 yield idx, onnx_output205 