codekingpro/portable-devtools
114k
1import warnings2from dataclasses import dataclass3from pathlib import Path4from typing import Any, Generic, Iterable, Sequence, Type, TypeVar5 6import numpy as np7import onnxruntime as ort8 9from numpy.typing import NDArray10from tokenizers import Tokenizer11 12from fastembed.common.types import OnnxProvider, NumpyArray, Device13from fastembed.parallel_processor import Worker14 15# Holds type of the embedding result16T = TypeVar("T")17 18 19@dataclass20class OnnxOutputContext:21 model_output: NumpyArray22 attention_mask: NDArray[np.int64] | None = None23 input_ids: NDArray[np.int64] | None = None24 metadata: dict[str, Any] | None = None25 26 27class OnnxModel(Generic[T]):28 EXPOSED_SESSION_OPTIONS = ("enable_cpu_mem_arena",)29 30 @classmethod31 def _get_worker_class(cls) -> Type["EmbeddingWorker[T]"]:32 raise NotImplementedError("Subclasses must implement this method")33 34 def _post_process_onnx_output(self, output: OnnxOutputContext, **kwargs: Any) -> Iterable[T]:35 """Post-process the ONNX model output to convert it into a usable format.36 37 Args:38 output (OnnxOutputContext): The raw output from the ONNX model.39 **kwargs: Additional keyword arguments that may be needed by specific implementations.40 41 Returns:42 Iterable[T]: Post-processed output as an iterable of type T.43 """44 raise NotImplementedError("Subclasses must implement this method")45 46 def __init__(self) -> None:47 self.model: ort.InferenceSession | None = None48 self.tokenizer: Tokenizer | None = None49 50 def _preprocess_onnx_input(51 self, onnx_input: dict[str, NumpyArray], **kwargs: Any52 ) -> dict[str, NumpyArray]:53 """54 Preprocess the onnx input.55 """56 return onnx_input57 58 def _load_onnx_model(59 self,60 model_dir: Path,61 model_file: str,62 threads: int | None,63 providers: Sequence[OnnxProvider] | None = None,64 cuda: bool | Device = Device.AUTO,65 device_id: int | None = None,66 extra_session_options: dict[str, Any] | None = None,67 ) -> None:68 model_path = model_dir / model_file69 # List of Execution Providers: https://onnxruntime.ai/docs/execution-providers70 available_providers = ort.get_available_providers()71 cuda_available = "CUDAExecutionProvider" in available_providers72 explicit_cuda = cuda is True or cuda == Device.CUDA73 74 if explicit_cuda and providers is not None:75 warnings.warn(76 f"`cuda` and `providers` are mutually exclusive parameters, "77 f"cuda: {cuda}, providers: {providers}. If you'd like to use providers, cuda should be one of "78 f"[False, Device.CPU, Device.AUTO].",79 category=UserWarning,80 stacklevel=6,81 )82 83 if providers is not None:84 onnx_providers = list(providers)85 elif explicit_cuda or (cuda == Device.AUTO and cuda_available):86 if device_id is None:87 onnx_providers = ["CUDAExecutionProvider"]88 else:89 onnx_providers = [("CUDAExecutionProvider", {"device_id": device_id})]90 else:91 onnx_providers = ["CPUExecutionProvider"]92 93 requested_provider_names: list[str] = []94 for provider in onnx_providers:95 # check providers available96 provider_name = provider if isinstance(provider, str) else provider[0]97 requested_provider_names.append(provider_name)98 if provider_name not in available_providers:99 raise ValueError(100 f"Provider {provider_name} is not available. Available providers: {available_providers}"101 )102 103 so = ort.SessionOptions()104 so.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL105 106 if threads is not None:107 so.intra_op_num_threads = threads108 so.inter_op_num_threads = threads109 110 if extra_session_options is not None:111 self.add_extra_session_options(so, extra_session_options)112 113 self.model = ort.InferenceSession(114 str(model_path), providers=onnx_providers, sess_options=so115 )116 if "CUDAExecutionProvider" in requested_provider_names:117 assert self.model is not None118 current_providers = self.model.get_providers()119 if "CUDAExecutionProvider" not in current_providers:120 warnings.warn(121 f"Attempt to set CUDAExecutionProvider failed. Current providers: {current_providers}."122 "If you are using CUDA 12.x, install onnxruntime-gpu via "123 "`pip install onnxruntime-gpu --extra-index-url https://aiinfra.pkgs.visualstudio.com/PublicPackages/_packaging/onnxruntime-cuda-12/pypi/simple/`",124 RuntimeWarning,125 )126 127 @classmethod128 def _select_exposed_session_options(cls, model_kwargs: dict[str, Any]) -> dict[str, Any]:129 """A convenience method to select the exposed session options in models130 131 Args:132 model_kwargs (dict[str, Any]): The model kwargs.133 134 Returns:135 dict[str, Any]: a dict with filtered exposed session options.136 """137 return {k: v for k, v in model_kwargs.items() if k in cls.EXPOSED_SESSION_OPTIONS}138 139 @classmethod140 def add_extra_session_options(141 cls, session_options: ort.SessionOptions, extra_options: dict[str, Any]142 ) -> None:143 """Add extra session options to the existing options object in-place144 145 Args:146 session_options (ort.SessionOptions): The existing session options object.147 extra_options (dict[str, Any]): The extra session options available in cls.EXPOSED_SESSION_OPTIONS.148 149 Returns:150 None151 """152 for option in extra_options:153 assert (154 option in cls.EXPOSED_SESSION_OPTIONS155 ), f"{option} is unknown or not exposed (exposed options: {cls.EXPOSED_SESSION_OPTIONS})"156 if "enable_cpu_mem_arena" in extra_options:157 session_options.enable_cpu_mem_arena = extra_options["enable_cpu_mem_arena"]158 159 def load_onnx_model(self) -> None:160 raise NotImplementedError("Subclasses must implement this method")161 162 def onnx_embed(self, *args: Any, **kwargs: Any) -> OnnxOutputContext:163 raise NotImplementedError("Subclasses must implement this method")164 165 166class EmbeddingWorker(Worker, Generic[T]):167 def init_embedding(168 self,169 model_name: str,170 cache_dir: str,171 **kwargs: Any,172 ) -> OnnxModel[T]:173 raise NotImplementedError()174 175 def __init__(176 self,177 model_name: str,178 cache_dir: str,179 **kwargs: Any,180 ):181 self.model = self.init_embedding(model_name, cache_dir, **kwargs)182 183 @classmethod184 def start(cls, model_name: str, cache_dir: str, **kwargs: Any) -> "EmbeddingWorker[T]":185 return cls(model_name=model_name, cache_dir=cache_dir, **kwargs)186 187 def process(self, items: Iterable[tuple[int, Any]]) -> Iterable[tuple[int, Any]]:188 raise NotImplementedError("Subclasses must implement this method")189 