From 0e17bea681d4ef104112052f1713d6c91603df7c Mon Sep 17 00:00:00 2001 From: George Panchuk Date: Mon, 24 Nov 2025 19:14:01 +0700 Subject: [PATCH 1/4] new: expose some onnx session options --- fastembed/common/onnx_model.py | 37 +++++++++++++++++++ fastembed/image/onnx_embedding.py | 2 + fastembed/image/onnx_image_model.py | 4 ++ fastembed/late_interaction/colbert.py | 3 ++ .../late_interaction_multimodal/colpali.py | 4 ++ .../onnx_multimodal_model.py | 6 +++ .../cross_encoder/onnx_text_cross_encoder.py | 4 ++ .../rerank/cross_encoder/onnx_text_model.py | 4 ++ fastembed/sparse/bm42.py | 3 ++ fastembed/sparse/minicoil.py | 2 + fastembed/sparse/splade_pp.py | 3 ++ fastembed/text/onnx_embedding.py | 4 +- fastembed/text/onnx_text_model.py | 4 ++ 13 files changed, 79 insertions(+), 1 deletion(-) diff --git a/fastembed/common/onnx_model.py b/fastembed/common/onnx_model.py index 1d6115417..00dbb9646 100644 --- a/fastembed/common/onnx_model.py +++ b/fastembed/common/onnx_model.py @@ -24,6 +24,8 @@ class OnnxOutputContext: class OnnxModel(Generic[T]): + EXPOSED_SESSION_OPTIONS = ("enable_cpu_mem_arena",) + @classmethod def _get_worker_class(cls) -> Type["EmbeddingWorker[T]"]: raise NotImplementedError("Subclasses must implement this method") @@ -60,6 +62,7 @@ def _load_onnx_model( providers: Optional[Sequence[OnnxProvider]] = None, cuda: bool = False, device_id: Optional[int] = None, + extra_session_options: Optional[dict[str, Any]] = None, ) -> None: model_path = model_dir / model_file # List of Execution Providers: https://onnxruntime.ai/docs/execution-providers @@ -99,6 +102,8 @@ def _load_onnx_model( so.intra_op_num_threads = threads so.inter_op_num_threads = threads + self.add_extra_session_options(so, extra_session_options) + self.model = ort.InferenceSession( str(model_path), providers=onnx_providers, sess_options=so ) @@ -113,6 +118,38 @@ def _load_onnx_model( RuntimeWarning, ) + @classmethod + def _select_exposed_session_options(cls, model_kwargs: dict[str, Any]) -> dict[str, Any]: + """A convenience method to select the exposed session options in models + + Args: + model_kwargs (dict[str, Any]): The model kwargs. + + Returns: + dict[str, Any]: a dict with filtered exposed session options. + """ + return {k: v for k, v in model_kwargs.items() if k in cls.EXPOSED_SESSION_OPTIONS} + + @classmethod + def add_extra_session_options( + cls, session_options: ort.SessionOptions, extra_options: dict[str, Any] + ) -> None: + """Add extra session options to the existing options object in-place + + Args: + session_options (ort.SessionOptions): The existing session options object. + extra_options (dict[str, Any]): The extra session options available in cls.EXPOSED_SESSION_OPTIONS. + + Returns: + None + """ + for option in extra_options: + assert ( + option in cls.EXPOSED_SESSION_OPTIONS + ), f"{option} is unknown or not exposed (exposed options: {cls.EXPOSED_SESSION_OPTIONS})" + if "enable_cpu_mem_arena" in extra_options: + session_options.enable_cpu_mem_arena = extra_options["enable_cpu_mem_arena"] + def load_onnx_model(self) -> None: raise NotImplementedError("Subclasses must implement this method") diff --git a/fastembed/image/onnx_embedding.py b/fastembed/image/onnx_embedding.py index 3b83b2483..36d0d6449 100644 --- a/fastembed/image/onnx_embedding.py +++ b/fastembed/image/onnx_embedding.py @@ -98,6 +98,7 @@ def __init__( super().__init__(model_name, cache_dir, threads, **kwargs) self.providers = providers self.lazy_load = lazy_load + self._extra_session_options = self._select_exposed_session_options(kwargs) # List of device ids, that can be used for data parallel processing in workers self.device_ids = device_ids @@ -134,6 +135,7 @@ def load_onnx_model(self) -> None: providers=self.providers, cuda=self.cuda, device_id=self.device_id, + extra_session_options=self._extra_session_options, ) @classmethod diff --git a/fastembed/image/onnx_image_model.py b/fastembed/image/onnx_image_model.py index a345f024c..03db4c554 100644 --- a/fastembed/image/onnx_image_model.py +++ b/fastembed/image/onnx_image_model.py @@ -55,6 +55,7 @@ def _load_onnx_model( providers: Optional[Sequence[OnnxProvider]] = None, cuda: bool = False, device_id: Optional[int] = None, + extra_session_options: Optional[dict[str, Any]] = None, ) -> None: super()._load_onnx_model( model_dir=model_dir, @@ -63,6 +64,7 @@ def _load_onnx_model( providers=providers, cuda=cuda, device_id=device_id, + extra_session_options=extra_session_options, ) self.processor = load_preprocessor(model_dir=model_dir) @@ -99,6 +101,7 @@ def _embed_images( device_ids: Optional[list[int]] = None, local_files_only: bool = False, specific_model_path: Optional[str] = None, + extra_session_options: Optional[dict[str, Any]] = None, **kwargs: Any, ) -> Iterable[T]: is_small = False @@ -127,6 +130,7 @@ def _embed_images( "providers": providers, "local_files_only": local_files_only, "specific_model_path": specific_model_path, + **extra_session_options, **kwargs, } diff --git a/fastembed/late_interaction/colbert.py b/fastembed/late_interaction/colbert.py index 3aa105ae6..4dfc2a058 100644 --- a/fastembed/late_interaction/colbert.py +++ b/fastembed/late_interaction/colbert.py @@ -143,6 +143,7 @@ def __init__( super().__init__(model_name, cache_dir, threads, **kwargs) self.providers = providers self.lazy_load = lazy_load + self._extra_session_options = self._select_exposed_session_options(kwargs) # List of device ids, that can be used for data parallel processing in workers self.device_ids = device_ids @@ -182,6 +183,7 @@ def load_onnx_model(self) -> None: providers=self.providers, cuda=self.cuda, device_id=self.device_id, + extra_session_options=self._extra_session_options, ) self.query_tokenizer, _ = load_tokenizer(model_dir=self._model_dir) @@ -235,6 +237,7 @@ def embed( device_ids=self.device_ids, local_files_only=self._local_files_only, specific_model_path=self._specific_model_path, + extra_session_options=self._extra_session_options, **kwargs, ) diff --git a/fastembed/late_interaction_multimodal/colpali.py b/fastembed/late_interaction_multimodal/colpali.py index 9fab95359..059f79716 100644 --- a/fastembed/late_interaction_multimodal/colpali.py +++ b/fastembed/late_interaction_multimodal/colpali.py @@ -80,6 +80,7 @@ def __init__( super().__init__(model_name, cache_dir, threads, **kwargs) self.providers = providers self.lazy_load = lazy_load + self._extra_session_options = self._select_exposed_session_options(kwargs) # List of device ids, that can be used for data parallel processing in workers self.device_ids = device_ids @@ -125,6 +126,7 @@ def load_onnx_model(self) -> None: providers=self.providers, cuda=self.cuda, device_id=self.device_id, + extra_session_options=self._extra_session_options, ) def _post_process_onnx_image_output( @@ -238,6 +240,7 @@ def embed_text( device_ids=self.device_ids, local_files_only=self._local_files_only, specific_model_path=self._specific_model_path, + extra_session_options=self._extra_session_options, **kwargs, ) @@ -273,6 +276,7 @@ def embed_image( device_ids=self.device_ids, local_files_only=self._local_files_only, specific_model_path=self._specific_model_path, + extra_session_options=self._extra_session_options, **kwargs, ) diff --git a/fastembed/late_interaction_multimodal/onnx_multimodal_model.py b/fastembed/late_interaction_multimodal/onnx_multimodal_model.py index 83706a2b4..9266f8b4f 100644 --- a/fastembed/late_interaction_multimodal/onnx_multimodal_model.py +++ b/fastembed/late_interaction_multimodal/onnx_multimodal_model.py @@ -64,6 +64,7 @@ def _load_onnx_model( providers: Optional[Sequence[OnnxProvider]] = None, cuda: bool = False, device_id: Optional[int] = None, + extra_session_options: Optional[dict[str, Any]] = None, ) -> None: super()._load_onnx_model( model_dir=model_dir, @@ -72,6 +73,7 @@ def _load_onnx_model( providers=providers, cuda=cuda, device_id=device_id, + extra_session_options=extra_session_options, ) self.tokenizer, self.special_token_to_id = load_tokenizer(model_dir=model_dir) assert self.tokenizer is not None @@ -122,6 +124,7 @@ def _embed_documents( device_ids: Optional[list[int]] = None, local_files_only: bool = False, specific_model_path: Optional[str] = None, + extra_session_options: Optional[dict[str, Any]] = None, **kwargs: Any, ) -> Iterable[T]: is_small = False @@ -150,6 +153,7 @@ def _embed_documents( "providers": providers, "local_files_only": local_files_only, "specific_model_path": specific_model_path, + **extra_session_options, **kwargs, } @@ -189,6 +193,7 @@ def _embed_images( device_ids: Optional[list[int]] = None, local_files_only: bool = False, specific_model_path: Optional[str] = None, + extra_session_options: Optional[dict[str, Any]] = None, **kwargs: Any, ) -> Iterable[T]: is_small = False @@ -217,6 +222,7 @@ def _embed_images( "providers": providers, "local_files_only": local_files_only, "specific_model_path": specific_model_path, + **extra_session_options, **kwargs, } diff --git a/fastembed/rerank/cross_encoder/onnx_text_cross_encoder.py b/fastembed/rerank/cross_encoder/onnx_text_cross_encoder.py index a171afa40..fdb4298c7 100644 --- a/fastembed/rerank/cross_encoder/onnx_text_cross_encoder.py +++ b/fastembed/rerank/cross_encoder/onnx_text_cross_encoder.py @@ -85,6 +85,7 @@ def __init__( lazy_load: bool = False, device_id: Optional[int] = None, specific_model_path: Optional[str] = None, + extra_session_options: Optional[dict[str, Any]] = None, **kwargs: Any, ): """ @@ -111,6 +112,7 @@ def __init__( super().__init__(model_name, cache_dir, threads, **kwargs) self.providers = providers self.lazy_load = lazy_load + self._extra_session_options = self._select_exposed_session_options(kwargs) # List of device ids, that can be used for data parallel processing in workers self.device_ids = device_ids @@ -150,6 +152,7 @@ def load_onnx_model(self) -> None: providers=self.providers, cuda=self.cuda, device_id=self.device_id, + extra_session_options=self._extra_session_options, ) def rerank( @@ -192,6 +195,7 @@ def rerank_pairs( device_ids=self.device_ids, local_files_only=self._local_files_only, specific_model_path=self._specific_model_path, + extra_session_options=self._extra_session_options, **kwargs, ) diff --git a/fastembed/rerank/cross_encoder/onnx_text_model.py b/fastembed/rerank/cross_encoder/onnx_text_model.py index 3fc4e81c4..8bc8fd691 100644 --- a/fastembed/rerank/cross_encoder/onnx_text_model.py +++ b/fastembed/rerank/cross_encoder/onnx_text_model.py @@ -33,6 +33,7 @@ def _load_onnx_model( providers: Optional[Sequence[OnnxProvider]] = None, cuda: bool = False, device_id: Optional[int] = None, + extra_session_options: Optional[dict[str, Any]] = None, ) -> None: super()._load_onnx_model( model_dir=model_dir, @@ -41,6 +42,7 @@ def _load_onnx_model( providers=providers, cuda=cuda, device_id=device_id, + extra_session_options=extra_session_options, ) self.tokenizer, _ = load_tokenizer(model_dir=model_dir) assert self.tokenizer is not None @@ -96,6 +98,7 @@ def _rerank_pairs( device_ids: Optional[list[int]] = None, local_files_only: bool = False, specific_model_path: Optional[str] = None, + extra_session_options: Optional[dict[str, Any]] = None, **kwargs: Any, ) -> Iterable[float]: is_small = False @@ -124,6 +127,7 @@ def _rerank_pairs( "providers": providers, "local_files_only": local_files_only, "specific_model_path": specific_model_path, + **extra_session_options, **kwargs, } diff --git a/fastembed/sparse/bm42.py b/fastembed/sparse/bm42.py index 3e51404f6..848b17531 100644 --- a/fastembed/sparse/bm42.py +++ b/fastembed/sparse/bm42.py @@ -103,6 +103,7 @@ def __init__( super().__init__(model_name, cache_dir, threads, **kwargs) self.providers = providers self.lazy_load = lazy_load + self._extra_session_options = self._select_exposed_session_options(kwargs) # List of device ids, that can be used for data parallel processing in workers self.device_ids = device_ids @@ -146,6 +147,7 @@ def load_onnx_model(self) -> None: providers=self.providers, cuda=self.cuda, device_id=self.device_id, + extra_session_options=self._extra_session_options, ) for token, idx in self.tokenizer.get_vocab().items(): # type: ignore[union-attr] @@ -312,6 +314,7 @@ def embed( alpha=self.alpha, local_files_only=self._local_files_only, specific_model_path=self._specific_model_path, + extra_session_options=self._extra_session_options, ) @classmethod diff --git a/fastembed/sparse/minicoil.py b/fastembed/sparse/minicoil.py index efaa9abbd..47a01f3fc 100644 --- a/fastembed/sparse/minicoil.py +++ b/fastembed/sparse/minicoil.py @@ -117,6 +117,8 @@ def __init__( self.device_ids = device_ids self.cuda = cuda self.device_id = device_id + self._extra_session_options = self._select_exposed_session_options(kwargs) + self.k = k self.b = b self.avg_len = avg_len diff --git a/fastembed/sparse/splade_pp.py b/fastembed/sparse/splade_pp.py index d2c4af38a..95e43bb2f 100644 --- a/fastembed/sparse/splade_pp.py +++ b/fastembed/sparse/splade_pp.py @@ -99,6 +99,7 @@ def __init__( super().__init__(model_name, cache_dir, threads, **kwargs) self.providers = providers self.lazy_load = lazy_load + self._extra_session_options = self._select_exposed_session_options(kwargs) # List of device ids, that can be used for data parallel processing in workers self.device_ids = device_ids @@ -133,6 +134,7 @@ def load_onnx_model(self) -> None: providers=self.providers, cuda=self.cuda, device_id=self.device_id, + extra_session_options=self._extra_session_options, ) def embed( @@ -168,6 +170,7 @@ def embed( device_ids=self.device_ids, local_files_only=self._local_files_only, specific_model_path=self._specific_model_path, + extra_session_options=self._extra_session_options, **kwargs, ) diff --git a/fastembed/text/onnx_embedding.py b/fastembed/text/onnx_embedding.py index 4cc892f59..d76db8bf4 100644 --- a/fastembed/text/onnx_embedding.py +++ b/fastembed/text/onnx_embedding.py @@ -233,7 +233,7 @@ def __init__( super().__init__(model_name, cache_dir, threads, **kwargs) self.providers = providers self.lazy_load = lazy_load - + self._extra_session_options = self._select_exposed_session_options(kwargs) # List of device ids, that can be used for data parallel processing in workers self.device_ids = device_ids self.cuda = cuda @@ -291,6 +291,7 @@ def embed( device_ids=self.device_ids, local_files_only=self._local_files_only, specific_model_path=self._specific_model_path, + extra_session_options=self._extra_session_options, **kwargs, ) @@ -327,6 +328,7 @@ def load_onnx_model(self) -> None: providers=self.providers, cuda=self.cuda, device_id=self.device_id, + extra_session_options=self._extra_session_options, ) diff --git a/fastembed/text/onnx_text_model.py b/fastembed/text/onnx_text_model.py index c939b21d5..46579376f 100644 --- a/fastembed/text/onnx_text_model.py +++ b/fastembed/text/onnx_text_model.py @@ -54,6 +54,7 @@ def _load_onnx_model( providers: Optional[Sequence[OnnxProvider]] = None, cuda: bool = False, device_id: Optional[int] = None, + extra_session_options: Optional[dict[str, Any]] = None, ) -> None: super()._load_onnx_model( model_dir=model_dir, @@ -62,6 +63,7 @@ def _load_onnx_model( providers=providers, cuda=cuda, device_id=device_id, + extra_session_options=extra_session_options, ) self.tokenizer, self.special_token_to_id = load_tokenizer(model_dir=model_dir) @@ -110,6 +112,7 @@ def _embed_documents( device_ids: Optional[list[int]] = None, local_files_only: bool = False, specific_model_path: Optional[str] = None, + extra_session_options: Optional[dict[str, Any]] = None, **kwargs: Any, ) -> Iterable[T]: is_small = False @@ -140,6 +143,7 @@ def _embed_documents( "providers": providers, "local_files_only": local_files_only, "specific_model_path": specific_model_path, + **extra_session_options, **kwargs, } From f4585626545b070c4b51da98c8c50e9f5485eebb Mon Sep 17 00:00:00 2001 From: George Panchuk Date: Mon, 24 Nov 2025 19:22:57 +0700 Subject: [PATCH 2/4] fix: fix extra session options is None case --- fastembed/common/onnx_model.py | 3 ++- fastembed/image/onnx_image_model.py | 4 +++- .../late_interaction_multimodal/onnx_multimodal_model.py | 8 ++++++-- fastembed/rerank/cross_encoder/onnx_text_model.py | 4 +++- fastembed/text/onnx_text_model.py | 4 +++- 5 files changed, 17 insertions(+), 6 deletions(-) diff --git a/fastembed/common/onnx_model.py b/fastembed/common/onnx_model.py index 00dbb9646..1b589e182 100644 --- a/fastembed/common/onnx_model.py +++ b/fastembed/common/onnx_model.py @@ -102,7 +102,8 @@ def _load_onnx_model( so.intra_op_num_threads = threads so.inter_op_num_threads = threads - self.add_extra_session_options(so, extra_session_options) + 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 diff --git a/fastembed/image/onnx_image_model.py b/fastembed/image/onnx_image_model.py index 03db4c554..2f4de833f 100644 --- a/fastembed/image/onnx_image_model.py +++ b/fastembed/image/onnx_image_model.py @@ -130,10 +130,12 @@ def _embed_images( "providers": providers, "local_files_only": local_files_only, "specific_model_path": specific_model_path, - **extra_session_options, **kwargs, } + if extra_session_options is not None: + params.update(extra_session_options) + pool = ParallelWorkerPool( num_workers=parallel or 1, worker=self._get_worker_class(), diff --git a/fastembed/late_interaction_multimodal/onnx_multimodal_model.py b/fastembed/late_interaction_multimodal/onnx_multimodal_model.py index 9266f8b4f..50a5d8b48 100644 --- a/fastembed/late_interaction_multimodal/onnx_multimodal_model.py +++ b/fastembed/late_interaction_multimodal/onnx_multimodal_model.py @@ -153,10 +153,12 @@ def _embed_documents( "providers": providers, "local_files_only": local_files_only, "specific_model_path": specific_model_path, - **extra_session_options, **kwargs, } + if extra_session_options is not None: + params.update(extra_session_options) + pool = ParallelWorkerPool( num_workers=parallel or 1, worker=self._get_text_worker_class(), @@ -222,10 +224,12 @@ def _embed_images( "providers": providers, "local_files_only": local_files_only, "specific_model_path": specific_model_path, - **extra_session_options, **kwargs, } + if extra_session_options is not None: + params.update(extra_session_options) + pool = ParallelWorkerPool( num_workers=parallel or 1, worker=self._get_image_worker_class(), diff --git a/fastembed/rerank/cross_encoder/onnx_text_model.py b/fastembed/rerank/cross_encoder/onnx_text_model.py index 8bc8fd691..5c85d27ec 100644 --- a/fastembed/rerank/cross_encoder/onnx_text_model.py +++ b/fastembed/rerank/cross_encoder/onnx_text_model.py @@ -127,10 +127,12 @@ def _rerank_pairs( "providers": providers, "local_files_only": local_files_only, "specific_model_path": specific_model_path, - **extra_session_options, **kwargs, } + if extra_session_options is not None: + params.update(extra_session_options) + pool = ParallelWorkerPool( num_workers=parallel or 1, worker=self._get_worker_class(), diff --git a/fastembed/text/onnx_text_model.py b/fastembed/text/onnx_text_model.py index 46579376f..6cb491781 100644 --- a/fastembed/text/onnx_text_model.py +++ b/fastembed/text/onnx_text_model.py @@ -143,10 +143,12 @@ def _embed_documents( "providers": providers, "local_files_only": local_files_only, "specific_model_path": specific_model_path, - **extra_session_options, **kwargs, } + if extra_session_options is not None: + params.update(extra_session_options) + pool = ParallelWorkerPool( num_workers=parallel or 1, worker=self._get_worker_class(), From 298912a38160e5c178a2f4992c8cc8bfdbc0937f Mon Sep 17 00:00:00 2001 From: George Panchuk Date: Mon, 24 Nov 2025 19:37:02 +0700 Subject: [PATCH 3/4] fix: fix missing params --- fastembed/image/onnx_embedding.py | 1 + fastembed/rerank/cross_encoder/onnx_text_cross_encoder.py | 1 - fastembed/sparse/minicoil.py | 2 ++ 3 files changed, 3 insertions(+), 1 deletion(-) diff --git a/fastembed/image/onnx_embedding.py b/fastembed/image/onnx_embedding.py index 36d0d6449..ae0ce848c 100644 --- a/fastembed/image/onnx_embedding.py +++ b/fastembed/image/onnx_embedding.py @@ -182,6 +182,7 @@ def embed( device_ids=self.device_ids, local_files_only=self._local_files_only, specific_model_path=self._specific_model_path, + extra_session_options=self._extra_session_options, **kwargs, ) diff --git a/fastembed/rerank/cross_encoder/onnx_text_cross_encoder.py b/fastembed/rerank/cross_encoder/onnx_text_cross_encoder.py index fdb4298c7..56f2b86c1 100644 --- a/fastembed/rerank/cross_encoder/onnx_text_cross_encoder.py +++ b/fastembed/rerank/cross_encoder/onnx_text_cross_encoder.py @@ -85,7 +85,6 @@ def __init__( lazy_load: bool = False, device_id: Optional[int] = None, specific_model_path: Optional[str] = None, - extra_session_options: Optional[dict[str, Any]] = None, **kwargs: Any, ): """ diff --git a/fastembed/sparse/minicoil.py b/fastembed/sparse/minicoil.py index 47a01f3fc..dde52d90d 100644 --- a/fastembed/sparse/minicoil.py +++ b/fastembed/sparse/minicoil.py @@ -155,6 +155,7 @@ def load_onnx_model(self) -> None: providers=self.providers, cuda=self.cuda, device_id=self.device_id, + extra_session_options=self._extra_session_options, ) assert self.tokenizer is not None @@ -223,6 +224,7 @@ def embed( is_query=False, local_files_only=self._local_files_only, specific_model_path=self._specific_model_path, + extra_session_options=self._extra_session_options, **kwargs, ) From 2ef4d17c062d75e59c2b48d464ba725727bd6283 Mon Sep 17 00:00:00 2001 From: George Panchuk Date: Tue, 25 Nov 2025 12:03:28 +0700 Subject: [PATCH 4/4] new: add tests --- tests/test_image_onnx_embeddings.py | 10 ++++++++++ tests/test_late_interaction_embeddings.py | 10 ++++++++++ tests/test_sparse_embeddings.py | 24 ++++++++++++++++++++++- tests/test_text_cross_encoder.py | 10 ++++++++++ tests/test_text_onnx_embeddings.py | 10 ++++++++++ 5 files changed, 63 insertions(+), 1 deletion(-) diff --git a/tests/test_image_onnx_embeddings.py b/tests/test_image_onnx_embeddings.py index 8369702ec..5ac5e44f1 100644 --- a/tests/test_image_onnx_embeddings.py +++ b/tests/test_image_onnx_embeddings.py @@ -163,3 +163,13 @@ def test_embedding_size() -> None: assert model.embedding_size == 512 if is_ci: delete_model_cache(model.model._model_dir) + + +@pytest.mark.parametrize("model_name", ["Qdrant/clip-ViT-B-32-vision"]) +def test_session_options(model_cache, model_name) -> None: + with model_cache(model_name) as default_model: + default_session_options = default_model.model.model.get_session_options() + assert default_session_options.enable_cpu_mem_arena is True + model = ImageEmbedding(model_name=model_name, enable_cpu_mem_arena=False) + session_options = model.model.model.get_session_options() + assert session_options.enable_cpu_mem_arena is False diff --git a/tests/test_late_interaction_embeddings.py b/tests/test_late_interaction_embeddings.py index f89882f43..f2499db86 100644 --- a/tests/test_late_interaction_embeddings.py +++ b/tests/test_late_interaction_embeddings.py @@ -308,3 +308,13 @@ def test_embedding_size(): assert model.embedding_size == 96 if is_ci: delete_model_cache(model.model._model_dir) + + +@pytest.mark.parametrize("model_name", ["answerdotai/answerai-ColBERT-small-v1"]) +def test_session_options(model_cache, model_name) -> None: + with model_cache(model_name) as default_model: + default_session_options = default_model.model.model.get_session_options() + assert default_session_options.enable_cpu_mem_arena is True + model = LateInteractionTextEmbedding(model_name=model_name, enable_cpu_mem_arena=False) + session_options = model.model.model.get_session_options() + assert session_options.enable_cpu_mem_arena is False diff --git a/tests/test_sparse_embeddings.py b/tests/test_sparse_embeddings.py index 90514967c..4c02a683b 100644 --- a/tests/test_sparse_embeddings.py +++ b/tests/test_sparse_embeddings.py @@ -77,7 +77,12 @@ } -_MODELS_TO_CACHE = ("prithivida/Splade_PP_en_v1", "Qdrant/minicoil-v1", "Qdrant/bm25") +_MODELS_TO_CACHE = ( + "prithivida/Splade_PP_en_v1", + "Qdrant/minicoil-v1", + "Qdrant/bm25", + "Qdrant/bm42-all-minilm-l6-v2-attentions", +) MODELS_TO_CACHE = tuple([x.lower() for x in _MODELS_TO_CACHE]) @@ -276,3 +281,20 @@ def test_lazy_load(model_name: str) -> None: if is_ci: delete_model_cache(model.model._model_dir) + + +@pytest.mark.parametrize( + "model_name", + [ + "prithivida/Splade_PP_en_v1", + "Qdrant/minicoil-v1", + "Qdrant/bm42-all-minilm-l6-v2-attentions", + ], +) +def test_session_options(model_cache, model_name) -> None: + with model_cache(model_name) as default_model: + default_session_options = default_model.model.model.get_session_options() + assert default_session_options.enable_cpu_mem_arena is True + model = SparseTextEmbedding(model_name=model_name, enable_cpu_mem_arena=False) + session_options = model.model.model.get_session_options() + assert session_options.enable_cpu_mem_arena is False diff --git a/tests/test_text_cross_encoder.py b/tests/test_text_cross_encoder.py index 925ae8a30..d23ee8ef4 100644 --- a/tests/test_text_cross_encoder.py +++ b/tests/test_text_cross_encoder.py @@ -122,3 +122,13 @@ def test_rerank_pairs_parallel(model_cache, model_name: str) -> None: assert np.allclose( scores_parallel[: len(canonical_scores)], canonical_scores, atol=1e-3 ), f"Model: {model_name}, Scores (Parallel): {scores_parallel}, Expected: {canonical_scores}" + + +@pytest.mark.parametrize("model_name", ["Xenova/ms-marco-MiniLM-L-6-v2"]) +def test_session_options(model_cache, model_name) -> None: + with model_cache(model_name) as default_model: + default_session_options = default_model.model.model.get_session_options() + assert default_session_options.enable_cpu_mem_arena is True + model = TextCrossEncoder(model_name=model_name, enable_cpu_mem_arena=False) + session_options = model.model.model.get_session_options() + assert session_options.enable_cpu_mem_arena is False diff --git a/tests/test_text_onnx_embeddings.py b/tests/test_text_onnx_embeddings.py index 46ce6554d..43e88ca85 100644 --- a/tests/test_text_onnx_embeddings.py +++ b/tests/test_text_onnx_embeddings.py @@ -193,3 +193,13 @@ def test_embedding_size() -> None: if is_ci: delete_model_cache(model.model._model_dir) + + +@pytest.mark.parametrize("model_name", ["sentence-transformers/all-MiniLM-L6-v2"]) +def test_session_options(model_cache, model_name) -> None: + with model_cache(model_name) as default_model: + default_session_options = default_model.model.model.get_session_options() + assert default_session_options.enable_cpu_mem_arena is True + model = TextEmbedding(model_name=model_name, enable_cpu_mem_arena=False) + session_options = model.model.model.get_session_options() + assert session_options.enable_cpu_mem_arena is False