Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
pooled_embedding.py137 linesDownload Raw Back to text
1from typing import Any, Iterable, Type2 3import numpy as np4from numpy.typing import NDArray5 6from fastembed.common.types import NumpyArray7from fastembed.common.onnx_model import OnnxOutputContext8from fastembed.common.utils import mean_pooling9from fastembed.text.onnx_embedding import OnnxTextEmbedding, OnnxTextEmbeddingWorker10from fastembed.common.model_description import DenseModelDescription, ModelSource11 12supported_pooled_models: list[DenseModelDescription] = [13    DenseModelDescription(14        model="nomic-ai/nomic-embed-text-v1.5",15        dim=768,16        description=(17            "Text embeddings, Multimodal (text, image), English, 8192 input tokens truncation, "18            "Prefixes for queries/documents: necessary, 2024 year."19        ),20        license="apache-2.0",21        size_in_GB=0.52,22        sources=ModelSource(hf="nomic-ai/nomic-embed-text-v1.5"),23        model_file="onnx/model.onnx",24    ),25    DenseModelDescription(26        model="nomic-ai/nomic-embed-text-v1.5-Q",27        dim=768,28        description=(29            "Text embeddings, Multimodal (text, image), English, 8192 input tokens truncation, "30            "Prefixes for queries/documents: necessary, 2024 year."31        ),32        license="apache-2.0",33        size_in_GB=0.13,34        sources=ModelSource(hf="nomic-ai/nomic-embed-text-v1.5"),35        model_file="onnx/model_quantized.onnx",36    ),37    DenseModelDescription(38        model="nomic-ai/nomic-embed-text-v1",39        dim=768,40        description=(41            "Text embeddings, Multimodal (text, image), English, 8192 input tokens truncation, "42            "Prefixes for queries/documents: necessary, 2024 year."43        ),44        license="apache-2.0",45        size_in_GB=0.52,46        sources=ModelSource(hf="nomic-ai/nomic-embed-text-v1"),47        model_file="onnx/model.onnx",48    ),49    DenseModelDescription(50        model="sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2",51        dim=384,52        description=(53            "Text embeddings, Unimodal (text), Multilingual (~50 languages), 512 input tokens truncation, "54            "Prefixes for queries/documents: not necessary, 2019 year."55        ),56        license="apache-2.0",57        size_in_GB=0.22,58        sources=ModelSource(hf="qdrant/paraphrase-multilingual-MiniLM-L12-v2-onnx-Q"),59        model_file="model_optimized.onnx",60    ),61    DenseModelDescription(62        model="sentence-transformers/paraphrase-multilingual-mpnet-base-v2",63        dim=768,64        description=(65            "Text embeddings, Unimodal (text), Multilingual (~50 languages), 384 input tokens truncation, "66            "Prefixes for queries/documents: not necessary, 2021 year."67        ),68        license="apache-2.0",69        size_in_GB=1.00,70        sources=ModelSource(hf="xenova/paraphrase-multilingual-mpnet-base-v2"),71        model_file="onnx/model.onnx",72    ),73    DenseModelDescription(74        model="intfloat/multilingual-e5-large",75        dim=1024,76        description=(77            "Text embeddings, Unimodal (text), Multilingual (~100 languages), 512 input tokens truncation, "78            "Prefixes for queries/documents: necessary, 2024 year."79        ),80        license="mit",81        size_in_GB=2.24,82        sources=ModelSource(83            hf="qdrant/multilingual-e5-large-onnx",84            url="https://storage.googleapis.com/qdrant-fastembed/fast-multilingual-e5-large.tar.gz",85            _deprecated_tar_struct=True,86        ),87        model_file="model.onnx",88        additional_files=["model.onnx_data"],89    ),90]91 92 93class PooledEmbedding(OnnxTextEmbedding):94    @classmethod95    def _get_worker_class(cls) -> Type[OnnxTextEmbeddingWorker]:96        return PooledEmbeddingWorker97 98    @classmethod99    def mean_pooling(100        cls, model_output: NumpyArray, attention_mask: NDArray[np.int64]101    ) -> NumpyArray:102        return mean_pooling(model_output, attention_mask)103 104    @classmethod105    def _list_supported_models(cls) -> list[DenseModelDescription]:106        """Lists the supported models.107 108        Returns:109            list[DenseModelDescription]: A list of DenseModelDescription objects containing the model information.110        """111        return supported_pooled_models112 113    def _post_process_onnx_output(114        self, output: OnnxOutputContext, **kwargs: Any115    ) -> Iterable[NumpyArray]:116        if output.attention_mask is None:117            raise ValueError("attention_mask must be provided for document post-processing")118 119        embeddings = output.model_output120        attn_mask = output.attention_mask121        return self.mean_pooling(embeddings, attn_mask)122 123 124class PooledEmbeddingWorker(OnnxTextEmbeddingWorker):125    def init_embedding(126        self,127        model_name: str,128        cache_dir: str,129        **kwargs: Any,130    ) -> OnnxTextEmbedding:131        return PooledEmbedding(132            model_name=model_name,133            cache_dir=cache_dir,134            threads=1,135            **kwargs,136        )137 
codekingpro/portable-devtools · Team Ai