mirror of
https://github.com/modelscope/modelscope.git
synced 2026-09-01 19:49:03 +02:00
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:
@@ -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'
|
||||
|
||||
@@ -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']
|
||||
|
||||
@@ -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'],
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user