diff --git a/modelscope/outputs/outputs.py b/modelscope/outputs/outputs.py index 4b89853e..d8b95ab4 100644 --- a/modelscope/outputs/outputs.py +++ b/modelscope/outputs/outputs.py @@ -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' diff --git a/modelscope/pipelines/audio/asr_inference_pipeline.py b/modelscope/pipelines/audio/asr_inference_pipeline.py index 20405333..5a339a08 100644 --- a/modelscope/pipelines/audio/asr_inference_pipeline.py +++ b/modelscope/pipelines/audio/asr_inference_pipeline.py @@ -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'] diff --git a/modelscope/pipelines/audio/punctuation_processing_pipeline.py b/modelscope/pipelines/audio/punctuation_processing_pipeline.py index 7008b403..072f9e85 100644 --- a/modelscope/pipelines/audio/punctuation_processing_pipeline.py +++ b/modelscope/pipelines/audio/punctuation_processing_pipeline.py @@ -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'], diff --git a/modelscope/pipelines/audio/speaker_verification_pipeline.py b/modelscope/pipelines/audio/speaker_verification_pipeline.py index 6939cf60..ed63dbcd 100644 --- a/modelscope/pipelines/audio/speaker_verification_pipeline.py +++ b/modelscope/pipelines/audio/speaker_verification_pipeline.py @@ -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 diff --git a/modelscope/trainers/audio/asr_trainer.py b/modelscope/trainers/audio/asr_trainer.py index 5f63bfd5..4ea25863 100644 --- a/modelscope/trainers/audio/asr_trainer.py +++ b/modelscope/trainers/audio/asr_trainer.py @@ -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) diff --git a/requirements/audio.txt b/requirements/audio.txt index b9e4c04a..983fd70f 100644 --- a/requirements/audio.txt +++ b/requirements/audio.txt @@ -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