Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
onnx_embedding.py354 linesDownload Raw Back to text
1from typing import Any, Iterable, Sequence, Type2 3from fastembed.common.types import NumpyArray, OnnxProvider, Device4from fastembed.common.onnx_model import OnnxOutputContext5from fastembed.common.utils import define_cache_dir, normalize6from fastembed.text.onnx_text_model import OnnxTextModel, TextEmbeddingWorker7from fastembed.text.text_embedding_base import TextEmbeddingBase8from fastembed.common.model_description import DenseModelDescription, ModelSource9 10supported_onnx_models: list[DenseModelDescription] = [11    DenseModelDescription(12        model="BAAI/bge-base-en",13        dim=768,14        description=(15            "Text embeddings, Unimodal (text), English, 512 input tokens truncation, "16            "Prefixes for queries/documents: necessary, 2023 year."17        ),18        license="mit",19        size_in_GB=0.42,20        sources=ModelSource(21            hf="Qdrant/fast-bge-base-en",22            url="https://storage.googleapis.com/qdrant-fastembed/fast-bge-base-en.tar.gz",23            _deprecated_tar_struct=True,24        ),25        model_file="model_optimized.onnx",26    ),27    DenseModelDescription(28        model="BAAI/bge-base-en-v1.5",29        dim=768,30        description=(31            "Text embeddings, Unimodal (text), English, 512 input tokens truncation, "32            "Prefixes for queries/documents: not so necessary, 2023 year."33        ),34        license="mit",35        size_in_GB=0.21,36        sources=ModelSource(37            hf="qdrant/bge-base-en-v1.5-onnx-q",38            url="https://storage.googleapis.com/qdrant-fastembed/fast-bge-base-en-v1.5.tar.gz",39            _deprecated_tar_struct=True,40        ),41        model_file="model_optimized.onnx",42    ),43    DenseModelDescription(44        model="BAAI/bge-large-en-v1.5",45        dim=1024,46        description=(47            "Text embeddings, Unimodal (text), English, 512 input tokens truncation, "48            "Prefixes for queries/documents: not so necessary, 2023 year."49        ),50        license="mit",51        size_in_GB=1.20,52        sources=ModelSource(hf="qdrant/bge-large-en-v1.5-onnx"),53        model_file="model.onnx",54    ),55    DenseModelDescription(56        model="BAAI/bge-small-en",57        dim=384,58        description=(59            "Text embeddings, Unimodal (text), English, 512 input tokens truncation, "60            "Prefixes for queries/documents: necessary, 2023 year."61        ),62        license="mit",63        size_in_GB=0.13,64        sources=ModelSource(65            hf="Qdrant/bge-small-en",66            url="https://storage.googleapis.com/qdrant-fastembed/BAAI-bge-small-en.tar.gz",67            _deprecated_tar_struct=True,68        ),69        model_file="model_optimized.onnx",70    ),71    DenseModelDescription(72        model="BAAI/bge-small-en-v1.5",73        dim=384,74        description=(75            "Text embeddings, Unimodal (text), English, 512 input tokens truncation, "76            "Prefixes for queries/documents: not so necessary, 2023 year."77        ),78        license="mit",79        size_in_GB=0.067,80        sources=ModelSource(hf="qdrant/bge-small-en-v1.5-onnx-q"),81        model_file="model_optimized.onnx",82    ),83    DenseModelDescription(84        model="BAAI/bge-small-zh-v1.5",85        dim=512,86        description=(87            "Text embeddings, Unimodal (text), Chinese, 512 input tokens truncation, "88            "Prefixes for queries/documents: not so necessary, 2023 year."89        ),90        license="mit",91        size_in_GB=0.09,92        sources=ModelSource(93            hf="Qdrant/bge-small-zh-v1.5",94            url="https://storage.googleapis.com/qdrant-fastembed/fast-bge-small-zh-v1.5.tar.gz",95            _deprecated_tar_struct=True,96        ),97        model_file="model_optimized.onnx",98    ),99    DenseModelDescription(100        model="mixedbread-ai/mxbai-embed-large-v1",101        dim=1024,102        description=(103            "Text embeddings, Unimodal (text), English, 512 input tokens truncation, "104            "Prefixes for queries/documents: necessary, 2024 year."105        ),106        license="apache-2.0",107        size_in_GB=0.64,108        sources=ModelSource(hf="mixedbread-ai/mxbai-embed-large-v1"),109        model_file="onnx/model.onnx",110    ),111    DenseModelDescription(112        model="snowflake/snowflake-arctic-embed-xs",113        dim=384,114        description=(115            "Text embeddings, Unimodal (text), English, 512 input tokens truncation, "116            "Prefixes for queries/documents: necessary, 2024 year."117        ),118        license="apache-2.0",119        size_in_GB=0.09,120        sources=ModelSource(hf="snowflake/snowflake-arctic-embed-xs"),121        model_file="onnx/model.onnx",122    ),123    DenseModelDescription(124        model="snowflake/snowflake-arctic-embed-s",125        dim=384,126        description=(127            "Text embeddings, Unimodal (text), English, 512 input tokens truncation, "128            "Prefixes for queries/documents: necessary, 2024 year."129        ),130        license="apache-2.0",131        size_in_GB=0.13,132        sources=ModelSource(hf="snowflake/snowflake-arctic-embed-s"),133        model_file="onnx/model.onnx",134    ),135    DenseModelDescription(136        model="snowflake/snowflake-arctic-embed-m",137        dim=768,138        description=(139            "Text embeddings, Unimodal (text), English, 512 input tokens truncation, "140            "Prefixes for queries/documents: necessary, 2024 year."141        ),142        license="apache-2.0",143        size_in_GB=0.43,144        sources=ModelSource(hf="Snowflake/snowflake-arctic-embed-m"),145        model_file="onnx/model.onnx",146    ),147    DenseModelDescription(148        model="snowflake/snowflake-arctic-embed-m-long",149        dim=768,150        description=(151            "Text embeddings, Unimodal (text), English, 2048 input tokens truncation, "152            "Prefixes for queries/documents: necessary, 2024 year."153        ),154        license="apache-2.0",155        size_in_GB=0.54,156        sources=ModelSource(hf="snowflake/snowflake-arctic-embed-m-long"),157        model_file="onnx/model.onnx",158    ),159    DenseModelDescription(160        model="snowflake/snowflake-arctic-embed-l",161        dim=1024,162        description=(163            "Text embeddings, Unimodal (text), English, 512 input tokens truncation, "164            "Prefixes for queries/documents: necessary, 2024 year."165        ),166        license="apache-2.0",167        size_in_GB=1.02,168        sources=ModelSource(hf="snowflake/snowflake-arctic-embed-l"),169        model_file="onnx/model.onnx",170    ),171    DenseModelDescription(172        model="jinaai/jina-clip-v1",173        dim=768,174        description=(175            "Text embeddings, Multimodal (text&image), English, Prefixes for queries/documents: "176            "not necessary, 2024 year"177        ),178        license="apache-2.0",179        size_in_GB=0.55,180        sources=ModelSource(hf="jinaai/jina-clip-v1"),181        model_file="onnx/text_model.onnx",182    ),183]184 185 186class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]):187    """Implementation of the Flag Embedding model."""188 189    @classmethod190    def _list_supported_models(cls) -> list[DenseModelDescription]:191        """192        Lists the supported models.193 194        Returns:195            list[DenseModelDescription]: A list of DenseModelDescription objects containing the model information.196        """197        return supported_onnx_models198 199    def __init__(200        self,201        model_name: str = "BAAI/bge-small-en-v1.5",202        cache_dir: str | None = None,203        threads: int | None = None,204        providers: Sequence[OnnxProvider] | None = None,205        cuda: bool | Device = Device.AUTO,206        device_ids: list[int] | None = None,207        lazy_load: bool = False,208        device_id: int | None = None,209        specific_model_path: str | None = None,210        **kwargs: Any,211    ):212        """213        Args:214            model_name (str): The name of the model to use.215            cache_dir (str, optional): The path to the cache directory.216                                       Can be set using the `FASTEMBED_CACHE_PATH` env variable.217                                       Defaults to `fastembed_cache` in the system's temp directory.218            threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.219            providers (Optional[Sequence[OnnxProvider]], optional): The list of onnxruntime providers to use.220                Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.221            cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`222                Defaults to Device.AUTO.223            device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in224                workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, mutually exclusive225                with `providers`. Defaults to None.226            lazy_load (bool, optional): Whether to load the model during class initialization or on demand.227                Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.228            device_id (Optional[int], optional): The device id to use for loading the model in the worker process.229            specific_model_path (Optional[str], optional): The specific path to the onnx model dir if it should be imported from somewhere else230 231        Raises:232            ValueError: If the model_name is not in the format <org>/<model> e.g. BAAI/bge-base-en.233        """234        super().__init__(model_name, cache_dir, threads, **kwargs)235        self.providers = providers236        self.lazy_load = lazy_load237        self._extra_session_options = self._select_exposed_session_options(kwargs)238        # List of device ids, that can be used for data parallel processing in workers239        self.device_ids = device_ids240        self.cuda = cuda241 242        # This device_id will be used if we need to load model in current process243        self.device_id: int | None = None244        if device_id is not None:245            self.device_id = device_id246        elif self.device_ids is not None:247            self.device_id = self.device_ids[0]248 249        self.model_description = self._get_model_description(model_name)250        self.cache_dir = str(define_cache_dir(cache_dir))251        self._specific_model_path = specific_model_path252        self._model_dir = self.download_model(253            self.model_description,254            self.cache_dir,255            local_files_only=self._local_files_only,256            specific_model_path=self._specific_model_path,257        )258 259        if not self.lazy_load:260            self.load_onnx_model()261 262    def embed(263        self,264        documents: str | Iterable[str],265        batch_size: int = 256,266        parallel: int | None = None,267        **kwargs: Any,268    ) -> Iterable[NumpyArray]:269        """270        Encode a list of documents into list of embeddings.271        We use mean pooling with attention so that the model can handle variable-length inputs.272 273        Args:274            documents: Iterator of documents or single document to embed275            batch_size: Batch size for encoding -- higher values will use more memory, but be faster276            parallel:277                If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.278                If 0, use all available cores.279                If None, don't use data-parallel processing, use default onnxruntime threading instead.280 281        Returns:282            List of embeddings, one per document283        """284        yield from self._embed_documents(285            model_name=self.model_name,286            cache_dir=str(self.cache_dir),287            documents=documents,288            batch_size=batch_size,289            parallel=parallel,290            providers=self.providers,291            cuda=self.cuda,292            device_ids=self.device_ids,293            local_files_only=self._local_files_only,294            specific_model_path=self._specific_model_path,295            extra_session_options=self._extra_session_options,296            **kwargs,297        )298 299    @classmethod300    def _get_worker_class(cls) -> Type["TextEmbeddingWorker[NumpyArray]"]:301        return OnnxTextEmbeddingWorker302 303    def _preprocess_onnx_input(304        self, onnx_input: dict[str, NumpyArray], **kwargs: Any305    ) -> dict[str, NumpyArray]:306        """307        Preprocess the onnx input.308        """309        return onnx_input310 311    def _post_process_onnx_output(312        self, output: OnnxOutputContext, **kwargs: Any313    ) -> Iterable[NumpyArray]:314        embeddings = output.model_output315 316        if embeddings.ndim == 3:  # (batch_size, seq_len, embedding_dim)317            processed_embeddings = embeddings[:, 0]318        elif embeddings.ndim == 2:  # (batch_size, embedding_dim)319            processed_embeddings = embeddings320        else:321            raise ValueError(f"Unsupported embedding shape: {embeddings.shape}")322        return normalize(processed_embeddings)323 324    def load_onnx_model(self) -> None:325        self._load_onnx_model(326            model_dir=self._model_dir,327            model_file=self.model_description.model_file,328            threads=self.threads,329            providers=self.providers,330            cuda=self.cuda,331            device_id=self.device_id,332            extra_session_options=self._extra_session_options,333        )334 335    def token_count(336        self, texts: str | Iterable[str], batch_size: int = 1024, **kwargs: Any337    ) -> int:338        return self._token_count(texts, batch_size=batch_size, **kwargs)339 340 341class OnnxTextEmbeddingWorker(TextEmbeddingWorker[NumpyArray]):342    def init_embedding(343        self,344        model_name: str,345        cache_dir: str,346        **kwargs: Any,347    ) -> OnnxTextEmbedding:348        return OnnxTextEmbedding(349            model_name=model_name,350            cache_dir=cache_dir,351            threads=1,352            **kwargs,353        )354 
codekingpro/portable-devtools · Team Ai