Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
token_embeddings.py84 linesDownload Raw Back to late_interaction
1from dataclasses import asdict2from typing import Iterable, Any, Type3 4from fastembed.common.model_description import DenseModelDescription, ModelSource5from fastembed.common.onnx_model import OnnxOutputContext6from fastembed.common.types import NumpyArray7from fastembed.late_interaction.late_interaction_embedding_base import (8    LateInteractionTextEmbeddingBase,9)10from fastembed.text.onnx_embedding import OnnxTextEmbedding11from fastembed.text.onnx_text_model import TextEmbeddingWorker12 13 14supported_token_embeddings_models = [15    DenseModelDescription(16        model="jinaai/jina-embeddings-v2-small-en-tokens",17        dim=512,18        description="Text embeddings, Unimodal (text), English, 8192 input tokens truncation,"19        " Prefixes for queries/documents: not necessary, 2023 year.",20        license="apache-2.0",21        size_in_GB=0.12,22        sources=ModelSource(hf="xenova/jina-embeddings-v2-small-en"),23        model_file="onnx/model.onnx",24    ),25]26 27 28class TokenEmbeddingsModel(OnnxTextEmbedding, LateInteractionTextEmbeddingBase):29    @classmethod30    def _list_supported_models(cls) -> list[DenseModelDescription]:31        """Lists the supported models.32 33        Returns:34            list[DenseModelDescription]: A list of DenseModelDescription objects containing the model information.35        """36        return supported_token_embeddings_models37 38    @classmethod39    def list_supported_models(cls) -> list[dict[str, Any]]:40        """Lists the supported models.41 42        Returns:43            list[dict[str, Any]]: A list of dictionaries containing the model information.44        """45        return [asdict(model) for model in cls._list_supported_models()]46 47    @classmethod48    def _get_worker_class(cls) -> Type[TextEmbeddingWorker[NumpyArray]]:49        return TokensEmbeddingWorker50 51    def _post_process_onnx_output(52        self, output: OnnxOutputContext, **kwargs: Any53    ) -> Iterable[NumpyArray]:54        # Size: (batch_size, sequence_length, hidden_size)55        embeddings = output.model_output56        # Size: (batch_size, sequence_length)57        assert output.attention_mask is not None58        masks = output.attention_mask59 60        # For each document we only select those embeddings that are not masked out61        for i in range(embeddings.shape[0]):62            yield embeddings[i, masks[i] == 1]63 64    def embed(65        self,66        documents: str | Iterable[str],67        batch_size: int = 256,68        parallel: int | None = None,69        **kwargs: Any,70    ) -> Iterable[NumpyArray]:71        yield from super().embed(documents, batch_size=batch_size, parallel=parallel, **kwargs)72 73 74class TokensEmbeddingWorker(TextEmbeddingWorker[NumpyArray]):75    def init_embedding(76        self, model_name: str, cache_dir: str, **kwargs: Any77    ) -> TokenEmbeddingsModel:78        return TokenEmbeddingsModel(79            model_name=model_name,80            cache_dir=cache_dir,81            threads=1,82            **kwargs,83        )84 
codekingpro/portable-devtools · Team Ai