Zipformer.axera / ax_pretrained_infer.py
HY-2012's picture
First commit
09bcf49 verified
Raw
History Blame
20.2 kB
#!/usr/bin/env python3
# Copyright 2023 Xiaomi Corp. (authors: Fangjun Kuang)
"""
This script loads ONNX models exported by ./export-onnx.py
and uses them to decode waves.
We use the pre-trained model from
https://huggingface.co/Zengwei/icefall-asr-librispeech-pruned-transducer-stateless7-streaming-2022-12-29
as an example to show how to use this file.
1. Download the pre-trained model
cd egs/librispeech/ASR
repo_url=https://huggingface.co/Zengwei/icefall-asr-librispeech-pruned-transducer-stateless7-streaming-2022-12-29
GIT_LFS_SKIP_SMUDGE=1 git clone $repo_url
repo=$(basename $repo_url)
pushd $repo
git lfs pull --include "data/lang_bpe_500/bpe.model"
git lfs pull --include "exp/pretrained.pt"
cd exp
ln -s pretrained.pt epoch-99.pt
popd
2. Export the model to ONNX
./pruned_transducer_stateless7_streaming/export-onnx.py \
--bpe-model $repo/data/lang_bpe_500/bpe.model \
--use-averaged-model 0 \
--epoch 99 \
--avg 1 \
--decode-chunk-len 32 \
--exp-dir $repo/exp/
It will generate the following 3 files in $repo/exp
- encoder-epoch-99-avg-1.onnx
- decoder-epoch-99-avg-1.onnx
- joiner-epoch-99-avg-1.onnx
3. Run this file with the exported ONNX models
./pruned_transducer_stateless7_streaming/onnx_pretrained.py \
--encoder-model-filename $repo/exp/encoder-epoch-99-avg-1.onnx \
--decoder-model-filename $repo/exp/decoder-epoch-99-avg-1.onnx \
--joiner-model-filename $repo/exp/joiner-epoch-99-avg-1.onnx \
--tokens $repo/data/lang_bpe_500/tokens.txt \
$repo/test_wavs/1089-134686-0001.wav
Note: Even though this script only supports decoding a single file,
the exported ONNX models do support batch processing.
"""
import argparse
import logging
from typing import Dict, List, Optional, Tuple
import k2
import numpy as np
# import onnxruntime as ort
import torch
# import torchaudio
from kaldifeat import FbankOptions, OnlineFbank, OnlineFeature
from axengine import InferenceSession
def get_parser():
parser = argparse.ArgumentParser(
formatter_class=argparse.ArgumentDefaultsHelpFormatter
)
parser.add_argument(
"--encoder-model-filename",
type=str,
required=True,
help="Path to the encoder onnx model. ",
)
parser.add_argument(
"--decoder-model-filename",
type=str,
required=True,
help="Path to the decoder onnx model. ",
)
parser.add_argument(
"--joiner-model-filename",
type=str,
required=True,
help="Path to the joiner onnx model. ",
)
parser.add_argument(
"--tokens",
type=str,
help="""Path to tokens.txt.""",
)
parser.add_argument(
"--sound-dir",
type=str,
help="The input sound file to transcribe. "
"Supported formats are those supported by torchaudio.load(). "
"For example, wav and flac are supported. "
"The sample rate has to be 16kHz.",
)
return parser
class OnnxModel:
def __init__(
self,
encoder_model_filename: str,
decoder_model_filename: str,
joiner_model_filename: str,
):
# session_opts = ort.SessionOptions()
# session_opts.inter_op_num_threads = 1
# session_opts.intra_op_num_threads = 1
# self.session_opts = session_opts
self.init_encoder(encoder_model_filename)
self.init_decoder(decoder_model_filename)
self.init_joiner(joiner_model_filename)
def init_encoder(self, encoder_model_filename: str):
self.encoder = InferenceSession(
encoder_model_filename,
#sess_options=self.session_opts,
)
print("==========Encoder input==========")
for i in self.encoder.get_inputs():
print(i)
print("==========Encoder output==========")
for i in self.encoder.get_outputs():
print(i)
self.init_encoder_states()
def init_encoder_states(self, batch_size: int = 1):
# encoder_meta = self.encoder.get_modelmeta().custom_metadata_map
# print(encoder_meta)
# model_type = encoder_meta["model_type"]
# assert model_type == "zipformer", model_type
# decode_chunk_len = int(encoder_meta["decode_chunk_len"])
# T = int(encoder_meta["T"])
# num_encoder_layers = encoder_meta["num_encoder_layers"]
# encoder_dims = encoder_meta["encoder_dims"]
# attention_dims = encoder_meta["attention_dims"]
# cnn_module_kernels = encoder_meta["cnn_module_kernels"]
# left_context_len = encoder_meta["left_context_len"]
model_type = "zipformer"
assert model_type == "zipformer", model_type
decode_chunk_len = 96
T = 103
num_encoder_layers = "2,2,2,2,2"
encoder_dims = "256,256,256,256,256"
attention_dims = "192,192,192,192,192"
cnn_module_kernels = "31,31,31,31,31"
left_context_len = "192,96,48,24,96"
def to_int_list(s):
return list(map(int, s.split(",")))
num_encoder_layers = to_int_list(num_encoder_layers)
encoder_dims = to_int_list(encoder_dims)
attention_dims = to_int_list(attention_dims)
cnn_module_kernels = to_int_list(cnn_module_kernels)
left_context_len = to_int_list(left_context_len)
# logging.info(f"decode_chunk_len: {decode_chunk_len}")
# logging.info(f"T: {T}")
# logging.info(f"num_encoder_layers: {num_encoder_layers}")
# logging.info(f"encoder_dims: {encoder_dims}")
# logging.info(f"attention_dims: {attention_dims}")
# logging.info(f"cnn_module_kernels: {cnn_module_kernels}")
# logging.info(f"left_context_len: {left_context_len}")
num_encoders = len(num_encoder_layers)
cached_len = []
cached_avg = []
cached_key = []
cached_val = []
cached_val2 = []
cached_conv1 = []
cached_conv2 = []
N = batch_size
for i in range(num_encoders):
cached_len.append(torch.zeros(num_encoder_layers[i], N, dtype=torch.int64))
cached_avg.append(torch.zeros(num_encoder_layers[i], N, encoder_dims[i]))
cached_key.append(
torch.zeros(
num_encoder_layers[i], left_context_len[i], N, attention_dims[i]
)
)
cached_val.append(
torch.zeros(
num_encoder_layers[i],
left_context_len[i],
N,
attention_dims[i] // 2,
)
)
cached_val2.append(
torch.zeros(
num_encoder_layers[i],
left_context_len[i],
N,
attention_dims[i] // 2,
)
)
cached_conv1.append(
torch.zeros(
num_encoder_layers[i], N, encoder_dims[i], cnn_module_kernels[i] - 1
)
)
cached_conv2.append(
torch.zeros(
num_encoder_layers[i], N, encoder_dims[i], cnn_module_kernels[i] - 1
)
)
self.cached_len = cached_len
self.cached_avg = cached_avg
self.cached_key = cached_key
self.cached_val = cached_val
self.cached_val2 = cached_val2
self.cached_conv1 = cached_conv1
self.cached_conv2 = cached_conv2
self.num_encoders = num_encoders
self.segment = T
self.offset = decode_chunk_len
def init_decoder(self, decoder_model_filename: str):
self.decoder = InferenceSession(
decoder_model_filename,
# sess_options=self.session_opts,
)
print("==========Decoder input==========")
for i in self.decoder.get_inputs():
print(i)
print("==========Decoder output==========")
for i in self.decoder.get_outputs():
print(i)
# decoder_meta = self.decoder.get_modelmeta().custom_metadata_map
# self.context_size = int(decoder_meta["context_size"])
# self.vocab_size = int(decoder_meta["vocab_size"])
self.context_size = 2
self.vocab_size = 6254
logging.info(f"context_size: {self.context_size}")
logging.info(f"vocab_size: {self.vocab_size}")
def init_joiner(self, joiner_model_filename: str):
self.joiner = InferenceSession(
joiner_model_filename,
# sess_options=self.session_opts,
)
print("==========Joiner input==========")
for i in self.joiner.get_inputs():
print(i)
print("==========Joiner output==========")
for i in self.joiner.get_outputs():
print(i)
# joiner_meta = self.joiner.get_modelmeta().custom_metadata_map
# self.joiner_dim = int(joiner_meta["joiner_dim"])
self.joiner_dim = 512
logging.info(f"joiner_dim: {self.joiner_dim}")
def _build_encoder_input_output(
self,
x: torch.Tensor,
) -> Tuple[Dict[str, np.ndarray], List[str]]:
encoder_input = {"x": x.numpy()}
encoder_output = ["encoder_out"]
def build_states_input(states: List[torch.Tensor], name: str):
for i, s in enumerate(states):
if isinstance(s, torch.Tensor):
# Fix dtype for cached_len (int32)
if name == "cached_len":
encoder_input[f"{name}_{i}"] = s.to(torch.int32).numpy()
else:
encoder_input[f"{name}_{i}"] = s.numpy()
else:
encoder_input[f"{name}_{i}"] = s
encoder_output.append(f"new_{name}_{i}")
build_states_input(self.cached_len, "cached_len")
build_states_input(self.cached_avg, "cached_avg")
build_states_input(self.cached_key, "cached_key")
build_states_input(self.cached_val, "cached_val")
build_states_input(self.cached_val2, "cached_val2")
build_states_input(self.cached_conv1, "cached_conv1")
build_states_input(self.cached_conv2, "cached_conv2")
return encoder_input, encoder_output
def _update_states(self, states: List[np.ndarray]):
num_encoders = self.num_encoders
self.cached_len = states[num_encoders * 0 : num_encoders * 1]
self.cached_avg = states[num_encoders * 1 : num_encoders * 2]
self.cached_key = states[num_encoders * 2 : num_encoders * 3]
self.cached_val = states[num_encoders * 3 : num_encoders * 4]
self.cached_val2 = states[num_encoders * 4 : num_encoders * 5]
self.cached_conv1 = states[num_encoders * 5 : num_encoders * 6]
self.cached_conv2 = states[num_encoders * 6 : num_encoders * 7]
def run_encoder(self, x: torch.Tensor) -> torch.Tensor:
"""
Args:
x:
A 3-D tensor of shape (N, T, C)
Returns:
Return a 3-D tensor of shape (N, T', joiner_dim) where
T' is usually equal to ((T-7)//2+1)//2
"""
encoder_input, encoder_output_names = self._build_encoder_input_output(x)
#
out = self.encoder.run(encoder_output_names, encoder_input)
self._update_states(out[1:])
return torch.from_numpy(out[0])
def run_decoder(self, decoder_input: torch.Tensor) -> torch.Tensor:
"""
Args:
decoder_input:
A 2-D tensor of shape (N, context_size)
Returns:
Return a 2-D tensor of shape (N, joiner_dim)
"""
# Fix dtype for decoder input (int32)
decoder_input_np = decoder_input.to(torch.int32).numpy()
out = self.decoder.run(
[self.decoder.get_outputs()[0].name],
{self.decoder.get_inputs()[0].name: decoder_input_np},
)[0]
return torch.from_numpy(out)
def run_joiner(
self, encoder_out: torch.Tensor, decoder_out: torch.Tensor
) -> torch.Tensor:
"""
Args:
encoder_out:
A 2-D tensor of shape (N, joiner_dim)
decoder_out:
A 2-D tensor of shape (N, joiner_dim)
Returns:
Return a 2-D tensor of shape (N, vocab_size)
"""
out = self.joiner.run(
[self.joiner.get_outputs()[0].name],
{
self.joiner.get_inputs()[0].name: encoder_out.numpy(),
self.joiner.get_inputs()[1].name: decoder_out.numpy(),
},
)[0]
return torch.from_numpy(out)
def read_sound_files(
filenames: List[str], expected_sample_rate: float
) -> List[torch.Tensor]:
"""Read a list of sound files into a list 1-D float32 torch tensors.
Args:
filenames:
A list of sound filenames.
expected_sample_rate:
The expected sample rate of the sound files.
Returns:
Return a list of 1-D float32 torch tensors.
"""
try:
import soundfile as sf
except ImportError:
raise ImportError("Please install soundfile: pip install soundfile")
try:
import librosa
except ImportError:
librosa = None
ans = []
for f in filenames:
data, sample_rate = sf.read(f)
if data.ndim > 1:
data = data[:, 0]
if sample_rate != expected_sample_rate:
if librosa is None:
raise ImportError("Please install librosa: pip install librosa for resampling audio.")
# librosa.resample expects float32
data = data.astype(np.float32)
data = librosa.resample(data, orig_sr=sample_rate, target_sr=expected_sample_rate)
sample_rate = expected_sample_rate
wave = torch.from_numpy(data.astype(np.float32)).contiguous()
ans.append(wave)
return ans
def create_streaming_feature_extractor() -> OnlineFeature:
"""Create a CPU streaming feature extractor.
At present, we assume it returns a fbank feature extractor with
fixed options. In the future, we will support passing in the options
from outside.
Returns:
Return a CPU streaming feature extractor.
"""
opts = FbankOptions()
opts.device = "cpu"
opts.frame_opts.dither = 0
opts.frame_opts.snip_edges = False
opts.frame_opts.samp_freq = 16000
opts.mel_opts.num_bins = 80
opts.mel_opts.high_freq = -400
return OnlineFbank(opts)
def greedy_search(
model: OnnxModel,
encoder_out: torch.Tensor,
context_size: int,
decoder_out: Optional[torch.Tensor] = None,
hyp: Optional[List[int]] = None,
) -> List[int]:
"""Greedy search in batch mode. It hardcodes --max-sym-per-frame=1.
Args:
model:
The transducer model.
encoder_out:
A 3-D tensor of shape (1, T, joiner_dim)
context_size:
The context size of the decoder model.
decoder_out:
Optional. Decoder output of the previous chunk.
hyp:
Decoding results for previous chunks.
Returns:
Return the decoded results so far.
"""
blank_id = 0
if decoder_out is None:
assert hyp is None, hyp
hyp = [blank_id] * context_size
decoder_input = torch.tensor([hyp], dtype=torch.int64)
decoder_out = model.run_decoder(decoder_input)
else:
assert hyp is not None, hyp
encoder_out = encoder_out.squeeze(0)
T = encoder_out.size(0)
for t in range(T):
cur_encoder_out = encoder_out[t : t + 1]
joiner_out = model.run_joiner(cur_encoder_out, decoder_out).squeeze(0)
y = joiner_out.argmax(dim=0).item()
if y != blank_id:
hyp.append(y)
decoder_input = hyp[-context_size:]
decoder_input = torch.tensor([decoder_input], dtype=torch.int64)
decoder_out = model.run_decoder(decoder_input)
return hyp, decoder_out
@torch.no_grad()
def main():
import time
parser = get_parser()
args = parser.parse_args()
logging.info(vars(args))
t0 = time.time()
model = OnnxModel(
encoder_model_filename=args.encoder_model_filename,
decoder_model_filename=args.decoder_model_filename,
joiner_model_filename=args.joiner_model_filename,
)
t1 = time.time()
print(f" Model loading time: {t1 - t0:.3f} s")
sample_rate = 16000
# logging.info("Constructing Fbank computer")
# online_fbank = create_streaming_feature_extractor()
import os
from glob import glob
# 判断sound_dir是文件还是目录
input_path = args.sound_dir
if os.path.isdir(input_path):
# 支持常见音频格式
audio_exts = ["*.wav", "*.flac", "*.mp3", "*.ogg", "*.m4a"]
audio_files = []
for ext in audio_exts:
audio_files.extend(glob(os.path.join(input_path, ext)))
audio_files.sort()
else:
audio_files = [input_path]
logging.info(f"Found {len(audio_files)} audio files.")
rtf_list = []
for idx, sound_file in enumerate(audio_files):
model.init_encoder_states() # 重置encoder缓存,保证每个音频独立识别
t2 = time.time()
waves = read_sound_files(
filenames=[sound_file],
expected_sample_rate=sample_rate,
)[0]
t3 = time.time()
print(f"[{idx+1}/{len(audio_files)}] {os.path.basename(sound_file)} Audio loading time: {t3 - t2:.3f} s")
tail_padding = torch.zeros(int(0.3 * sample_rate), dtype=torch.float32)
wave_samples = torch.cat([waves, tail_padding])
online_fbank = create_streaming_feature_extractor()
num_processed_frames = 0
segment = model.segment
offset = model.offset
context_size = model.context_size
hyp = None
decoder_out = None
chunk = int(1 * sample_rate) # 1 second
start = 0
t4 = time.time()
while start < wave_samples.numel():
end = min(start + chunk, wave_samples.numel())
samples = wave_samples[start:end]
start += chunk
online_fbank.accept_waveform(
sampling_rate=sample_rate,
waveform=samples,
)
while online_fbank.num_frames_ready - num_processed_frames >= segment:
frames = []
for i in range(segment):
frames.append(online_fbank.get_frame(num_processed_frames + i))
num_processed_frames += offset
frames = torch.cat(frames, dim=0)
frames = frames.unsqueeze(0)
encoder_out = model.run_encoder(frames)
hyp, decoder_out = greedy_search(
model,
encoder_out,
context_size,
decoder_out,
hyp,
)
t5 = time.time()
print(f" Audio processing time: {t5 - t4:.3f} s")
symbol_table = k2.SymbolTable.from_file(args.tokens)
text = ""
for i in hyp[context_size:]:
text += symbol_table[i]
text = text.replace("▁", " ").strip()
logging.info(sound_file)
logging.info(text)
audio_duration = float(waves.shape[0]) / sample_rate
total_time = (t3 - t2) + (t5 - t4)
rtf = total_time / audio_duration if audio_duration > 0 else float('inf')
print(f" Audio duration: {audio_duration:.3f} s")
print(f" Total (load+process) time: {total_time:.3f} s")
print(f" RTF (total_time/audio_duration): {rtf:.3f}")
rtf_list.append(rtf)
# logging.info("Decoding Done")
if rtf_list:
avg_rtf = sum(rtf_list) / len(rtf_list)
print(f"\nAverage RTF over {len(rtf_list)} files: {avg_rtf:.3f}")
else:
print("No audio files processed.")
if __name__ == "__main__":
formatter = "%(asctime)s %(levelname)s [%(filename)s:%(lineno)d] %(message)s"
logging.basicConfig(format=formatter, level=logging.INFO)
main()