Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
pooled_normalized_embedding.py165 linesDownload Raw Back to text
1from typing import Any, Iterable, Type2 3 4from fastembed.common.types import NumpyArray5from fastembed.common.onnx_model import OnnxOutputContext6from fastembed.common.utils import normalize7from fastembed.text.onnx_embedding import OnnxTextEmbedding, OnnxTextEmbeddingWorker8from fastembed.text.pooled_embedding import PooledEmbedding9from fastembed.common.model_description import DenseModelDescription, ModelSource10 11supported_pooled_normalized_models: list[DenseModelDescription] = [12    DenseModelDescription(13        model="sentence-transformers/all-MiniLM-L6-v2",14        dim=384,15        description=(16            "Text embeddings, Unimodal (text), English, 256 input tokens truncation, "17            "Prefixes for queries/documents: not necessary, 2021 year."18        ),19        license="apache-2.0",20        size_in_GB=0.09,21        sources=ModelSource(22            url="https://storage.googleapis.com/qdrant-fastembed/sentence-transformers-all-MiniLM-L6-v2.tar.gz",23            hf="qdrant/all-MiniLM-L6-v2-onnx",24            _deprecated_tar_struct=True,25        ),26        model_file="model.onnx",27    ),28    DenseModelDescription(29        model="jinaai/jina-embeddings-v2-base-en",30        dim=768,31        description=(32            "Text embeddings, Unimodal (text), English, 8192 input tokens truncation, "33            "Prefixes for queries/documents: not necessary, 2023 year."34        ),35        license="apache-2.0",36        size_in_GB=0.52,37        sources=ModelSource(hf="xenova/jina-embeddings-v2-base-en"),38        model_file="onnx/model.onnx",39    ),40    DenseModelDescription(41        model="jinaai/jina-embeddings-v2-small-en",42        dim=512,43        description=(44            "Text embeddings, Unimodal (text), English, 8192 input tokens truncation, "45            "Prefixes for queries/documents: not necessary, 2023 year."46        ),47        license="apache-2.0",48        size_in_GB=0.12,49        sources=ModelSource(hf="xenova/jina-embeddings-v2-small-en"),50        model_file="onnx/model.onnx",51    ),52    DenseModelDescription(53        model="jinaai/jina-embeddings-v2-base-de",54        dim=768,55        description=(56            "Text embeddings, Unimodal (text), Multilingual (German, English), 8192 input tokens truncation, "57            "Prefixes for queries/documents: not necessary, 2024 year."58        ),59        license="apache-2.0",60        size_in_GB=0.32,61        sources=ModelSource(hf="jinaai/jina-embeddings-v2-base-de"),62        model_file="onnx/model_fp16.onnx",63    ),64    DenseModelDescription(65        model="jinaai/jina-embeddings-v2-base-code",66        dim=768,67        description=(68            "Text embeddings, Unimodal (text), Multilingual (English, 30 programming languages), "69            "8192 input tokens truncation, Prefixes for queries/documents: not necessary, 2024 year."70        ),71        license="apache-2.0",72        size_in_GB=0.64,73        sources=ModelSource(hf="jinaai/jina-embeddings-v2-base-code"),74        model_file="onnx/model.onnx",75    ),76    DenseModelDescription(77        model="jinaai/jina-embeddings-v2-base-zh",78        dim=768,79        description=(80            "Text embeddings, Unimodal (text), supports mixed Chinese-English input text, "81            "8192 input tokens truncation, Prefixes for queries/documents: not necessary, 2024 year."82        ),83        license="apache-2.0",84        size_in_GB=0.64,85        sources=ModelSource(hf="jinaai/jina-embeddings-v2-base-zh"),86        model_file="onnx/model.onnx",87    ),88    DenseModelDescription(89        model="jinaai/jina-embeddings-v2-base-es",90        dim=768,91        description=(92            "Text embeddings, Unimodal (text), supports mixed Spanish-English input text, "93            "8192 input tokens truncation, Prefixes for queries/documents: not necessary, 2024 year."94        ),95        license="apache-2.0",96        size_in_GB=0.64,97        sources=ModelSource(hf="jinaai/jina-embeddings-v2-base-es"),98        model_file="onnx/model.onnx",99    ),100    DenseModelDescription(101        model="thenlper/gte-base",102        dim=768,103        description=(104            "General text embeddings, Unimodal (text), supports English only input text, "105            "512 input tokens truncation, Prefixes for queries/documents: not necessary, 2024 year."106        ),107        license="mit",108        size_in_GB=0.44,109        sources=ModelSource(hf="thenlper/gte-base"),110        model_file="onnx/model.onnx",111    ),112    DenseModelDescription(113        model="thenlper/gte-large",114        dim=1024,115        description=(116            "Text embeddings, Unimodal (text), English, 512 input tokens truncation, "117            "Prefixes for queries/documents: not necessary, 2023 year."118        ),119        license="mit",120        size_in_GB=1.20,121        sources=ModelSource(hf="qdrant/gte-large-onnx"),122        model_file="model.onnx",123    ),124]125 126 127class PooledNormalizedEmbedding(PooledEmbedding):128    @classmethod129    def _get_worker_class(cls) -> Type[OnnxTextEmbeddingWorker]:130        return PooledNormalizedEmbeddingWorker131 132    @classmethod133    def _list_supported_models(cls) -> list[DenseModelDescription]:134        """Lists the supported models.135 136        Returns:137            list[DenseModelDescription]: A list of DenseModelDescription objects containing the model information.138        """139        return supported_pooled_normalized_models140 141    def _post_process_onnx_output(142        self, output: OnnxOutputContext, **kwargs: Any143    ) -> Iterable[NumpyArray]:144        if output.attention_mask is None:145            raise ValueError("attention_mask must be provided for document post-processing")146 147        embeddings = output.model_output148        attn_mask = output.attention_mask149        return normalize(self.mean_pooling(embeddings, attn_mask))150 151 152class PooledNormalizedEmbeddingWorker(OnnxTextEmbeddingWorker):153    def init_embedding(154        self,155        model_name: str,156        cache_dir: str,157        **kwargs: Any,158    ) -> OnnxTextEmbedding:159        return PooledNormalizedEmbedding(160            model_name=model_name,161            cache_dir=cache_dir,162            threads=1,163            **kwargs,164        )165 
codekingpro/portable-devtools · Team Ai