Team Ai
Datasetpublic

codekingpro/portable-devtools

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