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