Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
onnx_model.py189 linesDownload Raw Back to common
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 
codekingpro/portable-devtools · Team Ai