Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
custom_text_embedding.py98 linesDownload Raw Back to text
1from typing import Sequence, Any, Iterable2from dataclasses import dataclass3 4import numpy as np5from numpy.typing import NDArray6 7from fastembed.common import OnnxProvider8from fastembed.common.model_description import (9    PoolingType,10    DenseModelDescription,11)12from fastembed.common.onnx_model import OnnxOutputContext13from fastembed.common.types import NumpyArray, Device14from fastembed.common.utils import normalize, mean_pooling15from fastembed.text.onnx_embedding import OnnxTextEmbedding16 17 18@dataclass(frozen=True)19class PostprocessingConfig:20    pooling: PoolingType21    normalization: bool22 23 24class CustomTextEmbedding(OnnxTextEmbedding):25    SUPPORTED_MODELS: list[DenseModelDescription] = []26    POSTPROCESSING_MAPPING: dict[str, PostprocessingConfig] = {}27 28    def __init__(29        self,30        model_name: str,31        cache_dir: str | None = None,32        threads: int | None = None,33        providers: Sequence[OnnxProvider] | None = None,34        cuda: bool | Device = Device.AUTO,35        device_ids: list[int] | None = None,36        lazy_load: bool = False,37        device_id: int | None = None,38        specific_model_path: str | None = None,39        **kwargs: Any,40    ):41        super().__init__(42            model_name=model_name,43            cache_dir=cache_dir,44            threads=threads,45            providers=providers,46            cuda=cuda,47            device_ids=device_ids,48            lazy_load=lazy_load,49            device_id=device_id,50            specific_model_path=specific_model_path,51            **kwargs,52        )53        self._pooling = self.POSTPROCESSING_MAPPING[model_name].pooling54        self._normalization = self.POSTPROCESSING_MAPPING[model_name].normalization55 56    @classmethod57    def _list_supported_models(cls) -> list[DenseModelDescription]:58        return cls.SUPPORTED_MODELS59 60    def _post_process_onnx_output(61        self, output: OnnxOutputContext, **kwargs: Any62    ) -> Iterable[NumpyArray]:63        return self._normalize(self._pool(output.model_output, output.attention_mask))64 65    def _pool(66        self, embeddings: NumpyArray, attention_mask: NDArray[np.int64] | None = None67    ) -> NumpyArray:68        if self._pooling == PoolingType.CLS:69            return embeddings[:, 0]70 71        if self._pooling == PoolingType.MEAN:72            if attention_mask is None:73                raise ValueError("attention_mask must be provided for mean pooling")74            return mean_pooling(embeddings, attention_mask)75 76        if self._pooling == PoolingType.DISABLED:77            return embeddings78 79        raise ValueError(80            f"Unsupported pooling type {self._pooling}. "81            f"Supported types are: {PoolingType.CLS}, {PoolingType.MEAN}, {PoolingType.DISABLED}."82        )83 84    def _normalize(self, embeddings: NumpyArray) -> NumpyArray:85        return normalize(embeddings) if self._normalization else embeddings86 87    @classmethod88    def add_model(89        cls,90        model_description: DenseModelDescription,91        pooling: PoolingType,92        normalization: bool,93    ) -> None:94        cls.SUPPORTED_MODELS.append(model_description)95        cls.POSTPROCESSING_MAPPING[model_description.model] = PostprocessingConfig(96            pooling=pooling, normalization=normalization97        )98 
codekingpro/portable-devtools · Team Ai