diff --git a/fastembed/common/model_management.py b/fastembed/common/model_management.py index ea42687..bf490be 100644 --- a/fastembed/common/model_management.py +++ b/fastembed/common/model_management.py @@ -24,6 +24,7 @@ from huggingface_hub.utils import ( from loguru import logger from tqdm import tqdm from fastembed.common.model_description import BaseModelDescription +from fastembed.common.onnx_external_data import link_external_data T = TypeVar("T", bound=BaseModelDescription) @@ -584,6 +585,19 @@ class ModelManagement(Generic[T]): return model_dir + @staticmethod + def _link_onnx_external_data(model: BaseModelDescription, model_dir: Path) -> Path: + """Links the external data of `model` right after its download, see link_external_data. + + Loading the model links it as well, but linking it here also covers lazy_load=True, e.g. + in a Docker build step, where links made by a later step would copy the files into a new + image layer. A failure, e.g. in a read-only cache, is reported when the model is loaded. + """ + if model.model_file.endswith(".onnx"): + with contextlib.suppress(OSError): + link_external_data(model_dir, model.model_file, model.additional_files) + return model_dir + @classmethod def download_model(cls, model: T, cache_dir: str, retries: int = 3, **kwargs: Any) -> Path: """ @@ -645,7 +659,7 @@ class ModelManagement(Generic[T]): if (resolved_path / model.model_file).exists() and all( (resolved_path / file).exists() for file in extra_patterns ): - return resolved_path + return cls._link_onnx_external_data(model, resolved_path) except CorruptedCacheError: force_download = True except Exception: @@ -667,7 +681,7 @@ class ModelManagement(Generic[T]): attempt_kwargs = {**kwargs, "force_download": True} if force_download else kwargs force_download = False try: - return Path( + model_dir = Path( cls.download_files_from_huggingface( hf_source, cache_dir=cache_dir, @@ -675,6 +689,7 @@ class ModelManagement(Generic[T]): **attempt_kwargs, ) ) + return cls._link_onnx_external_data(model, model_dir) except _HF_DOWNLOAD_ERRORS as e: logger.error( f"Could not download model from HuggingFace: {e} " diff --git a/fastembed/common/onnx_external_data.py b/fastembed/common/onnx_external_data.py new file mode 100644 index 0000000..47fb763 --- /dev/null +++ b/fastembed/common/onnx_external_data.py @@ -0,0 +1,91 @@ +"""Loading ONNX models with external data from a huggingface_hub cache. + +onnxruntime>=1.24 refuses external data that, once symlinks are resolved, is located outside of +the directory of the model file, and of the directory the model file resolves to (a fallback of +1.24.2, pyproject.toml excludes 1.24.0 and 1.24.1). Since huggingface_hub 1.32, a snapshot is made +of symlinks into a blob store shared by the whole cache and sharded by hash, where a model and its +data may resolve into different directories. Then the model and its data are hardlinked into +`ONNX_SNAPSHOTS_DIR` of the repo cache, laid out as in the snapshot, which takes no extra space. +""" + +import contextlib +import os +import shutil +from pathlib import Path + +# Next to `snapshots` in the cache of a huggingface_hub repo: `onnx_snapshots//...` +ONNX_SNAPSHOTS_DIR = "onnx_snapshots" + + +def link_external_data(model_dir: Path, model_file: str, additional_files: list[str]) -> Path: + """Returns a path of `model_file` from which onnxruntime can load its external data. + + Args: + model_dir (Path): The directory with the model files, e.g. a huggingface_hub snapshot. + model_file (str): The path of the ONNX file, relative to `model_dir`. + additional_files (list[str]): Other files of the model, relative to `model_dir`, + among which its external data. + + Returns: + Path: The ONNX file to load, `model_dir / model_file` unless it had to be linked. + + Raises: + OSError: If the files had to be linked, but couldn't be, e.g. in a read-only cache or + on a filesystem without hardlinks. + """ + model_path = model_dir / model_file + snapshots_dir = model_dir.parent + repo_dir = snapshots_dir.parent + if snapshots_dir.name != "snapshots" or not repo_dir.name.startswith("models--"): + return model_path # not a huggingface_hub cache, there is no place of ours to link into + if not model_path.exists(): + return model_path # onnxruntime reports it more clearly than a failed hardlink would + + # external data is located relative to the model file, so it can't be anywhere else + data_paths = [ + path + for path in (model_dir / file for file in additional_files) + if path.parent.is_relative_to(model_path.parent) and path.exists() + ] + # onnxruntime accepts data within these, as in the caches of older huggingface_hub versions + real_model_dirs = ( + os.path.realpath(model_path.parent), + os.path.dirname(os.path.realpath(model_path)), + ) + if all( + any(Path(os.path.realpath(path)).is_relative_to(d) for d in real_model_dirs) + for path in data_paths + ): + return model_path + + links_dir = repo_dir / ONNX_SNAPSHOTS_DIR + for path in (model_path, *data_paths): + _link_file(path, links_dir / model_dir.name / path.relative_to(model_dir)) + + # the links keep the blobs on disk, so drop those of revisions deleted from the cache + for revision_dir in links_dir.iterdir(): + if not (snapshots_dir / revision_dir.name).is_dir(): + shutil.rmtree(revision_dir, ignore_errors=True) + + return links_dir / model_dir.name / model_file + + +def _link_file(source: Path, link: Path) -> None: + """Makes `link` a hardlink to the file `source` resolves to, unless it already is one.""" + target = os.path.realpath(source) + try: + if os.path.samefile(link, target): + return + # a link to a blob downloaded again since, or a copy of the cache without its hardlinks + link.unlink() + except FileNotFoundError: + pass + except OSError: + # e.g. such a copy on a read-only filesystem, which works all the same + if os.path.getsize(link) != os.path.getsize(target): + raise + return + link.parent.mkdir(parents=True, exist_ok=True) + # os.link is atomic, so if the link exists, another process loading the model just made it + with contextlib.suppress(FileExistsError): + os.link(target, link) diff --git a/fastembed/common/onnx_model.py b/fastembed/common/onnx_model.py index aa80dfc..5057d95 100644 --- a/fastembed/common/onnx_model.py +++ b/fastembed/common/onnx_model.py @@ -6,9 +6,11 @@ from typing import Any, Generic, Iterable, Sequence, Type, TypeVar import numpy as np import onnxruntime as ort +from loguru import logger from numpy.typing import NDArray from tokenizers import Tokenizer +from fastembed.common.onnx_external_data import ONNX_SNAPSHOTS_DIR, link_external_data from fastembed.common.types import OnnxProvider, NumpyArray, Device from fastembed.parallel_processor import Worker @@ -78,8 +80,8 @@ class OnnxModel(Generic[T]): cuda: bool | Device = Device.AUTO, device_id: int | None = None, extra_session_options: dict[str, Any] | None = None, + additional_files: list[str] | None = None, ) -> None: - model_path = model_dir / model_file # List of Execution Providers: https://onnxruntime.ai/docs/execution-providers available_providers = ort.get_available_providers() cuda_available = "CUDAExecutionProvider" in available_providers @@ -124,9 +126,29 @@ class OnnxModel(Generic[T]): if extra_session_options is not None: self.add_extra_session_options(so, extra_session_options) - self.model = ort.InferenceSession( - str(model_path), providers=onnx_providers, sess_options=so - ) + model_path = model_dir / model_file + link_error: OSError | None = None + try: + model_path = link_external_data(model_dir, model_file, additional_files or []) + except OSError as e: + # e.g. a read-only cache, which onnxruntime<1.24 loads from all the same + link_error = e + + try: + self.model = ort.InferenceSession( + str(model_path), providers=onnx_providers, sess_options=so + ) + except Exception: + if link_error is not None: + logger.warning( + f"Could not link the files of {model_path} into {ONNX_SNAPSHOTS_DIR}: " + f"{link_error}. onnxruntime>=1.24 refuses external data that resolves " + "outside of the model directory, as in a huggingface_hub>=1.32 cache. " + "Download the model with fastembed into a writable cache_dir on a " + "filesystem with hardlinks, or delete it from the cache and download it " + "again with HF_HUB_DISABLE_SHARED_BLOBS=1." + ) + raise if "CUDAExecutionProvider" in requested_provider_names: assert self.model is not None current_providers = self.model.get_providers() diff --git a/fastembed/image/onnx_embedding.py b/fastembed/image/onnx_embedding.py index 4fd813a..fba875f 100644 --- a/fastembed/image/onnx_embedding.py +++ b/fastembed/image/onnx_embedding.py @@ -137,6 +137,7 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]): cuda=self.cuda, device_id=self.device_id, extra_session_options=self._extra_session_options, + additional_files=self.model_description.additional_files, ) @classmethod diff --git a/fastembed/image/onnx_image_model.py b/fastembed/image/onnx_image_model.py index 68a7251..9f002fa 100644 --- a/fastembed/image/onnx_image_model.py +++ b/fastembed/image/onnx_image_model.py @@ -58,6 +58,7 @@ class OnnxImageModel(OnnxModel[T]): cuda: bool | Device = Device.AUTO, device_id: int | None = None, extra_session_options: dict[str, Any] | None = None, + additional_files: list[str] | None = None, ) -> None: super()._load_onnx_model( model_dir=model_dir, @@ -67,6 +68,7 @@ class OnnxImageModel(OnnxModel[T]): cuda=cuda, device_id=device_id, extra_session_options=extra_session_options, + additional_files=additional_files, ) self.processor = load_preprocessor(model_dir=model_dir) diff --git a/fastembed/late_interaction/colbert.py b/fastembed/late_interaction/colbert.py index 81ebed6..9115115 100644 --- a/fastembed/late_interaction/colbert.py +++ b/fastembed/late_interaction/colbert.py @@ -221,6 +221,7 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]): cuda=self.cuda, device_id=self.device_id, extra_session_options=self._extra_session_options, + additional_files=self.model_description.additional_files, ) def _load_tokenizer(self, model_dir: Path) -> None: diff --git a/fastembed/late_interaction_multimodal/colmodernvbert.py b/fastembed/late_interaction_multimodal/colmodernvbert.py index 8ed6b27..f8fe015 100644 --- a/fastembed/late_interaction_multimodal/colmodernvbert.py +++ b/fastembed/late_interaction_multimodal/colmodernvbert.py @@ -133,6 +133,7 @@ class ColModernVBERT(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel cuda=self.cuda, device_id=self.device_id, extra_session_options=self._extra_session_options, + additional_files=self.model_description.additional_files, ) # Load image processing configuration diff --git a/fastembed/late_interaction_multimodal/colpali.py b/fastembed/late_interaction_multimodal/colpali.py index 24a4399..017906e 100644 --- a/fastembed/late_interaction_multimodal/colpali.py +++ b/fastembed/late_interaction_multimodal/colpali.py @@ -128,6 +128,7 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA cuda=self.cuda, device_id=self.device_id, extra_session_options=self._extra_session_options, + additional_files=self.model_description.additional_files, ) def _post_process_onnx_image_output( diff --git a/fastembed/late_interaction_multimodal/onnx_multimodal_model.py b/fastembed/late_interaction_multimodal/onnx_multimodal_model.py index cc1b1ae..bc76364 100644 --- a/fastembed/late_interaction_multimodal/onnx_multimodal_model.py +++ b/fastembed/late_interaction_multimodal/onnx_multimodal_model.py @@ -65,6 +65,7 @@ class OnnxMultimodalModel(OnnxModel[T]): cuda: bool | Device = Device.AUTO, device_id: int | None = None, extra_session_options: dict[str, Any] | None = None, + additional_files: list[str] | None = None, ) -> None: super()._load_onnx_model( model_dir=model_dir, @@ -74,6 +75,7 @@ class OnnxMultimodalModel(OnnxModel[T]): cuda=cuda, device_id=device_id, extra_session_options=extra_session_options, + additional_files=additional_files, ) self._ensure_tokenizer() assert self.tokenizer is not None diff --git a/fastembed/rerank/cross_encoder/onnx_text_cross_encoder.py b/fastembed/rerank/cross_encoder/onnx_text_cross_encoder.py index 63f8fb6..10af1f0 100644 --- a/fastembed/rerank/cross_encoder/onnx_text_cross_encoder.py +++ b/fastembed/rerank/cross_encoder/onnx_text_cross_encoder.py @@ -171,6 +171,7 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel): cuda=self.cuda, device_id=self.device_id, extra_session_options=self._extra_session_options, + additional_files=self.model_description.additional_files, ) def rerank( diff --git a/fastembed/rerank/cross_encoder/onnx_text_model.py b/fastembed/rerank/cross_encoder/onnx_text_model.py index 3ed09fc..6768d00 100644 --- a/fastembed/rerank/cross_encoder/onnx_text_model.py +++ b/fastembed/rerank/cross_encoder/onnx_text_model.py @@ -34,6 +34,7 @@ class OnnxCrossEncoderModel(OnnxModel[float]): cuda: bool | Device = Device.AUTO, device_id: int | None = None, extra_session_options: dict[str, Any] | None = None, + additional_files: list[str] | None = None, ) -> None: super()._load_onnx_model( model_dir=model_dir, @@ -43,6 +44,7 @@ class OnnxCrossEncoderModel(OnnxModel[float]): cuda=cuda, device_id=device_id, extra_session_options=extra_session_options, + additional_files=additional_files, ) self._ensure_tokenizer() assert self.tokenizer is not None diff --git a/fastembed/sparse/bm42.py b/fastembed/sparse/bm42.py index 9f99bc1..b4685a8 100644 --- a/fastembed/sparse/bm42.py +++ b/fastembed/sparse/bm42.py @@ -160,6 +160,7 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]): cuda=self.cuda, device_id=self.device_id, extra_session_options=self._extra_session_options, + additional_files=self.model_description.additional_files, ) self.stopwords = set(self._load_stopwords(self._model_dir)) diff --git a/fastembed/sparse/if_splade.py b/fastembed/sparse/if_splade.py index 9947d92..4e1c5bd 100644 --- a/fastembed/sparse/if_splade.py +++ b/fastembed/sparse/if_splade.py @@ -160,6 +160,7 @@ class IfSplade(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]): cuda=self.cuda, device_id=self.device_id, extra_session_options=self._extra_session_options, + additional_files=self.model_description.additional_files, ) def _load_idf(self) -> dict[int, float]: diff --git a/fastembed/sparse/minicoil.py b/fastembed/sparse/minicoil.py index 109d8a1..0aabad7 100644 --- a/fastembed/sparse/minicoil.py +++ b/fastembed/sparse/minicoil.py @@ -158,6 +158,7 @@ class MiniCOIL(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]): cuda=self.cuda, device_id=self.device_id, extra_session_options=self._extra_session_options, + additional_files=self.model_description.additional_files, ) def _load_tokenizer(self, model_dir: Path) -> None: diff --git a/fastembed/sparse/splade_pp.py b/fastembed/sparse/splade_pp.py index 562ebcd..1183f21 100644 --- a/fastembed/sparse/splade_pp.py +++ b/fastembed/sparse/splade_pp.py @@ -142,6 +142,7 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]): cuda=self.cuda, device_id=self.device_id, extra_session_options=self._extra_session_options, + additional_files=self.model_description.additional_files, ) def embed( diff --git a/fastembed/text/onnx_embedding.py b/fastembed/text/onnx_embedding.py index 3187673..c112218 100644 --- a/fastembed/text/onnx_embedding.py +++ b/fastembed/text/onnx_embedding.py @@ -388,6 +388,7 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]): cuda=self.cuda, device_id=self.device_id, extra_session_options=self._extra_session_options, + additional_files=self.model_description.additional_files, ) def token_count( diff --git a/fastembed/text/onnx_text_model.py b/fastembed/text/onnx_text_model.py index 74e9ab6..289d599 100644 --- a/fastembed/text/onnx_text_model.py +++ b/fastembed/text/onnx_text_model.py @@ -67,6 +67,7 @@ class OnnxTextModel(OnnxModel[T]): cuda: bool | Device = Device.AUTO, device_id: int | None = None, extra_session_options: dict[str, Any] | None = None, + additional_files: list[str] | None = None, ) -> None: super()._load_onnx_model( model_dir=model_dir, @@ -76,6 +77,7 @@ class OnnxTextModel(OnnxModel[T]): cuda=cuda, device_id=device_id, extra_session_options=extra_session_options, + additional_files=additional_files, ) self._ensure_tokenizer() diff --git a/tests/test_text_onnx_embeddings.py b/tests/test_text_onnx_embeddings.py index 7a2220b..5110a30 100644 --- a/tests/test_text_onnx_embeddings.py +++ b/tests/test_text_onnx_embeddings.py @@ -226,6 +226,17 @@ def test_query_embedding(model_cache) -> None: ), model_desc.model +def test_external_data_model(model_cache) -> None: + # Its weights are ONNX external data, which onnxruntime>=1.24 doesn't load straight from a + # huggingface_hub>=1.32 cache, see fastembed.common.onnx_external_data + model_name = "ibm-granite/granite-embedding-small-english-r2" + with model_cache(model_name) as model: + embedding = next(iter(model.embed(["hello world"]))) + + canonical_vector = CANONICAL_VECTOR_VALUES[model_name] + assert np.allclose(embedding[: canonical_vector.shape[0]], canonical_vector, atol=1e-3) + + def test_quantized_model_reports_onnxruntime_requirement(monkeypatch) -> None: """Old onnxruntime only implements 4-bit MatMulNBits, the error should say so.""" monkeypatch.setattr(