Skip to content

Commit 68ab888

Browse files
authored
Merge pull request #128 from solidDoWant/chore/minor-qwen3-asr-perf-improvements-1
Implement minor Qwen3-ASR performance improvements
2 parents 7de77a4 + f3c6232 commit 68ab888

1 file changed

Lines changed: 41 additions & 10 deletions

File tree

src/engine/openvino/qwen3_asr/qwen3_asr.py

Lines changed: 41 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)