model_bin.py 5.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145
  1. #!/usr/bin/env python3
  2. # -*- encoding: utf-8 -*-
  3. # Copyright FunASR (https://github.com/FunAudioLLM/SenseVoice). All Rights Reserved.
  4. # MIT License (https://opensource.org/licenses/MIT)
  5. import os.path
  6. from pathlib import Path
  7. from typing import List, Union, Tuple
  8. import torch
  9. import librosa
  10. import numpy as np
  11. from utils.infer_utils import (
  12. CharTokenizer,
  13. Hypothesis,
  14. ONNXRuntimeError,
  15. OrtInferSession,
  16. TokenIDConverter,
  17. get_logger,
  18. read_yaml,
  19. )
  20. from utils.frontend import WavFrontend
  21. from utils.infer_utils import pad_list
  22. logging = get_logger()
  23. class SenseVoiceSmallONNX:
  24. """
  25. Author: Speech Lab of DAMO Academy, Alibaba Group
  26. Paraformer: Fast and Accurate Parallel Transformer for Non-autoregressive End-to-End Speech Recognition
  27. https://arxiv.org/abs/2206.08317
  28. """
  29. def __init__(
  30. self,
  31. model_dir: Union[str, Path] = None,
  32. batch_size: int = 1,
  33. device_id: Union[str, int] = "-1",
  34. plot_timestamp_to: str = "",
  35. quantize: bool = False,
  36. intra_op_num_threads: int = 4,
  37. cache_dir: str = None,
  38. **kwargs,
  39. ):
  40. if quantize:
  41. model_file = os.path.join(model_dir, "model_quant.onnx")
  42. else:
  43. model_file = os.path.join(model_dir, "model.onnx")
  44. config_file = os.path.join(model_dir, "config.yaml")
  45. cmvn_file = os.path.join(model_dir, "am.mvn")
  46. config = read_yaml(config_file)
  47. # token_list = os.path.join(model_dir, "tokens.json")
  48. # with open(token_list, "r", encoding="utf-8") as f:
  49. # token_list = json.load(f)
  50. # self.converter = TokenIDConverter(token_list)
  51. self.tokenizer = CharTokenizer()
  52. config["frontend_conf"]['cmvn_file'] = cmvn_file
  53. self.frontend = WavFrontend(**config["frontend_conf"])
  54. self.ort_infer = OrtInferSession(
  55. model_file, device_id, intra_op_num_threads=intra_op_num_threads
  56. )
  57. self.batch_size = batch_size
  58. self.blank_id = 0
  59. def __call__(self,
  60. wav_content: Union[str, np.ndarray, List[str]],
  61. language: List,
  62. textnorm: List,
  63. tokenizer=None,
  64. **kwargs) -> List:
  65. waveform_list = self.load_data(wav_content, self.frontend.opts.frame_opts.samp_freq)
  66. waveform_nums = len(waveform_list)
  67. asr_res = []
  68. for beg_idx in range(0, waveform_nums, self.batch_size):
  69. end_idx = min(waveform_nums, beg_idx + self.batch_size)
  70. feats, feats_len = self.extract_feat(waveform_list[beg_idx:end_idx])
  71. ctc_logits, encoder_out_lens = self.infer(feats,
  72. feats_len,
  73. np.array(language, dtype=np.int32),
  74. np.array(textnorm, dtype=np.int32)
  75. )
  76. # back to torch.Tensor
  77. ctc_logits = torch.from_numpy(ctc_logits).float()
  78. # support batch_size=1 only currently
  79. x = ctc_logits[0, : encoder_out_lens[0].item(), :]
  80. yseq = x.argmax(dim=-1)
  81. yseq = torch.unique_consecutive(yseq, dim=-1)
  82. mask = yseq != self.blank_id
  83. token_int = yseq[mask].tolist()
  84. if tokenizer is not None:
  85. asr_res.append(tokenizer.tokens2text(token_int))
  86. else:
  87. asr_res.append(token_int)
  88. return asr_res
  89. def load_data(self, wav_content: Union[str, np.ndarray, List[str]], fs: int = None) -> List:
  90. def load_wav(path: str) -> np.ndarray:
  91. waveform, _ = librosa.load(path, sr=fs)
  92. return waveform
  93. if isinstance(wav_content, np.ndarray):
  94. return [wav_content]
  95. if isinstance(wav_content, str):
  96. return [load_wav(wav_content)]
  97. if isinstance(wav_content, list):
  98. return [load_wav(path) for path in wav_content]
  99. raise TypeError(f"The type of {wav_content} is not in [str, np.ndarray, list]")
  100. def extract_feat(self, waveform_list: List[np.ndarray]) -> Tuple[np.ndarray, np.ndarray]:
  101. feats, feats_len = [], []
  102. for waveform in waveform_list:
  103. speech, _ = self.frontend.fbank(waveform)
  104. feat, feat_len = self.frontend.lfr_cmvn(speech)
  105. feats.append(feat)
  106. feats_len.append(feat_len)
  107. feats = self.pad_feats(feats, np.max(feats_len))
  108. feats_len = np.array(feats_len).astype(np.int32)
  109. return feats, feats_len
  110. @staticmethod
  111. def pad_feats(feats: List[np.ndarray], max_feat_len: int) -> np.ndarray:
  112. def pad_feat(feat: np.ndarray, cur_len: int) -> np.ndarray:
  113. pad_width = ((0, max_feat_len - cur_len), (0, 0))
  114. return np.pad(feat, pad_width, "constant", constant_values=0)
  115. feat_res = [pad_feat(feat, feat.shape[0]) for feat in feats]
  116. feats = np.array(feat_res).astype(np.float32)
  117. return feats
  118. def infer(self,
  119. feats: np.ndarray,
  120. feats_len: np.ndarray,
  121. language: np.ndarray,
  122. textnorm: np.ndarray,) -> Tuple[np.ndarray, np.ndarray]:
  123. outputs = self.ort_infer([feats, feats_len, language, textnorm])
  124. return outputs