1import numpy as np
2import onnxruntime as ort
3import sentencepiece as spm
4import soundfile as sf
5import torch
6import torchaudio
7
8# --- config ---
9ENCODER = "onnx_int8/encoder.int8.onnx" # or onnx_fp32/encoder.onnx
10DECODER = "onnx_int8/decoder_joint.int8.onnx"
11TOKENIZER = "onnx_int8/tokenizer.model"
12
13SR = 16000
14BLANK = 1024
15
16# --- load audio ---
17audio, sr = sf.read("audio.wav", dtype="float32")
18wav = torch.from_numpy(audio).unsqueeze(0) # (1, samples)
19
20# --- mel features (must match NeMo preprocessor) ---
21mel_xform = torchaudio.transforms.MelSpectrogram(
22 sample_rate=SR, n_fft=512, win_length=400, hop_length=160,
23 n_mels=128, window_fn=torch.hann_window, power=2.0,
24 norm="slaney", mel_scale="slaney", center=True,
25)
26mel = torch.log(mel_xform(wav) + 2**-24)
27mel = (mel - mel.mean(dim=-1, keepdim=True)) / (mel.std(dim=-1, keepdim=True) + 1e-5)
28feats = mel.numpy().astype(np.float32)
29
30# --- encoder ---
31enc = ort.InferenceSession(ENCODER, providers=["CPUExecutionProvider"])
32enc_out, enc_len = enc.run(None, {"audio_signal": feats, "length": np.array([feats.shape[-1]], dtype=np.int64)})
33
34# --- greedy RNN-T decode ---
35dec = ort.InferenceSession(DECODER, providers=["CPUExecutionProvider"])
36state1 = np.zeros((2, 1, 640), dtype=np.float32)
37state2 = np.zeros((2, 1, 640), dtype=np.float32)
38last_token, tokens = BLANK, []
39
40for t in range(int(enc_len[0])):
41 enc_t = enc_out[:, :, t:t+1]
42 for _ in range(10):
43 logits, _, s1, s2 = dec.run(None, {
44 "encoder_outputs": enc_t,
45 "targets": np.array([[last_token]], dtype=np.int32),
46 "target_length": np.array([1], dtype=np.int32),
47 "input_states_1": state1, "input_states_2": state2,
48 })
49 idx = int(np.argmax(logits[0, 0, 0]))
50 if idx == BLANK:
51 break
52 tokens.append(idx)
53 last_token = idx
54 state1, state2 = s1, s2
55
56sp = spm.SentencePieceProcessor(model_file=TOKENIZER)
57print(sp.decode(tokens))