mirror of
https://github.com/qdrant/fastembed.git
synced 2026-07-23 11:20:51 -05:00
new: expose some onnx session options (#578)
* new: expose some onnx session options * fix: fix extra session options is None case * fix: fix missing params * new: add tests
This commit is contained in:
@@ -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 @@ class OnnxModel(Generic[T]):
|
||||
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,9 @@ class OnnxModel(Generic[T]):
|
||||
so.intra_op_num_threads = threads
|
||||
so.inter_op_num_threads = threads
|
||||
|
||||
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
|
||||
)
|
||||
@@ -113,6 +119,38 @@ class OnnxModel(Generic[T]):
|
||||
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")
|
||||
|
||||
|
||||
@@ -98,6 +98,7 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
|
||||
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 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_id=self.device_id,
|
||||
extra_session_options=self._extra_session_options,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@@ -180,6 +182,7 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
|
||||
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,
|
||||
)
|
||||
|
||||
|
||||
@@ -55,6 +55,7 @@ class OnnxImageModel(OnnxModel[T]):
|
||||
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 @@ class OnnxImageModel(OnnxModel[T]):
|
||||
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 @@ class OnnxImageModel(OnnxModel[T]):
|
||||
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
|
||||
@@ -130,6 +133,9 @@ class OnnxImageModel(OnnxModel[T]):
|
||||
**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(),
|
||||
|
||||
@@ -143,6 +143,7 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
|
||||
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 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
|
||||
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 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
|
||||
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,
|
||||
)
|
||||
|
||||
|
||||
@@ -80,6 +80,7 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
|
||||
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 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
|
||||
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 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
|
||||
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 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
|
||||
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,
|
||||
)
|
||||
|
||||
|
||||
@@ -64,6 +64,7 @@ class OnnxMultimodalModel(OnnxModel[T]):
|
||||
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 @@ class OnnxMultimodalModel(OnnxModel[T]):
|
||||
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 @@ class OnnxMultimodalModel(OnnxModel[T]):
|
||||
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
|
||||
@@ -153,6 +156,9 @@ class OnnxMultimodalModel(OnnxModel[T]):
|
||||
**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(),
|
||||
@@ -189,6 +195,7 @@ class OnnxMultimodalModel(OnnxModel[T]):
|
||||
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
|
||||
@@ -220,6 +227,9 @@ class OnnxMultimodalModel(OnnxModel[T]):
|
||||
**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(),
|
||||
|
||||
@@ -111,6 +111,7 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
|
||||
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 +151,7 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_id=self.device_id,
|
||||
extra_session_options=self._extra_session_options,
|
||||
)
|
||||
|
||||
def rerank(
|
||||
@@ -192,6 +194,7 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
|
||||
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,
|
||||
)
|
||||
|
||||
|
||||
@@ -33,6 +33,7 @@ class OnnxCrossEncoderModel(OnnxModel[float]):
|
||||
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 @@ class OnnxCrossEncoderModel(OnnxModel[float]):
|
||||
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 @@ class OnnxCrossEncoderModel(OnnxModel[float]):
|
||||
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
|
||||
@@ -127,6 +130,9 @@ class OnnxCrossEncoderModel(OnnxModel[float]):
|
||||
**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(),
|
||||
|
||||
@@ -103,6 +103,7 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
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 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
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 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
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
|
||||
|
||||
@@ -117,6 +117,8 @@ class MiniCOIL(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
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
|
||||
@@ -153,6 +155,7 @@ class MiniCOIL(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_id=self.device_id,
|
||||
extra_session_options=self._extra_session_options,
|
||||
)
|
||||
|
||||
assert self.tokenizer is not None
|
||||
@@ -221,6 +224,7 @@ class MiniCOIL(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
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,
|
||||
)
|
||||
|
||||
|
||||
@@ -99,6 +99,7 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
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 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_id=self.device_id,
|
||||
extra_session_options=self._extra_session_options,
|
||||
)
|
||||
|
||||
def embed(
|
||||
@@ -168,6 +170,7 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
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,
|
||||
)
|
||||
|
||||
|
||||
@@ -233,7 +233,7 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]):
|
||||
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 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]):
|
||||
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 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]):
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_id=self.device_id,
|
||||
extra_session_options=self._extra_session_options,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -54,6 +54,7 @@ class OnnxTextModel(OnnxModel[T]):
|
||||
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 @@ class OnnxTextModel(OnnxModel[T]):
|
||||
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 @@ class OnnxTextModel(OnnxModel[T]):
|
||||
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
|
||||
@@ -143,6 +146,9 @@ class OnnxTextModel(OnnxModel[T]):
|
||||
**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(),
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -77,7 +77,12 @@ CANONICAL_QUERY_VALUES = {
|
||||
}
|
||||
|
||||
|
||||
_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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user