codekingpro/portable-devtools
115k
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 