@@ -127,13 +127,13 @@ def compute_mel_filters():
127127 return fb .astype (np .float32 )
128128
129129 @staticmethod
130- def compute_mel_spectrogram (audio_np , mel_filters_np ):
130+ def compute_mel_spectrogram (audio_np , mel_filters_t , hann_window_t ):
131+ """Compute log-mel spectrogram. Caller supplies pre-built torch hann window
132+ and pre-converted mel_filters tensor so we don't rebuild them per chunk."""
131133 audio = torch .from_numpy (audio_np ).float ()
132- mel_filters = torch .from_numpy (mel_filters_np ).float ()
133- window = torch .hann_window (WINDOW_SIZE )
134- stft = torch .stft (audio , WINDOW_SIZE , HOP_LENGTH , window = window , return_complex = True )
134+ stft = torch .stft (audio , WINDOW_SIZE , HOP_LENGTH , window = hann_window_t , return_complex = True )
135135 mag2 = stft [..., :- 1 ].abs () ** 2
136- mel_spec = mel_filters .T @ mag2
136+ mel_spec = mel_filters_t .T @ mag2
137137 log_spec = torch .clamp (mel_spec , min = 1e-10 ).log10 ()
138138 log_spec = torch .maximum (log_spec , log_spec .max () - 8.0 )
139139 log_spec = (log_spec + 4.0 ) / 4.0
@@ -165,13 +165,17 @@ def bytes_to_unicode():
165165 return dict (zip (bs , [chr (c ) for c in cs ]))
166166
167167 @staticmethod
168- def decode_tokens (token_ids , tokenizer_dir : str ) -> str :
168+ def load_tokenizer_state (tokenizer_dir : str ) -> tuple [dict [Any , Any ], set [int ], dict [str , int ]]:
169+ """Load tokenizer state needed by detokenization. Returns (id_to_token,
170+ special_token_ids, byte_decoder). Intended to be called once per model load
171+ and cached on the inference instance — re-reading vocab.json (~3 MB) per
172+ chunk is otherwise a substantial overhead on long-form audio."""
169173 vocab_path = os .path .join (tokenizer_dir , "vocab.json" )
170174 with open (vocab_path , "r" , encoding = "utf-8" ) as f :
171175 vocab = json .load (f )
172176 id_to_token = {v : k for k , v in vocab .items ()}
173177
174- special_tokens = set ()
178+ special_tokens : set [ int ] = set ()
175179 tc_path = os .path .join (tokenizer_dir , "tokenizer_config.json" )
176180 if os .path .exists (tc_path ):
177181 with open (tc_path ) as f :
@@ -181,7 +185,15 @@ def decode_tokens(token_ids, tokenizer_dir: str) -> str:
181185
182186 byte_enc = Qwen3ASRHelpers .bytes_to_unicode ()
183187 byte_dec = {v : k for k , v in byte_enc .items ()}
188+ return id_to_token , special_tokens , byte_dec
184189
190+ @staticmethod
191+ def decode_tokens_cached (
192+ token_ids ,
193+ id_to_token : dict ,
194+ special_tokens : set ,
195+ byte_dec : dict ,
196+ ) -> str :
185197 pieces = []
186198 for tid in token_ids :
187199 if tid in special_tokens :
@@ -204,12 +216,19 @@ def __init__(self, load_config: ModelLoadConfig):
204216 self .runtime_cfg = Qwen3ASRHelpers .hf_config (self .ov_dir / "config.json" )
205217 self .chunk_size = self .runtime_cfg ["enc_n_window" ] * 2
206218 self .mel_filters = Qwen3ASRHelpers .compute_mel_filters ()
219+ # Cache torch tensors used every chunk by the mel spectrogram path.
220+ self ._mel_filters_t = torch .from_numpy (self .mel_filters ).float ()
221+ self ._hann_window = torch .hann_window (WINDOW_SIZE )
207222 self .core = ov .Core ()
208223 self .t_model_load = 0.0
209224 self .enc_model = None
210225 self .emb_model = None
211226 self .dec_model = None
212227 self .dec_request = None
228+ # Cached tokenizer state, populated in load_model.
229+ self ._tok_id_to_token : dict [Any , Any ] | None = None
230+ self ._tok_special : set [int ] | None = None
231+ self ._tok_byte_dec : dict [str , int ] | None = None
213232
214233 def load_model (self , load_config : ModelLoadConfig ) -> None :
215234 self .load_config = load_config
@@ -229,12 +248,18 @@ def load_model(self, load_config: ModelLoadConfig) -> None:
229248 load_config .device ,
230249 )
231250 self .dec_request = self .dec_model .create_infer_request ()
251+ (
252+ self ._tok_id_to_token ,
253+ self ._tok_special ,
254+ self ._tok_byte_dec ,
255+ ) = Qwen3ASRHelpers .load_tokenizer_state (str (self .ov_dir ))
232256 self .t_model_load = time .perf_counter () - t_load_start
233257
234258 def _embed_tokens (self , token_ids ):
235259 ids = np .asarray (token_ids , dtype = np .int64 )
236260 if ids .ndim == 1 :
237261 ids = ids [np .newaxis , :]
262+ assert self .emb_model , "Model not loaded"
238263 out = self .emb_model ([ids ])
239264 return out [self .emb_model .output (0 )]
240265
@@ -267,7 +292,7 @@ def collect_metrics(
267292
268293 def audio_chunks (self , chunk_audio : np .ndarray , max_tokens : int ):
269294 t_feature_start = time .perf_counter ()
270- mel = Qwen3ASRHelpers .compute_mel_spectrogram (chunk_audio , self .mel_filters )
295+ mel = Qwen3ASRHelpers .compute_mel_spectrogram (chunk_audio , self ._mel_filters_t , self . _hann_window )
271296 t_feature = time .perf_counter () - t_feature_start
272297 total_frames = mel .shape [1 ]
273298 expected_tokens = Qwen3ASRHelpers .count_encoder_tokens (total_frames , self .chunk_size )
@@ -278,6 +303,7 @@ def audio_chunks(self, chunk_audio: np.ndarray, max_tokens: int):
278303 mel_input = mel [np .newaxis , :, :].astype (np .float32 )
279304
280305 t_encoder_start = time .perf_counter ()
306+ assert self .enc_model , "Model not loaded"
281307 enc_out = self .enc_model ([mel_input ])
282308 t_encoder = time .perf_counter () - t_encoder_start
283309 audio_embeds = enc_out [self .enc_model .output (0 )]
@@ -293,8 +319,8 @@ def audio_chunks(self, chunk_audio: np.ndarray, max_tokens: int):
293319 prompt_len = len (input_ids )
294320 position_ids = np .arange (prompt_len , dtype = np .int64 )[np .newaxis , :]
295321
296-
297322 t_prefill_start = time .perf_counter ()
323+ assert self .dec_request , "Model not loaded"
298324 self .dec_request .reset_state ()
299325 self .dec_request .set_input_tensor (0 , ov .Tensor (input_embeds ))
300326 self .dec_request .set_input_tensor (1 , ov .Tensor (position_ids ))
@@ -326,7 +352,8 @@ def audio_chunks(self, chunk_audio: np.ndarray, max_tokens: int):
326352 generated .pop ()
327353
328354 t_detok_start = time .perf_counter ()
329- raw = Qwen3ASRHelpers .decode_tokens (generated , str (self .ov_dir ))
355+ assert self ._tok_id_to_token and self ._tok_special and self ._tok_byte_dec , "Tokenizer state not loaded"
356+ raw = Qwen3ASRHelpers .decode_tokens_cached (generated , self ._tok_id_to_token , self ._tok_special , self ._tok_byte_dec )
330357 t_detok = time .perf_counter () - t_detok_start
331358
332359 metrics = self .collect_metrics (
@@ -344,6 +371,7 @@ def audio_chunks(self, chunk_audio: np.ndarray, max_tokens: int):
344371 async def transcribe (self , gen_config : OV_Qwen3ASRGenConfig ) -> AsyncIterator [Union [Dict [str , Any ], str ]]:
345372 t_transcribe_start = time .perf_counter ()
346373 audio_input = gen_config .audio_base64
374+ assert audio_input , "audio_base64 is required"
347375 if not audio_input .startswith ("data:audio" ):
348376 audio_input = f"data:audio/wav;base64,{ audio_input } "
349377 audio_array = (await asyncio .to_thread (normalize_audios , audio_input ))[0 ]
@@ -434,6 +462,9 @@ async def unload_model(self, registry: ModelRegistry, model_name: str) -> bool:
434462 self .emb_model = None
435463 self .dec_model = None
436464 self .dec_request = None
465+ self ._tok_id_to_token = None
466+ self ._tok_special = None
467+ self ._tok_byte_dec = None
437468 gc .collect ()
438469 logger .info (f"[{ self .load_config .model_name } ] unloaded and memory cleaned up" )
439470 return removed
0 commit comments