Skip to content

Commit f6fd794

Browse files
fix: reduce code duplication in _load_onnx_model and add parallel execution support (#520)
Co-authored-by: google-labs-jules[bot] <161369871+google-labs-jules[bot]@users.noreply.github.com>
1 parent 4439a16 commit f6fd794

6 files changed

Lines changed: 62 additions & 67 deletions

File tree

qwen3_embed/common/onnx_model.py

Lines changed: 43 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@
1010
from numpy.typing import NDArray
1111
from tokenizers import Tokenizer
1212

13+
from qwen3_embed.common.preprocessor_utils import load_tokenizer
1314
from qwen3_embed.common.types import Device, NumpyArray, OnnxProvider
1415
from qwen3_embed.parallel_processor import Worker
1516

@@ -48,6 +49,7 @@ def __init__(self) -> None:
4849
self.model: ort.InferenceSession | None = None
4950
self.model_input_names: set[str] | None = None
5051
self.tokenizer: Tokenizer | None = None
52+
self.special_token_to_id: dict[str, int] = {}
5153

5254
def _preprocess_onnx_input(
5355
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
@@ -108,9 +110,14 @@ def _validate_providers(
108110
return requested_provider_names
109111

110112
def _create_session_options(
111-
self, threads: int | None, extra_session_options: dict[str, Any] | None
113+
self,
114+
threads: int | None,
115+
extra_session_options: dict[str, Any] | None,
116+
parallel_execution: bool = False,
112117
) -> Any:
113118
so = ort.SessionOptions() # type: ignore[possibly-missing-attribute]
119+
if parallel_execution:
120+
so.execution_mode = ort.ExecutionMode.ORT_PARALLEL # type: ignore[possibly-missing-attribute]
114121
so.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL # type: ignore[possibly-missing-attribute]
115122
# Disable memory pattern optimization to prevent ORT from retaining
116123
# peak-sized buffers across inferences with varying sequence lengths.
@@ -124,32 +131,28 @@ def _create_session_options(
124131
self.add_extra_session_options(so, extra_session_options)
125132
return so
126133

127-
def _load_onnx_model(
134+
def _instantiate_onnx_session(
128135
self,
129-
model_dir: Path,
130-
model_file: str,
136+
model_path: Path,
131137
threads: int | None,
132138
providers: Sequence[OnnxProvider] | None = None,
133139
cuda: bool | Device = Device.AUTO,
134140
device_id: int | None = None,
141+
parallel_execution: bool = False,
135142
extra_session_options: dict[str, Any] | None = None,
136-
) -> None:
137-
model_path = model_dir / model_file
143+
) -> tuple[ort.InferenceSession, list[str]]:
138144
# List of Execution Providers: https://onnxruntime.ai/docs/execution-providers
139145
available_providers = ort.get_available_providers() # type: ignore[possibly-missing-attribute]
140146

141147
onnx_providers = self._determine_providers(providers, cuda, device_id, available_providers)
142148
requested_provider_names = self._validate_providers(onnx_providers, available_providers)
143-
so = self._create_session_options(threads, extra_session_options)
149+
so = self._create_session_options(threads, extra_session_options, parallel_execution)
144150

145-
self.model = ort.InferenceSession(
146-
str(model_path), providers=onnx_providers, sess_options=so
147-
)
148-
self.model_input_names = {node.name for node in self.model.get_inputs()}
149-
logger.info(f"ONNX session created with providers: {self.model.get_providers()}")
151+
session = ort.InferenceSession(str(model_path), providers=onnx_providers, sess_options=so)
152+
input_names = [node.name for node in session.get_inputs()]
153+
logger.info(f"ONNX session created with providers: {session.get_providers()}")
150154
if "CUDAExecutionProvider" in requested_provider_names:
151-
assert self.model is not None
152-
current_providers = self.model.get_providers()
155+
current_providers = session.get_providers()
153156
if "CUDAExecutionProvider" not in current_providers:
154157
warnings.warn(
155158
f"Attempt to set CUDAExecutionProvider failed. Current providers: {current_providers}."
@@ -158,6 +161,32 @@ def _load_onnx_model(
158161
RuntimeWarning,
159162
stacklevel=2,
160163
)
164+
return session, input_names
165+
166+
def _load_onnx_model(
167+
self,
168+
model_dir: Path,
169+
model_file: str,
170+
threads: int | None,
171+
providers: Sequence[OnnxProvider] | None = None,
172+
cuda: bool | Device = Device.AUTO,
173+
device_id: int | None = None,
174+
parallel_execution: bool = False,
175+
extra_session_options: dict[str, Any] | None = None,
176+
) -> tuple[ort.InferenceSession, list[str]]:
177+
model_path = model_dir / model_file
178+
self.model, input_names = self._instantiate_onnx_session(
179+
model_path=model_path,
180+
threads=threads,
181+
providers=providers,
182+
cuda=cuda,
183+
device_id=device_id,
184+
parallel_execution=parallel_execution,
185+
extra_session_options=extra_session_options,
186+
)
187+
self.model_input_names = set(input_names)
188+
self.tokenizer, self.special_token_to_id = load_tokenizer(model_dir=model_dir)
189+
return self.model, input_names
161190

162191
@classmethod
163192
def _select_exposed_session_options(cls, model_kwargs: dict[str, Any]) -> dict[str, Any]:

qwen3_embed/rerank/cross_encoder/onnx_text_model.py

Lines changed: 0 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,6 @@
11
import os
22
from collections.abc import Iterable, Sequence
33
from multiprocessing import get_all_start_methods
4-
from pathlib import Path
54
from typing import Any
65

76
import numpy as np
@@ -13,7 +12,6 @@
1312
OnnxOutputContext,
1413
OnnxProvider,
1514
)
16-
from qwen3_embed.common.preprocessor_utils import load_tokenizer
1715
from qwen3_embed.common.types import Device, NumpyArray
1816
from qwen3_embed.common.utils import iter_batch
1917
from qwen3_embed.parallel_processor import ParallelWorkerPool
@@ -26,28 +24,6 @@ class OnnxCrossEncoderModel(OnnxModel[float]):
2624
def _get_worker_class(cls) -> type["TextRerankerWorker"]:
2725
raise NotImplementedError("Subclasses must implement this method")
2826

29-
def _load_onnx_model(
30-
self,
31-
model_dir: Path,
32-
model_file: str,
33-
threads: int | None,
34-
providers: Sequence[OnnxProvider] | None = None,
35-
cuda: bool | Device = Device.AUTO,
36-
device_id: int | None = None,
37-
extra_session_options: dict[str, Any] | None = None,
38-
) -> None:
39-
super()._load_onnx_model(
40-
model_dir=model_dir,
41-
model_file=model_file,
42-
threads=threads,
43-
providers=providers,
44-
cuda=cuda,
45-
device_id=device_id,
46-
extra_session_options=extra_session_options,
47-
)
48-
self.tokenizer, _ = load_tokenizer(model_dir=model_dir)
49-
assert self.tokenizer is not None
50-
5127
def tokenize(self, pairs: list[tuple[str, str]], **_: Any) -> list[Encoding]:
5228
assert self.tokenizer is not None
5329
return self.tokenizer.encode_batch(pairs)

qwen3_embed/text/onnx_text_model.py

Lines changed: 1 addition & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -1,15 +1,13 @@
11
import os
22
from collections.abc import Iterable, Sequence
33
from multiprocessing import get_all_start_methods
4-
from pathlib import Path
54
from typing import Any
65

76
import numpy as np
87
from numpy.typing import NDArray
9-
from tokenizers import Encoding, Tokenizer
8+
from tokenizers import Encoding
109

1110
from qwen3_embed.common.onnx_model import EmbeddingWorker, OnnxModel, OnnxOutputContext, T
12-
from qwen3_embed.common.preprocessor_utils import load_tokenizer
1311
from qwen3_embed.common.types import Device, NumpyArray, OnnxProvider
1412
from qwen3_embed.common.utils import iter_batch
1513
from qwen3_embed.parallel_processor import ParallelWorkerPool
@@ -36,8 +34,6 @@ def _post_process_onnx_output(self, output: OnnxOutputContext, **kwargs: Any) ->
3634

3735
def __init__(self) -> None:
3836
super().__init__()
39-
self.tokenizer: Tokenizer | None = None
40-
self.special_token_to_id: dict[str, int] = {}
4137

4238
def _preprocess_onnx_input(
4339
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
@@ -47,27 +43,6 @@ def _preprocess_onnx_input(
4743
"""
4844
return onnx_input
4945

50-
def _load_onnx_model(
51-
self,
52-
model_dir: Path,
53-
model_file: str,
54-
threads: int | None,
55-
providers: Sequence[OnnxProvider] | None = None,
56-
cuda: bool | Device = Device.AUTO,
57-
device_id: int | None = None,
58-
extra_session_options: dict[str, Any] | None = None,
59-
) -> None:
60-
super()._load_onnx_model(
61-
model_dir=model_dir,
62-
model_file=model_file,
63-
threads=threads,
64-
providers=providers,
65-
cuda=cuda,
66-
device_id=device_id,
67-
extra_session_options=extra_session_options,
68-
)
69-
self.tokenizer, self.special_token_to_id = load_tokenizer(model_dir=model_dir)
70-
7146
def load_onnx_model(self) -> None:
7247
raise NotImplementedError("Subclasses must implement this method")
7348

tests/test_cross_encoder_onnx.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -559,7 +559,7 @@ def test_wires_tokenizer_after_super(self, tmp_path: Path) -> None:
559559
with (
560560
patch("qwen3_embed.common.onnx_model.ort") as mock_ort,
561561
patch(
562-
"qwen3_embed.rerank.cross_encoder.onnx_text_model.load_tokenizer",
562+
"qwen3_embed.common.onnx_model.load_tokenizer",
563563
return_value=(mock_tokenizer, {}),
564564
),
565565
):

tests/test_onnx_model.py

Lines changed: 16 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -38,15 +38,20 @@ def model():
3838

3939
@pytest.fixture
4040
def mock_ort():
41-
with patch("qwen3_embed.common.onnx_model.ort") as mock:
41+
with (
42+
patch("qwen3_embed.common.onnx_model.ort") as mock,
43+
patch("qwen3_embed.common.onnx_model.load_tokenizer", return_value=(MagicMock(), {})),
44+
):
4245
# Default behavior: CPU provider available, no CUDA
4346
mock.get_available_providers.return_value = ["CPUExecutionProvider"]
4447
mock.SessionOptions.return_value = MagicMock()
4548
mock.GraphOptimizationLevel.ORT_ENABLE_ALL = 99
49+
mock.ExecutionMode.ORT_PARALLEL = 1
4650

4751
# Default session mock
4852
session_mock = MagicMock()
4953
session_mock.get_providers.return_value = ["CPUExecutionProvider"]
54+
session_mock.get_inputs.return_value = []
5055
mock.InferenceSession.return_value = session_mock
5156

5257
yield mock
@@ -196,6 +201,16 @@ def test_load_threads(model: ConcreteOnnxModel, mock_ort):
196201
assert so.inter_op_num_threads == 4
197202

198203

204+
def test_load_parallel_execution(model: ConcreteOnnxModel, mock_ort):
205+
"""Test parallel execution configuration."""
206+
mock_ort.get_available_providers.return_value = ["CPUExecutionProvider"]
207+
208+
model._load_onnx_model(Path("dummy"), "model.onnx", threads=None, parallel_execution=True)
209+
210+
so = mock_ort.SessionOptions.return_value
211+
assert so.execution_mode == mock_ort.ExecutionMode.ORT_PARALLEL
212+
213+
199214
def test_load_extra_session_options(model: ConcreteOnnxModel, mock_ort):
200215
"""Test extra session options."""
201216
mock_ort.get_available_providers.return_value = ["CPUExecutionProvider"]

tests/test_text_onnx_text_model.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -267,7 +267,7 @@ def test_wires_tokenizer_after_super(self, tmp_path: Path) -> None:
267267
with (
268268
patch("qwen3_embed.common.onnx_model.ort") as mock_ort,
269269
patch(
270-
"qwen3_embed.text.onnx_text_model.load_tokenizer",
270+
"qwen3_embed.common.onnx_model.load_tokenizer",
271271
return_value=(mock_tokenizer, mock_special),
272272
),
273273
):

0 commit comments

Comments
 (0)