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:
George
2025-11-25 17:49:02 +07:00
committed by GitHub
parent 44e332999c
commit ec0e3128ee
18 changed files with 155 additions and 2 deletions

View File

@@ -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")

View File

@@ -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,
)

View File

@@ -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(),

View File

@@ -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,
)

View File

@@ -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,
)

View File

@@ -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(),

View File

@@ -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,
)

View File

@@ -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(),

View File

@@ -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

View File

@@ -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,
)

View File

@@ -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,
)

View File

@@ -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,
)

View File

@@ -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(),

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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