codekingpro/portable-devtools
115k
1from enum import Enum2from typing import Any, Type, Iterable3 4import numpy as np5 6from fastembed.common.onnx_model import OnnxOutputContext7from fastembed.common.types import NumpyArray8from fastembed.text.pooled_normalized_embedding import PooledNormalizedEmbedding9from fastembed.text.onnx_embedding import OnnxTextEmbeddingWorker10from fastembed.common.model_description import DenseModelDescription, ModelSource11 12supported_multitask_models: list[DenseModelDescription] = [13 DenseModelDescription(14 model="jinaai/jina-embeddings-v3",15 dim=1024,16 tasks={17 "retrieval.query": 0,18 "retrieval.passage": 1,19 "separation": 2,20 "classification": 3,21 "text-matching": 4,22 },23 description=(24 "Multi-task unimodal (text) embedding model, multi-lingual (~100), "25 "1024 tokens truncation, and 8192 sequence length. Prefixes for queries/documents: not necessary, 2024 year."26 ),27 license="cc-by-nc-4.0",28 size_in_GB=2.29,29 sources=ModelSource(hf="jinaai/jina-embeddings-v3"),30 model_file="onnx/model.onnx",31 additional_files=["onnx/model.onnx_data"],32 ),33]34 35 36class Task(int, Enum):37 RETRIEVAL_QUERY = 038 RETRIEVAL_PASSAGE = 139 SEPARATION = 240 CLASSIFICATION = 341 TEXT_MATCHING = 442 43 44class JinaEmbeddingV3(PooledNormalizedEmbedding):45 PASSAGE_TASK = Task.RETRIEVAL_PASSAGE46 QUERY_TASK = Task.RETRIEVAL_QUERY47 48 def __init__(self, *args: Any, task_id: int | None = None, **kwargs: Any):49 super().__init__(*args, **kwargs)50 self.default_task_id: Task | int = task_id if task_id is not None else self.PASSAGE_TASK51 52 @classmethod53 def _get_worker_class(cls) -> Type[OnnxTextEmbeddingWorker]:54 return JinaEmbeddingV3Worker55 56 @classmethod57 def _list_supported_models(cls) -> list[DenseModelDescription]:58 return supported_multitask_models59 60 def _preprocess_onnx_input(61 self,62 onnx_input: dict[str, NumpyArray],63 task_id: int | Task | None = None,64 **kwargs: Any,65 ) -> dict[str, NumpyArray]:66 if task_id is None:67 raise ValueError(f"task_id must be provided for JinaEmbeddingV3, got <{task_id}>")68 onnx_input["task_id"] = np.array(task_id, dtype=np.int64)69 return onnx_input70 71 def embed(72 self,73 documents: str | Iterable[str],74 batch_size: int = 256,75 parallel: int | None = None,76 task_id: int | None = None,77 **kwargs: Any,78 ) -> Iterable[NumpyArray]:79 task_id = (80 task_id if task_id is not None else self.default_task_id81 ) # required for multiprocessing82 yield from super().embed(documents, batch_size, parallel, task_id=task_id, **kwargs)83 84 def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:85 yield from super().embed(query, task_id=self.QUERY_TASK, **kwargs)86 87 def passage_embed(self, texts: Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:88 yield from super().embed(texts, task_id=self.PASSAGE_TASK, **kwargs)89 90 91class JinaEmbeddingV3Worker(OnnxTextEmbeddingWorker):92 def init_embedding(93 self,94 model_name: str,95 cache_dir: str,96 **kwargs: Any,97 ) -> JinaEmbeddingV3:98 return JinaEmbeddingV3(99 model_name=model_name,100 cache_dir=cache_dir,101 threads=1,102 **kwargs,103 )104 105 def process(self, items: Iterable[tuple[int, Any]]) -> Iterable[tuple[int, OnnxOutputContext]]:106 self.model: JinaEmbeddingV3 # mypy complaints `self.model` does not have `default_task_id`107 for idx, batch in items:108 onnx_output = self.model.onnx_embed(batch, task_id=self.model.default_task_id)109 yield idx, onnx_output110 