sv inference & asr trainer: add new inputs

Link: https://code.alibaba-inc.com/Ali-MaaS/MaaS-lib/codereview/11424993

* add sv_inference for speaker embedding extraction

* add raw inputs support for speaker_verification

* add train meta_params
This commit is contained in:
jiangyu.xzy
2023-01-13 13:52:52 +00:00
committed by wenmeng.zwm
parent 883bbd0164
commit 742ad4b355
6 changed files with 106 additions and 16 deletions

View File

@@ -31,6 +31,7 @@ class OutputKeys(object):
OUTPUT_PCM_LIST = 'output_pcm_list'
OUTPUT_WAV = 'output_wav'
IMG_EMBEDDING = 'img_embedding'
SPK_EMBEDDING = 'spk_embedding'
SPO_LIST = 'spo_list'
TEXT_EMBEDDING = 'text_embedding'
TRANSLATION = 'translation'

View File

@@ -39,7 +39,7 @@ class AutomaticSpeechRecognitionPipeline(Pipeline):
self.output_dir = None
if 'output_dir' in kwargs:
self.output_dir = kwargs['output_dir']
self.cmd = self.get_cmd()
self.cmd = self.get_cmd(kwargs)
if self.cmd['code_base'] == 'funasr':
from funasr.bin import asr_inference_launch
self.funasr_infer_modelscope = asr_inference_launch.inference_launch(
@@ -133,7 +133,7 @@ class AutomaticSpeechRecognitionPipeline(Pipeline):
rst = self.postprocess(output)
return rst
def get_cmd(self) -> Dict[str, Any]:
def get_cmd(self, extra_args) -> Dict[str, Any]:
if self.preprocessor is None:
self.preprocessor = WavToScp()
@@ -224,6 +224,28 @@ class AutomaticSpeechRecognitionPipeline(Pipeline):
cmd['punc_model_config'] = outputs['punc_model_config']
else:
cmd['punc_model_config'] = None
if 'batch_size' in extra_args:
cmd['batch_size'] = extra_args['batch_size']
if 'mode' in extra_args:
cmd['mode'] = extra_args['mode']
if 'ngpu' in extra_args:
cmd['ngpu'] = extra_args['ngpu']
if 'beam_size' in extra_args:
cmd['beam_size'] = extra_args['beam_size']
if 'decoding_ind' in extra_args:
cmd['decoding_ind'] = extra_args['decoding_ind']
if 'decoding_mode' in extra_args:
cmd['decoding_mode'] = extra_args['decoding_mode']
if 'vad_model_file' in extra_args:
cmd['vad_model_name'] = extra_args['vad_model_file']
if 'vad_infer_config' in extra_args:
cmd['vad_model_config'] = extra_args['vad_infer_config']
if 'vad_cmvn_file' in extra_args:
cmd['vad_mvn_file'] = extra_args['vad_cmvn_file']
if 'punc_model_file' in extra_args:
cmd['punc_model_name'] = extra_args['punc_model_file']
if 'punc_infer_config' in extra_args:
cmd['punc_model_config'] = extra_args['punc_infer_config']
elif self.framework == Frameworks.tf:
cmd['fs']['model_fs'] = outputs['model_config']['fs']

View File

@@ -44,7 +44,9 @@ class PunctuationProcessingPipeline(Pipeline):
super().__init__(model=model, **kwargs)
self.model_cfg = self.model.forward()
self.cmd = self.get_cmd()
self.output_dir = None
if 'output_dir' in kwargs:
self.output_dir = kwargs['output_dir']
from funasr.bin import punc_inference_launch
self.funasr_infer_modelscope = punc_inference_launch.inference_launch(
mode=self.cmd['mode'],
@@ -52,7 +54,7 @@ class PunctuationProcessingPipeline(Pipeline):
log_level=self.cmd['log_level'],
dtype=self.cmd['dtype'],
seed=self.cmd['seed'],
output_dir=self.cmd['output_dir'],
output_dir=self.output_dir,
batch_size=self.cmd['batch_size'],
num_workers=self.cmd['num_workers'],
key_file=self.cmd['key_file'],

View File

@@ -10,7 +10,8 @@ from modelscope.models import Model
from modelscope.outputs import OutputKeys
from modelscope.pipelines.base import Pipeline
from modelscope.pipelines.builder import PIPELINES
from modelscope.utils.audio.audio_utils import generate_sv_scp_from_url
from modelscope.utils.audio.audio_utils import (generate_scp_for_sv,
generate_sv_scp_from_url)
from modelscope.utils.constant import Frameworks, Tasks
from modelscope.utils.logger import get_logger
@@ -60,7 +61,8 @@ class SpeakerVerificationPipeline(Pipeline):
key_file=self.cmd['key_file'],
model_tag=self.cmd['model_tag'])
def __call__(self, audio_in: tuple = None) -> Dict[str, Any]:
def __call__(self,
audio_in: Union[tuple, str, Any] = None) -> Dict[str, Any]:
if len(audio_in) == 0:
raise ValueError('The input of ITN should not be null.')
else:
@@ -76,8 +78,12 @@ class SpeakerVerificationPipeline(Pipeline):
rst = {}
for i in range(len(inputs)):
if i == 0:
score = inputs[0]['value']
rst[OutputKeys.SCORES] = score
if isinstance(self.audio_in, tuple):
score = inputs[0]['value']
rst[OutputKeys.SCORES] = score
else:
embedding = inputs[0]['value']
rst[OutputKeys.SPK_EMBEDDING] = embedding
else:
rst[inputs[i]['key']] = inputs[i]['value']
return rst
@@ -105,19 +111,42 @@ class SpeakerVerificationPipeline(Pipeline):
}
return cmd
def forward(self, audio_in: tuple = None) -> list:
def forward(self, audio_in: Union[tuple, str, Any] = None) -> list:
"""Decoding
"""
logger.info(
'Speaker Verification Processing: {0} ...'.format(audio_in))
# generate audio_scp
audio_scp_1, audio_scp_2 = generate_sv_scp_from_url(audio_in)
data_cmd = [(audio_scp_1, 'speech', 'sound'),
(audio_scp_2, 'ref_speech', 'sound')]
data_cmd, raw_inputs = None, None
if isinstance(audio_in, tuple):
# generate audio_scp
if isinstance(audio_in[0], str):
audio_scp_1, audio_scp_2 = generate_sv_scp_from_url(audio_in)
data_cmd = [(audio_scp_1, 'speech', 'sound'),
(audio_scp_2, 'ref_speech', 'sound')]
elif isinstance(audio_in[0], bytes):
data_cmd = [(audio_in[0], 'speech', 'bytes'),
(audio_in[1], 'ref_speech', 'bytes')]
else:
raise TypeError('Unsupported data type.')
else:
if isinstance(audio_in, str):
audio_scp = generate_scp_for_sv(audio_in)
data_cmd = [(audio_scp, 'speech', 'sound')]
elif isinstance(audio_in[0], bytes):
data_cmd = [(audio_in, 'speech', 'bytes')]
else:
import torch
import numpy as np
if isinstance(audio_in, torch.Tensor):
raw_inputs = audio_in
elif isinstance(audio_in, np.ndarray):
raw_inputs = audio_in
else:
raise TypeError('Unsupported data type.')
self.cmd['name_and_type'] = data_cmd
self.cmd['raw_inputs'] = None
self.cmd['raw_inputs'] = raw_inputs
punc_result = self.run_inference(self.cmd)
return punc_result

View File

@@ -31,7 +31,39 @@ class ASRTrainer(BaseTrainer):
dataset_type: str = 'small',
data_dir: Optional[Union[MsDataset, str]] = None,
model_revision: Optional[str] = DEFAULT_MODEL_REVISION,
batch_bins: Optional[int] = None,
max_epoch: Optional[int] = None,
lr: Optional[float] = None,
mate_params: Optional[dict] = None,
**kwargs):
"""ASR Trainer.
Args:
model (str) : model name
work_dir (str): output dir for saving results
distributed (bool): whether to enable DDP training
dataset_type (str): choose which dataset type to use
data_dir (str): the path of data
model_revision (str): set model version
batch_bins (str): batch size
max_epoch (int): the maximum epoch number for training
lr (float): learning rate
mate_params (dict): for saving other training args
Examples:
>>> import os
>>> from modelscope.metainfo import Trainers
>>> from modelscope.msdatasets import MsDataset
>>> from modelscope.trainers import build_trainer
>>> ds_dict = MsDataset.load('speech_asr_aishell1_trainsets')
>>> kwargs = dict(
>>> model='damo/speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-pytorch',
>>> data_dir=ds_dict,
>>> work_dir="./checkpoint")
>>> trainer = build_trainer(
>>> Trainers.speech_asr_trainer, default_args=kwargs)
>>> trainer.train()
"""
if not work_dir:
self.work_dir = tempfile.TemporaryDirectory().name
if not os.path.exists(self.work_dir):
@@ -71,7 +103,11 @@ class ASRTrainer(BaseTrainer):
data_dir=self.data_dir,
output_dir=self.work_dir,
distributed=self.distributed,
dataset_type=self.dataset_type)
dataset_type=self.dataset_type,
batch_bins=batch_bins,
max_epoch=max_epoch,
lr=lr,
mate_params=mate_params)
def parse_cfg(self, cfg_file):
cur_dir = os.path.dirname(cfg_file)

View File

@@ -1,7 +1,7 @@
bitstring
easyasr>=0.0.2
espnet==202204
funasr>=0.1.5
funasr>=0.1.6
funtextprocessing>=0.1.1
greenlet>=1.1.2
h5py