1010from numpy .typing import NDArray
1111from tokenizers import Tokenizer
1212
13+ from qwen3_embed .common .preprocessor_utils import load_tokenizer
1314from qwen3_embed .common .types import Device , NumpyArray , OnnxProvider
1415from 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 ]:
0 commit comments