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 numpy.typing import NDArray8from tokenizers import Encoding, Tokenizer9 10from fastembed.common.types import NumpyArray, OnnxProvider, Device11from fastembed.common.onnx_model import EmbeddingWorker, OnnxModel, OnnxOutputContext, T12from fastembed.common.preprocessor_utils import load_tokenizer13from fastembed.common.utils import iter_batch14from fastembed.parallel_processor import ParallelWorkerPool15 16 17class OnnxTextModel(OnnxModel[T]):18 ONNX_OUTPUT_NAMES: list[str] | None = None19 20 @classmethod21 def _get_worker_class(cls) -> Type["TextEmbeddingWorker[T]"]:22 raise NotImplementedError("Subclasses must implement this method")23 24 def _post_process_onnx_output(self, output: OnnxOutputContext, **kwargs: Any) -> Iterable[T]:25 """Post-process the ONNX model output to convert it into a usable format.26 27 Args:28 output (OnnxOutputContext): The raw output from the ONNX model.29 **kwargs: Additional keyword arguments that may be needed by specific implementations.30 31 Returns:32 Iterable[T]: Post-processed output as an iterable of type T.33 """34 raise NotImplementedError("Subclasses must implement this method")35 36 def __init__(self) -> None:37 super().__init__()38 self.tokenizer: Tokenizer | None = None39 self.special_token_to_id: dict[str, int] = {}40 41 def _preprocess_onnx_input(42 self, onnx_input: dict[str, NumpyArray], **kwargs: Any43 ) -> dict[str, NumpyArray | NDArray[np.int64]]:44 """45 Preprocess the onnx input.46 """47 return onnx_input48 49 def _load_onnx_model(50 self,51 model_dir: Path,52 model_file: str,53 threads: int | None,54 providers: Sequence[OnnxProvider] | None = None,55 cuda: bool | Device = Device.AUTO,56 device_id: int | None = None,57 extra_session_options: dict[str, Any] | None = None,58 ) -> None:59 super()._load_onnx_model(60 model_dir=model_dir,61 model_file=model_file,62 threads=threads,63 providers=providers,64 cuda=cuda,65 device_id=device_id,66 extra_session_options=extra_session_options,67 )68 self.tokenizer, self.special_token_to_id = load_tokenizer(model_dir=model_dir)69 70 def load_onnx_model(self) -> None:71 raise NotImplementedError("Subclasses must implement this method")72 73 def tokenize(self, documents: list[str], **kwargs: Any) -> list[Encoding]:74 return self.tokenizer.encode_batch(documents) # type: ignore[union-attr]75 76 def onnx_embed(77 self,78 documents: list[str],79 **kwargs: Any,80 ) -> OnnxOutputContext:81 encoded = self.tokenize(documents, **kwargs)82 input_ids = np.array([e.ids for e in encoded])83 attention_mask = np.array([e.attention_mask for e in encoded])84 input_names = {node.name for node in self.model.get_inputs()} # type: ignore[union-attr]85 onnx_input: dict[str, NumpyArray] = {86 "input_ids": np.array(input_ids, dtype=np.int64),87 }88 if "attention_mask" in input_names:89 onnx_input["attention_mask"] = np.array(attention_mask, dtype=np.int64)90 if "token_type_ids" in input_names:91 onnx_input["token_type_ids"] = np.array(92 [np.zeros(len(e), dtype=np.int64) for e in input_ids], dtype=np.int6493 )94 onnx_input = self._preprocess_onnx_input(onnx_input, **kwargs)95 96 model_output = self.model.run(self.ONNX_OUTPUT_NAMES, onnx_input) # type: ignore[union-attr]97 return OnnxOutputContext(98 model_output=model_output[0],99 attention_mask=onnx_input.get("attention_mask", attention_mask),100 input_ids=onnx_input.get("input_ids", input_ids),101 )102 103 def _embed_documents(104 self,105 model_name: str,106 cache_dir: str,107 documents: str | Iterable[str],108 batch_size: int = 256,109 parallel: int | None = None,110 providers: Sequence[OnnxProvider] | None = None,111 cuda: bool | Device = Device.AUTO,112 device_ids: list[int] | None = None,113 local_files_only: bool = False,114 specific_model_path: str | None = None,115 extra_session_options: dict[str, Any] | None = None,116 **kwargs: Any,117 ) -> Iterable[T]:118 is_small = False119 120 if isinstance(documents, str):121 documents = [documents]122 is_small = True123 124 if isinstance(documents, list):125 if len(documents) < batch_size:126 is_small = True127 128 if parallel is None or is_small:129 if not hasattr(self, "model") or self.model is None:130 self.load_onnx_model()131 for batch in iter_batch(documents, batch_size):132 yield from self._post_process_onnx_output(133 self.onnx_embed(batch, **kwargs), **kwargs134 )135 else:136 if parallel == 0:137 parallel = os.cpu_count()138 139 start_method = "forkserver" if "forkserver" in get_all_start_methods() else "spawn"140 params = {141 "model_name": model_name,142 "cache_dir": cache_dir,143 "providers": providers,144 "local_files_only": local_files_only,145 "specific_model_path": specific_model_path,146 **kwargs,147 }148 149 if extra_session_options is not None:150 params.update(extra_session_options)151 152 pool = ParallelWorkerPool(153 num_workers=parallel or 1,154 worker=self._get_worker_class(),155 cuda=cuda,156 device_ids=device_ids,157 start_method=start_method,158 )159 for batch in pool.ordered_map(iter_batch(documents, batch_size), **params):160 yield from self._post_process_onnx_output(batch, **kwargs) # type: ignore161 162 def _token_count(self, texts: str | Iterable[str], batch_size: int = 1024, **_: Any) -> int:163 if not hasattr(self, "model") or self.model is None:164 self.load_onnx_model() # loads the tokenizer as well165 166 token_num = 0167 assert self.tokenizer is not None168 texts = [texts] if isinstance(texts, str) else texts169 for batch in iter_batch(texts, batch_size):170 for tokens in self.tokenizer.encode_batch(batch):171 token_num += sum(tokens.attention_mask)172 173 return token_num174 175 176class TextEmbeddingWorker(EmbeddingWorker[T]):177 def process(self, items: Iterable[tuple[int, Any]]) -> Iterable[tuple[int, OnnxOutputContext]]:178 for idx, batch in items:179 onnx_output = self.model.onnx_embed(batch)180 yield idx, onnx_output181 