Team Ai
Datasetpublic

codekingpro/portable-devtools

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