codekingpro/portable-devtools
114k
1from typing import Any, Iterable, Sequence, Type2 3import numpy as np4from tokenizers import Encoding5 6from fastembed.common import OnnxProvider, ImageInput7from fastembed.common.onnx_model import OnnxOutputContext8from fastembed.common.types import NumpyArray, Device9from fastembed.common.utils import define_cache_dir, iter_batch10from fastembed.late_interaction_multimodal.late_interaction_multimodal_embedding_base import (11 LateInteractionMultimodalEmbeddingBase,12)13from fastembed.late_interaction_multimodal.onnx_multimodal_model import (14 OnnxMultimodalModel,15 TextEmbeddingWorker,16 ImageEmbeddingWorker,17)18from fastembed.common.model_description import DenseModelDescription, ModelSource19 20supported_colpali_models: list[DenseModelDescription] = [21 DenseModelDescription(22 model="Qdrant/colpali-v1.3-fp16",23 dim=128,24 description="Text embeddings, Multimodal (text&image), English, 50 tokens query length truncation, 2024.",25 license="mit",26 size_in_GB=6.5,27 sources=ModelSource(hf="Qdrant/colpali-v1.3-fp16"),28 additional_files=["model.onnx_data"],29 model_file="model.onnx",30 ),31]32 33 34class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyArray]):35 QUERY_PREFIX = "Query: "36 BOS_TOKEN = "<s>"37 PAD_TOKEN = "<pad>"38 QUERY_MARKER_TOKEN_ID = [2, 5098]39 IMAGE_PLACEHOLDER_SIZE = (3, 448, 448)40 EMPTY_TEXT_PLACEHOLDER = np.array(41 [257152] * 1024 + [2, 50721, 573, 2416, 235265, 108]42 ) # This is a tokenization of '<image>' * 1024 + '<bos>Describe the image.\n' line which is used as placeholder43 # while processing an image44 EVEN_ATTENTION_MASK = np.array([1] * 1030)45 46 def __init__(47 self,48 model_name: str,49 cache_dir: str | None = None,50 threads: int | None = None,51 providers: Sequence[OnnxProvider] | None = None,52 cuda: bool | Device = Device.AUTO,53 device_ids: list[int] | None = None,54 lazy_load: bool = False,55 device_id: int | None = None,56 specific_model_path: str | None = None,57 **kwargs: Any,58 ):59 """60 Args:61 model_name (str): The name of the model to use.62 cache_dir (str, optional): The path to the cache directory.63 Can be set using the `FASTEMBED_CACHE_PATH` env variable.64 Defaults to `fastembed_cache` in the system's temp directory.65 threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.66 providers (Optional[Sequence[OnnxProvider]], optional): The list of onnxruntime providers to use.67 Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.68 cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`69 Defaults to Device.AUTO.70 device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in71 workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, mutually exclusive72 with `providers`. Defaults to None.73 lazy_load (bool, optional): Whether to load the model during class initialization or on demand.74 Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.75 device_id (Optional[int], optional): The device id to use for loading the model in the worker process.76 77 Raises:78 ValueError: If the model_name is not in the format <org>/<model> e.g. BAAI/bge-base-en.79 """80 81 super().__init__(model_name, cache_dir, threads, **kwargs)82 self.providers = providers83 self.lazy_load = lazy_load84 self._extra_session_options = self._select_exposed_session_options(kwargs)85 86 # List of device ids, that can be used for data parallel processing in workers87 self.device_ids = device_ids88 self.cuda = cuda89 90 # This device_id will be used if we need to load model in current process91 self.device_id: int | None = None92 if device_id is not None:93 self.device_id = device_id94 elif self.device_ids is not None:95 self.device_id = self.device_ids[0]96 97 self.model_description = self._get_model_description(model_name)98 self.cache_dir = str(define_cache_dir(cache_dir))99 100 self._specific_model_path = specific_model_path101 self._model_dir = self.download_model(102 self.model_description,103 self.cache_dir,104 local_files_only=self._local_files_only,105 specific_model_path=self._specific_model_path,106 )107 self.mask_token_id = None108 self.pad_token_id = None109 110 if not self.lazy_load:111 self.load_onnx_model()112 113 @classmethod114 def _list_supported_models(cls) -> list[DenseModelDescription]:115 """Lists the supported models.116 117 Returns:118 list[DenseModelDescription]: A list of DenseModelDescription objects containing the model information.119 """120 return supported_colpali_models121 122 def load_onnx_model(self) -> None:123 self._load_onnx_model(124 model_dir=self._model_dir,125 model_file=self.model_description.model_file,126 threads=self.threads,127 providers=self.providers,128 cuda=self.cuda,129 device_id=self.device_id,130 extra_session_options=self._extra_session_options,131 )132 133 def _post_process_onnx_image_output(134 self,135 output: OnnxOutputContext,136 ) -> Iterable[NumpyArray]:137 """138 Post-process the ONNX model output to convert it into a usable format.139 140 Args:141 output (OnnxOutputContext): The raw output from the ONNX model.142 143 Returns:144 Iterable[NumpyArray]: Post-processed output as NumPy arrays.145 """146 assert self.model_description.dim is not None, "Model dim is not defined"147 return output.model_output.reshape(148 output.model_output.shape[0], -1, self.model_description.dim149 )150 151 def _post_process_onnx_text_output(152 self,153 output: OnnxOutputContext,154 ) -> Iterable[NumpyArray]:155 """156 Post-process the ONNX model output to convert it into a usable format.157 158 Args:159 output (OnnxOutputContext): The raw output from the ONNX model.160 161 Returns:162 Iterable[NumpyArray]: Post-processed output as NumPy arrays.163 """164 return output.model_output165 166 def tokenize(self, documents: list[str], **kwargs: Any) -> list[Encoding]:167 texts_query: list[str] = []168 for query in documents:169 query = self.BOS_TOKEN + self.QUERY_PREFIX + query + self.PAD_TOKEN * 10170 query += "\n"171 172 texts_query.append(query)173 encoded = self.tokenizer.encode_batch(texts_query) # type: ignore[union-attr]174 return encoded175 176 def token_count(177 self,178 texts: str | Iterable[str],179 batch_size: int = 1024,180 include_extension: bool = False,181 **kwargs: Any,182 ) -> int:183 if not hasattr(self, "model") or self.model is None:184 self.load_onnx_model() # loads the tokenizer as well185 token_num = 0186 texts = [texts] if isinstance(texts, str) else texts187 assert self.tokenizer is not None188 tokenize_func = self.tokenize if include_extension else self.tokenizer.encode_batch189 for batch in iter_batch(texts, batch_size):190 token_num += sum([sum(encoding.attention_mask) for encoding in tokenize_func(batch)])191 return token_num192 193 def _preprocess_onnx_text_input(194 self, onnx_input: dict[str, NumpyArray], **kwargs: Any195 ) -> dict[str, NumpyArray]:196 onnx_input["input_ids"] = np.array(197 [198 self.QUERY_MARKER_TOKEN_ID + input_ids[2:].tolist() # type: ignore[index]199 for input_ids in onnx_input["input_ids"]200 ]201 )202 empty_image_placeholder: NumpyArray = np.zeros(203 self.IMAGE_PLACEHOLDER_SIZE, dtype=np.float32204 )205 onnx_input["pixel_values"] = np.array(206 [empty_image_placeholder for _ in onnx_input["input_ids"]],207 )208 return onnx_input209 210 def _preprocess_onnx_image_input(211 self, onnx_input: dict[str, np.ndarray], **kwargs: Any212 ) -> dict[str, NumpyArray]:213 """214 Add placeholders for text input when processing image data for ONNX.215 Args:216 onnx_input (Dict[str, NumpyArray]): Preprocessed image inputs.217 **kwargs: Additional arguments.218 Returns:219 Dict[str, NumpyArray]: ONNX input with text placeholders.220 """221 onnx_input["input_ids"] = np.array(222 [self.EMPTY_TEXT_PLACEHOLDER for _ in onnx_input["pixel_values"]]223 )224 onnx_input["attention_mask"] = np.array(225 [self.EVEN_ATTENTION_MASK for _ in onnx_input["pixel_values"]]226 )227 return onnx_input228 229 def embed_text(230 self,231 documents: str | Iterable[str],232 batch_size: int = 256,233 parallel: int | None = None,234 **kwargs: Any,235 ) -> Iterable[NumpyArray]:236 """237 Encode a list of documents into list of embeddings.238 239 Args:240 documents: Iterator of documents or single document to embed241 batch_size: Batch size for encoding -- higher values will use more memory, but be faster242 parallel:243 If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.244 If 0, use all available cores.245 If None, don't use data-parallel processing, use default onnxruntime threading instead.246 247 Returns:248 List of embeddings, one per document249 """250 yield from self._embed_documents(251 model_name=self.model_name,252 cache_dir=str(self.cache_dir),253 documents=documents,254 batch_size=batch_size,255 parallel=parallel,256 providers=self.providers,257 cuda=self.cuda,258 device_ids=self.device_ids,259 local_files_only=self._local_files_only,260 specific_model_path=self._specific_model_path,261 extra_session_options=self._extra_session_options,262 **kwargs,263 )264 265 def embed_image(266 self,267 images: ImageInput | Iterable[ImageInput],268 batch_size: int = 16,269 parallel: int | None = None,270 **kwargs: Any,271 ) -> Iterable[NumpyArray]:272 """273 Encode a list of images into list of embeddings.274 275 Args:276 images: Iterator of image paths or single image path to embed277 batch_size: Batch size for encoding -- higher values will use more memory, but be faster278 parallel:279 If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.280 If 0, use all available cores.281 If None, don't use data-parallel processing, use default onnxruntime threading instead.282 283 Returns:284 List of embeddings, one per document285 """286 yield from self._embed_images(287 model_name=self.model_name,288 cache_dir=str(self.cache_dir),289 images=images,290 batch_size=batch_size,291 parallel=parallel,292 providers=self.providers,293 cuda=self.cuda,294 device_ids=self.device_ids,295 local_files_only=self._local_files_only,296 specific_model_path=self._specific_model_path,297 extra_session_options=self._extra_session_options,298 **kwargs,299 )300 301 @classmethod302 def _get_text_worker_class(cls) -> Type[TextEmbeddingWorker[NumpyArray]]:303 return ColPaliTextEmbeddingWorker304 305 @classmethod306 def _get_image_worker_class(cls) -> Type[ImageEmbeddingWorker[NumpyArray]]:307 return ColPaliImageEmbeddingWorker308 309 310class ColPaliTextEmbeddingWorker(TextEmbeddingWorker[NumpyArray]):311 def init_embedding(self, model_name: str, cache_dir: str, **kwargs: Any) -> ColPali:312 return ColPali(313 model_name=model_name,314 cache_dir=cache_dir,315 threads=1,316 **kwargs,317 )318 319 320class ColPaliImageEmbeddingWorker(ImageEmbeddingWorker[NumpyArray]):321 def init_embedding(self, model_name: str, cache_dir: str, **kwargs: Any) -> ColPali:322 return ColPali(323 model_name=model_name,324 cache_dir=cache_dir,325 threads=1,326 **kwargs,327 )328 