From f9fed09e96f43e7eab88378fc444c4987933badb Mon Sep 17 00:00:00 2001
From: zhifu gao <zhifu.gzf@alibaba-inc.com>
Date: 星期五, 09 十二月 2022 23:57:51 +0800
Subject: [PATCH] Merge pull request #10 from alibaba-damo-academy/dev
---
egs_modelscope/common/utils/apply_lfr_and_cmvn.sh | 38
egs_modelscope/wenetspeech/paraformer/utils/combine_cmvn_file.py | 73
egs_modelscope/common_uniasr/utils/apply_cmvn.sh | 29
egs_modelscope/aishell2/paraformer/utils/split_data.py | 60
egs_modelscope/aishell2/paraformer/utils/compute_wer.py | 157
egs/aishell/paraformer/utils/text2token.py | 135
egs/aishell/conformer/utils/apply_lfr_and_cmvn.py | 143
egs_modelscope/aishell/paraformer/utils/compute_wer.py | 157
egs_modelscope/common/modelscope_common_infer.sh | 7
egs/aishell/paraformerbert/utils/text_tokenize.sh | 35
egs_modelscope/aishell/paraformer/utils/print_args.py | 45
egs_modelscope/aishell2/paraformer/utils/compute_cmvn.py | 74
egs_modelscope/aishell2/paraformer/utils/compute_fbank.sh | 51
egs/aishell/paraformer/utils/apply_lfr_and_cmvn.sh | 38
egs_modelscope/aishell2/paraformer/utils/apply_cmvn.py | 79
egs_modelscope/aishell/paraformer/utils/apply_cmvn.sh | 29
egs_modelscope/aishell/paraformer/utils/split_scp.pl | 246
egs_modelscope/common/utils/filter_scp.pl | 87
egs_modelscope/speechio/paraformer/paraformer_large_infer.sh | 2
egs_modelscope/speechio/paraformer/utils/text_tokenize.sh | 35
egs_modelscope/aishell2/paraformer/utils/split_scp.pl | 246
egs_modelscope/common/utils/run.pl | 356
egs/aishell/conformer/utils/error_rate_zh | 370
egs_modelscope/common_uniasr/utils/run.pl | 356
egs_modelscope/speechio/paraformer/utils/extract_embeds.py | 47
egs_modelscope/common_uniasr/utils/text_tokenize.py | 106
egs_modelscope/common_uniasr/modelscope_common_infer.sh | 76
egs/aishell/conformer/utils/gen_ark_list.sh | 22
egs_modelscope/common_uniasr/utils/print_args.py | 45
egs_modelscope/aishell/paraformer/utils/fix_data.sh | 35
egs/aishell/paraformer/utils/subset_data_dir_tr_cv.sh | 30
egs/aishell/paraformerbert/utils/text2token.py | 135
egs_modelscope/aishell2/paraformer/utils/shuffle_list.pl | 44
egs_modelscope/speechio/paraformer/utils/compute_cmvn.sh | 25
funasr/bin/asr_inference_uniasr.py | 215
egs_modelscope/aishell2/paraformer/utils/apply_lfr_and_cmvn.py | 143
egs_modelscope/speechio/paraformer/utils/print_args.py | 45
egs/aishell/paraformer/utils/extract_embeds.py | 47
egs_modelscope/common_uniasr/utils/split_scp.pl | 246
egs_modelscope/common/utils/fix_data.sh | 35
egs/aishell/paraformer/utils/text_tokenize.sh | 35
egs/aishell/paraformerbert/utils/compute_wer.py | 157
egs/aishell/paraformer/utils/parse_options.sh | 97
egs/aishell/tranformer/utils/compute_cmvn.sh | 5
egs/aishell/conformer/utils/text2token.py | 135
egs/aishell/paraformer/utils/proce_text.py | 31
egs_modelscope/wenetspeech/paraformer/utils/filter_scp.pl | 87
egs_modelscope/common_uniasr/modelscope_utils/update_config.py | 41
egs_modelscope/aishell/paraformer/utils/proc_conf_oss.py | 35
egs_modelscope/aishell2/paraformer/conf/train_asr_paraformer_sanm_50e_16d_2048_512_lfr6.yaml | 1
egs_modelscope/wenetspeech/paraformer/modelscope_utils/download_model.py | 25
egs_modelscope/common_uniasr/utils/apply_cmvn.py | 79
egs_modelscope/aishell/paraformer/utils/split_data.py | 60
egs_modelscope/wenetspeech/paraformer/utils/run.pl | 356
egs_modelscope/wenetspeech/paraformer/utils/textnorm_zh.py | 834 +
egs/aishell/paraformerbert/utils/parse_options.sh | 97
egs_modelscope/common/utils/apply_lfr_and_cmvn.py | 143
docs/images/wechat.png | 0
egs/aishell/paraformer/utils/compute_fbank.sh | 51
egs_modelscope/aishell2/paraformer/utils/compute_cmvn.sh | 25
egs/aishell/paraformer/utils/apply_lfr_and_cmvn.py | 143
egs_modelscope/aishell2/paraformer/utils/apply_cmvn.sh | 29
egs/aishell/conformer/utils/split_data.py | 60
egs_modelscope/speechio/paraformer/utils/parse_options.sh | 97
egs_modelscope/speechio/paraformer/modelscope_utils/update_config.py | 41
egs_modelscope/aishell2/paraformer/utils/subset_data_dir_tr_cv.sh | 30
egs_modelscope/aishell2/paraformer/utils/apply_lfr_and_cmvn.sh | 38
egs/aishell/conformer/utils/compute_fbank.sh | 51
egs_modelscope/aishell2/paraformer/utils/__init__.py | 0
egs_modelscope/common_uniasr/utils/shuffle_list.pl | 44
egs_modelscope/aishell/paraformer/utils/apply_lfr_and_cmvn.sh | 38
egs/aishell/paraformerbert/utils/fix_data_feat.sh | 52
egs_modelscope/speechio/paraformer/utils/split_data.py | 60
egs/aishell/conformer/utils/compute_wer.py | 157
egs_modelscope/speechio/paraformer/utils/apply_cmvn.sh | 29
egs_modelscope/speechio/paraformer/utils/fix_data_feat.sh | 52
egs_modelscope/speechio/paraformer/utils/compute_cmvn.py | 74
egs_modelscope/aishell/paraformer/utils/compute_cmvn.sh | 25
egs_modelscope/common/utils/combine_cmvn_file.py | 73
egs_modelscope/common_uniasr/utils/textnorm_zh.py | 834 +
setup.py | 4
egs_modelscope/aishell2/paraformer/utils/run.pl | 356
funasr/bin/modelscope_infer.py | 2
egs_modelscope/speechio/paraformer/utils/shuffle_list.pl | 44
egs/aishell/paraformerbert/utils/split_data.py | 60
egs_modelscope/wenetspeech/paraformer/modelscope_utils/update_config.py | 41
egs/aishell/paraformerbert/utils/gen_ark_list.sh | 22
egs_modelscope/aishell/paraformer/utils/filter_scp.pl | 87
egs/aishell/paraformer/utils/filter_scp.pl | 87
egs_modelscope/common_uniasr/utils/compute_wer.py | 157
egs_modelscope/speechio/paraformer/utils/run.pl | 356
egs/aishell/paraformer/utils/gen_ark_list.sh | 22
egs_modelscope/aishell/paraformer/utils/text_tokenize.sh | 35
egs/aishell/conformer/utils/subset_data_dir_tr_cv.sh | 30
egs_modelscope/wenetspeech/paraformer/utils/compute_wer.py | 157
egs_modelscope/wenetspeech/paraformer/utils/split_data.py | 60
egs_modelscope/common/utils/apply_cmvn.sh | 29
egs_modelscope/common/utils/gen_ark_list.sh | 22
egs/aishell/paraformerbert/utils/fix_data.sh | 35
egs/aishell/paraformer/utils/textnorm_zh.py | 834 +
egs_modelscope/common/utils/text2token.py | 135
egs_modelscope/aishell/paraformer/utils/error_rate_zh | 370
egs_modelscope/wenetspeech/paraformer/utils/shuffle_list.pl | 44
egs/aishell/conformer/utils/compute_cmvn.sh | 25
egs_modelscope/speechio/paraformer/utils/text_tokenize.py | 106
egs_modelscope/aishell2/paraformer/utils/text2token.py | 135
egs_modelscope/wenetspeech/paraformer/modelscope_utils/modelscope_infer.sh | 88
egs/aishell/conformer/utils/parse_options.sh | 97
egs/aishell/tranformer/utils/gen_ark_list.sh | 14
egs_modelscope/speechio/paraformer/utils/filter_scp.pl | 87
egs_modelscope/wenetspeech/paraformer/utils/proce_text.py | 31
egs/aishell/paraformer/utils/split_data.py | 60
egs_modelscope/wenetspeech/paraformer/utils/proc_conf_oss.py | 35
egs_modelscope/aishell/paraformer/utils/subset_data_dir_tr_cv.sh | 30
egs_modelscope/common/utils/proc_conf_oss.py | 35
egs_modelscope/common_uniasr/utils/compute_cmvn.sh | 25
egs_modelscope/common_uniasr/utils/compute_fbank.py | 153
funasr/bin/asr_inference.py | 225
egs_modelscope/aishell2/paraformer/paraformer_large_finetune.sh | 87
egs/aishell/paraformerbert/utils/compute_fbank.sh | 51
egs_modelscope/speechio/paraformer/utils/gen_ark_list.sh | 22
egs_modelscope/aishell2/paraformer/utils/text_tokenize.sh | 35
egs_modelscope/aishell/paraformer/utils/text_tokenize.py | 106
egs_modelscope/aishell/paraformer/utils/apply_cmvn.py | 79
egs_modelscope/aishell2/paraformer/utils/fix_data.sh | 35
egs_modelscope/aishell2/paraformer/utils/print_args.py | 45
egs_modelscope/common_uniasr/utils/compute_fbank.sh | 51
egs/aishell/paraformerbert/utils/textnorm_zh.py | 834 +
egs_modelscope/aishell/paraformer/utils/apply_lfr_and_cmvn.py | 143
egs/aishell/conformer/utils/combine_cmvn_file.py | 73
egs/aishell/paraformerbert/utils/extract_embeds.py | 47
egs_modelscope/wenetspeech/paraformer/utils/parse_options.sh | 97
egs_modelscope/aishell2/paraformer/modelscope_utils/update_config.py | 41
egs_modelscope/speechio/paraformer/utils/apply_cmvn.py | 79
egs_modelscope/aishell2/paraformer/utils/text_tokenize.py | 106
egs_modelscope/common/utils/print_args.py | 45
egs_modelscope/aishell/paraformer/utils/compute_fbank.py | 153
egs_modelscope/common/utils/parse_options.sh | 97
egs_modelscope/aishell2/paraformer/utils/compute_fbank.py | 153
egs_modelscope/common_uniasr/utils/__init__.py | 0
egs_modelscope/wenetspeech/paraformer/utils/fix_data_feat.sh | 52
egs/aishell/paraformerbert/utils/shuffle_list.pl | 44
egs_modelscope/aishell/paraformer/utils/shuffle_list.pl | 44
egs_modelscope/speechio/paraformer/modelscope_utils/download_model.py | 25
egs_modelscope/common/utils/fix_data_feat.sh | 52
egs_modelscope/common_uniasr/modelscope_utils/modelscope_infer.sh | 88
egs_modelscope/aishell2/paraformer/utils/error_rate_zh | 370
egs/aishell/paraformer/utils/error_rate_zh | 370
egs_modelscope/aishell/paraformer/utils/run.pl | 356
egs_modelscope/speechio/paraformer/utils/proc_conf_oss.py | 35
egs_modelscope/common/utils/error_rate_zh | 370
docs/images/dingding.jpg | 0
egs_modelscope/common_uniasr/path.sh | 5
egs_modelscope/common_uniasr/utils/compute_cmvn.py | 74
egs_modelscope/common/utils/apply_cmvn.py | 79
egs_modelscope/common/utils/compute_cmvn.py | 74
egs/aishell/paraformerbert/utils/proce_text.py | 31
egs/aishell/conformer/utils/apply_lfr_and_cmvn.sh | 38
egs_modelscope/common_uniasr/utils/fix_data.sh | 35
egs_modelscope/speechio/paraformer/utils/error_rate_zh | 370
egs/aishell/conformer/utils/shuffle_list.pl | 44
egs_modelscope/wenetspeech/paraformer/utils/extract_embeds.py | 47
egs/aishell/paraformer/utils/compute_fbank.py | 153
egs_modelscope/common_uniasr/utils/combine_cmvn_file.py | 73
egs_modelscope/common_uniasr/utils/extract_embeds.py | 47
egs_modelscope/aishell/paraformer/utils/gen_ark_list.sh | 22
egs/aishell/tranformer/utils/compute_cmvn.py | 11
egs/aishell/conformer/utils/proce_text.py | 31
egs/aishell/paraformerbert/utils/split_scp.pl | 246
egs_modelscope/common_uniasr/utils/text_tokenize.sh | 35
egs_modelscope/aishell2/paraformer/paraformer_large_infer.sh | 2
egs/aishell/conformer/utils/run.pl | 356
egs/aishell/paraformer/utils/text_tokenize.py | 106
funasr/models/predictor/cif.py | 56
egs_modelscope/aishell/paraformer/utils/__init__.py | 0
egs_modelscope/common_uniasr/modelscope_common_infer_after_finetune.sh | 66
egs_modelscope/speechio/paraformer/utils/__init__.py | 0
egs_modelscope/aishell2/paraformer/utils/proce_text.py | 31
egs_modelscope/common/utils/split_data.py | 60
egs_modelscope/aishell/paraformer/utils/textnorm_zh.py | 834 +
egs_modelscope/common_uniasr/utils/text2token.py | 135
egs_modelscope/aishell/paraformer/utils/extract_embeds.py | 47
egs/aishell/paraformer/utils/print_args.py | 45
egs/aishell/conformer/run.sh | 92
egs_modelscope/aishell/paraformer/utils/combine_cmvn_file.py | 73
egs/aishell/paraformer/utils/apply_cmvn.sh | 29
egs_modelscope/aishell2/paraformer/utils/textnorm_zh.py | 834 +
egs_modelscope/speechio/paraformer/utils/text2token.py | 135
funasr/bin/asr_inference_paraformer.py | 321
egs_modelscope/wenetspeech/paraformer/utils/fix_data.sh | 35
egs_modelscope/aishell/paraformer/paraformer_large_finetune.sh | 88
egs/aishell/paraformer/utils/split_scp.pl | 246
egs_modelscope/wenetspeech/paraformer/utils/__init__.py | 0
egs_modelscope/aishell2/paraformer/utils/filter_scp.pl | 87
egs_modelscope/common_uniasr/utils/gen_ark_list.sh | 22
egs/aishell/paraformer/utils/compute_cmvn.sh | 25
egs_modelscope/common/utils/shuffle_list.pl | 44
egs/aishell/conformer/utils/textnorm_zh.py | 834 +
egs/aishell/conformer/utils/fix_data_feat.sh | 52
egs/aishell/paraformerbert/utils/compute_cmvn.sh | 25
egs/aishell/conformer/utils/fix_data.sh | 35
egs_modelscope/common/utils/compute_cmvn.sh | 25
egs/aishell/paraformer/utils/shuffle_list.pl | 44
egs/aishell/paraformerbert/utils/subset_data_dir_tr_cv.sh | 30
egs/aishell/paraformerbert/utils/apply_cmvn.sh | 29
egs_modelscope/common/utils/compute_fbank.py | 153
egs_modelscope/common_uniasr/conf/decode_asr_uniasr.yaml | 9
egs_modelscope/common_uniasr/conf/train_asr_uniasr_40e1_12d1_20e2_12d2_1280_320_lfr6.yaml | 192
egs_modelscope/common_uniasr/utils/proc_conf_oss.py | 35
egs/aishell/paraformerbert/utils/combine_cmvn_file.py | 73
egs_modelscope/aishell/paraformer/utils/compute_fbank.sh | 51
egs_modelscope/aishell/paraformer/conf/train_asr_paraformer_sanm_50e_16d_2048_512_lfr6.yaml | 2
egs_modelscope/wenetspeech/paraformer/utils/compute_fbank.py | 153
egs_modelscope/aishell/paraformer/utils/text2token.py | 135
egs/aishell/conformer/utils/filter_scp.pl | 87
egs_modelscope/common_uniasr/modelscope_utils/download_model.py | 25
egs_modelscope/wenetspeech/paraformer/utils/text_tokenize.py | 106
egs/aishell/paraformer/utils/apply_cmvn.py | 79
egs_modelscope/common/utils/text_tokenize.py | 106
egs/aishell/conformer/utils/__init__.py | 0
egs/aishell/conformer/utils/split_scp.pl | 246
egs_modelscope/wenetspeech/paraformer/utils/compute_cmvn.sh | 25
egs/aishell/conformer/utils/extract_embeds.py | 47
egs/aishell/paraformerbert/utils/compute_cmvn.py | 74
egs_modelscope/aishell/paraformer/utils/compute_cmvn.py | 74
egs_modelscope/aishell2/paraformer/utils/gen_ark_list.sh | 22
egs/aishell/paraformer/utils/compute_cmvn.py | 74
egs_modelscope/wenetspeech/paraformer/utils/gen_ark_list.sh | 22
egs_modelscope/common_uniasr/utils/apply_lfr_and_cmvn.sh | 38
egs_modelscope/common_uniasr/README.md | 27
egs_modelscope/speechio/paraformer/utils/combine_cmvn_file.py | 73
egs/aishell/paraformerbert/utils/__init__.py | 0
docs/images/funasr_logo.jpg | 0
egs_modelscope/aishell2/paraformer/modelscope_utils/download_model.py | 25
egs_modelscope/common_uniasr/utils/error_rate_zh | 370
egs/aishell/conformer/utils/apply_cmvn.sh | 29
egs_modelscope/aishell/paraformer/modelscope_utils/update_config.py | 41
funasr/models/e2e_asr_paraformer.py | 2
egs/aishell/paraformer/utils/proc_conf_oss.py | 35
egs/aishell/conformer/utils/proc_conf_oss.py | 35
egs_modelscope/speechio/paraformer/utils/compute_fbank.py | 153
egs_modelscope/speechio/paraformer/utils/compute_wer.py | 157
egs_modelscope/common_uniasr/utils/proce_text.py | 31
egs/aishell/paraformerbert/utils/compute_fbank.py | 153
egs/aishell/paraformer/utils/combine_cmvn_file.py | 73
egs_modelscope/speechio/paraformer/utils/fix_data.sh | 35
egs_modelscope/common/modelscope_common_finetune.sh | 97
egs/aishell/paraformerbert/utils/filter_scp.pl | 87
egs_modelscope/wenetspeech/paraformer/utils/error_rate_zh | 370
egs/aishell/conformer/utils/compute_cmvn.py | 74
egs_modelscope/common_uniasr/utils/subset_data_dir_tr_cv.sh | 30
egs/aishell/paraformer/run.sh | 94
egs_modelscope/common/utils/textnorm_zh.py | 834 +
egs_modelscope/aishell2/paraformer/utils/proc_conf_oss.py | 35
egs/aishell/paraformerbert/utils/apply_lfr_and_cmvn.py | 143
egs/aishell/paraformer/utils/__init__.py | 0
egs_modelscope/common_uniasr/modelscope_common_finetune.sh | 268
egs_modelscope/common_uniasr/utils/filter_scp.pl | 87
egs/aishell/paraformerbert/run.sh | 142
egs_modelscope/speechio/paraformer/utils/textnorm_zh.py | 834 +
egs_modelscope/common/utils/split_scp.pl | 246
egs/aishell/paraformer/utils/compute_wer.py | 157
egs_modelscope/aishell2/paraformer/utils/combine_cmvn_file.py | 73
egs/aishell/conformer/utils/apply_cmvn.py | 79
egs_modelscope/wenetspeech/paraformer/utils/text_tokenize.sh | 35
egs_modelscope/wenetspeech/paraformer/utils/apply_lfr_and_cmvn.py | 143
egs_modelscope/common_uniasr/utils/fix_data_feat.sh | 52
funasr/modules/streaming_utils/__init__.py | 0
egs/aishell/paraformerbert/utils/run.pl | 356
egs_modelscope/speechio/paraformer/utils/subset_data_dir_tr_cv.sh | 30
egs_modelscope/wenetspeech/paraformer/utils/compute_cmvn.py | 74
egs_modelscope/speechio/paraformer/utils/apply_lfr_and_cmvn.py | 143
egs_modelscope/wenetspeech/paraformer/utils/split_scp.pl | 246
egs/aishell/paraformerbert/utils/text_tokenize.py | 106
egs/aishell/tranformer/utils/combine_cmvn_file.py | 11
egs/aishell/conformer/utils/text_tokenize.sh | 35
README.md | 20
egs/aishell/conformer/utils/print_args.py | 45
egs_modelscope/common/utils/compute_wer.py | 157
egs/aishell/paraformerbert/utils/apply_lfr_and_cmvn.sh | 38
egs_modelscope/speechio/paraformer/utils/compute_fbank.sh | 51
egs_modelscope/aishell/paraformer/modelscope_utils/modelscope_infer.sh | 88
docs/modelscope_models.md | 34
egs_modelscope/wenetspeech/paraformer/utils/apply_cmvn.py | 79
egs_modelscope/aishell2/paraformer/modelscope_utils/modelscope_infer.sh | 88
egs_modelscope/speechio/paraformer/modelscope_utils/modelscope_infer.sh | 88
funasr/version.txt | 2
docs/images/.DS_Store | 0
egs/aishell/paraformerbert/utils/error_rate_zh | 370
egs_modelscope/wenetspeech/paraformer/utils/text2token.py | 135
egs_modelscope/aishell2/paraformer/utils/parse_options.sh | 97
egs_modelscope/speechio/paraformer/utils/apply_lfr_and_cmvn.sh | 38
egs_modelscope/wenetspeech/paraformer/utils/apply_lfr_and_cmvn.sh | 38
egs_modelscope/wenetspeech/paraformer/utils/apply_cmvn.sh | 29
egs_modelscope/aishell/paraformer/paraformer_large_infer.sh | 2
egs_modelscope/common/utils/subset_data_dir_tr_cv.sh | 30
egs/aishell/paraformerbert/utils/print_args.py | 45
egs/aishell/tranformer/utils/compute_fbank.sh | 4
egs_modelscope/common/utils/extract_embeds.py | 47
egs_modelscope/wenetspeech/paraformer/paraformer_large_infer.sh | 2
egs_modelscope/speechio/paraformer/utils/split_scp.pl | 246
egs/aishell/paraformerbert/utils/proc_conf_oss.py | 35
egs_modelscope/aishell2/paraformer/utils/extract_embeds.py | 47
egs/aishell/tranformer/run.sh | 93
egs_modelscope/aishell/paraformer/utils/parse_options.sh | 97
egs_modelscope/wenetspeech/paraformer/utils/print_args.py | 45
egs/aishell/conformer/utils/compute_fbank.py | 153
egs/aishell/paraformer/utils/fix_data_feat.sh | 52
egs_modelscope/speechio/paraformer/utils/proce_text.py | 31
egs_modelscope/aishell2/paraformer/utils/fix_data_feat.sh | 52
egs/aishell/paraformer/utils/run.pl | 356
egs_modelscope/common_uniasr/utils/split_data.py | 60
egs/aishell/paraformer/utils/fix_data.sh | 35
egs_modelscope/aishell/paraformer/modelscope_utils/download_model.py | 25
egs_modelscope/wenetspeech/paraformer/utils/compute_fbank.sh | 51
egs_modelscope/common/utils/compute_fbank.sh | 51
egs_modelscope/common_uniasr/utils/apply_lfr_and_cmvn.py | 143
egs_modelscope/wenetspeech/paraformer/utils/subset_data_dir_tr_cv.sh | 30
egs_modelscope/common/utils/__init__.py | 0
egs_modelscope/common_uniasr/utils/parse_options.sh | 97
/dev/null | 686 -
egs_modelscope/common/utils/text_tokenize.sh | 35
egs_modelscope/common/utils/proce_text.py | 31
egs/aishell/conformer/utils/text_tokenize.py | 106
egs/aishell/paraformerbert/utils/apply_cmvn.py | 79
egs_modelscope/aishell/paraformer/utils/fix_data_feat.sh | 52
funasr/bin/asr_inference_launch.py | 38
egs_modelscope/aishell/paraformer/utils/proce_text.py | 31
328 files changed, 34,008 insertions(+), 1,167 deletions(-)
diff --git a/README.md b/README.md
index b140a20..6dd38b2 100644
--- a/README.md
+++ b/README.md
@@ -1,8 +1,14 @@
-<div align="left"><img src="image/funasr_logo.jpg" width="400"/></div>
+<div align="left"><img src="docs/images/funasr_logo.jpg" width="400"/></div>
# FunASR: A Fundamental End-to-End Speech Recognition Toolkit
<strong>FunASR</strong> hopes to build a bridge between academic research and industrial applications on speech recognition. By supporting the training & finetuning of the industrial-grade speech recognition model released on [ModelScope](https://www.modelscope.cn/models?page=1&tasks=auto-speech-recognition), researchers and developers can conduct research and production of speech recognition models more conveniently, and promote the development of speech recognition ecology. ASR for Fun锛�
+
+## Highlights
+- FunASR supports many types of models, such as, Tranformer, Conformer, [Paraformer](https://arxiv.org/abs/2206.08317).
+- A large number of ASR models trained on academic datasets or industrial datasets are open sourced on [ModelScope](https://www.modelscope.cn/models?page=1&tasks=auto-speech-recognition),
+- The pretrained model [Paraformer-large](https://www.modelscope.cn/models/damo/speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-pytorch/summary) obtains the first place on many task in [SpeechIO leaderboard](https://github.com/SpeechColab/Leaderboard)
+- FunASR supports large-scale dataset dataloader and multi-GPU training.
## Installation(Training and Developing)
@@ -27,18 +33,24 @@
| 10.2 | conda install pytorch==1.8.0 torchvision==0.9.0 torchaudio==0.8.0 cudatoolkit=10.2 -c pytorch |
| 11.1 | conda install pytorch==1.8.0 torchvision==0.9.0 torchaudio==0.8.0 cudatoolkit=11.1 -c pytorch |
-For more versions, please see https://pytorch.org/get-started/locally/
+For more versions, please see [https://pytorch.org/get-started/locally](https://pytorch.org/get-started/locally)
- Install ModelScope:
``` sh
pip install "modelscope[audio]" -f https://modelscope.oss-cn-beijing.aliyuncs.com/releases/repo.html
```
-- Install other packages:
+For more details about modelscope, please see [modelscope installation](https://modelscope.cn/docs/%E7%8E%AF%E5%A2%83%E5%AE%89%E8%A3%85)
+
+- Install FunASR and other packages:
``` sh
pip install --editable ./
```
+
+## Pretrained model hub
+
+We have trained many academic and industrial models, [model hub](docs/modelscope_models.md)
## Contact
@@ -47,7 +59,7 @@
- email: [funasr@list.alibaba-inc.com](funasr@list.alibaba-inc.com)
- Dingding group:
-<div align="left"><img src="image/dingding.jpg" width="400"/></div>
+<div align="left"><img src="docs/images/dingding.jpg" width="400"/></div>
## Acknowledge
diff --git a/docs/images/.DS_Store b/docs/images/.DS_Store
new file mode 100644
index 0000000..5ef0f4c
--- /dev/null
+++ b/docs/images/.DS_Store
Binary files differ
diff --git a/image/dingding.jpg b/docs/images/dingding.jpg
similarity index 100%
rename from image/dingding.jpg
rename to docs/images/dingding.jpg
Binary files differ
diff --git a/image/funasr_logo.jpg b/docs/images/funasr_logo.jpg
similarity index 100%
rename from image/funasr_logo.jpg
rename to docs/images/funasr_logo.jpg
Binary files differ
diff --git a/image/wechat.png b/docs/images/wechat.png
similarity index 100%
rename from image/wechat.png
rename to docs/images/wechat.png
Binary files differ
diff --git a/docs/modelscope_models.md b/docs/modelscope_models.md
new file mode 100644
index 0000000..277d8e9
--- /dev/null
+++ b/docs/modelscope_models.md
@@ -0,0 +1,34 @@
+# Pretrained models on ModelScope
+
+## Model License
+- Apache License 2.0
+
+## Model Zoo
+Here we provided several pretrained models on different datasets. The details of models and datasets can be found on [ModelScope](https://www.modelscope.cn/models?page=1&tasks=auto-speech-recognition).
+
+| Datasets | Hours | Model | Online/Offline | Language | Framework | Checkpoint |
+|:-----:|:-----:|:--------------:|:--------------:| :---: | :---: | --- |
+| Alibaba Speech Data | 60000 | Paraformer | Offline | CN | Pytorch |[speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-pytorch](https://www.modelscope.cn/models/damo/speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-pytorch/summary) |
+| Alibaba Speech Data | 50000 | Paraformer | Offline | CN | Tensorflow |[speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8358-tensorflow1](https://www.modelscope.cn/models/damo/speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8358-tensorflow1/summary) |
+| Alibaba Speech Data | 50000 | Paraformer | Offline | CN | Tensorflow |[speech_paraformer_asr_nat-zh-cn-16k-common-vocab8358-tensorflow1](https://www.modelscope.cn/models/damo/speech_paraformer_asr_nat-zh-cn-16k-common-vocab8358-tensorflow1/summary) |
+| Alibaba Speech Data | 50000 | Paraformer | Online | CN | Tensorflow |[speech_paraformer_asr_nat-zh-cn-16k-common-vocab3444-tensorflow1-online](http://www.modelscope.cn/models/damo/speech_paraformer_asr_nat-zh-cn-16k-common-vocab3444-tensorflow1-online/summary) |
+| Alibaba Speech Data | 50000 | UniASR | Online | CN | Tensorflow |[speech_UniASR_asr_2pass-zh-cn-16k-common-vocab8358-tensorflow1-online](https://www.modelscope.cn/models/damo/speech_UniASR_asr_2pass-zh-cn-16k-common-vocab8358-tensorflow1-online/summary) |
+| Alibaba Speech Data | 50000 | UniASR | Offline | CN | Tensorflow |[speech_UniASR-large_asr_2pass-zh-cn-16k-common-vocab8358-tensorflow1-offline](https://www.modelscope.cn/models/damo/speech_UniASR-large_asr_2pass-zh-cn-16k-common-vocab8358-tensorflow1-offline/summary) |
+| Alibaba Speech Data | 50000 | UniASR | Online | CN&EN | Tensorflow |[speech_UniASR_asr_2pass-cn-en-moe-16k-vocab8358-tensorflow1-online](https://www.modelscope.cn/models/damo/speech_UniASR_asr_2pass-cn-en-moe-16k-vocab8358-tensorflow1-online/summary) |
+| Alibaba Speech Data | 50000 | UniASR | Offline | CN&EN | Tensorflow |[speech_UniASR_asr_2pass-cn-en-moe-16k-vocab8358-tensorflow1-offline](https://www.modelscope.cn/models/damo/speech_UniASR_asr_2pass-cn-en-moe-16k-vocab8358-tensorflow1-offline/summary) |
+| Alibaba Speech Data | 20000 | UniASR | Online | CN-Accent | Tensorflow |[speech_UniASR_asr_2pass-cn-dialect-16k-vocab8358-tensorflow1-online](https://www.modelscope.cn/models/damo/speech_UniASR_asr_2pass-cn-dialect-16k-vocab8358-tensorflow1-online/summary) |
+| Alibaba Speech Data | 20000 | UniASR | Offline | CN-Accent | Tensorflow |[speech_UniASR_asr_2pass-cn-dialect-16k-vocab8358-tensorflow1-offline](https://www.modelscope.cn/models/damo/speech_UniASR_asr_2pass-cn-dialect-16k-vocab8358-tensorflow1-offline/summary) |
+| Alibaba Speech Data | 30000 | Paraformer-8K | Online | CN | Tensorflow |[speech_paraformer_asr_nat-zh-cn-8k-common-vocab3444-tensorflow1-online](https://www.modelscope.cn/models/damo/speech_paraformer_asr_nat-zh-cn-8k-common-vocab3444-tensorflow1-online/summary) |
+| Alibaba Speech Data | 30000 | Paraformer-8K | Offline | CN | Tensorflow |[speech_paraformer_asr_nat-zh-cn-8k-common-vocab8358-tensorflow1](https://www.modelscope.cn/models/damo/speech_paraformer_asr_nat-zh-cn-8k-common-vocab8358-tensorflow1/summary) |
+| Alibaba Speech Data | 30000 | Paraformer-8K | Online | CN | Pytorch |[speech_UniASR_asr_2pass-zh-cn-8k-common-vocab3445-pytorch-online](https://www.modelscope.cn/models/damo/speech_UniASR_asr_2pass-zh-cn-8k-common-vocab3445-pytorch-online/summary) |
+| Alibaba Speech Data | 30000 | Paraformer-8K | Offline | CN | Pytorch |[speech_UniASR_asr_2pass-zh-cn-8k-common-vocab3445-pytorch-offline](https://www.modelscope.cn/models/damo/speech_UniASR_asr_2pass-zh-cn-8k-common-vocab3445-pytorch-offline/summary) |
+| Alibaba Speech Data | 30000 | UniASR-8K | Online | CN | Tensorflow |[speech_UniASR_asr_2pass-zh-cn-8k-common-vocab8358-tensorflow1-online](https://www.modelscope.cn/models/damo/speech_UniASR_asr_2pass-zh-cn-8k-common-vocab8358-tensorflow1-online/summary) |
+| Alibaba Speech Data | 30000 | UniASR-8K | Offline | CN | Tensorflow |[speech_UniASR_asr_2pass-zh-cn-8k-common-vocab8358-tensorflow1-offline](https://www.modelscope.cn/models/damo/speech_UniASR_asr_2pass-zh-cn-8k-common-vocab8358-tensorflow1-offline/summary) |
+| Alibaba Speech Data | 30000 | UniASR-8K | Online | CN | Pytorch |[speech_UniASR_asr_2pass-zh-cn-8k-common-vocab3445-pytorch-online](https://www.modelscope.cn/models/damo/speech_UniASR_asr_2pass-zh-cn-8k-common-vocab3445-pytorch-online/summary) |
+| Alibaba Speech Data | 30000 | UniASR-8K | Offline | CN | Pytorch |[speech_UniASR_asr_2pass-zh-cn-8k-common-vocab3445-pytorch-offline](https://www.modelscope.cn/models/damo/speech_UniASR_asr_2pass-zh-cn-8k-common-vocab3445-pytorch-offline/summary) |
+| AISHELL-1 | 178 | Paraformer | Offline | CN | Pytorch | [speech_paraformer_asr_nat-aishell1-pytorch](https://www.modelscope.cn/models/damo/speech_paraformer_asr_nat-aishell1-pytorch/summary) |
+| AISHELL-2 | 1000 | Paraformer | Offline | CN | Pytorch | [speech_paraformer_asr_nat-aishell2-pytorch](https://www.modelscope.cn/models/damo/speech_paraformer_asr_nat-aishell2-pytorch/summary) |
+| AISHELL-1 | 178 | ParaformerBert | Offline | CN | Pytorch | [speech_paraformerbert_asr_nat-zh-cn-16k-aishell1-vocab4234-pytorch](https://modelscope.cn/models/damo/speech_paraformerbert_asr_nat-zh-cn-16k-aishell1-vocab4234-pytorch/summary) |
+| AISHELL-2 | 1000 | ParaformerBert | Offline | CN | Pytorch | [speech_paraformerbert_asr_nat-zh-cn-16k-aishell2-vocab5212-pytorch](https://modelscope.cn/models/damo/speech_paraformerbert_asr_nat-zh-cn-16k-aishell2-vocab5212-pytorch/summary) |
+| AISHELL-1 | 178 | Conformer | Offline | CN | Pytorch | [speech_conformer_asr_nat-zh-cn-16k-aishell1-vocab4234-pytorch](https://modelscope.cn/models/damo/speech_conformer_asr_nat-zh-cn-16k-aishell1-vocab4234-pytorch/summary) |
+| AISHELL-2 | 1000 | Conformer | Offline | CN | Pytorch | [speech_conformer_asr_nat-zh-cn-16k-aishell2-vocab5212-pytorch](https://modelscope.cn/models/damo/speech_conformer_asr_nat-zh-cn-16k-aishell2-vocab5212-pytorch/summary) |
diff --git a/egs/aishell/conformer/run.sh b/egs/aishell/conformer/run.sh
index 16ebc67..d865982 100755
--- a/egs/aishell/conformer/run.sh
+++ b/egs/aishell/conformer/run.sh
@@ -10,9 +10,10 @@
# for gpu decoding, inference_nj=ngpu*njob; for cpu decoding, inference_nj=njob
njob=8
train_cmd=utils/run.pl
+infer_cmd=utils/run.pl
# general configuration
-feats_dir=".." #feature output dictionary, for large data
+feats_dir="../DATA" #feature output dictionary, for large data
exp_dir="."
lang=zh
dumpdir=dump/fbank
@@ -59,8 +60,10 @@
if ${gpu_inference}; then
inference_nj=$[${ngpu}*${njob}]
+ _ngpu=1
else
inference_nj=$njob
+ _ngpu=0
fi
if [ ${stage} -le 0 ] && [ ${stop_stage} -ge 0 ]; then
@@ -83,18 +86,18 @@
echo "stage 1: Feature Generation"
# compute fbank features
fbankdir=${feats_dir}/fbank
- utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj --speed_perturb ${speed_perturb} \
+ utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj --feats_dim ${feats_dim} --sample_frequency ${sample_frequency} --speed_perturb ${speed_perturb} \
${feats_dir}/data/train ${exp_dir}/exp/make_fbank/train ${fbankdir}/train
utils/fix_data_feat.sh ${fbankdir}/train
- utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj \
+ utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj --feats_dim ${feats_dim} --sample_frequency ${sample_frequency} \
${feats_dir}/data/dev ${exp_dir}/exp/make_fbank/dev ${fbankdir}/dev
utils/fix_data_feat.sh ${fbankdir}/dev
- utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj \
+ utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj --feats_dim ${feats_dim} --sample_frequency ${sample_frequency} \
${feats_dir}/data/test ${exp_dir}/exp/make_fbank/test ${fbankdir}/test
utils/fix_data_feat.sh ${fbankdir}/test
# compute global cmvn
- utils/compute_cmvn.sh --cmd "$train_cmd" --nj $nj \
+ utils/compute_cmvn.sh --cmd "$train_cmd" --nj $nj --feats_dim ${feats_dim} \
${fbankdir}/train ${exp_dir}/exp/make_fbank/train
# apply cmvn
@@ -112,6 +115,10 @@
utils/fix_data_feat.sh ${feat_train_dir}
utils/fix_data_feat.sh ${feat_dev_dir}
utils/fix_data_feat.sh ${feat_test_dir}
+
+ #generate ark list
+ utils/gen_ark_list.sh --cmd "$train_cmd" --nj $nj ${feat_train_dir} ${fbankdir}/train ${feat_train_dir}
+ utils/gen_ark_list.sh --cmd "$train_cmd" --nj $nj ${feat_dev_dir} ${fbankdir}/dev ${feat_dev_dir}
fi
token_list=${feats_dir}/data/${lang}_token_list/char/tokens.txt
@@ -140,9 +147,10 @@
# Training Stage
world_size=$gpu_num # run on one machine
if [ ${stage} -le 3 ] && [ ${stop_stage} -ge 3 ]; then
+ echo "stage 3: Training"
mkdir -p ${exp_dir}/exp/${model_dir}
mkdir -p ${exp_dir}/exp/${model_dir}/log
- INIT_FILE=$exp_dir/ddp_init
+ INIT_FILE=${exp_dir}/exp/${model_dir}/ddp_init
if [ -f $INIT_FILE ];then
rm -f $INIT_FILE
fi
@@ -184,25 +192,57 @@
# Testing Stage
if [ ${stage} -le 4 ] && [ ${stop_stage} -ge 4 ]; then
- utils/easy_asr_infer.sh \
- --lang zh \
- --datadir ${feats_dir} \
- --feats_type ${feats_type} \
- --feats_dim ${feats_dim} \
- --token_type ${token_type} \
- --gpu_inference ${gpu_inference} \
- --inference_config "${inference_config}" \
- --test_sets "${test_sets}" \
- --token_list $token_list \
- --asr_exp ${exp_dir}/${model_dir} \
- --stage 12 \
- --stop_stage 12 \
- --scp $scp \
- --text text \
- --inference_nj $inference_nj \
- --njob $njob \
- --inference_asr_model $inference_asr_model \
- --gpuid_list $gpuid_list \
- --mode asr
+ echo "stage 4: Inference"
+ for dset in ${test_sets}; do
+ asr_exp=${exp_dir}/exp/${model_dir}
+ inference_tag="$(basename "${inference_config}" .yaml)"
+ _dir="${asr_exp}/${inference_tag}/${inference_asr_model}/${dset}"
+ _logdir="${_dir}/logdir"
+ if [ -d ${_dir} ]; then
+ echo "${_dir} is already exists. if you want to decode again, please delete this dir first."
+ exit 0
+ fi
+ mkdir -p "${_logdir}"
+ _data="${feats_dir}/${dumpdir}/${dset}"
+ key_file=${_data}/${scp}
+ num_scp_file="$(<${key_file} wc -l)"
+ _nj=$([ $inference_nj -le $num_scp_file ] && echo "$inference_nj" || echo "$num_scp_file")
+ split_scps=
+ for n in $(seq "${_nj}"); do
+ split_scps+=" ${_logdir}/keys.${n}.scp"
+ done
+ # shellcheck disable=SC2086
+ utils/split_scp.pl "${key_file}" ${split_scps}
+ _opts=
+ if [ -n "${inference_config}" ]; then
+ _opts+="--config ${inference_config} "
+ fi
+ ${infer_cmd} --gpu "${_ngpu}" --max-jobs-run "${_nj}" JOB=1:"${_nj}" "${_logdir}"/asr_inference.JOB.log \
+ python -m funasr.bin.asr_inference_launch \
+ --batch_size 1 \
+ --ngpu "${_ngpu}" \
+ --njob ${njob} \
+ --gpuid_list ${gpuid_list} \
+ --data_path_and_name_and_type "${_data}/${scp},speech,${type}" \
+ --key_file "${_logdir}"/keys.JOB.scp \
+ --asr_train_config "${asr_exp}"/config.yaml \
+ --asr_model_file "${asr_exp}"/"${inference_asr_model}" \
+ --output_dir "${_logdir}"/output.JOB \
+ --mode asr \
+ ${_opts}
+
+ for f in token token_int score text; do
+ if [ -f "${_logdir}/output.1/1best_recog/${f}" ]; then
+ for i in $(seq "${_nj}"); do
+ cat "${_logdir}/output.${i}/1best_recog/${f}"
+ done | sort -k1 >"${_dir}/${f}"
+ fi
+ done
+ python utils/proce_text.py ${_dir}/text ${_dir}/text.proc
+ python utils/proce_text.py ${_data}/text ${_data}/text.proc
+ python utils/compute_wer.py ${_data}/text.proc ${_dir}/text.proc ${_dir}/text.cer
+ tail -n 3 ${_dir}/text.cer > ${_dir}/text.cer.txt
+ cat ${_dir}/text.cer.txt
+ done
fi
diff --git a/egs/aishell/conformer/utils b/egs/aishell/conformer/utils
deleted file mode 120000
index 40e14f5..0000000
--- a/egs/aishell/conformer/utils
+++ /dev/null
@@ -1 +0,0 @@
-../tranformer/utils
\ No newline at end of file
diff --git a/egs/aishell/conformer/utils/__init__.py b/egs/aishell/conformer/utils/__init__.py
new file mode 100644
index 0000000..e69de29
--- /dev/null
+++ b/egs/aishell/conformer/utils/__init__.py
diff --git a/egs/aishell/conformer/utils/apply_cmvn.py b/egs/aishell/conformer/utils/apply_cmvn.py
new file mode 100755
index 0000000..b5c5086
--- /dev/null
+++ b/egs/aishell/conformer/utils/apply_cmvn.py
@@ -0,0 +1,79 @@
+from kaldiio import ReadHelper
+from kaldiio import WriteHelper
+
+import argparse
+import json
+import math
+import numpy as np
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="apply cmvn",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--ark-file",
+ "-a",
+ default=False,
+ required=True,
+ type=str,
+ help="fbank ark file",
+ )
+ parser.add_argument(
+ "--cmvn-file",
+ "-c",
+ default=False,
+ required=True,
+ type=str,
+ help="cmvn file",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ ark_file = args.output_dir + "/feats." + str(args.ark_index) + ".ark"
+ scp_file = args.output_dir + "/feats." + str(args.ark_index) + ".scp"
+ ark_writer = WriteHelper('ark,scp:{},{}'.format(ark_file, scp_file))
+
+ with open(args.cmvn_file) as f:
+ cmvn_stats = json.load(f)
+
+ means = cmvn_stats['mean_stats']
+ vars = cmvn_stats['var_stats']
+ total_frames = cmvn_stats['total_frames']
+
+ for i in range(len(means)):
+ means[i] /= total_frames
+ vars[i] = vars[i] / total_frames - means[i] * means[i]
+ if vars[i] < 1.0e-20:
+ vars[i] = 1.0e-20
+ vars[i] = 1.0 / math.sqrt(vars[i])
+
+ with ReadHelper('ark:{}'.format(args.ark_file)) as ark_reader:
+ for key, mat in ark_reader:
+ mat = (mat - means) * vars
+ ark_writer(key, mat)
+
+
+if __name__ == '__main__':
+ main()
diff --git a/egs/aishell/conformer/utils/apply_cmvn.sh b/egs/aishell/conformer/utils/apply_cmvn.sh
new file mode 100755
index 0000000..f8fd1d1
--- /dev/null
+++ b/egs/aishell/conformer/utils/apply_cmvn.sh
@@ -0,0 +1,29 @@
+#!/usr/bin/env bash
+
+. ./path.sh || exit 1;
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+fbankdir=$1
+cmvn_file=$2
+logdir=$3
+output_dir=$4
+
+dump_dir=${output_dir}/ark; mkdir -p ${dump_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/apply_cmvn.JOB.log \
+ python utils/apply_cmvn.py -a $fbankdir/ark/feats.JOB.ark \
+ -c $cmvn_file -i JOB -o ${dump_dir} \
+ || exit 1;
+
+for n in $(seq $nj); do
+ cat ${dump_dir}/feats.$n.scp || exit 1
+done > ${output_dir}/feats.scp || exit 1
+
+echo "$0: Succeeded apply cmvn"
diff --git a/egs/aishell/conformer/utils/apply_lfr_and_cmvn.py b/egs/aishell/conformer/utils/apply_lfr_and_cmvn.py
new file mode 100755
index 0000000..50d18d1
--- /dev/null
+++ b/egs/aishell/conformer/utils/apply_lfr_and_cmvn.py
@@ -0,0 +1,143 @@
+from kaldiio import ReadHelper, WriteHelper
+
+import argparse
+import numpy as np
+
+
+def build_LFR_features(inputs, m=7, n=6):
+ LFR_inputs = []
+ T = inputs.shape[0]
+ T_lfr = int(np.ceil(T / n))
+ left_padding = np.tile(inputs[0], ((m - 1) // 2, 1))
+ inputs = np.vstack((left_padding, inputs))
+ T = T + (m - 1) // 2
+ for i in range(T_lfr):
+ if m <= T - i * n:
+ LFR_inputs.append(np.hstack(inputs[i * n:i * n + m]))
+ else:
+ num_padding = m - (T - i * n)
+ frame = np.hstack(inputs[i * n:])
+ for _ in range(num_padding):
+ frame = np.hstack((frame, inputs[-1]))
+ LFR_inputs.append(frame)
+ return np.vstack(LFR_inputs)
+
+
+def build_CMVN_features(inputs, mvn_file): # noqa
+ with open(mvn_file, 'r', encoding='utf-8') as f:
+ lines = f.readlines()
+
+ add_shift_list = []
+ rescale_list = []
+ for i in range(len(lines)):
+ line_item = lines[i].split()
+ if line_item[0] == '<AddShift>':
+ line_item = lines[i + 1].split()
+ if line_item[0] == '<LearnRateCoef>':
+ add_shift_line = line_item[3:(len(line_item) - 1)]
+ add_shift_list = list(add_shift_line)
+ continue
+ elif line_item[0] == '<Rescale>':
+ line_item = lines[i + 1].split()
+ if line_item[0] == '<LearnRateCoef>':
+ rescale_line = line_item[3:(len(line_item) - 1)]
+ rescale_list = list(rescale_line)
+ continue
+
+ for j in range(inputs.shape[0]):
+ for k in range(inputs.shape[1]):
+ add_shift_value = add_shift_list[k]
+ rescale_value = rescale_list[k]
+ inputs[j, k] = float(inputs[j, k]) + float(add_shift_value)
+ inputs[j, k] = float(inputs[j, k]) * float(rescale_value)
+
+ return inputs
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="apply low_frame_rate and cmvn",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--ark-file",
+ "-a",
+ default=False,
+ required=True,
+ type=str,
+ help="fbank ark file",
+ )
+ parser.add_argument(
+ "--lfr",
+ "-f",
+ default=True,
+ type=str,
+ help="low frame rate",
+ )
+ parser.add_argument(
+ "--lfr-m",
+ "-m",
+ default=7,
+ type=int,
+ help="number of frames to stack",
+ )
+ parser.add_argument(
+ "--lfr-n",
+ "-n",
+ default=6,
+ type=int,
+ help="number of frames to skip",
+ )
+ parser.add_argument(
+ "--cmvn-file",
+ "-c",
+ default=False,
+ required=True,
+ type=str,
+ help="global cmvn file",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ dump_ark_file = args.output_dir + "/feats." + str(args.ark_index) + ".ark"
+ dump_scp_file = args.output_dir + "/feats." + str(args.ark_index) + ".scp"
+ shape_file = args.output_dir + "/len." + str(args.ark_index)
+ ark_writer = WriteHelper('ark,scp:{},{}'.format(dump_ark_file, dump_scp_file))
+
+ shape_writer = open(shape_file, 'w')
+ with ReadHelper('ark:{}'.format(args.ark_file)) as ark_reader:
+ for key, mat in ark_reader:
+ if args.lfr:
+ lfr = build_LFR_features(mat, args.lfr_m, args.lfr_n)
+ else:
+ lfr = mat
+ cmvn = build_CMVN_features(lfr, args.cmvn_file)
+ dims = cmvn.shape[1]
+ lens = cmvn.shape[0]
+ shape_writer.write(key + " " + str(lens) + "," + str(dims) + '\n')
+ ark_writer(key, cmvn)
+
+
+if __name__ == '__main__':
+ main()
+
diff --git a/egs/aishell/conformer/utils/apply_lfr_and_cmvn.sh b/egs/aishell/conformer/utils/apply_lfr_and_cmvn.sh
new file mode 100755
index 0000000..3119fdb
--- /dev/null
+++ b/egs/aishell/conformer/utils/apply_lfr_and_cmvn.sh
@@ -0,0 +1,38 @@
+#!/usr/bin/env bash
+
+
+# Begin configuration section.
+nj=32
+cmd=utils/run.pl
+
+# feature configuration
+lfr=True
+lfr_m=7
+lfr_n=6
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+fbankdir=$1
+cmvn_file=$2
+logdir=$3
+output_dir=$4
+
+dump_dir=${output_dir}/ark; mkdir -p ${dump_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/apply_lfr_and_cmvn.JOB.log \
+ python utils/apply_lfr_and_cmvn.py -a $fbankdir/ark/feats.JOB.ark \
+ -f $lfr -m $lfr_m -n $lfr_n -c $cmvn_file -i JOB -o ${dump_dir} \
+ || exit 1;
+
+for n in $(seq $nj); do
+ cat ${dump_dir}/feats.$n.scp || exit 1
+done > ${output_dir}/feats.scp || exit 1
+
+for n in $(seq $nj); do
+ cat ${dump_dir}/len.$n || exit 1
+done > ${output_dir}/speech_shape || exit 1
+
+echo "$0: Succeeded apply low frame rate and cmvn"
diff --git a/egs/aishell/conformer/utils/combine_cmvn_file.py b/egs/aishell/conformer/utils/combine_cmvn_file.py
new file mode 100755
index 0000000..b2974a4
--- /dev/null
+++ b/egs/aishell/conformer/utils/combine_cmvn_file.py
@@ -0,0 +1,73 @@
+import argparse
+import json
+import numpy as np
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="combine cmvn file",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--dims",
+ "-d",
+ default=80,
+ type=int,
+ help="feature dims",
+ )
+ parser.add_argument(
+ "--cmvn-dir",
+ "-c",
+ default=False,
+ required=True,
+ type=str,
+ help="cmvn dir",
+ )
+
+ parser.add_argument(
+ "--nj",
+ "-n",
+ default=1,
+ required=True,
+ type=int,
+ help="num of cmvn file",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ total_means = np.zeros(args.dims)
+ total_vars = np.zeros(args.dims)
+ total_frames = 0
+
+ cmvn_file = args.output_dir + "/cmvn.json"
+
+ for i in range(1, args.nj+1):
+ with open(args.cmvn_dir + "/cmvn." + str(i) + ".json", "r") as fin:
+ cmvn_stats = json.load(fin)
+
+ total_means += np.array(cmvn_stats["mean_stats"])
+ total_vars += np.array(cmvn_stats["var_stats"])
+ total_frames += cmvn_stats["total_frames"]
+
+ cmvn_info = {
+ 'mean_stats': list(total_means.tolist()),
+ 'var_stats': list(total_vars.tolist()),
+ 'total_frames': total_frames
+ }
+ with open(cmvn_file, 'w') as fout:
+ fout.write(json.dumps(cmvn_info))
+
+
+if __name__ == '__main__':
+ main()
diff --git a/egs/aishell/conformer/utils/compute_cmvn.py b/egs/aishell/conformer/utils/compute_cmvn.py
new file mode 100755
index 0000000..2b96e26
--- /dev/null
+++ b/egs/aishell/conformer/utils/compute_cmvn.py
@@ -0,0 +1,74 @@
+from kaldiio import ReadHelper
+
+import argparse
+import numpy as np
+import json
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="computer global cmvn",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--dims",
+ "-d",
+ default=80,
+ type=int,
+ help="feature dims",
+ )
+ parser.add_argument(
+ "--ark-file",
+ "-a",
+ default=False,
+ required=True,
+ type=str,
+ help="fbank ark file",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ ark_file = args.ark_file + "/feats." + str(args.ark_index) + ".ark"
+ cmvn_file = args.output_dir + "/cmvn." + str(args.ark_index) + ".json"
+
+ mean_stats = np.zeros(args.dims)
+ var_stats = np.zeros(args.dims)
+ total_frames = 0
+
+ with ReadHelper('ark:{}'.format(ark_file)) as ark_reader:
+ for key, mat in ark_reader:
+ mean_stats += np.sum(mat, axis=0)
+ var_stats += np.sum(np.square(mat), axis=0)
+ total_frames += mat.shape[0]
+
+ cmvn_info = {
+ 'mean_stats': list(mean_stats.tolist()),
+ 'var_stats': list(var_stats.tolist()),
+ 'total_frames': total_frames
+ }
+ with open(cmvn_file, 'w') as fout:
+ fout.write(json.dumps(cmvn_info))
+
+
+if __name__ == '__main__':
+ main()
diff --git a/egs/aishell/conformer/utils/compute_cmvn.sh b/egs/aishell/conformer/utils/compute_cmvn.sh
new file mode 100755
index 0000000..12173ee
--- /dev/null
+++ b/egs/aishell/conformer/utils/compute_cmvn.sh
@@ -0,0 +1,25 @@
+#!/usr/bin/env bash
+
+. ./path.sh || exit 1;
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+feats_dim=80
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+fbankdir=$1
+logdir=$2
+
+output_dir=${fbankdir}/cmvn; mkdir -p ${output_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/cmvn.JOB.log \
+ python utils/compute_cmvn.py -d ${feats_dim} -a $fbankdir/ark -i JOB -o ${output_dir} \
+ || exit 1;
+
+python utils/combine_cmvn_file.py -d ${feats_dim} -c ${output_dir} -n $nj -o $fbankdir
+
+echo "$0: Succeeded compute global cmvn"
diff --git a/egs/aishell/conformer/utils/compute_fbank.py b/egs/aishell/conformer/utils/compute_fbank.py
new file mode 100755
index 0000000..d03b5a8
--- /dev/null
+++ b/egs/aishell/conformer/utils/compute_fbank.py
@@ -0,0 +1,153 @@
+from kaldiio import WriteHelper
+
+import argparse
+import numpy as np
+import json
+import torch
+import torchaudio
+import torchaudio.compliance.kaldi as kaldi
+
+
+def compute_fbank(wav_file,
+ num_mel_bins=80,
+ frame_length=25,
+ frame_shift=10,
+ dither=0.0,
+ resample_rate=16000,
+ speed=1.0):
+
+ waveform, sample_rate = torchaudio.load(wav_file)
+ if resample_rate != sample_rate:
+ waveform = torchaudio.transforms.Resample(orig_freq=sample_rate,
+ new_freq=resample_rate)(waveform)
+ if speed != 1.0:
+ waveform, _ = torchaudio.sox_effects.apply_effects_tensor(
+ waveform, resample_rate,
+ [['speed', str(speed)], ['rate', str(resample_rate)]]
+ )
+
+ waveform = waveform * (1 << 15)
+ mat = kaldi.fbank(waveform,
+ num_mel_bins=num_mel_bins,
+ frame_length=frame_length,
+ frame_shift=frame_shift,
+ dither=dither,
+ energy_floor=0.0,
+ window_type='hamming',
+ sample_frequency=resample_rate)
+
+ return mat.numpy()
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="computer features",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--wav-lists",
+ "-w",
+ default=False,
+ required=True,
+ type=str,
+ help="input wav lists",
+ )
+ parser.add_argument(
+ "--text-files",
+ "-t",
+ default=False,
+ required=True,
+ type=str,
+ help="input text files",
+ )
+ parser.add_argument(
+ "--dims",
+ "-d",
+ default=80,
+ type=int,
+ help="feature dims",
+ )
+ parser.add_argument(
+ "--sample-frequency",
+ "-s",
+ default=16000,
+ type=int,
+ help="sample frequency",
+ )
+ parser.add_argument(
+ "--speed-perturb",
+ "-p",
+ default="1.0",
+ type=str,
+ help="speed perturb",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-a",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ ark_file = args.output_dir + "/ark/feats." + str(args.ark_index) + ".ark"
+ scp_file = args.output_dir + "/ark/feats." + str(args.ark_index) + ".scp"
+ text_file = args.output_dir + "/txt/text." + str(args.ark_index) + ".txt"
+ feats_shape_file = args.output_dir + "/ark/len." + str(args.ark_index)
+ text_shape_file = args.output_dir + "/txt/len." + str(args.ark_index)
+
+ ark_writer = WriteHelper('ark,scp:{},{}'.format(ark_file, scp_file))
+ text_writer = open(text_file, 'w')
+ feats_shape_writer = open(feats_shape_file, 'w')
+ text_shape_writer = open(text_shape_file, 'w')
+
+ speed_perturb_list = args.speed_perturb.split(',')
+
+ for speed in speed_perturb_list:
+ with open(args.wav_lists, 'r', encoding='utf-8') as wavfile:
+ with open(args.text_files, 'r', encoding='utf-8') as textfile:
+ for wav, text in zip(wavfile, textfile):
+ s_w = wav.strip().split()
+ wav_id = s_w[0]
+ wav_file = s_w[1]
+
+ s_t = text.strip().split()
+ text_id = s_t[0]
+ txt = s_t[1:]
+ fbank = compute_fbank(wav_file,
+ num_mel_bins=args.dims,
+ resample_rate=args.sample_frequency,
+ speed=float(speed)
+ )
+ feats_dims = fbank.shape[1]
+ feats_lens = fbank.shape[0]
+ txt_lens = len(txt)
+ if speed == "1.0":
+ wav_id_sp = wav_id
+ else:
+ wav_id_sp = wav_id + "_sp" + speed
+
+ feats_shape_writer.write(wav_id_sp + " " + str(feats_lens) + "," + str(feats_dims) + '\n')
+ text_shape_writer.write(wav_id_sp + " " + str(txt_lens) + '\n')
+
+ text_writer.write(wav_id_sp + " " + " ".join(txt) + '\n')
+ ark_writer(wav_id_sp, fbank)
+
+
+if __name__ == '__main__':
+ main()
+
diff --git a/egs/aishell/conformer/utils/compute_fbank.sh b/egs/aishell/conformer/utils/compute_fbank.sh
new file mode 100755
index 0000000..92a4fe6
--- /dev/null
+++ b/egs/aishell/conformer/utils/compute_fbank.sh
@@ -0,0 +1,51 @@
+#!/usr/bin/env bash
+
+. ./path.sh || exit 1;
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+
+# feature configuration
+feats_dim=80
+sample_frequency=16000
+speed_perturb="1.0"
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+data=$1
+logdir=$2
+fbankdir=$3
+
+[ ! -f $data/wav.scp ] && echo "$0: no such file $data/wav.scp" && exit 1;
+[ ! -f $data/text ] && echo "$0: no such file $data/text" && exit 1;
+
+python utils/split_data.py $data $data $nj
+
+ark_dir=${fbankdir}/ark; mkdir -p ${ark_dir}
+text_dir=${fbankdir}/txt; mkdir -p ${text_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/make_fbank.JOB.log \
+ python utils/compute_fbank.py -w $data/split${nj}/JOB/wav.scp -t $data/split${nj}/JOB/text \
+ -d $feats_dim -s $sample_frequency -p ${speed_perturb} -a JOB -o ${fbankdir} \
+ || exit 1;
+
+for n in $(seq $nj); do
+ cat ${ark_dir}/feats.$n.scp || exit 1
+done > $fbankdir/feats.scp || exit 1
+
+for n in $(seq $nj); do
+ cat ${text_dir}/text.$n.txt || exit 1
+done > $fbankdir/text || exit 1
+
+for n in $(seq $nj); do
+ cat ${ark_dir}/len.$n || exit 1
+done > $fbankdir/speech_shape || exit 1
+
+for n in $(seq $nj); do
+ cat ${text_dir}/len.$n || exit 1
+done > $fbankdir/text_shape || exit 1
+
+echo "$0: Succeeded compute FBANK features"
diff --git a/egs/aishell/conformer/utils/compute_wer.py b/egs/aishell/conformer/utils/compute_wer.py
new file mode 100755
index 0000000..349a3f6
--- /dev/null
+++ b/egs/aishell/conformer/utils/compute_wer.py
@@ -0,0 +1,157 @@
+import os
+import numpy as np
+import sys
+
+def compute_wer(ref_file,
+ hyp_file,
+ cer_detail_file):
+ rst = {
+ 'Wrd': 0,
+ 'Corr': 0,
+ 'Ins': 0,
+ 'Del': 0,
+ 'Sub': 0,
+ 'Snt': 0,
+ 'Err': 0.0,
+ 'S.Err': 0.0,
+ 'wrong_words': 0,
+ 'wrong_sentences': 0
+ }
+
+ hyp_dict = {}
+ ref_dict = {}
+ with open(hyp_file, 'r') as hyp_reader:
+ for line in hyp_reader:
+ key = line.strip().split()[0]
+ value = line.strip().split()[1:]
+ hyp_dict[key] = value
+ with open(ref_file, 'r') as ref_reader:
+ for line in ref_reader:
+ key = line.strip().split()[0]
+ value = line.strip().split()[1:]
+ ref_dict[key] = value
+
+ cer_detail_writer = open(cer_detail_file, 'w')
+ for hyp_key in hyp_dict:
+ if hyp_key in ref_dict:
+ out_item = compute_wer_by_line(hyp_dict[hyp_key], ref_dict[hyp_key])
+ rst['Wrd'] += out_item['nwords']
+ rst['Corr'] += out_item['cor']
+ rst['wrong_words'] += out_item['wrong']
+ rst['Ins'] += out_item['ins']
+ rst['Del'] += out_item['del']
+ rst['Sub'] += out_item['sub']
+ rst['Snt'] += 1
+ if out_item['wrong'] > 0:
+ rst['wrong_sentences'] += 1
+ cer_detail_writer.write(hyp_key + print_cer_detail(out_item) + '\n')
+ cer_detail_writer.write("ref:" + '\t' + "".join(ref_dict[hyp_key]) + '\n')
+ cer_detail_writer.write("hyp:" + '\t' + "".join(hyp_dict[hyp_key]) + '\n')
+
+ if rst['Wrd'] > 0:
+ rst['Err'] = round(rst['wrong_words'] * 100 / rst['Wrd'], 2)
+ if rst['Snt'] > 0:
+ rst['S.Err'] = round(rst['wrong_sentences'] * 100 / rst['Snt'], 2)
+
+ cer_detail_writer.write('\n')
+ cer_detail_writer.write("%WER " + str(rst['Err']) + " [ " + str(rst['wrong_words'])+ " / " + str(rst['Wrd']) +
+ ", " + str(rst['Ins']) + " ins, " + str(rst['Del']) + " del, " + str(rst['Sub']) + " sub ]" + '\n')
+ cer_detail_writer.write("%SER " + str(rst['S.Err']) + " [ " + str(rst['wrong_sentences']) + " / " + str(rst['Snt']) + " ]" + '\n')
+ cer_detail_writer.write("Scored " + str(len(hyp_dict)) + " sentences, " + str(len(hyp_dict) - rst['Snt']) + " not present in hyp." + '\n')
+
+
+def compute_wer_by_line(hyp,
+ ref):
+ hyp = list(map(lambda x: x.lower(), hyp))
+ ref = list(map(lambda x: x.lower(), ref))
+
+ len_hyp = len(hyp)
+ len_ref = len(ref)
+
+ cost_matrix = np.zeros((len_hyp + 1, len_ref + 1), dtype=np.int16)
+
+ ops_matrix = np.zeros((len_hyp + 1, len_ref + 1), dtype=np.int8)
+
+ for i in range(len_hyp + 1):
+ cost_matrix[i][0] = i
+ for j in range(len_ref + 1):
+ cost_matrix[0][j] = j
+
+ for i in range(1, len_hyp + 1):
+ for j in range(1, len_ref + 1):
+ if hyp[i - 1] == ref[j - 1]:
+ cost_matrix[i][j] = cost_matrix[i - 1][j - 1]
+ else:
+ substitution = cost_matrix[i - 1][j - 1] + 1
+ insertion = cost_matrix[i - 1][j] + 1
+ deletion = cost_matrix[i][j - 1] + 1
+
+ compare_val = [substitution, insertion, deletion]
+
+ min_val = min(compare_val)
+ operation_idx = compare_val.index(min_val) + 1
+ cost_matrix[i][j] = min_val
+ ops_matrix[i][j] = operation_idx
+
+ match_idx = []
+ i = len_hyp
+ j = len_ref
+ rst = {
+ 'nwords': len_ref,
+ 'cor': 0,
+ 'wrong': 0,
+ 'ins': 0,
+ 'del': 0,
+ 'sub': 0
+ }
+ while i >= 0 or j >= 0:
+ i_idx = max(0, i)
+ j_idx = max(0, j)
+
+ if ops_matrix[i_idx][j_idx] == 0: # correct
+ if i - 1 >= 0 and j - 1 >= 0:
+ match_idx.append((j - 1, i - 1))
+ rst['cor'] += 1
+
+ i -= 1
+ j -= 1
+
+ elif ops_matrix[i_idx][j_idx] == 2: # insert
+ i -= 1
+ rst['ins'] += 1
+
+ elif ops_matrix[i_idx][j_idx] == 3: # delete
+ j -= 1
+ rst['del'] += 1
+
+ elif ops_matrix[i_idx][j_idx] == 1: # substitute
+ i -= 1
+ j -= 1
+ rst['sub'] += 1
+
+ if i < 0 and j >= 0:
+ rst['del'] += 1
+ elif j < 0 and i >= 0:
+ rst['ins'] += 1
+
+ match_idx.reverse()
+ wrong_cnt = cost_matrix[len_hyp][len_ref]
+ rst['wrong'] = wrong_cnt
+
+ return rst
+
+def print_cer_detail(rst):
+ return ("(" + "nwords=" + str(rst['nwords']) + ",cor=" + str(rst['cor'])
+ + ",ins=" + str(rst['ins']) + ",del=" + str(rst['del']) + ",sub="
+ + str(rst['sub']) + ") corr:" + '{:.2%}'.format(rst['cor']/rst['nwords'])
+ + ",cer:" + '{:.2%}'.format(rst['wrong']/rst['nwords']))
+
+if __name__ == '__main__':
+ if len(sys.argv) != 4:
+ print("usage : python compute-wer.py test.ref test.hyp test.wer")
+ sys.exit(0)
+
+ ref_file = sys.argv[1]
+ hyp_file = sys.argv[2]
+ cer_detail_file = sys.argv[3]
+ compute_wer(ref_file, hyp_file, cer_detail_file)
diff --git a/egs/aishell/conformer/utils/error_rate_zh b/egs/aishell/conformer/utils/error_rate_zh
new file mode 100755
index 0000000..6871a07
--- /dev/null
+++ b/egs/aishell/conformer/utils/error_rate_zh
@@ -0,0 +1,370 @@
+#!/usr/bin/env python3
+# coding=utf8
+
+# Copyright 2021 Jiayu DU
+
+import sys
+import argparse
+import json
+import logging
+logging.basicConfig(stream=sys.stderr, level=logging.INFO, format='[%(levelname)s] %(message)s')
+
+DEBUG = None
+
+def GetEditType(ref_token, hyp_token):
+ if ref_token == None and hyp_token != None:
+ return 'I'
+ elif ref_token != None and hyp_token == None:
+ return 'D'
+ elif ref_token == hyp_token:
+ return 'C'
+ elif ref_token != hyp_token:
+ return 'S'
+ else:
+ raise RuntimeError
+
+class AlignmentArc:
+ def __init__(self, src, dst, ref, hyp):
+ self.src = src
+ self.dst = dst
+ self.ref = ref
+ self.hyp = hyp
+ self.edit_type = GetEditType(ref, hyp)
+
+def similarity_score_function(ref_token, hyp_token):
+ return 0 if (ref_token == hyp_token) else -1.0
+
+def insertion_score_function(token):
+ return -1.0
+
+def deletion_score_function(token):
+ return -1.0
+
+def EditDistance(
+ ref,
+ hyp,
+ similarity_score_function = similarity_score_function,
+ insertion_score_function = insertion_score_function,
+ deletion_score_function = deletion_score_function):
+ assert(len(ref) != 0)
+ class DPState:
+ def __init__(self):
+ self.score = -float('inf')
+ # backpointer
+ self.prev_r = None
+ self.prev_h = None
+
+ def print_search_grid(S, R, H, fstream):
+ print(file=fstream)
+ for r in range(R):
+ for h in range(H):
+ print(F'[{r},{h}]:{S[r][h].score:4.3f}:({S[r][h].prev_r},{S[r][h].prev_h}) ', end='', file=fstream)
+ print(file=fstream)
+
+ R = len(ref) + 1
+ H = len(hyp) + 1
+
+ # Construct DP search space, a (R x H) grid
+ S = [ [] for r in range(R) ]
+ for r in range(R):
+ S[r] = [ DPState() for x in range(H) ]
+
+ # initialize DP search grid origin, S(r = 0, h = 0)
+ S[0][0].score = 0.0
+ S[0][0].prev_r = None
+ S[0][0].prev_h = None
+
+ # initialize REF axis
+ for r in range(1, R):
+ S[r][0].score = S[r-1][0].score + deletion_score_function(ref[r-1])
+ S[r][0].prev_r = r-1
+ S[r][0].prev_h = 0
+
+ # initialize HYP axis
+ for h in range(1, H):
+ S[0][h].score = S[0][h-1].score + insertion_score_function(hyp[h-1])
+ S[0][h].prev_r = 0
+ S[0][h].prev_h = h-1
+
+ best_score = S[0][0].score
+ best_state = (0, 0)
+
+ for r in range(1, R):
+ for h in range(1, H):
+ sub_or_cor_score = similarity_score_function(ref[r-1], hyp[h-1])
+ new_score = S[r-1][h-1].score + sub_or_cor_score
+ if new_score >= S[r][h].score:
+ S[r][h].score = new_score
+ S[r][h].prev_r = r-1
+ S[r][h].prev_h = h-1
+
+ del_score = deletion_score_function(ref[r-1])
+ new_score = S[r-1][h].score + del_score
+ if new_score >= S[r][h].score:
+ S[r][h].score = new_score
+ S[r][h].prev_r = r - 1
+ S[r][h].prev_h = h
+
+ ins_score = insertion_score_function(hyp[h-1])
+ new_score = S[r][h-1].score + ins_score
+ if new_score >= S[r][h].score:
+ S[r][h].score = new_score
+ S[r][h].prev_r = r
+ S[r][h].prev_h = h-1
+
+ best_score = S[R-1][H-1].score
+ best_state = (R-1, H-1)
+
+ if DEBUG:
+ print_search_grid(S, R, H, sys.stderr)
+
+ # Backtracing best alignment path, i.e. a list of arcs
+ # arc = (src, dst, ref, hyp, edit_type)
+ # src/dst = (r, h), where r/h refers to search grid state-id along Ref/Hyp axis
+ best_path = []
+ r, h = best_state[0], best_state[1]
+ prev_r, prev_h = S[r][h].prev_r, S[r][h].prev_h
+ score = S[r][h].score
+ # loop invariant:
+ # 1. (prev_r, prev_h) -> (r, h) is a "forward arc" on best alignment path
+ # 2. score is the value of point(r, h) on DP search grid
+ while prev_r != None or prev_h != None:
+ src = (prev_r, prev_h)
+ dst = (r, h)
+ if (r == prev_r + 1 and h == prev_h + 1): # Substitution or correct
+ arc = AlignmentArc(src, dst, ref[prev_r], hyp[prev_h])
+ elif (r == prev_r + 1 and h == prev_h): # Deletion
+ arc = AlignmentArc(src, dst, ref[prev_r], None)
+ elif (r == prev_r and h == prev_h + 1): # Insertion
+ arc = AlignmentArc(src, dst, None, hyp[prev_h])
+ else:
+ raise RuntimeError
+ best_path.append(arc)
+ r, h = prev_r, prev_h
+ prev_r, prev_h = S[r][h].prev_r, S[r][h].prev_h
+ score = S[r][h].score
+
+ best_path.reverse()
+ return (best_path, best_score)
+
+def PrettyPrintAlignment(alignment, stream = sys.stderr):
+ def get_token_str(token):
+ if token == None:
+ return "*"
+ return token
+
+ def is_double_width_char(ch):
+ if (ch >= '\u4e00') and (ch <= '\u9fa5'): # codepoint ranges for Chinese chars
+ return True
+ # TODO: support other double-width-char language such as Japanese, Korean
+ else:
+ return False
+
+ def display_width(token_str):
+ m = 0
+ for c in token_str:
+ if is_double_width_char(c):
+ m += 2
+ else:
+ m += 1
+ return m
+
+ R = ' REF : '
+ H = ' HYP : '
+ E = ' EDIT : '
+ for arc in alignment:
+ r = get_token_str(arc.ref)
+ h = get_token_str(arc.hyp)
+ e = arc.edit_type if arc.edit_type != 'C' else ''
+
+ nr, nh, ne = display_width(r), display_width(h), display_width(e)
+ n = max(nr, nh, ne) + 1
+
+ R += r + ' ' * (n-nr)
+ H += h + ' ' * (n-nh)
+ E += e + ' ' * (n-ne)
+
+ print(R, file=stream)
+ print(H, file=stream)
+ print(E, file=stream)
+
+def CountEdits(alignment):
+ c, s, i, d = 0, 0, 0, 0
+ for arc in alignment:
+ if arc.edit_type == 'C':
+ c += 1
+ elif arc.edit_type == 'S':
+ s += 1
+ elif arc.edit_type == 'I':
+ i += 1
+ elif arc.edit_type == 'D':
+ d += 1
+ else:
+ raise RuntimeError
+ return (c, s, i, d)
+
+def ComputeTokenErrorRate(c, s, i, d):
+ return 100.0 * (s + d + i) / (s + d + c)
+
+def ComputeSentenceErrorRate(num_err_utts, num_utts):
+ assert(num_utts != 0)
+ return 100.0 * num_err_utts / num_utts
+
+
+class EvaluationResult:
+ def __init__(self):
+ self.num_ref_utts = 0
+ self.num_hyp_utts = 0
+ self.num_eval_utts = 0 # seen in both ref & hyp
+ self.num_hyp_without_ref = 0
+
+ self.C = 0
+ self.S = 0
+ self.I = 0
+ self.D = 0
+ self.token_error_rate = 0.0
+
+ self.num_utts_with_error = 0
+ self.sentence_error_rate = 0.0
+
+ def to_json(self):
+ return json.dumps(self.__dict__)
+
+ def to_kaldi(self):
+ info = (
+ F'%WER {self.token_error_rate:.2f} [ {self.S + self.D + self.I} / {self.C + self.S + self.D}, {self.I} ins, {self.D} del, {self.S} sub ]\n'
+ F'%SER {self.sentence_error_rate:.2f} [ {self.num_utts_with_error} / {self.num_eval_utts} ]\n'
+ )
+ return info
+
+ def to_sclite(self):
+ return "TODO"
+
+ def to_espnet(self):
+ return "TODO"
+
+ def to_summary(self):
+ #return json.dumps(self.__dict__, indent=4)
+ summary = (
+ '==================== Overall Statistics ====================\n'
+ F'num_ref_utts: {self.num_ref_utts}\n'
+ F'num_hyp_utts: {self.num_hyp_utts}\n'
+ F'num_hyp_without_ref: {self.num_hyp_without_ref}\n'
+ F'num_eval_utts: {self.num_eval_utts}\n'
+ F'sentence_error_rate: {self.sentence_error_rate:.2f}%\n'
+ F'token_error_rate: {self.token_error_rate:.2f}%\n'
+ F'token_stats:\n'
+ F' - tokens:{self.C + self.S + self.D:>7}\n'
+ F' - edits: {self.S + self.I + self.D:>7}\n'
+ F' - cor: {self.C:>7}\n'
+ F' - sub: {self.S:>7}\n'
+ F' - ins: {self.I:>7}\n'
+ F' - del: {self.D:>7}\n'
+ '============================================================\n'
+ )
+ return summary
+
+
+class Utterance:
+ def __init__(self, uid, text):
+ self.uid = uid
+ self.text = text
+
+
+def LoadUtterances(filepath, format):
+ utts = {}
+ if format == 'text': # utt_id word1 word2 ...
+ with open(filepath, 'r', encoding='utf8') as f:
+ for line in f:
+ line = line.strip()
+ if line:
+ cols = line.split(maxsplit=1)
+ assert(len(cols) == 2 or len(cols) == 1)
+ uid = cols[0]
+ text = cols[1] if len(cols) == 2 else ''
+ if utts.get(uid) != None:
+ raise RuntimeError(F'Found duplicated utterence id {uid}')
+ utts[uid] = Utterance(uid, text)
+ else:
+ raise RuntimeError(F'Unsupported text format {format}')
+ return utts
+
+
+def tokenize_text(text, tokenizer):
+ if tokenizer == 'whitespace':
+ return text.split()
+ elif tokenizer == 'char':
+ return [ ch for ch in ''.join(text.split()) ]
+ else:
+ raise RuntimeError(F'ERROR: Unsupported tokenizer {tokenizer}')
+
+
+if __name__ == '__main__':
+ parser = argparse.ArgumentParser()
+ # optional
+ parser.add_argument('--tokenizer', choices=['whitespace', 'char'], default='whitespace', help='whitespace for WER, char for CER')
+ parser.add_argument('--ref-format', choices=['text'], default='text', help='reference format, first col is utt_id, the rest is text')
+ parser.add_argument('--hyp-format', choices=['text'], default='text', help='hypothesis format, first col is utt_id, the rest is text')
+ # required
+ parser.add_argument('--ref', type=str, required=True, help='input reference file')
+ parser.add_argument('--hyp', type=str, required=True, help='input hypothesis file')
+
+ parser.add_argument('result_file', type=str)
+ args = parser.parse_args()
+ logging.info(args)
+
+ ref_utts = LoadUtterances(args.ref, args.ref_format)
+ hyp_utts = LoadUtterances(args.hyp, args.hyp_format)
+
+ r = EvaluationResult()
+
+ # check valid utterances in hyp that have matched non-empty reference
+ eval_utts = []
+ r.num_hyp_without_ref = 0
+ for uid in sorted(hyp_utts.keys()):
+ if uid in ref_utts.keys(): # TODO: efficiency
+ if ref_utts[uid].text.strip(): # non-empty reference
+ eval_utts.append(uid)
+ else:
+ logging.warn(F'Found {uid} with empty reference, skipping...')
+ else:
+ logging.warn(F'Found {uid} without reference, skipping...')
+ r.num_hyp_without_ref += 1
+
+ r.num_hyp_utts = len(hyp_utts)
+ r.num_ref_utts = len(ref_utts)
+ r.num_eval_utts = len(eval_utts)
+
+ with open(args.result_file, 'w+', encoding='utf8') as fo:
+ for uid in eval_utts:
+ ref = ref_utts[uid]
+ hyp = hyp_utts[uid]
+
+ alignment, score = EditDistance(
+ tokenize_text(ref.text, args.tokenizer),
+ tokenize_text(hyp.text, args.tokenizer)
+ )
+
+ c, s, i, d = CountEdits(alignment)
+ utt_ter = ComputeTokenErrorRate(c, s, i, d)
+
+ # utt-level evaluation result
+ print(F'{{"uid":{uid}, "score":{score}, "ter":{utt_ter:.2f}, "cor":{c}, "sub":{s}, "ins":{i}, "del":{d}}}', file=fo)
+ PrettyPrintAlignment(alignment, fo)
+
+ r.C += c
+ r.S += s
+ r.I += i
+ r.D += d
+
+ if utt_ter > 0:
+ r.num_utts_with_error += 1
+
+ # corpus level evaluation result
+ r.sentence_error_rate = ComputeSentenceErrorRate(r.num_utts_with_error, r.num_eval_utts)
+ r.token_error_rate = ComputeTokenErrorRate(r.C, r.S, r.I, r.D)
+
+ print(r.to_summary(), file=fo)
+
+ print(r.to_json())
+ print(r.to_kaldi())
diff --git a/egs/aishell/conformer/utils/extract_embeds.py b/egs/aishell/conformer/utils/extract_embeds.py
new file mode 100755
index 0000000..7b817d8
--- /dev/null
+++ b/egs/aishell/conformer/utils/extract_embeds.py
@@ -0,0 +1,47 @@
+from transformers import AutoTokenizer, AutoModel, pipeline
+import numpy as np
+import sys
+import os
+import torch
+from kaldiio import WriteHelper
+import re
+text_file_json = sys.argv[1]
+out_ark = sys.argv[2]
+out_scp = sys.argv[3]
+out_shape = sys.argv[4]
+device = int(sys.argv[5])
+model_path = sys.argv[6]
+
+model = AutoModel.from_pretrained(model_path)
+tokenizer = AutoTokenizer.from_pretrained(model_path)
+extractor = pipeline(task="feature-extraction", model=model, tokenizer=tokenizer, device=device)
+
+with open(text_file_json, 'r') as f:
+ js = f.readlines()
+
+
+f_shape = open(out_shape, "w")
+with WriteHelper('ark,scp:{},{}'.format(out_ark, out_scp)) as writer:
+ with torch.no_grad():
+ for idx, line in enumerate(js):
+ id, tokens = line.strip().split(" ", 1)
+ tokens = re.sub(" ", "", tokens.strip())
+ tokens = ' '.join([j for j in tokens])
+ token_num = len(tokens.split(" "))
+ outputs = extractor(tokens)
+ outputs = np.array(outputs)
+ embeds = outputs[0, 1:-1, :]
+
+ token_num_embeds, dim = embeds.shape
+ if token_num == token_num_embeds:
+ writer(id, embeds)
+ shape_line = "{} {},{}\n".format(id, token_num_embeds, dim)
+ f_shape.write(shape_line)
+ else:
+ print("{}, size has changed, {}, {}, {}".format(id, token_num, token_num_embeds, tokens))
+
+
+
+f_shape.close()
+
+
diff --git a/egs/aishell/conformer/utils/filter_scp.pl b/egs/aishell/conformer/utils/filter_scp.pl
new file mode 100755
index 0000000..003530d
--- /dev/null
+++ b/egs/aishell/conformer/utils/filter_scp.pl
@@ -0,0 +1,87 @@
+#!/usr/bin/env perl
+# Copyright 2010-2012 Microsoft Corporation
+# Johns Hopkins University (author: Daniel Povey)
+
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+# This script takes a list of utterance-ids or any file whose first field
+# of each line is an utterance-id, and filters an scp
+# file (or any file whose "n-th" field is an utterance id), printing
+# out only those lines whose "n-th" field is in id_list. The index of
+# the "n-th" field is 1, by default, but can be changed by using
+# the -f <n> switch
+
+$exclude = 0;
+$field = 1;
+$shifted = 0;
+
+do {
+ $shifted=0;
+ if ($ARGV[0] eq "--exclude") {
+ $exclude = 1;
+ shift @ARGV;
+ $shifted=1;
+ }
+ if ($ARGV[0] eq "-f") {
+ $field = $ARGV[1];
+ shift @ARGV; shift @ARGV;
+ $shifted=1
+ }
+} while ($shifted);
+
+if(@ARGV < 1 || @ARGV > 2) {
+ die "Usage: filter_scp.pl [--exclude] [-f <field-to-filter-on>] id_list [in.scp] > out.scp \n" .
+ "Prints only the input lines whose f'th field (default: first) is in 'id_list'.\n" .
+ "Note: only the first field of each line in id_list matters. With --exclude, prints\n" .
+ "only the lines that were *not* in id_list.\n" .
+ "Caution: previously, the -f option was interpreted as a zero-based field index.\n" .
+ "If your older scripts (written before Oct 2014) stopped working and you used the\n" .
+ "-f option, add 1 to the argument.\n" .
+ "See also: scripts/filter_scp.pl .\n";
+}
+
+
+$idlist = shift @ARGV;
+open(F, "<$idlist") || die "Could not open id-list file $idlist";
+while(<F>) {
+ @A = split;
+ @A>=1 || die "Invalid id-list file line $_";
+ $seen{$A[0]} = 1;
+}
+
+if ($field == 1) { # Treat this as special case, since it is common.
+ while(<>) {
+ $_ =~ m/\s*(\S+)\s*/ || die "Bad line $_, could not get first field.";
+ # $1 is what we filter on.
+ if ((!$exclude && $seen{$1}) || ($exclude && !defined $seen{$1})) {
+ print $_;
+ }
+ }
+} else {
+ while(<>) {
+ @A = split;
+ @A > 0 || die "Invalid scp file line $_";
+ @A >= $field || die "Invalid scp file line $_";
+ if ((!$exclude && $seen{$A[$field-1]}) || ($exclude && !defined $seen{$A[$field-1]})) {
+ print $_;
+ }
+ }
+}
+
+# tests:
+# the following should print "foo 1"
+# ( echo foo 1; echo bar 2 ) | scripts/filter_scp.pl <(echo foo)
+# the following should print "bar 2".
+# ( echo foo 1; echo bar 2 ) | scripts/filter_scp.pl -f 2 <(echo 2)
diff --git a/egs/aishell/conformer/utils/fix_data.sh b/egs/aishell/conformer/utils/fix_data.sh
new file mode 100755
index 0000000..32cdde5
--- /dev/null
+++ b/egs/aishell/conformer/utils/fix_data.sh
@@ -0,0 +1,35 @@
+#!/usr/bin/env bash
+
+echo "$0 $@"
+data_dir=$1
+
+if [ ! -f ${data_dir}/wav.scp ]; then
+ echo "$0: wav.scp is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/text ]; then
+ echo "$0: text is not found"
+ exit 1;
+fi
+
+
+
+mkdir -p ${data_dir}/.backup
+
+awk '{print $1}' ${data_dir}/wav.scp > ${data_dir}/.backup/wav_id
+awk '{print $1}' ${data_dir}/text > ${data_dir}/.backup/text_id
+
+sort ${data_dir}/.backup/wav_id ${data_dir}/.backup/text_id | uniq -d > ${data_dir}/.backup/id
+
+cp ${data_dir}/wav.scp ${data_dir}/.backup/wav.scp
+cp ${data_dir}/text ${data_dir}/.backup/text
+
+mv ${data_dir}/wav.scp ${data_dir}/wav.scp.bak
+mv ${data_dir}/text ${data_dir}/text.bak
+
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/wav.scp.bak > ${data_dir}/wav.scp
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/text.bak > ${data_dir}/text
+
+rm ${data_dir}/wav.scp.bak
+rm ${data_dir}/text.bak
diff --git a/egs/aishell/conformer/utils/fix_data_feat.sh b/egs/aishell/conformer/utils/fix_data_feat.sh
new file mode 100755
index 0000000..2c92d7f
--- /dev/null
+++ b/egs/aishell/conformer/utils/fix_data_feat.sh
@@ -0,0 +1,52 @@
+#!/usr/bin/env bash
+
+echo "$0 $@"
+data_dir=$1
+
+if [ ! -f ${data_dir}/feats.scp ]; then
+ echo "$0: feats.scp is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/text ]; then
+ echo "$0: text is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/speech_shape ]; then
+ echo "$0: feature lengths is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/text_shape ]; then
+ echo "$0: text lengths is not found"
+ exit 1;
+fi
+
+mkdir -p ${data_dir}/.backup
+
+awk '{print $1}' ${data_dir}/feats.scp > ${data_dir}/.backup/wav_id
+awk '{print $1}' ${data_dir}/text > ${data_dir}/.backup/text_id
+
+sort ${data_dir}/.backup/wav_id ${data_dir}/.backup/text_id | uniq -d > ${data_dir}/.backup/id
+
+cp ${data_dir}/feats.scp ${data_dir}/.backup/feats.scp
+cp ${data_dir}/text ${data_dir}/.backup/text
+cp ${data_dir}/speech_shape ${data_dir}/.backup/speech_shape
+cp ${data_dir}/text_shape ${data_dir}/.backup/text_shape
+
+mv ${data_dir}/feats.scp ${data_dir}/feats.scp.bak
+mv ${data_dir}/text ${data_dir}/text.bak
+mv ${data_dir}/speech_shape ${data_dir}/speech_shape.bak
+mv ${data_dir}/text_shape ${data_dir}/text_shape.bak
+
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/feats.scp.bak > ${data_dir}/feats.scp
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/text.bak > ${data_dir}/text
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/speech_shape.bak > ${data_dir}/speech_shape
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/text_shape.bak > ${data_dir}/text_shape
+
+rm ${data_dir}/feats.scp.bak
+rm ${data_dir}/text.bak
+rm ${data_dir}/speech_shape.bak
+rm ${data_dir}/text_shape.bak
+
diff --git a/egs/aishell/conformer/utils/gen_ark_list.sh b/egs/aishell/conformer/utils/gen_ark_list.sh
new file mode 100755
index 0000000..aebf356
--- /dev/null
+++ b/egs/aishell/conformer/utils/gen_ark_list.sh
@@ -0,0 +1,22 @@
+#!/usr/bin/env bash
+
+
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+ark_dir=$1
+txt_dir=$2
+output_dir=$3
+
+[ ! -d ${ark_dir}/ark ] && echo "$0: ark data is required" && exit 1;
+[ ! -d ${txt_dir}/txt ] && echo "$0: txt data is required" && exit 1;
+
+for n in $(seq $nj); do
+ echo "${ark_dir}/ark/feats.$n.ark ${txt_dir}/txt/text.$n.txt" || exit 1
+done > ${output_dir}/ark_txt.scp || exit 1
+
diff --git a/egs/aishell/conformer/utils/parse_options.sh b/egs/aishell/conformer/utils/parse_options.sh
new file mode 100755
index 0000000..71fb9e5
--- /dev/null
+++ b/egs/aishell/conformer/utils/parse_options.sh
@@ -0,0 +1,97 @@
+#!/usr/bin/env bash
+
+# Copyright 2012 Johns Hopkins University (Author: Daniel Povey);
+# Arnab Ghoshal, Karel Vesely
+
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+# Parse command-line options.
+# To be sourced by another script (as in ". parse_options.sh").
+# Option format is: --option-name arg
+# and shell variable "option_name" gets set to value "arg."
+# The exception is --help, which takes no arguments, but prints the
+# $help_message variable (if defined).
+
+
+###
+### The --config file options have lower priority to command line
+### options, so we need to import them first...
+###
+
+# Now import all the configs specified by command-line, in left-to-right order
+for ((argpos=1; argpos<$#; argpos++)); do
+ if [ "${!argpos}" == "--config" ]; then
+ argpos_plus1=$((argpos+1))
+ config=${!argpos_plus1}
+ [ ! -r $config ] && echo "$0: missing config '$config'" && exit 1
+ . $config # source the config file.
+ fi
+done
+
+
+###
+### Now we process the command line options
+###
+while true; do
+ [ -z "${1:-}" ] && break; # break if there are no arguments
+ case "$1" in
+ # If the enclosing script is called with --help option, print the help
+ # message and exit. Scripts should put help messages in $help_message
+ --help|-h) if [ -z "$help_message" ]; then echo "No help found." 1>&2;
+ else printf "$help_message\n" 1>&2 ; fi;
+ exit 0 ;;
+ --*=*) echo "$0: options to scripts must be of the form --name value, got '$1'"
+ exit 1 ;;
+ # If the first command-line argument begins with "--" (e.g. --foo-bar),
+ # then work out the variable name as $name, which will equal "foo_bar".
+ --*) name=`echo "$1" | sed s/^--// | sed s/-/_/g`;
+ # Next we test whether the variable in question is undefned-- if so it's
+ # an invalid option and we die. Note: $0 evaluates to the name of the
+ # enclosing script.
+ # The test [ -z ${foo_bar+xxx} ] will return true if the variable foo_bar
+ # is undefined. We then have to wrap this test inside "eval" because
+ # foo_bar is itself inside a variable ($name).
+ eval '[ -z "${'$name'+xxx}" ]' && echo "$0: invalid option $1" 1>&2 && exit 1;
+
+ oldval="`eval echo \\$$name`";
+ # Work out whether we seem to be expecting a Boolean argument.
+ if [ "$oldval" == "true" ] || [ "$oldval" == "false" ]; then
+ was_bool=true;
+ else
+ was_bool=false;
+ fi
+
+ # Set the variable to the right value-- the escaped quotes make it work if
+ # the option had spaces, like --cmd "queue.pl -sync y"
+ eval $name=\"$2\";
+
+ # Check that Boolean-valued arguments are really Boolean.
+ if $was_bool && [[ "$2" != "true" && "$2" != "false" ]]; then
+ echo "$0: expected \"true\" or \"false\": $1 $2" 1>&2
+ exit 1;
+ fi
+ shift 2;
+ ;;
+ *) break;
+ esac
+done
+
+
+# Check for an empty argument to the --cmd option, which can easily occur as a
+# result of scripting errors.
+[ ! -z "${cmd+xxx}" ] && [ -z "$cmd" ] && echo "$0: empty argument to --cmd option" 1>&2 && exit 1;
+
+
+true; # so this script returns exit code 0.
diff --git a/egs/aishell/conformer/utils/print_args.py b/egs/aishell/conformer/utils/print_args.py
new file mode 100755
index 0000000..b0c61e5
--- /dev/null
+++ b/egs/aishell/conformer/utils/print_args.py
@@ -0,0 +1,45 @@
+#!/usr/bin/env python
+import sys
+
+
+def get_commandline_args(no_executable=True):
+ extra_chars = [
+ " ",
+ ";",
+ "&",
+ "|",
+ "<",
+ ">",
+ "?",
+ "*",
+ "~",
+ "`",
+ '"',
+ "'",
+ "\\",
+ "{",
+ "}",
+ "(",
+ ")",
+ ]
+
+ # Escape the extra characters for shell
+ argv = [
+ arg.replace("'", "'\\''")
+ if all(char not in arg for char in extra_chars)
+ else "'" + arg.replace("'", "'\\''") + "'"
+ for arg in sys.argv
+ ]
+
+ if no_executable:
+ return " ".join(argv[1:])
+ else:
+ return sys.executable + " " + " ".join(argv)
+
+
+def main():
+ print(get_commandline_args())
+
+
+if __name__ == "__main__":
+ main()
diff --git a/egs/aishell/conformer/utils/proc_conf_oss.py b/egs/aishell/conformer/utils/proc_conf_oss.py
new file mode 100755
index 0000000..c4a90c5
--- /dev/null
+++ b/egs/aishell/conformer/utils/proc_conf_oss.py
@@ -0,0 +1,35 @@
+from pathlib import Path
+
+import torch
+import yaml
+
+
+class NoAliasSafeDumper(yaml.SafeDumper):
+ # Disable anchor/alias in yaml because looks ugly
+ def ignore_aliases(self, data):
+ return True
+
+
+def yaml_no_alias_safe_dump(data, stream=None, **kwargs):
+ """Safe-dump in yaml with no anchor/alias"""
+ return yaml.dump(
+ data, stream, allow_unicode=True, Dumper=NoAliasSafeDumper, **kwargs
+ )
+
+
+def gen_conf(file, out_dir):
+ conf = torch.load(file)["config"]
+ conf["oss_bucket"] = "null"
+ print(conf)
+ output_dir = Path(out_dir)
+ output_dir.mkdir(parents=True, exist_ok=True)
+ with (output_dir / "config.yaml").open("w", encoding="utf-8") as f:
+ yaml_no_alias_safe_dump(conf, f, indent=4, sort_keys=False)
+
+
+if __name__ == "__main__":
+ import sys
+
+ in_f = sys.argv[1]
+ out_f = sys.argv[2]
+ gen_conf(in_f, out_f)
diff --git a/egs/aishell/conformer/utils/proce_text.py b/egs/aishell/conformer/utils/proce_text.py
new file mode 100755
index 0000000..9e517a4
--- /dev/null
+++ b/egs/aishell/conformer/utils/proce_text.py
@@ -0,0 +1,31 @@
+
+import sys
+import re
+
+in_f = sys.argv[1]
+out_f = sys.argv[2]
+
+
+with open(in_f, "r", encoding="utf-8") as f:
+ lines = f.readlines()
+
+with open(out_f, "w", encoding="utf-8") as f:
+ for line in lines:
+ outs = line.strip().split(" ", 1)
+ if len(outs) == 2:
+ idx, text = outs
+ text = re.sub("</s>", "", text)
+ text = re.sub("<s>", "", text)
+ text = re.sub("@@", "", text)
+ text = re.sub("@", "", text)
+ text = re.sub("<unk>", "", text)
+ text = re.sub(" ", "", text)
+ text = text.lower()
+ else:
+ idx = outs[0]
+ text = " "
+
+ text = [x for x in text]
+ text = " ".join(text)
+ out = "{} {}\n".format(idx, text)
+ f.write(out)
diff --git a/egs/aishell/conformer/utils/run.pl b/egs/aishell/conformer/utils/run.pl
new file mode 100755
index 0000000..483f95b
--- /dev/null
+++ b/egs/aishell/conformer/utils/run.pl
@@ -0,0 +1,356 @@
+#!/usr/bin/env perl
+use warnings; #sed replacement for -w perl parameter
+# In general, doing
+# run.pl some.log a b c is like running the command a b c in
+# the bash shell, and putting the standard error and output into some.log.
+# To run parallel jobs (backgrounded on the host machine), you can do (e.g.)
+# run.pl JOB=1:4 some.JOB.log a b c JOB is like running the command a b c JOB
+# and putting it in some.JOB.log, for each one. [Note: JOB can be any identifier].
+# If any of the jobs fails, this script will fail.
+
+# A typical example is:
+# run.pl some.log my-prog "--opt=foo bar" foo \| other-prog baz
+# and run.pl will run something like:
+# ( my-prog '--opt=foo bar' foo | other-prog baz ) >& some.log
+#
+# Basically it takes the command-line arguments, quotes them
+# as necessary to preserve spaces, and evaluates them with bash.
+# In addition it puts the command line at the top of the log, and
+# the start and end times of the command at the beginning and end.
+# The reason why this is useful is so that we can create a different
+# version of this program that uses a queueing system instead.
+
+#use Data::Dumper;
+
+@ARGV < 2 && die "usage: run.pl log-file command-line arguments...";
+
+#print STDERR "COMMAND-LINE: " . Dumper(\@ARGV) . "\n";
+$job_pick = 'all';
+$max_jobs_run = -1;
+$jobstart = 1;
+$jobend = 1;
+$ignored_opts = ""; # These will be ignored.
+
+# First parse an option like JOB=1:4, and any
+# options that would normally be given to
+# queue.pl, which we will just discard.
+
+for (my $x = 1; $x <= 2; $x++) { # This for-loop is to
+ # allow the JOB=1:n option to be interleaved with the
+ # options to qsub.
+ while (@ARGV >= 2 && $ARGV[0] =~ m:^-:) {
+ # parse any options that would normally go to qsub, but which will be ignored here.
+ my $switch = shift @ARGV;
+ if ($switch eq "-V") {
+ $ignored_opts .= "-V ";
+ } elsif ($switch eq "--max-jobs-run" || $switch eq "-tc") {
+ # we do support the option --max-jobs-run n, and its GridEngine form -tc n.
+ # if the command appears multiple times uses the smallest option.
+ if ( $max_jobs_run <= 0 ) {
+ $max_jobs_run = shift @ARGV;
+ } else {
+ my $new_constraint = shift @ARGV;
+ if ( ($new_constraint < $max_jobs_run) ) {
+ $max_jobs_run = $new_constraint;
+ }
+ }
+
+ if (! ($max_jobs_run > 0)) {
+ die "run.pl: invalid option --max-jobs-run $max_jobs_run";
+ }
+ } else {
+ my $argument = shift @ARGV;
+ if ($argument =~ m/^--/) {
+ print STDERR "run.pl: WARNING: suspicious argument '$argument' to $switch; starts with '-'\n";
+ }
+ if ($switch eq "-sync" && $argument =~ m/^[yY]/) {
+ $ignored_opts .= "-sync "; # Note: in the
+ # corresponding code in queue.pl it says instead, just "$sync = 1;".
+ } elsif ($switch eq "-pe") { # e.g. -pe smp 5
+ my $argument2 = shift @ARGV;
+ $ignored_opts .= "$switch $argument $argument2 ";
+ } elsif ($switch eq "--gpu") {
+ $using_gpu = $argument;
+ } elsif ($switch eq "--pick") {
+ if($argument =~ m/^(all|failed|incomplete)$/) {
+ $job_pick = $argument;
+ } else {
+ print STDERR "run.pl: ERROR: --pick argument must be one of 'all', 'failed' or 'incomplete'"
+ }
+ } else {
+ # Ignore option.
+ $ignored_opts .= "$switch $argument ";
+ }
+ }
+ }
+ if ($ARGV[0] =~ m/^([\w_][\w\d_]*)+=(\d+):(\d+)$/) { # e.g. JOB=1:20
+ $jobname = $1;
+ $jobstart = $2;
+ $jobend = $3;
+ if ($jobstart > $jobend) {
+ die "run.pl: invalid job range $ARGV[0]";
+ }
+ if ($jobstart <= 0) {
+ die "run.pl: invalid job range $ARGV[0], start must be strictly positive (this is required for GridEngine compatibility).";
+ }
+ shift;
+ } elsif ($ARGV[0] =~ m/^([\w_][\w\d_]*)+=(\d+)$/) { # e.g. JOB=1.
+ $jobname = $1;
+ $jobstart = $2;
+ $jobend = $2;
+ shift;
+ } elsif ($ARGV[0] =~ m/.+\=.*\:.*$/) {
+ print STDERR "run.pl: Warning: suspicious first argument to run.pl: $ARGV[0]\n";
+ }
+}
+
+# Users found this message confusing so we are removing it.
+# if ($ignored_opts ne "") {
+# print STDERR "run.pl: Warning: ignoring options \"$ignored_opts\"\n";
+# }
+
+if ($max_jobs_run == -1) { # If --max-jobs-run option not set,
+ # then work out the number of processors if possible,
+ # and set it based on that.
+ $max_jobs_run = 0;
+ if ($using_gpu) {
+ if (open(P, "nvidia-smi -L |")) {
+ $max_jobs_run++ while (<P>);
+ close(P);
+ }
+ if ($max_jobs_run == 0) {
+ $max_jobs_run = 1;
+ print STDERR "run.pl: Warning: failed to detect number of GPUs from nvidia-smi, using ${max_jobs_run}\n";
+ }
+ } elsif (open(P, "</proc/cpuinfo")) { # Linux
+ while (<P>) { if (m/^processor/) { $max_jobs_run++; } }
+ if ($max_jobs_run == 0) {
+ print STDERR "run.pl: Warning: failed to detect any processors from /proc/cpuinfo\n";
+ $max_jobs_run = 10; # reasonable default.
+ }
+ close(P);
+ } elsif (open(P, "sysctl -a |")) { # BSD/Darwin
+ while (<P>) {
+ if (m/hw\.ncpu\s*[:=]\s*(\d+)/) { # hw.ncpu = 4, or hw.ncpu: 4
+ $max_jobs_run = $1;
+ last;
+ }
+ }
+ close(P);
+ if ($max_jobs_run == 0) {
+ print STDERR "run.pl: Warning: failed to detect any processors from sysctl -a\n";
+ $max_jobs_run = 10; # reasonable default.
+ }
+ } else {
+ # allow at most 32 jobs at once, on non-UNIX systems; change this code
+ # if you need to change this default.
+ $max_jobs_run = 32;
+ }
+ # The just-computed value of $max_jobs_run is just the number of processors
+ # (or our best guess); and if it happens that the number of jobs we need to
+ # run is just slightly above $max_jobs_run, it will make sense to increase
+ # $max_jobs_run to equal the number of jobs, so we don't have a small number
+ # of leftover jobs.
+ $num_jobs = $jobend - $jobstart + 1;
+ if (!$using_gpu &&
+ $num_jobs > $max_jobs_run && $num_jobs < 1.4 * $max_jobs_run) {
+ $max_jobs_run = $num_jobs;
+ }
+}
+
+sub pick_or_exit {
+ # pick_or_exit ( $logfile )
+ # Invoked before each job is started helps to run jobs selectively.
+ #
+ # Given the name of the output logfile decides whether the job must be
+ # executed (by returning from the subroutine) or not (by terminating the
+ # process calling exit)
+ #
+ # PRE: $job_pick is a global variable set by command line switch --pick
+ # and indicates which class of jobs must be executed.
+ #
+ # 1) If a failed job is not executed the process exit code will indicate
+ # failure, just as if the task was just executed and failed.
+ #
+ # 2) If a task is incomplete it will be executed. Incomplete may be either
+ # a job whose log file does not contain the accounting notes in the end,
+ # or a job whose log file does not exist.
+ #
+ # 3) If the $job_pick is set to 'all' (default behavior) a task will be
+ # executed regardless of the result of previous attempts.
+ #
+ # This logic could have been implemented in the main execution loop
+ # but a subroutine to preserve the current level of readability of
+ # that part of the code.
+ #
+ # Alexandre Felipe, (o.alexandre.felipe@gmail.com) 14th of August of 2020
+ #
+ if($job_pick eq 'all'){
+ return; # no need to bother with the previous log
+ }
+ open my $fh, "<", $_[0] or return; # job not executed yet
+ my $log_line;
+ my $cur_line;
+ while ($cur_line = <$fh>) {
+ if( $cur_line =~ m/# Ended \(code .*/ ) {
+ $log_line = $cur_line;
+ }
+ }
+ close $fh;
+ if (! defined($log_line)){
+ return; # incomplete
+ }
+ if ( $log_line =~ m/# Ended \(code 0\).*/ ) {
+ exit(0); # complete
+ } elsif ( $log_line =~ m/# Ended \(code \d+(; signal \d+)?\).*/ ){
+ if ($job_pick !~ m/^(failed|all)$/) {
+ exit(1); # failed but not going to run
+ } else {
+ return; # failed
+ }
+ } elsif ( $log_line =~ m/.*\S.*/ ) {
+ return; # incomplete jobs are always run
+ }
+}
+
+
+$logfile = shift @ARGV;
+
+if (defined $jobname && $logfile !~ m/$jobname/ &&
+ $jobend > $jobstart) {
+ print STDERR "run.pl: you are trying to run a parallel job but "
+ . "you are putting the output into just one log file ($logfile)\n";
+ exit(1);
+}
+
+$cmd = "";
+
+foreach $x (@ARGV) {
+ if ($x =~ m/^\S+$/) { $cmd .= $x . " "; }
+ elsif ($x =~ m:\":) { $cmd .= "'$x' "; }
+ else { $cmd .= "\"$x\" "; }
+}
+
+#$Data::Dumper::Indent=0;
+$ret = 0;
+$numfail = 0;
+%active_pids=();
+
+use POSIX ":sys_wait_h";
+for ($jobid = $jobstart; $jobid <= $jobend; $jobid++) {
+ if (scalar(keys %active_pids) >= $max_jobs_run) {
+
+ # Lets wait for a change in any child's status
+ # Then we have to work out which child finished
+ $r = waitpid(-1, 0);
+ $code = $?;
+ if ($r < 0 ) { die "run.pl: Error waiting for child process"; } # should never happen.
+ if ( defined $active_pids{$r} ) {
+ $jid=$active_pids{$r};
+ $fail[$jid]=$code;
+ if ($code !=0) { $numfail++;}
+ delete $active_pids{$r};
+ # print STDERR "Finished: $r/$jid " . Dumper(\%active_pids) . "\n";
+ } else {
+ die "run.pl: Cannot find the PID of the child process that just finished.";
+ }
+
+ # In theory we could do a non-blocking waitpid over all jobs running just
+ # to find out if only one or more jobs finished during the previous waitpid()
+ # However, we just omit this and will reap the next one in the next pass
+ # through the for(;;) cycle
+ }
+ $childpid = fork();
+ if (!defined $childpid) { die "run.pl: Error forking in run.pl (writing to $logfile)"; }
+ if ($childpid == 0) { # We're in the child... this branch
+ # executes the job and returns (possibly with an error status).
+ if (defined $jobname) {
+ $cmd =~ s/$jobname/$jobid/g;
+ $logfile =~ s/$jobname/$jobid/g;
+ }
+ # exit if the job does not need to be executed
+ pick_or_exit( $logfile );
+
+ system("mkdir -p `dirname $logfile` 2>/dev/null");
+ open(F, ">$logfile") || die "run.pl: Error opening log file $logfile";
+ print F "# " . $cmd . "\n";
+ print F "# Started at " . `date`;
+ $starttime = `date +'%s'`;
+ print F "#\n";
+ close(F);
+
+ # Pipe into bash.. make sure we're not using any other shell.
+ open(B, "|bash") || die "run.pl: Error opening shell command";
+ print B "( " . $cmd . ") 2>>$logfile >> $logfile";
+ close(B); # If there was an error, exit status is in $?
+ $ret = $?;
+
+ $lowbits = $ret & 127;
+ $highbits = $ret >> 8;
+ if ($lowbits != 0) { $return_str = "code $highbits; signal $lowbits" }
+ else { $return_str = "code $highbits"; }
+
+ $endtime = `date +'%s'`;
+ open(F, ">>$logfile") || die "run.pl: Error opening log file $logfile (again)";
+ $enddate = `date`;
+ chop $enddate;
+ print F "# Accounting: time=" . ($endtime - $starttime) . " threads=1\n";
+ print F "# Ended ($return_str) at " . $enddate . ", elapsed time " . ($endtime-$starttime) . " seconds\n";
+ close(F);
+ exit($ret == 0 ? 0 : 1);
+ } else {
+ $pid[$jobid] = $childpid;
+ $active_pids{$childpid} = $jobid;
+ # print STDERR "Queued: " . Dumper(\%active_pids) . "\n";
+ }
+}
+
+# Now we have submitted all the jobs, lets wait until all the jobs finish
+foreach $child (keys %active_pids) {
+ $jobid=$active_pids{$child};
+ $r = waitpid($pid[$jobid], 0);
+ $code = $?;
+ if ($r == -1) { die "run.pl: Error waiting for child process"; } # should never happen.
+ if ($r != 0) { $fail[$jobid]=$code; $numfail++ if $code!=0; } # Completed successfully
+}
+
+# Some sanity checks:
+# The $fail array should not contain undefined codes
+# The number of non-zeros in that array should be equal to $numfail
+# We cannot do foreach() here, as the JOB ids do not start at zero
+$failed_jids=0;
+for ($jobid = $jobstart; $jobid <= $jobend; $jobid++) {
+ $job_return = $fail[$jobid];
+ if (not defined $job_return ) {
+ # print Dumper(\@fail);
+
+ die "run.pl: Sanity check failed: we have indication that some jobs are running " .
+ "even after we waited for all jobs to finish" ;
+ }
+ if ($job_return != 0 ){ $failed_jids++;}
+}
+if ($failed_jids != $numfail) {
+ die "run.pl: Sanity check failed: cannot find out how many jobs failed ($failed_jids x $numfail)."
+}
+if ($numfail > 0) { $ret = 1; }
+
+if ($ret != 0) {
+ $njobs = $jobend - $jobstart + 1;
+ if ($njobs == 1) {
+ if (defined $jobname) {
+ $logfile =~ s/$jobname/$jobstart/; # only one numbered job, so replace name with
+ # that job.
+ }
+ print STDERR "run.pl: job failed, log is in $logfile\n";
+ if ($logfile =~ m/JOB/) {
+ print STDERR "run.pl: probably you forgot to put JOB=1:\$nj in your script.";
+ }
+ }
+ else {
+ $logfile =~ s/$jobname/*/g;
+ print STDERR "run.pl: $numfail / $njobs failed, log is in $logfile\n";
+ }
+}
+
+
+exit ($ret);
diff --git a/egs/aishell/conformer/utils/shuffle_list.pl b/egs/aishell/conformer/utils/shuffle_list.pl
new file mode 100755
index 0000000..a116200
--- /dev/null
+++ b/egs/aishell/conformer/utils/shuffle_list.pl
@@ -0,0 +1,44 @@
+#!/usr/bin/env perl
+
+# Copyright 2013 Johns Hopkins University (author: Daniel Povey)
+
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+if ($ARGV[0] eq "--srand") {
+ $n = $ARGV[1];
+ $n =~ m/\d+/ || die "Bad argument to --srand option: \"$n\"";
+ srand($ARGV[1]);
+ shift;
+ shift;
+} else {
+ srand(0); # Gives inconsistent behavior if we don't seed.
+}
+
+if (@ARGV > 1 || $ARGV[0] =~ m/^-.+/) { # >1 args, or an option we
+ # don't understand.
+ print "Usage: shuffle_list.pl [--srand N] [input file] > output\n";
+ print "randomizes the order of lines of input.\n";
+ exit(1);
+}
+
+@lines;
+while (<>) {
+ push @lines, [ (rand(), $_)] ;
+}
+
+@lines = sort { $a->[0] cmp $b->[0] } @lines;
+foreach $l (@lines) {
+ print $l->[1];
+}
\ No newline at end of file
diff --git a/egs/aishell/conformer/utils/split_data.py b/egs/aishell/conformer/utils/split_data.py
new file mode 100755
index 0000000..060eae6
--- /dev/null
+++ b/egs/aishell/conformer/utils/split_data.py
@@ -0,0 +1,60 @@
+import os
+import sys
+import random
+
+
+in_dir = sys.argv[1]
+out_dir = sys.argv[2]
+num_split = sys.argv[3]
+
+
+def split_scp(scp, num):
+ assert len(scp) >= num
+ avg = len(scp) // num
+ out = []
+ begin = 0
+
+ for i in range(num):
+ if i == num - 1:
+ out.append(scp[begin:])
+ else:
+ out.append(scp[begin:begin+avg])
+ begin += avg
+
+ return out
+
+
+os.path.exists("{}/wav.scp".format(in_dir))
+os.path.exists("{}/text".format(in_dir))
+
+with open("{}/wav.scp".format(in_dir), 'r') as infile:
+ wav_list = infile.readlines()
+
+with open("{}/text".format(in_dir), 'r') as infile:
+ text_list = infile.readlines()
+
+assert len(wav_list) == len(text_list)
+
+x = list(zip(wav_list, text_list))
+random.shuffle(x)
+wav_shuffle_list, text_shuffle_list = zip(*x)
+
+num_split = int(num_split)
+wav_split_list = split_scp(wav_shuffle_list, num_split)
+text_split_list = split_scp(text_shuffle_list, num_split)
+
+for idx, wav_list in enumerate(wav_split_list, 1):
+ path = out_dir + "/split" + str(num_split) + "/" + str(idx)
+ if not os.path.exists(path):
+ os.makedirs(path)
+ with open("{}/wav.scp".format(path), 'w') as wav_writer:
+ for line in wav_list:
+ wav_writer.write(line)
+
+for idx, text_list in enumerate(text_split_list, 1):
+ path = out_dir + "/split" + str(num_split) + "/" + str(idx)
+ if not os.path.exists(path):
+ os.makedirs(path)
+ with open("{}/text".format(path), 'w') as text_writer:
+ for line in text_list:
+ text_writer.write(line)
diff --git a/egs/aishell/conformer/utils/split_scp.pl b/egs/aishell/conformer/utils/split_scp.pl
new file mode 100755
index 0000000..0876dcb
--- /dev/null
+++ b/egs/aishell/conformer/utils/split_scp.pl
@@ -0,0 +1,246 @@
+#!/usr/bin/env perl
+
+# Copyright 2010-2011 Microsoft Corporation
+
+# See ../../COPYING for clarification regarding multiple authors
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+# This program splits up any kind of .scp or archive-type file.
+# If there is no utt2spk option it will work on any text file and
+# will split it up with an approximately equal number of lines in
+# each but.
+# With the --utt2spk option it will work on anything that has the
+# utterance-id as the first entry on each line; the utt2spk file is
+# of the form "utterance speaker" (on each line).
+# It splits it into equal size chunks as far as it can. If you use the utt2spk
+# option it will make sure these chunks coincide with speaker boundaries. In
+# this case, if there are more chunks than speakers (and in some other
+# circumstances), some of the resulting chunks will be empty and it will print
+# an error message and exit with nonzero status.
+# You will normally call this like:
+# split_scp.pl scp scp.1 scp.2 scp.3 ...
+# or
+# split_scp.pl --utt2spk=utt2spk scp scp.1 scp.2 scp.3 ...
+# Note that you can use this script to split the utt2spk file itself,
+# e.g. split_scp.pl --utt2spk=utt2spk utt2spk utt2spk.1 utt2spk.2 ...
+
+# You can also call the scripts like:
+# split_scp.pl -j 3 0 scp scp.0
+# [note: with this option, it assumes zero-based indexing of the split parts,
+# i.e. the second number must be 0 <= n < num-jobs.]
+
+use warnings;
+
+$num_jobs = 0;
+$job_id = 0;
+$utt2spk_file = "";
+$one_based = 0;
+
+for ($x = 1; $x <= 3 && @ARGV > 0; $x++) {
+ if ($ARGV[0] eq "-j") {
+ shift @ARGV;
+ $num_jobs = shift @ARGV;
+ $job_id = shift @ARGV;
+ }
+ if ($ARGV[0] =~ /--utt2spk=(.+)/) {
+ $utt2spk_file=$1;
+ shift;
+ }
+ if ($ARGV[0] eq '--one-based') {
+ $one_based = 1;
+ shift @ARGV;
+ }
+}
+
+if ($num_jobs != 0 && ($num_jobs < 0 || $job_id - $one_based < 0 ||
+ $job_id - $one_based >= $num_jobs)) {
+ die "$0: Invalid job number/index values for '-j $num_jobs $job_id" .
+ ($one_based ? " --one-based" : "") . "'\n"
+}
+
+$one_based
+ and $job_id--;
+
+if(($num_jobs == 0 && @ARGV < 2) || ($num_jobs > 0 && (@ARGV < 1 || @ARGV > 2))) {
+ die
+"Usage: split_scp.pl [--utt2spk=<utt2spk_file>] in.scp out1.scp out2.scp ...
+ or: split_scp.pl -j num-jobs job-id [--one-based] [--utt2spk=<utt2spk_file>] in.scp [out.scp]
+ ... where 0 <= job-id < num-jobs, or 1 <= job-id <- num-jobs if --one-based.\n";
+}
+
+$error = 0;
+$inscp = shift @ARGV;
+if ($num_jobs == 0) { # without -j option
+ @OUTPUTS = @ARGV;
+} else {
+ for ($j = 0; $j < $num_jobs; $j++) {
+ if ($j == $job_id) {
+ if (@ARGV > 0) { push @OUTPUTS, $ARGV[0]; }
+ else { push @OUTPUTS, "-"; }
+ } else {
+ push @OUTPUTS, "/dev/null";
+ }
+ }
+}
+
+if ($utt2spk_file ne "") { # We have the --utt2spk option...
+ open($u_fh, '<', $utt2spk_file) || die "$0: Error opening utt2spk file $utt2spk_file: $!\n";
+ while(<$u_fh>) {
+ @A = split;
+ @A == 2 || die "$0: Bad line $_ in utt2spk file $utt2spk_file\n";
+ ($u,$s) = @A;
+ $utt2spk{$u} = $s;
+ }
+ close $u_fh;
+ open($i_fh, '<', $inscp) || die "$0: Error opening input scp file $inscp: $!\n";
+ @spkrs = ();
+ while(<$i_fh>) {
+ @A = split;
+ if(@A == 0) { die "$0: Empty or space-only line in scp file $inscp\n"; }
+ $u = $A[0];
+ $s = $utt2spk{$u};
+ defined $s || die "$0: No utterance $u in utt2spk file $utt2spk_file\n";
+ if(!defined $spk_count{$s}) {
+ push @spkrs, $s;
+ $spk_count{$s} = 0;
+ $spk_data{$s} = []; # ref to new empty array.
+ }
+ $spk_count{$s}++;
+ push @{$spk_data{$s}}, $_;
+ }
+ # Now split as equally as possible ..
+ # First allocate spks to files by allocating an approximately
+ # equal number of speakers.
+ $numspks = @spkrs; # number of speakers.
+ $numscps = @OUTPUTS; # number of output files.
+ if ($numspks < $numscps) {
+ die "$0: Refusing to split data because number of speakers $numspks " .
+ "is less than the number of output .scp files $numscps\n";
+ }
+ for($scpidx = 0; $scpidx < $numscps; $scpidx++) {
+ $scparray[$scpidx] = []; # [] is array reference.
+ }
+ for ($spkidx = 0; $spkidx < $numspks; $spkidx++) {
+ $scpidx = int(($spkidx*$numscps) / $numspks);
+ $spk = $spkrs[$spkidx];
+ push @{$scparray[$scpidx]}, $spk;
+ $scpcount[$scpidx] += $spk_count{$spk};
+ }
+
+ # Now will try to reassign beginning + ending speakers
+ # to different scp's and see if it gets more balanced.
+ # Suppose objf we're minimizing is sum_i (num utts in scp[i] - average)^2.
+ # We can show that if considering changing just 2 scp's, we minimize
+ # this by minimizing the squared difference in sizes. This is
+ # equivalent to minimizing the absolute difference in sizes. This
+ # shows this method is bound to converge.
+
+ $changed = 1;
+ while($changed) {
+ $changed = 0;
+ for($scpidx = 0; $scpidx < $numscps; $scpidx++) {
+ # First try to reassign ending spk of this scp.
+ if($scpidx < $numscps-1) {
+ $sz = @{$scparray[$scpidx]};
+ if($sz > 0) {
+ $spk = $scparray[$scpidx]->[$sz-1];
+ $count = $spk_count{$spk};
+ $nutt1 = $scpcount[$scpidx];
+ $nutt2 = $scpcount[$scpidx+1];
+ if( abs( ($nutt2+$count) - ($nutt1-$count))
+ < abs($nutt2 - $nutt1)) { # Would decrease
+ # size-diff by reassigning spk...
+ $scpcount[$scpidx+1] += $count;
+ $scpcount[$scpidx] -= $count;
+ pop @{$scparray[$scpidx]};
+ unshift @{$scparray[$scpidx+1]}, $spk;
+ $changed = 1;
+ }
+ }
+ }
+ if($scpidx > 0 && @{$scparray[$scpidx]} > 0) {
+ $spk = $scparray[$scpidx]->[0];
+ $count = $spk_count{$spk};
+ $nutt1 = $scpcount[$scpidx-1];
+ $nutt2 = $scpcount[$scpidx];
+ if( abs( ($nutt2-$count) - ($nutt1+$count))
+ < abs($nutt2 - $nutt1)) { # Would decrease
+ # size-diff by reassigning spk...
+ $scpcount[$scpidx-1] += $count;
+ $scpcount[$scpidx] -= $count;
+ shift @{$scparray[$scpidx]};
+ push @{$scparray[$scpidx-1]}, $spk;
+ $changed = 1;
+ }
+ }
+ }
+ }
+ # Now print out the files...
+ for($scpidx = 0; $scpidx < $numscps; $scpidx++) {
+ $scpfile = $OUTPUTS[$scpidx];
+ ($scpfile ne '-' ? open($f_fh, '>', $scpfile)
+ : open($f_fh, '>&', \*STDOUT)) ||
+ die "$0: Could not open scp file $scpfile for writing: $!\n";
+ $count = 0;
+ if(@{$scparray[$scpidx]} == 0) {
+ print STDERR "$0: eError: split_scp.pl producing empty .scp file " .
+ "$scpfile (too many splits and too few speakers?)\n";
+ $error = 1;
+ } else {
+ foreach $spk ( @{$scparray[$scpidx]} ) {
+ print $f_fh @{$spk_data{$spk}};
+ $count += $spk_count{$spk};
+ }
+ $count == $scpcount[$scpidx] || die "Count mismatch [code error]";
+ }
+ close($f_fh);
+ }
+} else {
+ # This block is the "normal" case where there is no --utt2spk
+ # option and we just break into equal size chunks.
+
+ open($i_fh, '<', $inscp) || die "$0: Error opening input scp file $inscp: $!\n";
+
+ $numscps = @OUTPUTS; # size of array.
+ @F = ();
+ while(<$i_fh>) {
+ push @F, $_;
+ }
+ $numlines = @F;
+ if($numlines == 0) {
+ print STDERR "$0: error: empty input scp file $inscp\n";
+ $error = 1;
+ }
+ $linesperscp = int( $numlines / $numscps); # the "whole part"..
+ $linesperscp >= 1 || die "$0: You are splitting into too many pieces! [reduce \$nj ($numscps) to be smaller than the number of lines ($numlines) in $inscp]\n";
+ $remainder = $numlines - ($linesperscp * $numscps);
+ ($remainder >= 0 && $remainder < $numlines) || die "bad remainder $remainder";
+ # [just doing int() rounds down].
+ $n = 0;
+ for($scpidx = 0; $scpidx < @OUTPUTS; $scpidx++) {
+ $scpfile = $OUTPUTS[$scpidx];
+ ($scpfile ne '-' ? open($o_fh, '>', $scpfile)
+ : open($o_fh, '>&', \*STDOUT)) ||
+ die "$0: Could not open scp file $scpfile for writing: $!\n";
+ for($k = 0; $k < $linesperscp + ($scpidx < $remainder ? 1 : 0); $k++) {
+ print $o_fh $F[$n++];
+ }
+ close($o_fh) || die "$0: Eror closing scp file $scpfile: $!\n";
+ }
+ $n == $numlines || die "$n != $numlines [code error]";
+}
+
+exit ($error);
diff --git a/egs/aishell/conformer/utils/subset_data_dir_tr_cv.sh b/egs/aishell/conformer/utils/subset_data_dir_tr_cv.sh
new file mode 100755
index 0000000..e16cebd
--- /dev/null
+++ b/egs/aishell/conformer/utils/subset_data_dir_tr_cv.sh
@@ -0,0 +1,30 @@
+#!/usr/bin/env bash
+
+dev_num_utt=1000
+
+echo "$0 $@"
+. utils/parse_options.sh || exit 1;
+
+train_data=$1
+out_dir=$2
+
+[ ! -f ${train_data}/wav.scp ] && echo "$0: no such file ${train_data}/wav.scp" && exit 1;
+[ ! -f ${train_data}/text ] && echo "$0: no such file ${train_data}/text" && exit 1;
+
+mkdir -p ${out_dir}/train && mkdir -p ${out_dir}/dev
+
+cp ${train_data}/wav.scp ${out_dir}/train/wav.scp.bak
+cp ${train_data}/text ${out_dir}/train/text.bak
+
+num_utt=$(wc -l <${out_dir}/train/wav.scp.bak)
+
+utils/shuffle_list.pl --srand 1 ${out_dir}/train/wav.scp.bak > ${out_dir}/train/wav.scp.shuf
+head -n ${dev_num_utt} ${out_dir}/train/wav.scp.shuf > ${out_dir}/dev/wav.scp
+tail -n $((${num_utt}-${dev_num_utt})) ${out_dir}/train/wav.scp.shuf > ${out_dir}/train/wav.scp
+
+utils/shuffle_list.pl --srand 1 ${out_dir}/train/text.bak > ${out_dir}/train/text.shuf
+head -n ${dev_num_utt} ${out_dir}/train/text.shuf > ${out_dir}/dev/text
+tail -n $((${num_utt}-${dev_num_utt})) ${out_dir}/train/text.shuf > ${out_dir}/train/text
+
+rm ${out_dir}/train/wav.scp.bak ${out_dir}/train/text.bak
+rm ${out_dir}/train/wav.scp.shuf ${out_dir}/train/text.shuf
diff --git a/egs/aishell/conformer/utils/text2token.py b/egs/aishell/conformer/utils/text2token.py
new file mode 100755
index 0000000..56c3913
--- /dev/null
+++ b/egs/aishell/conformer/utils/text2token.py
@@ -0,0 +1,135 @@
+#!/usr/bin/env python3
+
+# Copyright 2017 Johns Hopkins University (Shinji Watanabe)
+# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
+
+
+import argparse
+import codecs
+import re
+import sys
+
+is_python2 = sys.version_info[0] == 2
+
+
+def exist_or_not(i, match_pos):
+ start_pos = None
+ end_pos = None
+ for pos in match_pos:
+ if pos[0] <= i < pos[1]:
+ start_pos = pos[0]
+ end_pos = pos[1]
+ break
+
+ return start_pos, end_pos
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="convert raw text to tokenized text",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--nchar",
+ "-n",
+ default=1,
+ type=int,
+ help="number of characters to split, i.e., \
+ aabb -> a a b b with -n 1 and aa bb with -n 2",
+ )
+ parser.add_argument(
+ "--skip-ncols", "-s", default=0, type=int, help="skip first n columns"
+ )
+ parser.add_argument("--space", default="<space>", type=str, help="space symbol")
+ parser.add_argument(
+ "--non-lang-syms",
+ "-l",
+ default=None,
+ type=str,
+ help="list of non-linguistic symobles, e.g., <NOISE> etc.",
+ )
+ parser.add_argument("text", type=str, default=False, nargs="?", help="input text")
+ parser.add_argument(
+ "--trans_type",
+ "-t",
+ type=str,
+ default="char",
+ choices=["char", "phn"],
+ help="""Transcript type. char/phn. e.g., for TIMIT FADG0_SI1279 -
+ If trans_type is char,
+ read from SI1279.WRD file -> "bricks are an alternative"
+ Else if trans_type is phn,
+ read from SI1279.PHN file -> "sil b r ih sil k s aa r er n aa l
+ sil t er n ih sil t ih v sil" """,
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ rs = []
+ if args.non_lang_syms is not None:
+ with codecs.open(args.non_lang_syms, "r", encoding="utf-8") as f:
+ nls = [x.rstrip() for x in f.readlines()]
+ rs = [re.compile(re.escape(x)) for x in nls]
+
+ if args.text:
+ f = codecs.open(args.text, encoding="utf-8")
+ else:
+ f = codecs.getreader("utf-8")(sys.stdin if is_python2 else sys.stdin.buffer)
+
+ sys.stdout = codecs.getwriter("utf-8")(
+ sys.stdout if is_python2 else sys.stdout.buffer
+ )
+ line = f.readline()
+ n = args.nchar
+ while line:
+ x = line.split()
+ print(" ".join(x[: args.skip_ncols]), end=" ")
+ a = " ".join(x[args.skip_ncols :])
+
+ # get all matched positions
+ match_pos = []
+ for r in rs:
+ i = 0
+ while i >= 0:
+ m = r.search(a, i)
+ if m:
+ match_pos.append([m.start(), m.end()])
+ i = m.end()
+ else:
+ break
+
+ if args.trans_type == "phn":
+ a = a.split(" ")
+ else:
+ if len(match_pos) > 0:
+ chars = []
+ i = 0
+ while i < len(a):
+ start_pos, end_pos = exist_or_not(i, match_pos)
+ if start_pos is not None:
+ chars.append(a[start_pos:end_pos])
+ i = end_pos
+ else:
+ chars.append(a[i])
+ i += 1
+ a = chars
+
+ a = [a[j : j + n] for j in range(0, len(a), n)]
+
+ a_flat = []
+ for z in a:
+ a_flat.append("".join(z))
+
+ a_chars = [z.replace(" ", args.space) for z in a_flat]
+ if args.trans_type == "phn":
+ a_chars = [z.replace("sil", args.space) for z in a_chars]
+ print(" ".join(a_chars))
+ line = f.readline()
+
+
+if __name__ == "__main__":
+ main()
diff --git a/egs/aishell/conformer/utils/text_tokenize.py b/egs/aishell/conformer/utils/text_tokenize.py
new file mode 100755
index 0000000..962ea11
--- /dev/null
+++ b/egs/aishell/conformer/utils/text_tokenize.py
@@ -0,0 +1,106 @@
+import re
+import argparse
+
+
+def load_dict(seg_file):
+ seg_dict = {}
+ with open(seg_file, 'r') as infile:
+ for line in infile:
+ s = line.strip().split()
+ key = s[0]
+ value = s[1:]
+ seg_dict[key] = " ".join(value)
+ return seg_dict
+
+
+def forward_segment(text, dic):
+ word_list = []
+ i = 0
+ while i < len(text):
+ longest_word = text[i]
+ for j in range(i + 1, len(text) + 1):
+ word = text[i:j]
+ if word in dic:
+ if len(word) > len(longest_word):
+ longest_word = word
+ word_list.append(longest_word)
+ i += len(longest_word)
+ return word_list
+
+
+def tokenize(txt,
+ seg_dict):
+ out_txt = ""
+ pattern = re.compile(r"([\u4E00-\u9FA5A-Za-z0-9])")
+ for word in txt:
+ if pattern.match(word):
+ if word in seg_dict:
+ out_txt += seg_dict[word] + " "
+ else:
+ out_txt += "<unk>" + " "
+ else:
+ continue
+ return out_txt.strip()
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="text tokenize",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--text-file",
+ "-t",
+ default=False,
+ required=True,
+ type=str,
+ help="input text",
+ )
+ parser.add_argument(
+ "--seg-file",
+ "-s",
+ default=False,
+ required=True,
+ type=str,
+ help="seg file",
+ )
+ parser.add_argument(
+ "--txt-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="txt index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ txt_writer = open("{}/text.{}.txt".format(args.output_dir, args.txt_index), 'w')
+ shape_writer = open("{}/len.{}".format(args.output_dir, args.txt_index), 'w')
+ seg_dict = load_dict(args.seg_file)
+ with open(args.text_file, 'r') as infile:
+ for line in infile:
+ s = line.strip().split()
+ text_id = s[0]
+ text_list = forward_segment("".join(s[1:]).lower(), seg_dict)
+ text = tokenize(text_list, seg_dict)
+ lens = len(text.strip().split())
+ txt_writer.write(text_id + " " + text + '\n')
+ shape_writer.write(text_id + " " + str(lens) + '\n')
+
+
+if __name__ == '__main__':
+ main()
+
diff --git a/egs/aishell/conformer/utils/text_tokenize.sh b/egs/aishell/conformer/utils/text_tokenize.sh
new file mode 100755
index 0000000..6b74fef
--- /dev/null
+++ b/egs/aishell/conformer/utils/text_tokenize.sh
@@ -0,0 +1,35 @@
+#!/usr/bin/env bash
+
+
+# Begin configuration section.
+nj=32
+cmd=utils/run.pl
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+# tokenize configuration
+text_dir=$1
+seg_file=$2
+logdir=$3
+output_dir=$4
+
+txt_dir=${output_dir}/txt; mkdir -p ${output_dir}/txt
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/text_tokenize.JOB.log \
+ python utils/text_tokenize.py -t ${text_dir}/txt/text.JOB.txt \
+ -s ${seg_file} -i JOB -o ${txt_dir} \
+ || exit 1;
+
+# concatenate the text files together.
+for n in $(seq $nj); do
+ cat ${txt_dir}/text.$n.txt || exit 1
+done > ${output_dir}/text || exit 1
+
+for n in $(seq $nj); do
+ cat ${txt_dir}/len.$n || exit 1
+done > ${output_dir}/text_shape || exit 1
+
+echo "$0: Succeeded text tokenize"
diff --git a/egs/aishell/conformer/utils/textnorm_zh.py b/egs/aishell/conformer/utils/textnorm_zh.py
new file mode 100755
index 0000000..79feb83
--- /dev/null
+++ b/egs/aishell/conformer/utils/textnorm_zh.py
@@ -0,0 +1,834 @@
+#!/usr/bin/env python3
+# coding=utf-8
+
+# Authors:
+# 2019.5 Zhiyang Zhou (https://github.com/Joee1995/chn_text_norm.git)
+# 2019.9 Jiayu DU
+#
+# requirements:
+# - python 3.X
+# notes: python 2.X WILL fail or produce misleading results
+
+import sys, os, argparse, codecs, string, re
+
+# ================================================================================ #
+# basic constant
+# ================================================================================ #
+CHINESE_DIGIS = u'闆朵竴浜屼笁鍥涗簲鍏竷鍏節'
+BIG_CHINESE_DIGIS_SIMPLIFIED = u'闆跺9璐板弫鑲嗕紞闄嗘煉鎹岀帠'
+BIG_CHINESE_DIGIS_TRADITIONAL = u'闆跺9璨冲弮鑲嗕紞闄告煉鎹岀帠'
+SMALLER_BIG_CHINESE_UNITS_SIMPLIFIED = u'鍗佺櫨鍗冧竾'
+SMALLER_BIG_CHINESE_UNITS_TRADITIONAL = u'鎷句桨浠熻惉'
+LARGER_CHINESE_NUMERING_UNITS_SIMPLIFIED = u'浜垮厗浜灀绉┌娌熸锭姝h浇'
+LARGER_CHINESE_NUMERING_UNITS_TRADITIONAL = u'鍎勫厗浜灀绉┌婧濇緱姝h級'
+SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED = u'鍗佺櫨鍗冧竾'
+SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL = u'鎷句桨浠熻惉'
+
+ZERO_ALT = u'銆�'
+ONE_ALT = u'骞�'
+TWO_ALTS = [u'涓�', u'鍏�']
+
+POSITIVE = [u'姝�', u'姝�']
+NEGATIVE = [u'璐�', u'璨�']
+POINT = [u'鐐�', u'榛�']
+# PLUS = [u'鍔�', u'鍔�']
+# SIL = [u'鏉�', u'妲�']
+
+FILLER_CHARS = ['鍛�', '鍟�']
+ER_WHITELIST = '(鍎垮コ|鍎垮瓙|鍎垮瓩|濂冲効|鍎垮|濡诲効|' \
+ '鑳庡効|濠村効|鏂扮敓鍎縷濠村辜鍎縷骞煎効|灏戝効|灏忓効|鍎挎瓕|鍎跨|鍎跨|鎵樺効鎵�|瀛ゅ効|' \
+ '鍎挎垙|鍎垮寲|鍙板効搴剕楣垮効宀泑姝e効鍏粡|鍚婂効閮庡綋|鐢熷効鑲插コ|鎵樺効甯﹀コ|鍏诲効闃茶�亅鐥村効鍛嗗コ|' \
+ '浣冲効浣冲|鍎挎�滃吔鎵皘鍎挎棤甯哥埗|鍎夸笉瀚屾瘝涓憒鍎胯鍗冮噷姣嶆媴蹇鍎垮ぇ涓嶇敱鐖穦鑻忎篂鍎�)'
+
+# 涓枃鏁板瓧绯荤粺绫诲瀷
+NUMBERING_TYPES = ['low', 'mid', 'high']
+
+CURRENCY_NAMES = '(浜烘皯甯亅缇庡厓|鏃ュ厓|鑻遍晳|娆у厓|椹厠|娉曢儙|鍔犳嬁澶у厓|婢冲厓|娓竵|鍏堜护|鑺叞椹厠|鐖卞皵鍏伴晳|' \
+ '閲屾媺|鑽峰叞鐩緗鍩冩柉搴撳|姣斿濉攟鍗板凹鐩緗鏋楀悏鐗箌鏂拌タ鍏板厓|姣旂储|鍗㈠竷|鏂板姞鍧″厓|闊╁厓|娉伴摙)'
+CURRENCY_UNITS = '((浜縷鍗冧竾|鐧句竾|涓噟鍗億鐧�)|(浜縷鍗冧竾|鐧句竾|涓噟鍗億鐧緗)鍏億(浜縷鍗冧竾|鐧句竾|涓噟鍗億鐧緗)鍧梶瑙抾姣泑鍒�)'
+COM_QUANTIFIERS = '(鍖箌寮爘搴鍥瀨鍦簗灏緗鏉涓獆棣東闃檤闃祙缃憒鐐畖椤秥涓榺妫祙鍙獆鏀瘄琚瓅杈唡鎸憒鎷厊棰梶澹硘绐爘鏇瞸澧檤缇鑵攟' \
+ '鐮搴瀹璐瘄鎵巪鎹唡鍒�|浠鎵搢鎵媩缃梶鍧灞眧宀瓅姹焲婧獆閽焲闃焲鍗晐鍙寍瀵箌鍑簗鍙澶磡鑴殀鏉縷璺硘鏋潀浠秥璐磡' \
+ '閽坾绾縷绠鍚峾浣峾韬珅鍫倈璇緗鏈瑋椤祙瀹秥鎴穦灞倈涓潀姣珅鍘榺鍒唡閽眧涓鏂鎷厊閾鐭硘閽閿眧蹇絴(鍗億姣珅寰�)鍏媩' \
+ '姣珅鍘榺鍒唡瀵竱灏簗涓坾閲寍瀵粅甯竱閾簗绋媩(鍗億鍒唡鍘榺姣珅寰�)绫硘鎾畖鍕簗鍚坾鍗噟鏂梶鐭硘鐩榺纰梶纰焲鍙爘妗秥绗紎鐩唡' \
+ '鐩抾鏉瘄閽焲鏂泑閿厊绨媩绡畖鐩榺妗秥缃恷鐡秥澹秥鍗畖鐩弢绠﹟绠眧鐓瞸鍟東琚媩閽祙骞磡鏈坾鏃瀛鍒粅鏃秥鍛▅澶﹟绉抾鍒唡鏃瑋' \
+ '绾獆宀亅涓東鏇磡澶渱鏄澶弢绉媩鍐瑋浠浼弢杈坾涓竱娉绮抾棰梶骞鍫唡鏉鏍箌鏀瘄閬搢闈鐗噟寮爘棰梶鍧�)'
+
+# punctuation information are based on Zhon project (https://github.com/tsroten/zhon.git)
+CHINESE_PUNC_STOP = '锛侊紵锝°��'
+CHINESE_PUNC_NON_STOP = '锛傦純锛勶紖锛嗭紘锛堬級锛婏紜锛岋紞锛忥細锛涳紲锛濓紴锛狅蓟锛硷冀锛撅伎锝�锝涳綔锝濓綖锝燂綘锝剑锝ゃ�併�冦�嬨�屻�嶃�庛�忋�愩�戙�斻�曘�栥�椼�樸�欍�氥�涖�溿�濄�炪�熴�般�俱�库�撯�斺�樷�欌�涒�溾�濃�炩�熲�︹�э箯'
+CHINESE_PUNC_LIST = CHINESE_PUNC_STOP + CHINESE_PUNC_NON_STOP
+
+# ================================================================================ #
+# basic class
+# ================================================================================ #
+class ChineseChar(object):
+ """
+ 涓枃瀛楃
+ 姣忎釜瀛楃瀵瑰簲绠�浣撳拰绻佷綋,
+ e.g. 绠�浣� = '璐�', 绻佷綋 = '璨�'
+ 杞崲鏃跺彲杞崲涓虹畝浣撴垨绻佷綋
+ """
+
+ def __init__(self, simplified, traditional):
+ self.simplified = simplified
+ self.traditional = traditional
+ #self.__repr__ = self.__str__
+
+ def __str__(self):
+ return self.simplified or self.traditional or None
+
+ def __repr__(self):
+ return self.__str__()
+
+
+class ChineseNumberUnit(ChineseChar):
+ """
+ 涓枃鏁板瓧/鏁颁綅瀛楃
+ 姣忎釜瀛楃闄ょ箒绠�浣撳杩樻湁涓�涓澶栫殑澶у啓瀛楃
+ e.g. '闄�' 鍜� '闄�'
+ """
+
+ def __init__(self, power, simplified, traditional, big_s, big_t):
+ super(ChineseNumberUnit, self).__init__(simplified, traditional)
+ self.power = power
+ self.big_s = big_s
+ self.big_t = big_t
+
+ def __str__(self):
+ return '10^{}'.format(self.power)
+
+ @classmethod
+ def create(cls, index, value, numbering_type=NUMBERING_TYPES[1], small_unit=False):
+
+ if small_unit:
+ return ChineseNumberUnit(power=index + 1,
+ simplified=value[0], traditional=value[1], big_s=value[1], big_t=value[1])
+ elif numbering_type == NUMBERING_TYPES[0]:
+ return ChineseNumberUnit(power=index + 8,
+ simplified=value[0], traditional=value[1], big_s=value[0], big_t=value[1])
+ elif numbering_type == NUMBERING_TYPES[1]:
+ return ChineseNumberUnit(power=(index + 2) * 4,
+ simplified=value[0], traditional=value[1], big_s=value[0], big_t=value[1])
+ elif numbering_type == NUMBERING_TYPES[2]:
+ return ChineseNumberUnit(power=pow(2, index + 3),
+ simplified=value[0], traditional=value[1], big_s=value[0], big_t=value[1])
+ else:
+ raise ValueError(
+ 'Counting type should be in {0} ({1} provided).'.format(NUMBERING_TYPES, numbering_type))
+
+
+class ChineseNumberDigit(ChineseChar):
+ """
+ 涓枃鏁板瓧瀛楃
+ """
+
+ def __init__(self, value, simplified, traditional, big_s, big_t, alt_s=None, alt_t=None):
+ super(ChineseNumberDigit, self).__init__(simplified, traditional)
+ self.value = value
+ self.big_s = big_s
+ self.big_t = big_t
+ self.alt_s = alt_s
+ self.alt_t = alt_t
+
+ def __str__(self):
+ return str(self.value)
+
+ @classmethod
+ def create(cls, i, v):
+ return ChineseNumberDigit(i, v[0], v[1], v[2], v[3])
+
+
+class ChineseMath(ChineseChar):
+ """
+ 涓枃鏁颁綅瀛楃
+ """
+
+ def __init__(self, simplified, traditional, symbol, expression=None):
+ super(ChineseMath, self).__init__(simplified, traditional)
+ self.symbol = symbol
+ self.expression = expression
+ self.big_s = simplified
+ self.big_t = traditional
+
+
+CC, CNU, CND, CM = ChineseChar, ChineseNumberUnit, ChineseNumberDigit, ChineseMath
+
+
+class NumberSystem(object):
+ """
+ 涓枃鏁板瓧绯荤粺
+ """
+ pass
+
+
+class MathSymbol(object):
+ """
+ 鐢ㄤ簬涓枃鏁板瓧绯荤粺鐨勬暟瀛︾鍙� (绻�/绠�浣�), e.g.
+ positive = ['姝�', '姝�']
+ negative = ['璐�', '璨�']
+ point = ['鐐�', '榛�']
+ """
+
+ def __init__(self, positive, negative, point):
+ self.positive = positive
+ self.negative = negative
+ self.point = point
+
+ def __iter__(self):
+ for v in self.__dict__.values():
+ yield v
+
+
+# class OtherSymbol(object):
+# """
+# 鍏朵粬绗﹀彿
+# """
+#
+# def __init__(self, sil):
+# self.sil = sil
+#
+# def __iter__(self):
+# for v in self.__dict__.values():
+# yield v
+
+
+# ================================================================================ #
+# basic utils
+# ================================================================================ #
+def create_system(numbering_type=NUMBERING_TYPES[1]):
+ """
+ 鏍规嵁鏁板瓧绯荤粺绫诲瀷杩斿洖鍒涘缓鐩稿簲鐨勬暟瀛楃郴缁燂紝榛樿涓� mid
+ NUMBERING_TYPES = ['low', 'mid', 'high']: 涓枃鏁板瓧绯荤粺绫诲瀷
+ low: '鍏�' = '浜�' * '鍗�' = $10^{9}$, '浜�' = '鍏�' * '鍗�', etc.
+ mid: '鍏�' = '浜�' * '涓�' = $10^{12}$, '浜�' = '鍏�' * '涓�', etc.
+ high: '鍏�' = '浜�' * '浜�' = $10^{16}$, '浜�' = '鍏�' * '鍏�', etc.
+ 杩斿洖瀵瑰簲鐨勬暟瀛楃郴缁�
+ """
+
+ # chinese number units of '浜�' and larger
+ all_larger_units = zip(
+ LARGER_CHINESE_NUMERING_UNITS_SIMPLIFIED, LARGER_CHINESE_NUMERING_UNITS_TRADITIONAL)
+ larger_units = [CNU.create(i, v, numbering_type, False)
+ for i, v in enumerate(all_larger_units)]
+ # chinese number units of '鍗�, 鐧�, 鍗�, 涓�'
+ all_smaller_units = zip(
+ SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED, SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL)
+ smaller_units = [CNU.create(i, v, small_unit=True)
+ for i, v in enumerate(all_smaller_units)]
+ # digis
+ chinese_digis = zip(CHINESE_DIGIS, CHINESE_DIGIS,
+ BIG_CHINESE_DIGIS_SIMPLIFIED, BIG_CHINESE_DIGIS_TRADITIONAL)
+ digits = [CND.create(i, v) for i, v in enumerate(chinese_digis)]
+ digits[0].alt_s, digits[0].alt_t = ZERO_ALT, ZERO_ALT
+ digits[1].alt_s, digits[1].alt_t = ONE_ALT, ONE_ALT
+ digits[2].alt_s, digits[2].alt_t = TWO_ALTS[0], TWO_ALTS[1]
+
+ # symbols
+ positive_cn = CM(POSITIVE[0], POSITIVE[1], '+', lambda x: x)
+ negative_cn = CM(NEGATIVE[0], NEGATIVE[1], '-', lambda x: -x)
+ point_cn = CM(POINT[0], POINT[1], '.', lambda x,
+ y: float(str(x) + '.' + str(y)))
+ # sil_cn = CM(SIL[0], SIL[1], '-', lambda x, y: float(str(x) + '-' + str(y)))
+ system = NumberSystem()
+ system.units = smaller_units + larger_units
+ system.digits = digits
+ system.math = MathSymbol(positive_cn, negative_cn, point_cn)
+ # system.symbols = OtherSymbol(sil_cn)
+ return system
+
+
+def chn2num(chinese_string, numbering_type=NUMBERING_TYPES[1]):
+
+ def get_symbol(char, system):
+ for u in system.units:
+ if char in [u.traditional, u.simplified, u.big_s, u.big_t]:
+ return u
+ for d in system.digits:
+ if char in [d.traditional, d.simplified, d.big_s, d.big_t, d.alt_s, d.alt_t]:
+ return d
+ for m in system.math:
+ if char in [m.traditional, m.simplified]:
+ return m
+
+ def string2symbols(chinese_string, system):
+ int_string, dec_string = chinese_string, ''
+ for p in [system.math.point.simplified, system.math.point.traditional]:
+ if p in chinese_string:
+ int_string, dec_string = chinese_string.split(p)
+ break
+ return [get_symbol(c, system) for c in int_string], \
+ [get_symbol(c, system) for c in dec_string]
+
+ def correct_symbols(integer_symbols, system):
+ """
+ 涓�鐧惧叓 to 涓�鐧惧叓鍗�
+ 涓�浜夸竴鍗冧笁鐧句竾 to 涓�浜� 涓�鍗冧竾 涓夌櫨涓�
+ """
+
+ if integer_symbols and isinstance(integer_symbols[0], CNU):
+ if integer_symbols[0].power == 1:
+ integer_symbols = [system.digits[1]] + integer_symbols
+
+ if len(integer_symbols) > 1:
+ if isinstance(integer_symbols[-1], CND) and isinstance(integer_symbols[-2], CNU):
+ integer_symbols.append(
+ CNU(integer_symbols[-2].power - 1, None, None, None, None))
+
+ result = []
+ unit_count = 0
+ for s in integer_symbols:
+ if isinstance(s, CND):
+ result.append(s)
+ unit_count = 0
+ elif isinstance(s, CNU):
+ current_unit = CNU(s.power, None, None, None, None)
+ unit_count += 1
+
+ if unit_count == 1:
+ result.append(current_unit)
+ elif unit_count > 1:
+ for i in range(len(result)):
+ if isinstance(result[-i - 1], CNU) and result[-i - 1].power < current_unit.power:
+ result[-i - 1] = CNU(result[-i - 1].power +
+ current_unit.power, None, None, None, None)
+ return result
+
+ def compute_value(integer_symbols):
+ """
+ Compute the value.
+ When current unit is larger than previous unit, current unit * all previous units will be used as all previous units.
+ e.g. '涓ゅ崈涓�' = 2000 * 10000 not 2000 + 10000
+ """
+ value = [0]
+ last_power = 0
+ for s in integer_symbols:
+ if isinstance(s, CND):
+ value[-1] = s.value
+ elif isinstance(s, CNU):
+ value[-1] *= pow(10, s.power)
+ if s.power > last_power:
+ value[:-1] = list(map(lambda v: v *
+ pow(10, s.power), value[:-1]))
+ last_power = s.power
+ value.append(0)
+ return sum(value)
+
+ system = create_system(numbering_type)
+ int_part, dec_part = string2symbols(chinese_string, system)
+ int_part = correct_symbols(int_part, system)
+ int_str = str(compute_value(int_part))
+ dec_str = ''.join([str(d.value) for d in dec_part])
+ if dec_part:
+ return '{0}.{1}'.format(int_str, dec_str)
+ else:
+ return int_str
+
+
+def num2chn(number_string, numbering_type=NUMBERING_TYPES[1], big=False,
+ traditional=False, alt_zero=False, alt_one=False, alt_two=True,
+ use_zeros=True, use_units=True):
+
+ def get_value(value_string, use_zeros=True):
+
+ striped_string = value_string.lstrip('0')
+
+ # record nothing if all zeros
+ if not striped_string:
+ return []
+
+ # record one digits
+ elif len(striped_string) == 1:
+ if use_zeros and len(value_string) != len(striped_string):
+ return [system.digits[0], system.digits[int(striped_string)]]
+ else:
+ return [system.digits[int(striped_string)]]
+
+ # recursively record multiple digits
+ else:
+ result_unit = next(u for u in reversed(
+ system.units) if u.power < len(striped_string))
+ result_string = value_string[:-result_unit.power]
+ return get_value(result_string) + [result_unit] + get_value(striped_string[-result_unit.power:])
+
+ system = create_system(numbering_type)
+
+ int_dec = number_string.split('.')
+ if len(int_dec) == 1:
+ int_string = int_dec[0]
+ dec_string = ""
+ elif len(int_dec) == 2:
+ int_string = int_dec[0]
+ dec_string = int_dec[1]
+ else:
+ raise ValueError(
+ "invalid input num string with more than one dot: {}".format(number_string))
+
+ if use_units and len(int_string) > 1:
+ result_symbols = get_value(int_string)
+ else:
+ result_symbols = [system.digits[int(c)] for c in int_string]
+ dec_symbols = [system.digits[int(c)] for c in dec_string]
+ if dec_string:
+ result_symbols += [system.math.point] + dec_symbols
+
+ if alt_two:
+ liang = CND(2, system.digits[2].alt_s, system.digits[2].alt_t,
+ system.digits[2].big_s, system.digits[2].big_t)
+ for i, v in enumerate(result_symbols):
+ if isinstance(v, CND) and v.value == 2:
+ next_symbol = result_symbols[i +
+ 1] if i < len(result_symbols) - 1 else None
+ previous_symbol = result_symbols[i - 1] if i > 0 else None
+ if isinstance(next_symbol, CNU) and isinstance(previous_symbol, (CNU, type(None))):
+ if next_symbol.power != 1 and ((previous_symbol is None) or (previous_symbol.power != 1)):
+ result_symbols[i] = liang
+
+ # if big is True, '涓�' will not be used and `alt_two` has no impact on output
+ if big:
+ attr_name = 'big_'
+ if traditional:
+ attr_name += 't'
+ else:
+ attr_name += 's'
+ else:
+ if traditional:
+ attr_name = 'traditional'
+ else:
+ attr_name = 'simplified'
+
+ result = ''.join([getattr(s, attr_name) for s in result_symbols])
+
+ # if not use_zeros:
+ # result = result.strip(getattr(system.digits[0], attr_name))
+
+ if alt_zero:
+ result = result.replace(
+ getattr(system.digits[0], attr_name), system.digits[0].alt_s)
+
+ if alt_one:
+ result = result.replace(
+ getattr(system.digits[1], attr_name), system.digits[1].alt_s)
+
+ for i, p in enumerate(POINT):
+ if result.startswith(p):
+ return CHINESE_DIGIS[0] + result
+
+ # ^10, 11, .., 19
+ if len(result) >= 2 and result[1] in [SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED[0],
+ SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL[0]] and \
+ result[0] in [CHINESE_DIGIS[1], BIG_CHINESE_DIGIS_SIMPLIFIED[1], BIG_CHINESE_DIGIS_TRADITIONAL[1]]:
+ result = result[1:]
+
+ return result
+
+
+# ================================================================================ #
+# different types of rewriters
+# ================================================================================ #
+class Cardinal:
+ """
+ CARDINAL绫�
+ """
+
+ def __init__(self, cardinal=None, chntext=None):
+ self.cardinal = cardinal
+ self.chntext = chntext
+
+ def chntext2cardinal(self):
+ return chn2num(self.chntext)
+
+ def cardinal2chntext(self):
+ return num2chn(self.cardinal)
+
+class Digit:
+ """
+ DIGIT绫�
+ """
+
+ def __init__(self, digit=None, chntext=None):
+ self.digit = digit
+ self.chntext = chntext
+
+ # def chntext2digit(self):
+ # return chn2num(self.chntext)
+
+ def digit2chntext(self):
+ return num2chn(self.digit, alt_two=False, use_units=False)
+
+
+class TelePhone:
+ """
+ TELEPHONE绫�
+ """
+
+ def __init__(self, telephone=None, raw_chntext=None, chntext=None):
+ self.telephone = telephone
+ self.raw_chntext = raw_chntext
+ self.chntext = chntext
+
+ # def chntext2telephone(self):
+ # sil_parts = self.raw_chntext.split('<SIL>')
+ # self.telephone = '-'.join([
+ # str(chn2num(p)) for p in sil_parts
+ # ])
+ # return self.telephone
+
+ def telephone2chntext(self, fixed=False):
+
+ if fixed:
+ sil_parts = self.telephone.split('-')
+ self.raw_chntext = '<SIL>'.join([
+ num2chn(part, alt_two=False, use_units=False) for part in sil_parts
+ ])
+ self.chntext = self.raw_chntext.replace('<SIL>', '')
+ else:
+ sp_parts = self.telephone.strip('+').split()
+ self.raw_chntext = '<SP>'.join([
+ num2chn(part, alt_two=False, use_units=False) for part in sp_parts
+ ])
+ self.chntext = self.raw_chntext.replace('<SP>', '')
+ return self.chntext
+
+
+class Fraction:
+ """
+ FRACTION绫�
+ """
+
+ def __init__(self, fraction=None, chntext=None):
+ self.fraction = fraction
+ self.chntext = chntext
+
+ def chntext2fraction(self):
+ denominator, numerator = self.chntext.split('鍒嗕箣')
+ return chn2num(numerator) + '/' + chn2num(denominator)
+
+ def fraction2chntext(self):
+ numerator, denominator = self.fraction.split('/')
+ return num2chn(denominator) + '鍒嗕箣' + num2chn(numerator)
+
+
+class Date:
+ """
+ DATE绫�
+ """
+
+ def __init__(self, date=None, chntext=None):
+ self.date = date
+ self.chntext = chntext
+
+ # def chntext2date(self):
+ # chntext = self.chntext
+ # try:
+ # year, other = chntext.strip().split('骞�', maxsplit=1)
+ # year = Digit(chntext=year).digit2chntext() + '骞�'
+ # except ValueError:
+ # other = chntext
+ # year = ''
+ # if other:
+ # try:
+ # month, day = other.strip().split('鏈�', maxsplit=1)
+ # month = Cardinal(chntext=month).chntext2cardinal() + '鏈�'
+ # except ValueError:
+ # day = chntext
+ # month = ''
+ # if day:
+ # day = Cardinal(chntext=day[:-1]).chntext2cardinal() + day[-1]
+ # else:
+ # month = ''
+ # day = ''
+ # date = year + month + day
+ # self.date = date
+ # return self.date
+
+ def date2chntext(self):
+ date = self.date
+ try:
+ year, other = date.strip().split('骞�', 1)
+ year = Digit(digit=year).digit2chntext() + '骞�'
+ except ValueError:
+ other = date
+ year = ''
+ if other:
+ try:
+ month, day = other.strip().split('鏈�', 1)
+ month = Cardinal(cardinal=month).cardinal2chntext() + '鏈�'
+ except ValueError:
+ day = date
+ month = ''
+ if day:
+ day = Cardinal(cardinal=day[:-1]).cardinal2chntext() + day[-1]
+ else:
+ month = ''
+ day = ''
+ chntext = year + month + day
+ self.chntext = chntext
+ return self.chntext
+
+
+class Money:
+ """
+ MONEY绫�
+ """
+
+ def __init__(self, money=None, chntext=None):
+ self.money = money
+ self.chntext = chntext
+
+ # def chntext2money(self):
+ # return self.money
+
+ def money2chntext(self):
+ money = self.money
+ pattern = re.compile(r'(\d+(\.\d+)?)')
+ matchers = pattern.findall(money)
+ if matchers:
+ for matcher in matchers:
+ money = money.replace(matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext())
+ self.chntext = money
+ return self.chntext
+
+
+class Percentage:
+ """
+ PERCENTAGE绫�
+ """
+
+ def __init__(self, percentage=None, chntext=None):
+ self.percentage = percentage
+ self.chntext = chntext
+
+ def chntext2percentage(self):
+ return chn2num(self.chntext.strip().strip('鐧惧垎涔�')) + '%'
+
+ def percentage2chntext(self):
+ return '鐧惧垎涔�' + num2chn(self.percentage.strip().strip('%'))
+
+
+def remove_erhua(text, er_whitelist):
+ """
+ 鍘婚櫎鍎垮寲闊宠瘝涓殑鍎�:
+ 浠栧コ鍎垮湪閭h竟鍎� -> 浠栧コ鍎垮湪閭h竟
+ """
+
+ er_pattern = re.compile(er_whitelist)
+ new_str=''
+ while re.search('鍎�',text):
+ a = re.search('鍎�',text).span()
+ remove_er_flag = 0
+
+ if er_pattern.search(text):
+ b = er_pattern.search(text).span()
+ if b[0] <= a[0]:
+ remove_er_flag = 1
+
+ if remove_er_flag == 0 :
+ new_str = new_str + text[0:a[0]]
+ text = text[a[1]:]
+ else:
+ new_str = new_str + text[0:b[1]]
+ text = text[b[1]:]
+
+ text = new_str + text
+ return text
+
+# ================================================================================ #
+# NSW Normalizer
+# ================================================================================ #
+class NSWNormalizer:
+ def __init__(self, raw_text):
+ self.raw_text = '^' + raw_text + '$'
+ self.norm_text = ''
+
+ def _particular(self):
+ text = self.norm_text
+ pattern = re.compile(r"(([a-zA-Z]+)浜�([a-zA-Z]+))")
+ matchers = pattern.findall(text)
+ if matchers:
+ # print('particular')
+ for matcher in matchers:
+ text = text.replace(matcher[0], matcher[1]+'2'+matcher[2], 1)
+ self.norm_text = text
+ return self.norm_text
+
+ def normalize(self):
+ text = self.raw_text
+
+ # 瑙勮寖鍖栨棩鏈�
+ pattern = re.compile(r"\D+((([089]\d|(19|20)\d{2})骞�)?(\d{1,2}鏈�(\d{1,2}[鏃ュ彿])?)?)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('date')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Date(date=matcher[0]).date2chntext(), 1)
+
+ # 瑙勮寖鍖栭噾閽�
+ pattern = re.compile(r"\D+((\d+(\.\d+)?)[澶氫綑鍑燷?" + CURRENCY_UNITS + r"(\d" + CURRENCY_UNITS + r"?)?)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('money')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Money(money=matcher[0]).money2chntext(), 1)
+
+ # 瑙勮寖鍖栧浐璇�/鎵嬫満鍙风爜
+ # 鎵嬫満
+ # http://www.jihaoba.com/news/show/13680
+ # 绉诲姩锛�139銆�138銆�137銆�136銆�135銆�134銆�159銆�158銆�157銆�150銆�151銆�152銆�188銆�187銆�182銆�183銆�184銆�178銆�198
+ # 鑱旈�氾細130銆�131銆�132銆�156銆�155銆�186銆�185銆�176
+ # 鐢典俊锛�133銆�153銆�189銆�180銆�181銆�177
+ pattern = re.compile(r"\D((\+?86 ?)?1([38]\d|5[0-35-9]|7[678]|9[89])\d{8})\D")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('telephone')
+ for matcher in matchers:
+ text = text.replace(matcher[0], TelePhone(telephone=matcher[0]).telephone2chntext(), 1)
+ # 鍥鸿瘽
+ pattern = re.compile(r"\D((0(10|2[1-3]|[3-9]\d{2})-?)?[1-9]\d{6,7})\D")
+ matchers = pattern.findall(text)
+ if matchers:
+ # print('fixed telephone')
+ for matcher in matchers:
+ text = text.replace(matcher[0], TelePhone(telephone=matcher[0]).telephone2chntext(fixed=True), 1)
+
+ # 瑙勮寖鍖栧垎鏁�
+ pattern = re.compile(r"(\d+/\d+)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('fraction')
+ for matcher in matchers:
+ text = text.replace(matcher, Fraction(fraction=matcher).fraction2chntext(), 1)
+
+ # 瑙勮寖鍖栫櫨鍒嗘暟
+ text = text.replace('锛�', '%')
+ pattern = re.compile(r"(\d+(\.\d+)?%)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('percentage')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Percentage(percentage=matcher[0]).percentage2chntext(), 1)
+
+ # 瑙勮寖鍖栫函鏁�+閲忚瘝
+ pattern = re.compile(r"(\d+(\.\d+)?)[澶氫綑鍑燷?" + COM_QUANTIFIERS)
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('cardinal+quantifier')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext(), 1)
+
+ # 瑙勮寖鍖栨暟瀛楃紪鍙�
+ pattern = re.compile(r"(\d{4,32})")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('digit')
+ for matcher in matchers:
+ text = text.replace(matcher, Digit(digit=matcher).digit2chntext(), 1)
+
+ # 瑙勮寖鍖栫函鏁�
+ pattern = re.compile(r"(\d+(\.\d+)?)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('cardinal')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext(), 1)
+
+ self.norm_text = text
+ self._particular()
+
+ return self.norm_text.lstrip('^').rstrip('$')
+
+
+def nsw_test_case(raw_text):
+ print('I:' + raw_text)
+ print('O:' + NSWNormalizer(raw_text).normalize())
+ print('')
+
+
+def nsw_test():
+ nsw_test_case('鍥鸿瘽锛�0595-23865596鎴�23880880銆�')
+ nsw_test_case('鍥鸿瘽锛�0595-23865596鎴�23880880銆�')
+ nsw_test_case('鎵嬫満锛�+86 19859213959鎴�15659451527銆�')
+ nsw_test_case('鍒嗘暟锛�32477/76391銆�')
+ nsw_test_case('鐧惧垎鏁帮細80.03%銆�')
+ nsw_test_case('缂栧彿锛�31520181154418銆�')
+ nsw_test_case('绾暟锛�2983.07鍏嬫垨12345.60绫炽��')
+ nsw_test_case('鏃ユ湡锛�1999骞�2鏈�20鏃ユ垨09骞�3鏈�15鍙枫��')
+ nsw_test_case('閲戦挶锛�12鍧�5锛�34.5鍏冿紝20.1涓�')
+ nsw_test_case('鐗规畩锛歄2O鎴朆2C銆�')
+ nsw_test_case('3456涓囧惃')
+ nsw_test_case('2938涓�')
+ nsw_test_case('938')
+ nsw_test_case('浠婂ぉ鍚冧簡115涓皬绗煎寘231涓澶�')
+ nsw_test_case('鏈�62锛呯殑姒傜巼')
+
+
+if __name__ == '__main__':
+ #nsw_test()
+
+ p = argparse.ArgumentParser()
+ p.add_argument('ifile', help='input filename, assume utf-8 encoding')
+ p.add_argument('ofile', help='output filename')
+ p.add_argument('--to_upper', action='store_true', help='convert to upper case')
+ p.add_argument('--to_lower', action='store_true', help='convert to lower case')
+ p.add_argument('--has_key', action='store_true', help="input text has Kaldi's key as first field.")
+ p.add_argument('--remove_fillers', type=bool, default=True, help='remove filler chars such as "鍛�, 鍟�"')
+ p.add_argument('--remove_erhua', type=bool, default=True, help='remove erhua chars such as "杩欏効"')
+ p.add_argument('--log_interval', type=int, default=10000, help='log interval in number of processed lines')
+ args = p.parse_args()
+
+ ifile = codecs.open(args.ifile, 'r', 'utf8')
+ ofile = codecs.open(args.ofile, 'w+', 'utf8')
+
+ n = 0
+ for l in ifile:
+ key = ''
+ text = ''
+ if args.has_key:
+ cols = l.split(maxsplit=1)
+ key = cols[0]
+ if len(cols) == 2:
+ text = cols[1].strip()
+ else:
+ text = ''
+ else:
+ text = l.strip()
+
+ # cases
+ if args.to_upper and args.to_lower:
+ sys.stderr.write('text norm: to_upper OR to_lower?')
+ exit(1)
+ if args.to_upper:
+ text = text.upper()
+ if args.to_lower:
+ text = text.lower()
+
+ # Filler chars removal
+ if args.remove_fillers:
+ for ch in FILLER_CHARS:
+ text = text.replace(ch, '')
+
+ if args.remove_erhua:
+ text = remove_erhua(text, ER_WHITELIST)
+
+ # NSW(Non-Standard-Word) normalization
+ text = NSWNormalizer(text).normalize()
+
+ # Punctuations removal
+ old_chars = CHINESE_PUNC_LIST + string.punctuation # includes all CN and EN punctuations
+ new_chars = ' ' * len(old_chars)
+ del_chars = ''
+ text = text.translate(str.maketrans(old_chars, new_chars, del_chars))
+
+ #
+ if args.has_key:
+ ofile.write(key + '\t' + text + '\n')
+ else:
+ ofile.write(text + '\n')
+
+ n += 1
+ if n % args.log_interval == 0:
+ sys.stderr.write("text norm: {} lines done.\n".format(n))
+
+ sys.stderr.write("text norm: {} lines done in total.\n".format(n))
+
+ ifile.close()
+ ofile.close()
diff --git a/egs/aishell/paraformer/conf/train_asr_paraformer_conformer_20e_6d_1280_320.yaml b/egs/aishell/paraformer/conf/train_asr_paraformer_conformer_20e_6d_1280_320.yaml
deleted file mode 100644
index 2b5e2d1..0000000
--- a/egs/aishell/paraformer/conf/train_asr_paraformer_conformer_20e_6d_1280_320.yaml
+++ /dev/null
@@ -1,92 +0,0 @@
-# network architecture
-# encoder related
-encoder: conformer
-encoder_conf:
- output_size: 320 # dimension of attention
- attention_heads: 4
- linear_units: 1280 # the number of units of position-wise feed forward
- num_blocks: 20 # the number of encoder blocks
- dropout_rate: 0.1
- positional_dropout_rate: 0.1
- attention_dropout_rate: 0.0
- input_layer: conv2d # encoder architecture type
- normalize_before: true
- pos_enc_layer_type: rel_pos
- selfattention_layer_type: rel_selfattn
- activation_type: swish
- macaron_style: true
- use_cnn_module: true
- cnn_module_kernel: 15
-
-# decoder related
-decoder: paraformer_decoder_san
-decoder_conf:
- attention_heads: 4
- linear_units: 2048
- num_blocks: 6
- dropout_rate: 0.1
- positional_dropout_rate: 0.1
- self_attention_dropout_rate: 0.0
- src_attention_dropout_rate: 0.0
-
-# hybrid CTC/attention
-model: paraformer
-model_conf:
- ctc_weight: 0.3
- lsm_weight: 0.1 # label smoothing option
- length_normalized_loss: false
- predictor_weight: 1.0
- sampling_ratio: 0.4
-
-# minibatch related
-batch_type: length
-batch_bins: 25000
-num_workers: 16
-
-# optimization related
-accum_grad: 4
-grad_clip: 5
-max_epoch: 50
-val_scheduler_criterion:
- - valid
- - acc
-best_model_criterion:
-- - valid
- - acc
- - max
-keep_nbest_models: 10
-
-optim: adam
-optim_conf:
- lr: 0.0005
-scheduler: warmuplr
-scheduler_conf:
- warmup_steps: 30000
-
-specaug: specaug
-specaug_conf:
- apply_time_warp: true
- time_warp_window: 5
- time_warp_mode: bicubic
- apply_freq_mask: true
- freq_mask_width_range:
- - 0
- - 30
- num_freq_mask: 2
- apply_time_mask: true
- time_mask_width_range:
- - 0
- - 40
- num_time_mask: 2
-
-predictor: cif_predictor
-predictor_conf:
- idim: 256
- threshold: 1.0
- l_order: 1
- r_order: 1
- tail_threshold: 0.45
-
-
-log_interval: 50
-normalize: None
\ No newline at end of file
diff --git a/egs/aishell/paraformer/conf/train_asr_paraformer_sanm_tf_40e_12d_1280_320_lfr6.yaml b/egs/aishell/paraformer/conf/train_asr_paraformer_sanm_tf_40e_12d_1280_320_lfr6.yaml
deleted file mode 100644
index 8643507..0000000
--- a/egs/aishell/paraformer/conf/train_asr_paraformer_sanm_tf_40e_12d_1280_320_lfr6.yaml
+++ /dev/null
@@ -1,114 +0,0 @@
-# network architecture
-# encoder related
-encoder: sanm
-encoder_conf:
- output_size: 320 # dimension of attention
- attention_heads: 4
- linear_units: 1280 # the number of units of position-wise feed forward
- num_blocks: 40 # the number of encoder blocks
- dropout_rate: 0.1
- positional_dropout_rate: 0.1
- attention_dropout_rate: 0.1
- input_layer: pe # encoder architecture type
- pos_enc_class: SinusoidalPositionEncoder
- normalize_before: true
- kernel_size: 11
- sanm_shfit: 0
- selfattention_layer_type: sanm
-
-# decoder related
-decoder: paraformer_decoder_sanm
-decoder_conf:
- attention_heads: 4
- linear_units: 1280
- num_blocks: 12
- dropout_rate: 0.1
- positional_dropout_rate: 0.1
- self_attention_dropout_rate: 0.1
- src_attention_dropout_rate: 0.1
- att_layer_num: 6
- kernel_size: 11
- sanm_shfit: 0
-
-
-predictor: cif_predictor
-predictor_conf:
- idim: 320
- threshold: 1.0
- l_order: 1
- r_order: 1
- tail_threshold: 0.45
-
-# hybrid CTC/attention
-model: paraformer
-model_conf:
- ctc_weight: 0.0
- lsm_weight: 0.1 # label smoothing option
- length_normalized_loss: true
- predictor_weight: 1.0
- predictor_bias: 0
- sampling_ratio: 0.75
-
-
-# minibatch related
-# dataset_type: small
-batch_type: length
-batch_bins: 6000
-num_workers: 16
-# dataset_type: large
-dataset_conf:
- filter_conf:
- min_length: 10
- max_length: 250
- min_token_length: 1
- max_token_length: 200
- shuffle: True
- shuffle_conf:
- shuffle_size: 10240
- sort_size: 500
- batch_conf:
- batch_type: token
- batch_size: 6000
- num_workers: 16
-
-# optimization related
-accum_grad: 1
-grad_clip: 5
-max_epoch: 20
-val_scheduler_criterion:
- - valid
- - acc
-best_model_criterion:
-- - valid
- - acc
- - max
-keep_nbest_models: 5
-
-optim: adam
-optim_conf:
- lr: 0.0005
-scheduler: warmuplr
-scheduler_conf:
- warmup_steps: 15000
-
-specaug: specaug_lfr
-specaug_conf:
- apply_time_warp: false
- time_warp_window: 5
- time_warp_mode: bicubic
- apply_freq_mask: true
- freq_mask_width_range:
- - 0
- - 30
- lfr_rate: 6
- num_freq_mask: 1
- apply_time_mask: true
- time_mask_width_range:
- - 0
- - 12
- num_time_mask: 1
-
-unused_parameters: true
-log_interval: 50
-normalize: None
-split_with_space: true
\ No newline at end of file
diff --git a/egs/aishell/paraformer/conf/train_asr_paraformer_sanm_tf_50e_16d_2048_512_lfr6.yaml b/egs/aishell/paraformer/conf/train_asr_paraformer_sanm_tf_50e_16d_2048_512_lfr6.yaml
deleted file mode 100644
index 67983b3..0000000
--- a/egs/aishell/paraformer/conf/train_asr_paraformer_sanm_tf_50e_16d_2048_512_lfr6.yaml
+++ /dev/null
@@ -1,114 +0,0 @@
-# network architecture
-# encoder related
-encoder: sanm
-encoder_conf:
- output_size: 512 # dimension of attention
- attention_heads: 4
- linear_units: 2048 # the number of units of position-wise feed forward
- num_blocks: 50 # the number of encoder blocks
- dropout_rate: 0.1
- positional_dropout_rate: 0.1
- attention_dropout_rate: 0.1
- input_layer: pe # encoder architecture type
- pos_enc_class: SinusoidalPositionEncoder
- normalize_before: true
- kernel_size: 11
- sanm_shfit: 0
- selfattention_layer_type: sanm
-
-# decoder related
-decoder: paraformer_decoder_sanm
-decoder_conf:
- attention_heads: 4
- linear_units: 2048
- num_blocks: 16
- dropout_rate: 0.1
- positional_dropout_rate: 0.1
- self_attention_dropout_rate: 0.1
- src_attention_dropout_rate: 0.1
- att_layer_num: 16
- kernel_size: 11
- sanm_shfit: 0
-
-
-predictor: cif_predictor_v2
-predictor_conf:
- idim: 512
- threshold: 1.0
- l_order: 1
- r_order: 1
- tail_threshold: 0.45
-
-# hybrid CTC/attention
-model: paraformer
-model_conf:
- ctc_weight: 0.0
- lsm_weight: 0.1 # label smoothing option
- length_normalized_loss: true
- predictor_weight: 1.0
- predictor_bias: 1
- sampling_ratio: 0.75
-
-
-# minibatch related
-# dataset_type: small
-batch_type: length
-batch_bins: 10000
-num_workers: 16
-# dataset_type: large
-dataset_conf:
- filter_conf:
- min_length: 10
- max_length: 250
- min_token_length: 1
- max_token_length: 200
- shuffle: true
- shuffle_conf:
- shuffle_size: 10240
- sort_size: 500
- batch_conf:
- batch_type: 'token'
- batch_size: 6000
- num_workers: 16
-
-# optimization related
-accum_grad: 1
-grad_clip: 5
-max_epoch: 20
-val_scheduler_criterion:
- - valid
- - acc
-best_model_criterion:
-- - valid
- - acc
- - max
-keep_nbest_models: 5
-
-optim: adam
-optim_conf:
- lr: 0.0005
-scheduler: warmuplr
-scheduler_conf:
- warmup_steps: 30000
-
-specaug: specaug_lfr
-specaug_conf:
- apply_time_warp: false
- time_warp_window: 5
- time_warp_mode: bicubic
- apply_freq_mask: true
- freq_mask_width_range:
- - 0
- - 30
- lfr_rate: 6
- num_freq_mask: 1
- apply_time_mask: true
- time_mask_width_range:
- - 0
- - 12
- num_time_mask: 1
-
-unused_parameters: true
-log_interval: 50
-normalize: None
-split_with_space: true
diff --git a/egs/aishell/paraformer/conf/train_asr_paraformerbert_conformer_12e_6d_2048_256.yaml b/egs/aishell/paraformer/conf/train_asr_paraformerbert_conformer_12e_6d_2048_256.yaml
deleted file mode 100644
index f369b3d..0000000
--- a/egs/aishell/paraformer/conf/train_asr_paraformerbert_conformer_12e_6d_2048_256.yaml
+++ /dev/null
@@ -1,99 +0,0 @@
-# network architecture
-# encoder related
-encoder: conformer
-encoder_conf:
- output_size: 256 # dimension of attention
- attention_heads: 4
- linear_units: 2048 # the number of units of position-wise feed forward
- num_blocks: 12 # the number of encoder blocks
- dropout_rate: 0.1
- positional_dropout_rate: 0.1
- attention_dropout_rate: 0.0
- input_layer: conv2d # encoder architecture type
- normalize_before: true
- pos_enc_layer_type: rel_pos
- selfattention_layer_type: rel_selfattn
- activation_type: swish
- macaron_style: true
- use_cnn_module: true
- cnn_module_kernel: 15
-
-# decoder related
-decoder: paraformer_decoder_san
-decoder_conf:
- attention_heads: 4
- linear_units: 2048
- num_blocks: 6
- dropout_rate: 0.1
- positional_dropout_rate: 0.1
- self_attention_dropout_rate: 0.0
- src_attention_dropout_rate: 0.0
-
-# hybrid CTC/attention
-model: paraformer_bert
-model_conf:
- ctc_weight: 0.3
- lsm_weight: 0.1 # label smoothing option
- length_normalized_loss: false
- predictor_weight: 1.0
- sampling_ratio: 0.4
- embeds_id: 3
- embed_dims: 768
- embeds_loss_weight: 2.0
-
-
-
-# minibatch related
-#batch_type: length
-#batch_bins: 40000
-batch_type: numel
-batch_bins: 2000000
-num_workers: 16
-
-# optimization related
-accum_grad: 4
-grad_clip: 5
-max_epoch: 50
-val_scheduler_criterion:
- - valid
- - acc
-best_model_criterion:
-- - valid
- - acc
- - max
-keep_nbest_models: 10
-
-optim: adam
-optim_conf:
- lr: 0.0005
-scheduler: warmuplr
-scheduler_conf:
- warmup_steps: 30000
-
-specaug: specaug
-specaug_conf:
- apply_time_warp: true
- time_warp_window: 5
- time_warp_mode: bicubic
- apply_freq_mask: true
- freq_mask_width_range:
- - 0
- - 30
- num_freq_mask: 2
- apply_time_mask: true
- time_mask_width_range:
- - 0
- - 40
- num_time_mask: 2
-
-predictor: cif_predictor
-predictor_conf:
- idim: 256
- threshold: 1.0
- l_order: 1
- r_order: 1
- tail_threshold: 0.45
-
-
-log_interval: 50
-normalize: None
\ No newline at end of file
diff --git a/egs/aishell/paraformer/run.sh b/egs/aishell/paraformer/run.sh
index bebb646..06322ce 100755
--- a/egs/aishell/paraformer/run.sh
+++ b/egs/aishell/paraformer/run.sh
@@ -10,9 +10,10 @@
# for gpu decoding, inference_nj=ngpu*njob; for cpu decoding, inference_nj=njob
njob=8
train_cmd=utils/run.pl
+infer_cmd=utils/run.pl
# general configuration
-feats_dir=".." #feature output dictionary, for large data
+feats_dir="../DATA" #feature output dictionary, for large data
exp_dir="."
lang=zh
dumpdir=dump/fbank
@@ -59,8 +60,10 @@
if ${gpu_inference}; then
inference_nj=$[${ngpu}*${njob}]
+ _ngpu=1
else
inference_nj=$njob
+ _ngpu=0
fi
if [ ${stage} -le 0 ] && [ ${stop_stage} -ge 0 ]; then
@@ -83,18 +86,18 @@
echo "stage 1: Feature Generation"
# compute fbank features
fbankdir=${feats_dir}/fbank
- utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj --speed_perturb ${speed_perturb} \
+ utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj --feats_dim ${feats_dim} --sample_frequency ${sample_frequency} --speed_perturb ${speed_perturb} \
${feats_dir}/data/train ${exp_dir}/exp/make_fbank/train ${fbankdir}/train
utils/fix_data_feat.sh ${fbankdir}/train
- utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj \
+ utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj --feats_dim ${feats_dim} --sample_frequency ${sample_frequency} \
${feats_dir}/data/dev ${exp_dir}/exp/make_fbank/dev ${fbankdir}/dev
utils/fix_data_feat.sh ${fbankdir}/dev
- utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj \
+ utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj --feats_dim ${feats_dim} --sample_frequency ${sample_frequency} \
${feats_dir}/data/test ${exp_dir}/exp/make_fbank/test ${fbankdir}/test
utils/fix_data_feat.sh ${fbankdir}/test
# compute global cmvn
- utils/compute_cmvn.sh --cmd "$train_cmd" --nj $nj \
+ utils/compute_cmvn.sh --cmd "$train_cmd" --nj $nj --feats_dim ${feats_dim} \
${fbankdir}/train ${exp_dir}/exp/make_fbank/train
# apply cmvn
@@ -112,6 +115,10 @@
utils/fix_data_feat.sh ${feat_train_dir}
utils/fix_data_feat.sh ${feat_dev_dir}
utils/fix_data_feat.sh ${feat_test_dir}
+
+ #generate ark list
+ utils/gen_ark_list.sh --cmd "$train_cmd" --nj $nj ${feat_train_dir} ${fbankdir}/train ${feat_train_dir}
+ utils/gen_ark_list.sh --cmd "$train_cmd" --nj $nj ${feat_dev_dir} ${fbankdir}/dev ${feat_dev_dir}
fi
token_list=${feats_dir}/data/${lang}_token_list/char/tokens.txt
@@ -140,9 +147,10 @@
# Training Stage
world_size=$gpu_num # run on one machine
if [ ${stage} -le 3 ] && [ ${stop_stage} -ge 3 ]; then
+ echo "stage 3: Training"
mkdir -p ${exp_dir}/exp/${model_dir}
- mkdir -p ${exp_dir}/exp/log
- INIT_FILE=$exp_dir/ddp_init
+ mkdir -p ${exp_dir}/exp/${model_dir}/log
+ INIT_FILE=${exp_dir}/exp/${model_dir}/ddp_init
if [ -f $INIT_FILE ];then
rm -f $INIT_FILE
fi
@@ -184,25 +192,57 @@
# Testing Stage
if [ ${stage} -le 4 ] && [ ${stop_stage} -ge 4 ]; then
- utils/easy_asr_infer.sh \
- --lang zh \
- --datadir ${feats_dir} \
- --feats_type ${feats_type} \
- --feats_dim ${feats_dim} \
- --token_type ${token_type} \
- --gpu_inference ${gpu_inference} \
- --inference_config "${inference_config}" \
- --test_sets "${test_sets}" \
- --token_list $token_list \
- --asr_exp ${exp_dir}/${model_dir} \
- --stage 12 \
- --stop_stage 12 \
- --scp $scp \
- --text text \
- --inference_nj $inference_nj \
- --njob $njob \
- --inference_asr_model $inference_asr_model \
- --gpuid_list $gpuid_list \
- --mode paraformer
+ echo "stage 4: Inference"
+ for dset in ${test_sets}; do
+ asr_exp=${exp_dir}/exp/${model_dir}
+ inference_tag="$(basename "${inference_config}" .yaml)"
+ _dir="${asr_exp}/${inference_tag}/${inference_asr_model}/${dset}"
+ _logdir="${_dir}/logdir"
+ if [ -d ${_dir} ]; then
+ echo "${_dir} is already exists. if you want to decode again, please delete this dir first."
+ exit 0
+ fi
+ mkdir -p "${_logdir}"
+ _data="${feats_dir}/${dumpdir}/${dset}"
+ key_file=${_data}/${scp}
+ num_scp_file="$(<${key_file} wc -l)"
+ _nj=$([ $inference_nj -le $num_scp_file ] && echo "$inference_nj" || echo "$num_scp_file")
+ split_scps=
+ for n in $(seq "${_nj}"); do
+ split_scps+=" ${_logdir}/keys.${n}.scp"
+ done
+ # shellcheck disable=SC2086
+ utils/split_scp.pl "${key_file}" ${split_scps}
+ _opts=
+ if [ -n "${inference_config}" ]; then
+ _opts+="--config ${inference_config} "
+ fi
+ ${infer_cmd} --gpu "${_ngpu}" --max-jobs-run "${_nj}" JOB=1:"${_nj}" "${_logdir}"/asr_inference.JOB.log \
+ python -m funasr.bin.asr_inference_launch \
+ --batch_size 1 \
+ --ngpu "${_ngpu}" \
+ --njob ${njob} \
+ --gpuid_list ${gpuid_list} \
+ --data_path_and_name_and_type "${_data}/${scp},speech,${type}" \
+ --key_file "${_logdir}"/keys.JOB.scp \
+ --asr_train_config "${asr_exp}"/config.yaml \
+ --asr_model_file "${asr_exp}"/"${inference_asr_model}" \
+ --output_dir "${_logdir}"/output.JOB \
+ --mode paraformer \
+ ${_opts}
+
+ for f in token token_int score text; do
+ if [ -f "${_logdir}/output.1/1best_recog/${f}" ]; then
+ for i in $(seq "${_nj}"); do
+ cat "${_logdir}/output.${i}/1best_recog/${f}"
+ done | sort -k1 >"${_dir}/${f}"
+ fi
+ done
+ python utils/proce_text.py ${_dir}/text ${_dir}/text.proc
+ python utils/proce_text.py ${_data}/text ${_data}/text.proc
+ python utils/compute_wer.py ${_data}/text.proc ${_dir}/text.proc ${_dir}/text.cer
+ tail -n 3 ${_dir}/text.cer > ${_dir}/text.cer.txt
+ cat ${_dir}/text.cer.txt
+ done
fi
diff --git a/egs/aishell/paraformer/utils b/egs/aishell/paraformer/utils
deleted file mode 120000
index 40e14f5..0000000
--- a/egs/aishell/paraformer/utils
+++ /dev/null
@@ -1 +0,0 @@
-../tranformer/utils
\ No newline at end of file
diff --git a/egs/aishell/paraformer/utils/__init__.py b/egs/aishell/paraformer/utils/__init__.py
new file mode 100644
index 0000000..e69de29
--- /dev/null
+++ b/egs/aishell/paraformer/utils/__init__.py
diff --git a/egs/aishell/paraformer/utils/apply_cmvn.py b/egs/aishell/paraformer/utils/apply_cmvn.py
new file mode 100755
index 0000000..b5c5086
--- /dev/null
+++ b/egs/aishell/paraformer/utils/apply_cmvn.py
@@ -0,0 +1,79 @@
+from kaldiio import ReadHelper
+from kaldiio import WriteHelper
+
+import argparse
+import json
+import math
+import numpy as np
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="apply cmvn",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--ark-file",
+ "-a",
+ default=False,
+ required=True,
+ type=str,
+ help="fbank ark file",
+ )
+ parser.add_argument(
+ "--cmvn-file",
+ "-c",
+ default=False,
+ required=True,
+ type=str,
+ help="cmvn file",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ ark_file = args.output_dir + "/feats." + str(args.ark_index) + ".ark"
+ scp_file = args.output_dir + "/feats." + str(args.ark_index) + ".scp"
+ ark_writer = WriteHelper('ark,scp:{},{}'.format(ark_file, scp_file))
+
+ with open(args.cmvn_file) as f:
+ cmvn_stats = json.load(f)
+
+ means = cmvn_stats['mean_stats']
+ vars = cmvn_stats['var_stats']
+ total_frames = cmvn_stats['total_frames']
+
+ for i in range(len(means)):
+ means[i] /= total_frames
+ vars[i] = vars[i] / total_frames - means[i] * means[i]
+ if vars[i] < 1.0e-20:
+ vars[i] = 1.0e-20
+ vars[i] = 1.0 / math.sqrt(vars[i])
+
+ with ReadHelper('ark:{}'.format(args.ark_file)) as ark_reader:
+ for key, mat in ark_reader:
+ mat = (mat - means) * vars
+ ark_writer(key, mat)
+
+
+if __name__ == '__main__':
+ main()
diff --git a/egs/aishell/paraformer/utils/apply_cmvn.sh b/egs/aishell/paraformer/utils/apply_cmvn.sh
new file mode 100755
index 0000000..f8fd1d1
--- /dev/null
+++ b/egs/aishell/paraformer/utils/apply_cmvn.sh
@@ -0,0 +1,29 @@
+#!/usr/bin/env bash
+
+. ./path.sh || exit 1;
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+fbankdir=$1
+cmvn_file=$2
+logdir=$3
+output_dir=$4
+
+dump_dir=${output_dir}/ark; mkdir -p ${dump_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/apply_cmvn.JOB.log \
+ python utils/apply_cmvn.py -a $fbankdir/ark/feats.JOB.ark \
+ -c $cmvn_file -i JOB -o ${dump_dir} \
+ || exit 1;
+
+for n in $(seq $nj); do
+ cat ${dump_dir}/feats.$n.scp || exit 1
+done > ${output_dir}/feats.scp || exit 1
+
+echo "$0: Succeeded apply cmvn"
diff --git a/egs/aishell/paraformer/utils/apply_lfr_and_cmvn.py b/egs/aishell/paraformer/utils/apply_lfr_and_cmvn.py
new file mode 100755
index 0000000..50d18d1
--- /dev/null
+++ b/egs/aishell/paraformer/utils/apply_lfr_and_cmvn.py
@@ -0,0 +1,143 @@
+from kaldiio import ReadHelper, WriteHelper
+
+import argparse
+import numpy as np
+
+
+def build_LFR_features(inputs, m=7, n=6):
+ LFR_inputs = []
+ T = inputs.shape[0]
+ T_lfr = int(np.ceil(T / n))
+ left_padding = np.tile(inputs[0], ((m - 1) // 2, 1))
+ inputs = np.vstack((left_padding, inputs))
+ T = T + (m - 1) // 2
+ for i in range(T_lfr):
+ if m <= T - i * n:
+ LFR_inputs.append(np.hstack(inputs[i * n:i * n + m]))
+ else:
+ num_padding = m - (T - i * n)
+ frame = np.hstack(inputs[i * n:])
+ for _ in range(num_padding):
+ frame = np.hstack((frame, inputs[-1]))
+ LFR_inputs.append(frame)
+ return np.vstack(LFR_inputs)
+
+
+def build_CMVN_features(inputs, mvn_file): # noqa
+ with open(mvn_file, 'r', encoding='utf-8') as f:
+ lines = f.readlines()
+
+ add_shift_list = []
+ rescale_list = []
+ for i in range(len(lines)):
+ line_item = lines[i].split()
+ if line_item[0] == '<AddShift>':
+ line_item = lines[i + 1].split()
+ if line_item[0] == '<LearnRateCoef>':
+ add_shift_line = line_item[3:(len(line_item) - 1)]
+ add_shift_list = list(add_shift_line)
+ continue
+ elif line_item[0] == '<Rescale>':
+ line_item = lines[i + 1].split()
+ if line_item[0] == '<LearnRateCoef>':
+ rescale_line = line_item[3:(len(line_item) - 1)]
+ rescale_list = list(rescale_line)
+ continue
+
+ for j in range(inputs.shape[0]):
+ for k in range(inputs.shape[1]):
+ add_shift_value = add_shift_list[k]
+ rescale_value = rescale_list[k]
+ inputs[j, k] = float(inputs[j, k]) + float(add_shift_value)
+ inputs[j, k] = float(inputs[j, k]) * float(rescale_value)
+
+ return inputs
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="apply low_frame_rate and cmvn",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--ark-file",
+ "-a",
+ default=False,
+ required=True,
+ type=str,
+ help="fbank ark file",
+ )
+ parser.add_argument(
+ "--lfr",
+ "-f",
+ default=True,
+ type=str,
+ help="low frame rate",
+ )
+ parser.add_argument(
+ "--lfr-m",
+ "-m",
+ default=7,
+ type=int,
+ help="number of frames to stack",
+ )
+ parser.add_argument(
+ "--lfr-n",
+ "-n",
+ default=6,
+ type=int,
+ help="number of frames to skip",
+ )
+ parser.add_argument(
+ "--cmvn-file",
+ "-c",
+ default=False,
+ required=True,
+ type=str,
+ help="global cmvn file",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ dump_ark_file = args.output_dir + "/feats." + str(args.ark_index) + ".ark"
+ dump_scp_file = args.output_dir + "/feats." + str(args.ark_index) + ".scp"
+ shape_file = args.output_dir + "/len." + str(args.ark_index)
+ ark_writer = WriteHelper('ark,scp:{},{}'.format(dump_ark_file, dump_scp_file))
+
+ shape_writer = open(shape_file, 'w')
+ with ReadHelper('ark:{}'.format(args.ark_file)) as ark_reader:
+ for key, mat in ark_reader:
+ if args.lfr:
+ lfr = build_LFR_features(mat, args.lfr_m, args.lfr_n)
+ else:
+ lfr = mat
+ cmvn = build_CMVN_features(lfr, args.cmvn_file)
+ dims = cmvn.shape[1]
+ lens = cmvn.shape[0]
+ shape_writer.write(key + " " + str(lens) + "," + str(dims) + '\n')
+ ark_writer(key, cmvn)
+
+
+if __name__ == '__main__':
+ main()
+
diff --git a/egs/aishell/paraformer/utils/apply_lfr_and_cmvn.sh b/egs/aishell/paraformer/utils/apply_lfr_and_cmvn.sh
new file mode 100755
index 0000000..3119fdb
--- /dev/null
+++ b/egs/aishell/paraformer/utils/apply_lfr_and_cmvn.sh
@@ -0,0 +1,38 @@
+#!/usr/bin/env bash
+
+
+# Begin configuration section.
+nj=32
+cmd=utils/run.pl
+
+# feature configuration
+lfr=True
+lfr_m=7
+lfr_n=6
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+fbankdir=$1
+cmvn_file=$2
+logdir=$3
+output_dir=$4
+
+dump_dir=${output_dir}/ark; mkdir -p ${dump_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/apply_lfr_and_cmvn.JOB.log \
+ python utils/apply_lfr_and_cmvn.py -a $fbankdir/ark/feats.JOB.ark \
+ -f $lfr -m $lfr_m -n $lfr_n -c $cmvn_file -i JOB -o ${dump_dir} \
+ || exit 1;
+
+for n in $(seq $nj); do
+ cat ${dump_dir}/feats.$n.scp || exit 1
+done > ${output_dir}/feats.scp || exit 1
+
+for n in $(seq $nj); do
+ cat ${dump_dir}/len.$n || exit 1
+done > ${output_dir}/speech_shape || exit 1
+
+echo "$0: Succeeded apply low frame rate and cmvn"
diff --git a/egs/aishell/paraformer/utils/combine_cmvn_file.py b/egs/aishell/paraformer/utils/combine_cmvn_file.py
new file mode 100755
index 0000000..b2974a4
--- /dev/null
+++ b/egs/aishell/paraformer/utils/combine_cmvn_file.py
@@ -0,0 +1,73 @@
+import argparse
+import json
+import numpy as np
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="combine cmvn file",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--dims",
+ "-d",
+ default=80,
+ type=int,
+ help="feature dims",
+ )
+ parser.add_argument(
+ "--cmvn-dir",
+ "-c",
+ default=False,
+ required=True,
+ type=str,
+ help="cmvn dir",
+ )
+
+ parser.add_argument(
+ "--nj",
+ "-n",
+ default=1,
+ required=True,
+ type=int,
+ help="num of cmvn file",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ total_means = np.zeros(args.dims)
+ total_vars = np.zeros(args.dims)
+ total_frames = 0
+
+ cmvn_file = args.output_dir + "/cmvn.json"
+
+ for i in range(1, args.nj+1):
+ with open(args.cmvn_dir + "/cmvn." + str(i) + ".json", "r") as fin:
+ cmvn_stats = json.load(fin)
+
+ total_means += np.array(cmvn_stats["mean_stats"])
+ total_vars += np.array(cmvn_stats["var_stats"])
+ total_frames += cmvn_stats["total_frames"]
+
+ cmvn_info = {
+ 'mean_stats': list(total_means.tolist()),
+ 'var_stats': list(total_vars.tolist()),
+ 'total_frames': total_frames
+ }
+ with open(cmvn_file, 'w') as fout:
+ fout.write(json.dumps(cmvn_info))
+
+
+if __name__ == '__main__':
+ main()
diff --git a/egs/aishell/paraformer/utils/compute_cmvn.py b/egs/aishell/paraformer/utils/compute_cmvn.py
new file mode 100755
index 0000000..2b96e26
--- /dev/null
+++ b/egs/aishell/paraformer/utils/compute_cmvn.py
@@ -0,0 +1,74 @@
+from kaldiio import ReadHelper
+
+import argparse
+import numpy as np
+import json
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="computer global cmvn",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--dims",
+ "-d",
+ default=80,
+ type=int,
+ help="feature dims",
+ )
+ parser.add_argument(
+ "--ark-file",
+ "-a",
+ default=False,
+ required=True,
+ type=str,
+ help="fbank ark file",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ ark_file = args.ark_file + "/feats." + str(args.ark_index) + ".ark"
+ cmvn_file = args.output_dir + "/cmvn." + str(args.ark_index) + ".json"
+
+ mean_stats = np.zeros(args.dims)
+ var_stats = np.zeros(args.dims)
+ total_frames = 0
+
+ with ReadHelper('ark:{}'.format(ark_file)) as ark_reader:
+ for key, mat in ark_reader:
+ mean_stats += np.sum(mat, axis=0)
+ var_stats += np.sum(np.square(mat), axis=0)
+ total_frames += mat.shape[0]
+
+ cmvn_info = {
+ 'mean_stats': list(mean_stats.tolist()),
+ 'var_stats': list(var_stats.tolist()),
+ 'total_frames': total_frames
+ }
+ with open(cmvn_file, 'w') as fout:
+ fout.write(json.dumps(cmvn_info))
+
+
+if __name__ == '__main__':
+ main()
diff --git a/egs/aishell/paraformer/utils/compute_cmvn.sh b/egs/aishell/paraformer/utils/compute_cmvn.sh
new file mode 100755
index 0000000..12173ee
--- /dev/null
+++ b/egs/aishell/paraformer/utils/compute_cmvn.sh
@@ -0,0 +1,25 @@
+#!/usr/bin/env bash
+
+. ./path.sh || exit 1;
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+feats_dim=80
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+fbankdir=$1
+logdir=$2
+
+output_dir=${fbankdir}/cmvn; mkdir -p ${output_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/cmvn.JOB.log \
+ python utils/compute_cmvn.py -d ${feats_dim} -a $fbankdir/ark -i JOB -o ${output_dir} \
+ || exit 1;
+
+python utils/combine_cmvn_file.py -d ${feats_dim} -c ${output_dir} -n $nj -o $fbankdir
+
+echo "$0: Succeeded compute global cmvn"
diff --git a/egs/aishell/paraformer/utils/compute_fbank.py b/egs/aishell/paraformer/utils/compute_fbank.py
new file mode 100755
index 0000000..d03b5a8
--- /dev/null
+++ b/egs/aishell/paraformer/utils/compute_fbank.py
@@ -0,0 +1,153 @@
+from kaldiio import WriteHelper
+
+import argparse
+import numpy as np
+import json
+import torch
+import torchaudio
+import torchaudio.compliance.kaldi as kaldi
+
+
+def compute_fbank(wav_file,
+ num_mel_bins=80,
+ frame_length=25,
+ frame_shift=10,
+ dither=0.0,
+ resample_rate=16000,
+ speed=1.0):
+
+ waveform, sample_rate = torchaudio.load(wav_file)
+ if resample_rate != sample_rate:
+ waveform = torchaudio.transforms.Resample(orig_freq=sample_rate,
+ new_freq=resample_rate)(waveform)
+ if speed != 1.0:
+ waveform, _ = torchaudio.sox_effects.apply_effects_tensor(
+ waveform, resample_rate,
+ [['speed', str(speed)], ['rate', str(resample_rate)]]
+ )
+
+ waveform = waveform * (1 << 15)
+ mat = kaldi.fbank(waveform,
+ num_mel_bins=num_mel_bins,
+ frame_length=frame_length,
+ frame_shift=frame_shift,
+ dither=dither,
+ energy_floor=0.0,
+ window_type='hamming',
+ sample_frequency=resample_rate)
+
+ return mat.numpy()
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="computer features",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--wav-lists",
+ "-w",
+ default=False,
+ required=True,
+ type=str,
+ help="input wav lists",
+ )
+ parser.add_argument(
+ "--text-files",
+ "-t",
+ default=False,
+ required=True,
+ type=str,
+ help="input text files",
+ )
+ parser.add_argument(
+ "--dims",
+ "-d",
+ default=80,
+ type=int,
+ help="feature dims",
+ )
+ parser.add_argument(
+ "--sample-frequency",
+ "-s",
+ default=16000,
+ type=int,
+ help="sample frequency",
+ )
+ parser.add_argument(
+ "--speed-perturb",
+ "-p",
+ default="1.0",
+ type=str,
+ help="speed perturb",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-a",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ ark_file = args.output_dir + "/ark/feats." + str(args.ark_index) + ".ark"
+ scp_file = args.output_dir + "/ark/feats." + str(args.ark_index) + ".scp"
+ text_file = args.output_dir + "/txt/text." + str(args.ark_index) + ".txt"
+ feats_shape_file = args.output_dir + "/ark/len." + str(args.ark_index)
+ text_shape_file = args.output_dir + "/txt/len." + str(args.ark_index)
+
+ ark_writer = WriteHelper('ark,scp:{},{}'.format(ark_file, scp_file))
+ text_writer = open(text_file, 'w')
+ feats_shape_writer = open(feats_shape_file, 'w')
+ text_shape_writer = open(text_shape_file, 'w')
+
+ speed_perturb_list = args.speed_perturb.split(',')
+
+ for speed in speed_perturb_list:
+ with open(args.wav_lists, 'r', encoding='utf-8') as wavfile:
+ with open(args.text_files, 'r', encoding='utf-8') as textfile:
+ for wav, text in zip(wavfile, textfile):
+ s_w = wav.strip().split()
+ wav_id = s_w[0]
+ wav_file = s_w[1]
+
+ s_t = text.strip().split()
+ text_id = s_t[0]
+ txt = s_t[1:]
+ fbank = compute_fbank(wav_file,
+ num_mel_bins=args.dims,
+ resample_rate=args.sample_frequency,
+ speed=float(speed)
+ )
+ feats_dims = fbank.shape[1]
+ feats_lens = fbank.shape[0]
+ txt_lens = len(txt)
+ if speed == "1.0":
+ wav_id_sp = wav_id
+ else:
+ wav_id_sp = wav_id + "_sp" + speed
+
+ feats_shape_writer.write(wav_id_sp + " " + str(feats_lens) + "," + str(feats_dims) + '\n')
+ text_shape_writer.write(wav_id_sp + " " + str(txt_lens) + '\n')
+
+ text_writer.write(wav_id_sp + " " + " ".join(txt) + '\n')
+ ark_writer(wav_id_sp, fbank)
+
+
+if __name__ == '__main__':
+ main()
+
diff --git a/egs/aishell/paraformer/utils/compute_fbank.sh b/egs/aishell/paraformer/utils/compute_fbank.sh
new file mode 100755
index 0000000..92a4fe6
--- /dev/null
+++ b/egs/aishell/paraformer/utils/compute_fbank.sh
@@ -0,0 +1,51 @@
+#!/usr/bin/env bash
+
+. ./path.sh || exit 1;
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+
+# feature configuration
+feats_dim=80
+sample_frequency=16000
+speed_perturb="1.0"
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+data=$1
+logdir=$2
+fbankdir=$3
+
+[ ! -f $data/wav.scp ] && echo "$0: no such file $data/wav.scp" && exit 1;
+[ ! -f $data/text ] && echo "$0: no such file $data/text" && exit 1;
+
+python utils/split_data.py $data $data $nj
+
+ark_dir=${fbankdir}/ark; mkdir -p ${ark_dir}
+text_dir=${fbankdir}/txt; mkdir -p ${text_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/make_fbank.JOB.log \
+ python utils/compute_fbank.py -w $data/split${nj}/JOB/wav.scp -t $data/split${nj}/JOB/text \
+ -d $feats_dim -s $sample_frequency -p ${speed_perturb} -a JOB -o ${fbankdir} \
+ || exit 1;
+
+for n in $(seq $nj); do
+ cat ${ark_dir}/feats.$n.scp || exit 1
+done > $fbankdir/feats.scp || exit 1
+
+for n in $(seq $nj); do
+ cat ${text_dir}/text.$n.txt || exit 1
+done > $fbankdir/text || exit 1
+
+for n in $(seq $nj); do
+ cat ${ark_dir}/len.$n || exit 1
+done > $fbankdir/speech_shape || exit 1
+
+for n in $(seq $nj); do
+ cat ${text_dir}/len.$n || exit 1
+done > $fbankdir/text_shape || exit 1
+
+echo "$0: Succeeded compute FBANK features"
diff --git a/egs/aishell/paraformer/utils/compute_wer.py b/egs/aishell/paraformer/utils/compute_wer.py
new file mode 100755
index 0000000..349a3f6
--- /dev/null
+++ b/egs/aishell/paraformer/utils/compute_wer.py
@@ -0,0 +1,157 @@
+import os
+import numpy as np
+import sys
+
+def compute_wer(ref_file,
+ hyp_file,
+ cer_detail_file):
+ rst = {
+ 'Wrd': 0,
+ 'Corr': 0,
+ 'Ins': 0,
+ 'Del': 0,
+ 'Sub': 0,
+ 'Snt': 0,
+ 'Err': 0.0,
+ 'S.Err': 0.0,
+ 'wrong_words': 0,
+ 'wrong_sentences': 0
+ }
+
+ hyp_dict = {}
+ ref_dict = {}
+ with open(hyp_file, 'r') as hyp_reader:
+ for line in hyp_reader:
+ key = line.strip().split()[0]
+ value = line.strip().split()[1:]
+ hyp_dict[key] = value
+ with open(ref_file, 'r') as ref_reader:
+ for line in ref_reader:
+ key = line.strip().split()[0]
+ value = line.strip().split()[1:]
+ ref_dict[key] = value
+
+ cer_detail_writer = open(cer_detail_file, 'w')
+ for hyp_key in hyp_dict:
+ if hyp_key in ref_dict:
+ out_item = compute_wer_by_line(hyp_dict[hyp_key], ref_dict[hyp_key])
+ rst['Wrd'] += out_item['nwords']
+ rst['Corr'] += out_item['cor']
+ rst['wrong_words'] += out_item['wrong']
+ rst['Ins'] += out_item['ins']
+ rst['Del'] += out_item['del']
+ rst['Sub'] += out_item['sub']
+ rst['Snt'] += 1
+ if out_item['wrong'] > 0:
+ rst['wrong_sentences'] += 1
+ cer_detail_writer.write(hyp_key + print_cer_detail(out_item) + '\n')
+ cer_detail_writer.write("ref:" + '\t' + "".join(ref_dict[hyp_key]) + '\n')
+ cer_detail_writer.write("hyp:" + '\t' + "".join(hyp_dict[hyp_key]) + '\n')
+
+ if rst['Wrd'] > 0:
+ rst['Err'] = round(rst['wrong_words'] * 100 / rst['Wrd'], 2)
+ if rst['Snt'] > 0:
+ rst['S.Err'] = round(rst['wrong_sentences'] * 100 / rst['Snt'], 2)
+
+ cer_detail_writer.write('\n')
+ cer_detail_writer.write("%WER " + str(rst['Err']) + " [ " + str(rst['wrong_words'])+ " / " + str(rst['Wrd']) +
+ ", " + str(rst['Ins']) + " ins, " + str(rst['Del']) + " del, " + str(rst['Sub']) + " sub ]" + '\n')
+ cer_detail_writer.write("%SER " + str(rst['S.Err']) + " [ " + str(rst['wrong_sentences']) + " / " + str(rst['Snt']) + " ]" + '\n')
+ cer_detail_writer.write("Scored " + str(len(hyp_dict)) + " sentences, " + str(len(hyp_dict) - rst['Snt']) + " not present in hyp." + '\n')
+
+
+def compute_wer_by_line(hyp,
+ ref):
+ hyp = list(map(lambda x: x.lower(), hyp))
+ ref = list(map(lambda x: x.lower(), ref))
+
+ len_hyp = len(hyp)
+ len_ref = len(ref)
+
+ cost_matrix = np.zeros((len_hyp + 1, len_ref + 1), dtype=np.int16)
+
+ ops_matrix = np.zeros((len_hyp + 1, len_ref + 1), dtype=np.int8)
+
+ for i in range(len_hyp + 1):
+ cost_matrix[i][0] = i
+ for j in range(len_ref + 1):
+ cost_matrix[0][j] = j
+
+ for i in range(1, len_hyp + 1):
+ for j in range(1, len_ref + 1):
+ if hyp[i - 1] == ref[j - 1]:
+ cost_matrix[i][j] = cost_matrix[i - 1][j - 1]
+ else:
+ substitution = cost_matrix[i - 1][j - 1] + 1
+ insertion = cost_matrix[i - 1][j] + 1
+ deletion = cost_matrix[i][j - 1] + 1
+
+ compare_val = [substitution, insertion, deletion]
+
+ min_val = min(compare_val)
+ operation_idx = compare_val.index(min_val) + 1
+ cost_matrix[i][j] = min_val
+ ops_matrix[i][j] = operation_idx
+
+ match_idx = []
+ i = len_hyp
+ j = len_ref
+ rst = {
+ 'nwords': len_ref,
+ 'cor': 0,
+ 'wrong': 0,
+ 'ins': 0,
+ 'del': 0,
+ 'sub': 0
+ }
+ while i >= 0 or j >= 0:
+ i_idx = max(0, i)
+ j_idx = max(0, j)
+
+ if ops_matrix[i_idx][j_idx] == 0: # correct
+ if i - 1 >= 0 and j - 1 >= 0:
+ match_idx.append((j - 1, i - 1))
+ rst['cor'] += 1
+
+ i -= 1
+ j -= 1
+
+ elif ops_matrix[i_idx][j_idx] == 2: # insert
+ i -= 1
+ rst['ins'] += 1
+
+ elif ops_matrix[i_idx][j_idx] == 3: # delete
+ j -= 1
+ rst['del'] += 1
+
+ elif ops_matrix[i_idx][j_idx] == 1: # substitute
+ i -= 1
+ j -= 1
+ rst['sub'] += 1
+
+ if i < 0 and j >= 0:
+ rst['del'] += 1
+ elif j < 0 and i >= 0:
+ rst['ins'] += 1
+
+ match_idx.reverse()
+ wrong_cnt = cost_matrix[len_hyp][len_ref]
+ rst['wrong'] = wrong_cnt
+
+ return rst
+
+def print_cer_detail(rst):
+ return ("(" + "nwords=" + str(rst['nwords']) + ",cor=" + str(rst['cor'])
+ + ",ins=" + str(rst['ins']) + ",del=" + str(rst['del']) + ",sub="
+ + str(rst['sub']) + ") corr:" + '{:.2%}'.format(rst['cor']/rst['nwords'])
+ + ",cer:" + '{:.2%}'.format(rst['wrong']/rst['nwords']))
+
+if __name__ == '__main__':
+ if len(sys.argv) != 4:
+ print("usage : python compute-wer.py test.ref test.hyp test.wer")
+ sys.exit(0)
+
+ ref_file = sys.argv[1]
+ hyp_file = sys.argv[2]
+ cer_detail_file = sys.argv[3]
+ compute_wer(ref_file, hyp_file, cer_detail_file)
diff --git a/egs/aishell/paraformer/utils/error_rate_zh b/egs/aishell/paraformer/utils/error_rate_zh
new file mode 100755
index 0000000..6871a07
--- /dev/null
+++ b/egs/aishell/paraformer/utils/error_rate_zh
@@ -0,0 +1,370 @@
+#!/usr/bin/env python3
+# coding=utf8
+
+# Copyright 2021 Jiayu DU
+
+import sys
+import argparse
+import json
+import logging
+logging.basicConfig(stream=sys.stderr, level=logging.INFO, format='[%(levelname)s] %(message)s')
+
+DEBUG = None
+
+def GetEditType(ref_token, hyp_token):
+ if ref_token == None and hyp_token != None:
+ return 'I'
+ elif ref_token != None and hyp_token == None:
+ return 'D'
+ elif ref_token == hyp_token:
+ return 'C'
+ elif ref_token != hyp_token:
+ return 'S'
+ else:
+ raise RuntimeError
+
+class AlignmentArc:
+ def __init__(self, src, dst, ref, hyp):
+ self.src = src
+ self.dst = dst
+ self.ref = ref
+ self.hyp = hyp
+ self.edit_type = GetEditType(ref, hyp)
+
+def similarity_score_function(ref_token, hyp_token):
+ return 0 if (ref_token == hyp_token) else -1.0
+
+def insertion_score_function(token):
+ return -1.0
+
+def deletion_score_function(token):
+ return -1.0
+
+def EditDistance(
+ ref,
+ hyp,
+ similarity_score_function = similarity_score_function,
+ insertion_score_function = insertion_score_function,
+ deletion_score_function = deletion_score_function):
+ assert(len(ref) != 0)
+ class DPState:
+ def __init__(self):
+ self.score = -float('inf')
+ # backpointer
+ self.prev_r = None
+ self.prev_h = None
+
+ def print_search_grid(S, R, H, fstream):
+ print(file=fstream)
+ for r in range(R):
+ for h in range(H):
+ print(F'[{r},{h}]:{S[r][h].score:4.3f}:({S[r][h].prev_r},{S[r][h].prev_h}) ', end='', file=fstream)
+ print(file=fstream)
+
+ R = len(ref) + 1
+ H = len(hyp) + 1
+
+ # Construct DP search space, a (R x H) grid
+ S = [ [] for r in range(R) ]
+ for r in range(R):
+ S[r] = [ DPState() for x in range(H) ]
+
+ # initialize DP search grid origin, S(r = 0, h = 0)
+ S[0][0].score = 0.0
+ S[0][0].prev_r = None
+ S[0][0].prev_h = None
+
+ # initialize REF axis
+ for r in range(1, R):
+ S[r][0].score = S[r-1][0].score + deletion_score_function(ref[r-1])
+ S[r][0].prev_r = r-1
+ S[r][0].prev_h = 0
+
+ # initialize HYP axis
+ for h in range(1, H):
+ S[0][h].score = S[0][h-1].score + insertion_score_function(hyp[h-1])
+ S[0][h].prev_r = 0
+ S[0][h].prev_h = h-1
+
+ best_score = S[0][0].score
+ best_state = (0, 0)
+
+ for r in range(1, R):
+ for h in range(1, H):
+ sub_or_cor_score = similarity_score_function(ref[r-1], hyp[h-1])
+ new_score = S[r-1][h-1].score + sub_or_cor_score
+ if new_score >= S[r][h].score:
+ S[r][h].score = new_score
+ S[r][h].prev_r = r-1
+ S[r][h].prev_h = h-1
+
+ del_score = deletion_score_function(ref[r-1])
+ new_score = S[r-1][h].score + del_score
+ if new_score >= S[r][h].score:
+ S[r][h].score = new_score
+ S[r][h].prev_r = r - 1
+ S[r][h].prev_h = h
+
+ ins_score = insertion_score_function(hyp[h-1])
+ new_score = S[r][h-1].score + ins_score
+ if new_score >= S[r][h].score:
+ S[r][h].score = new_score
+ S[r][h].prev_r = r
+ S[r][h].prev_h = h-1
+
+ best_score = S[R-1][H-1].score
+ best_state = (R-1, H-1)
+
+ if DEBUG:
+ print_search_grid(S, R, H, sys.stderr)
+
+ # Backtracing best alignment path, i.e. a list of arcs
+ # arc = (src, dst, ref, hyp, edit_type)
+ # src/dst = (r, h), where r/h refers to search grid state-id along Ref/Hyp axis
+ best_path = []
+ r, h = best_state[0], best_state[1]
+ prev_r, prev_h = S[r][h].prev_r, S[r][h].prev_h
+ score = S[r][h].score
+ # loop invariant:
+ # 1. (prev_r, prev_h) -> (r, h) is a "forward arc" on best alignment path
+ # 2. score is the value of point(r, h) on DP search grid
+ while prev_r != None or prev_h != None:
+ src = (prev_r, prev_h)
+ dst = (r, h)
+ if (r == prev_r + 1 and h == prev_h + 1): # Substitution or correct
+ arc = AlignmentArc(src, dst, ref[prev_r], hyp[prev_h])
+ elif (r == prev_r + 1 and h == prev_h): # Deletion
+ arc = AlignmentArc(src, dst, ref[prev_r], None)
+ elif (r == prev_r and h == prev_h + 1): # Insertion
+ arc = AlignmentArc(src, dst, None, hyp[prev_h])
+ else:
+ raise RuntimeError
+ best_path.append(arc)
+ r, h = prev_r, prev_h
+ prev_r, prev_h = S[r][h].prev_r, S[r][h].prev_h
+ score = S[r][h].score
+
+ best_path.reverse()
+ return (best_path, best_score)
+
+def PrettyPrintAlignment(alignment, stream = sys.stderr):
+ def get_token_str(token):
+ if token == None:
+ return "*"
+ return token
+
+ def is_double_width_char(ch):
+ if (ch >= '\u4e00') and (ch <= '\u9fa5'): # codepoint ranges for Chinese chars
+ return True
+ # TODO: support other double-width-char language such as Japanese, Korean
+ else:
+ return False
+
+ def display_width(token_str):
+ m = 0
+ for c in token_str:
+ if is_double_width_char(c):
+ m += 2
+ else:
+ m += 1
+ return m
+
+ R = ' REF : '
+ H = ' HYP : '
+ E = ' EDIT : '
+ for arc in alignment:
+ r = get_token_str(arc.ref)
+ h = get_token_str(arc.hyp)
+ e = arc.edit_type if arc.edit_type != 'C' else ''
+
+ nr, nh, ne = display_width(r), display_width(h), display_width(e)
+ n = max(nr, nh, ne) + 1
+
+ R += r + ' ' * (n-nr)
+ H += h + ' ' * (n-nh)
+ E += e + ' ' * (n-ne)
+
+ print(R, file=stream)
+ print(H, file=stream)
+ print(E, file=stream)
+
+def CountEdits(alignment):
+ c, s, i, d = 0, 0, 0, 0
+ for arc in alignment:
+ if arc.edit_type == 'C':
+ c += 1
+ elif arc.edit_type == 'S':
+ s += 1
+ elif arc.edit_type == 'I':
+ i += 1
+ elif arc.edit_type == 'D':
+ d += 1
+ else:
+ raise RuntimeError
+ return (c, s, i, d)
+
+def ComputeTokenErrorRate(c, s, i, d):
+ return 100.0 * (s + d + i) / (s + d + c)
+
+def ComputeSentenceErrorRate(num_err_utts, num_utts):
+ assert(num_utts != 0)
+ return 100.0 * num_err_utts / num_utts
+
+
+class EvaluationResult:
+ def __init__(self):
+ self.num_ref_utts = 0
+ self.num_hyp_utts = 0
+ self.num_eval_utts = 0 # seen in both ref & hyp
+ self.num_hyp_without_ref = 0
+
+ self.C = 0
+ self.S = 0
+ self.I = 0
+ self.D = 0
+ self.token_error_rate = 0.0
+
+ self.num_utts_with_error = 0
+ self.sentence_error_rate = 0.0
+
+ def to_json(self):
+ return json.dumps(self.__dict__)
+
+ def to_kaldi(self):
+ info = (
+ F'%WER {self.token_error_rate:.2f} [ {self.S + self.D + self.I} / {self.C + self.S + self.D}, {self.I} ins, {self.D} del, {self.S} sub ]\n'
+ F'%SER {self.sentence_error_rate:.2f} [ {self.num_utts_with_error} / {self.num_eval_utts} ]\n'
+ )
+ return info
+
+ def to_sclite(self):
+ return "TODO"
+
+ def to_espnet(self):
+ return "TODO"
+
+ def to_summary(self):
+ #return json.dumps(self.__dict__, indent=4)
+ summary = (
+ '==================== Overall Statistics ====================\n'
+ F'num_ref_utts: {self.num_ref_utts}\n'
+ F'num_hyp_utts: {self.num_hyp_utts}\n'
+ F'num_hyp_without_ref: {self.num_hyp_without_ref}\n'
+ F'num_eval_utts: {self.num_eval_utts}\n'
+ F'sentence_error_rate: {self.sentence_error_rate:.2f}%\n'
+ F'token_error_rate: {self.token_error_rate:.2f}%\n'
+ F'token_stats:\n'
+ F' - tokens:{self.C + self.S + self.D:>7}\n'
+ F' - edits: {self.S + self.I + self.D:>7}\n'
+ F' - cor: {self.C:>7}\n'
+ F' - sub: {self.S:>7}\n'
+ F' - ins: {self.I:>7}\n'
+ F' - del: {self.D:>7}\n'
+ '============================================================\n'
+ )
+ return summary
+
+
+class Utterance:
+ def __init__(self, uid, text):
+ self.uid = uid
+ self.text = text
+
+
+def LoadUtterances(filepath, format):
+ utts = {}
+ if format == 'text': # utt_id word1 word2 ...
+ with open(filepath, 'r', encoding='utf8') as f:
+ for line in f:
+ line = line.strip()
+ if line:
+ cols = line.split(maxsplit=1)
+ assert(len(cols) == 2 or len(cols) == 1)
+ uid = cols[0]
+ text = cols[1] if len(cols) == 2 else ''
+ if utts.get(uid) != None:
+ raise RuntimeError(F'Found duplicated utterence id {uid}')
+ utts[uid] = Utterance(uid, text)
+ else:
+ raise RuntimeError(F'Unsupported text format {format}')
+ return utts
+
+
+def tokenize_text(text, tokenizer):
+ if tokenizer == 'whitespace':
+ return text.split()
+ elif tokenizer == 'char':
+ return [ ch for ch in ''.join(text.split()) ]
+ else:
+ raise RuntimeError(F'ERROR: Unsupported tokenizer {tokenizer}')
+
+
+if __name__ == '__main__':
+ parser = argparse.ArgumentParser()
+ # optional
+ parser.add_argument('--tokenizer', choices=['whitespace', 'char'], default='whitespace', help='whitespace for WER, char for CER')
+ parser.add_argument('--ref-format', choices=['text'], default='text', help='reference format, first col is utt_id, the rest is text')
+ parser.add_argument('--hyp-format', choices=['text'], default='text', help='hypothesis format, first col is utt_id, the rest is text')
+ # required
+ parser.add_argument('--ref', type=str, required=True, help='input reference file')
+ parser.add_argument('--hyp', type=str, required=True, help='input hypothesis file')
+
+ parser.add_argument('result_file', type=str)
+ args = parser.parse_args()
+ logging.info(args)
+
+ ref_utts = LoadUtterances(args.ref, args.ref_format)
+ hyp_utts = LoadUtterances(args.hyp, args.hyp_format)
+
+ r = EvaluationResult()
+
+ # check valid utterances in hyp that have matched non-empty reference
+ eval_utts = []
+ r.num_hyp_without_ref = 0
+ for uid in sorted(hyp_utts.keys()):
+ if uid in ref_utts.keys(): # TODO: efficiency
+ if ref_utts[uid].text.strip(): # non-empty reference
+ eval_utts.append(uid)
+ else:
+ logging.warn(F'Found {uid} with empty reference, skipping...')
+ else:
+ logging.warn(F'Found {uid} without reference, skipping...')
+ r.num_hyp_without_ref += 1
+
+ r.num_hyp_utts = len(hyp_utts)
+ r.num_ref_utts = len(ref_utts)
+ r.num_eval_utts = len(eval_utts)
+
+ with open(args.result_file, 'w+', encoding='utf8') as fo:
+ for uid in eval_utts:
+ ref = ref_utts[uid]
+ hyp = hyp_utts[uid]
+
+ alignment, score = EditDistance(
+ tokenize_text(ref.text, args.tokenizer),
+ tokenize_text(hyp.text, args.tokenizer)
+ )
+
+ c, s, i, d = CountEdits(alignment)
+ utt_ter = ComputeTokenErrorRate(c, s, i, d)
+
+ # utt-level evaluation result
+ print(F'{{"uid":{uid}, "score":{score}, "ter":{utt_ter:.2f}, "cor":{c}, "sub":{s}, "ins":{i}, "del":{d}}}', file=fo)
+ PrettyPrintAlignment(alignment, fo)
+
+ r.C += c
+ r.S += s
+ r.I += i
+ r.D += d
+
+ if utt_ter > 0:
+ r.num_utts_with_error += 1
+
+ # corpus level evaluation result
+ r.sentence_error_rate = ComputeSentenceErrorRate(r.num_utts_with_error, r.num_eval_utts)
+ r.token_error_rate = ComputeTokenErrorRate(r.C, r.S, r.I, r.D)
+
+ print(r.to_summary(), file=fo)
+
+ print(r.to_json())
+ print(r.to_kaldi())
diff --git a/egs/aishell/paraformer/utils/extract_embeds.py b/egs/aishell/paraformer/utils/extract_embeds.py
new file mode 100755
index 0000000..7b817d8
--- /dev/null
+++ b/egs/aishell/paraformer/utils/extract_embeds.py
@@ -0,0 +1,47 @@
+from transformers import AutoTokenizer, AutoModel, pipeline
+import numpy as np
+import sys
+import os
+import torch
+from kaldiio import WriteHelper
+import re
+text_file_json = sys.argv[1]
+out_ark = sys.argv[2]
+out_scp = sys.argv[3]
+out_shape = sys.argv[4]
+device = int(sys.argv[5])
+model_path = sys.argv[6]
+
+model = AutoModel.from_pretrained(model_path)
+tokenizer = AutoTokenizer.from_pretrained(model_path)
+extractor = pipeline(task="feature-extraction", model=model, tokenizer=tokenizer, device=device)
+
+with open(text_file_json, 'r') as f:
+ js = f.readlines()
+
+
+f_shape = open(out_shape, "w")
+with WriteHelper('ark,scp:{},{}'.format(out_ark, out_scp)) as writer:
+ with torch.no_grad():
+ for idx, line in enumerate(js):
+ id, tokens = line.strip().split(" ", 1)
+ tokens = re.sub(" ", "", tokens.strip())
+ tokens = ' '.join([j for j in tokens])
+ token_num = len(tokens.split(" "))
+ outputs = extractor(tokens)
+ outputs = np.array(outputs)
+ embeds = outputs[0, 1:-1, :]
+
+ token_num_embeds, dim = embeds.shape
+ if token_num == token_num_embeds:
+ writer(id, embeds)
+ shape_line = "{} {},{}\n".format(id, token_num_embeds, dim)
+ f_shape.write(shape_line)
+ else:
+ print("{}, size has changed, {}, {}, {}".format(id, token_num, token_num_embeds, tokens))
+
+
+
+f_shape.close()
+
+
diff --git a/egs/aishell/paraformer/utils/filter_scp.pl b/egs/aishell/paraformer/utils/filter_scp.pl
new file mode 100755
index 0000000..003530d
--- /dev/null
+++ b/egs/aishell/paraformer/utils/filter_scp.pl
@@ -0,0 +1,87 @@
+#!/usr/bin/env perl
+# Copyright 2010-2012 Microsoft Corporation
+# Johns Hopkins University (author: Daniel Povey)
+
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+# This script takes a list of utterance-ids or any file whose first field
+# of each line is an utterance-id, and filters an scp
+# file (or any file whose "n-th" field is an utterance id), printing
+# out only those lines whose "n-th" field is in id_list. The index of
+# the "n-th" field is 1, by default, but can be changed by using
+# the -f <n> switch
+
+$exclude = 0;
+$field = 1;
+$shifted = 0;
+
+do {
+ $shifted=0;
+ if ($ARGV[0] eq "--exclude") {
+ $exclude = 1;
+ shift @ARGV;
+ $shifted=1;
+ }
+ if ($ARGV[0] eq "-f") {
+ $field = $ARGV[1];
+ shift @ARGV; shift @ARGV;
+ $shifted=1
+ }
+} while ($shifted);
+
+if(@ARGV < 1 || @ARGV > 2) {
+ die "Usage: filter_scp.pl [--exclude] [-f <field-to-filter-on>] id_list [in.scp] > out.scp \n" .
+ "Prints only the input lines whose f'th field (default: first) is in 'id_list'.\n" .
+ "Note: only the first field of each line in id_list matters. With --exclude, prints\n" .
+ "only the lines that were *not* in id_list.\n" .
+ "Caution: previously, the -f option was interpreted as a zero-based field index.\n" .
+ "If your older scripts (written before Oct 2014) stopped working and you used the\n" .
+ "-f option, add 1 to the argument.\n" .
+ "See also: scripts/filter_scp.pl .\n";
+}
+
+
+$idlist = shift @ARGV;
+open(F, "<$idlist") || die "Could not open id-list file $idlist";
+while(<F>) {
+ @A = split;
+ @A>=1 || die "Invalid id-list file line $_";
+ $seen{$A[0]} = 1;
+}
+
+if ($field == 1) { # Treat this as special case, since it is common.
+ while(<>) {
+ $_ =~ m/\s*(\S+)\s*/ || die "Bad line $_, could not get first field.";
+ # $1 is what we filter on.
+ if ((!$exclude && $seen{$1}) || ($exclude && !defined $seen{$1})) {
+ print $_;
+ }
+ }
+} else {
+ while(<>) {
+ @A = split;
+ @A > 0 || die "Invalid scp file line $_";
+ @A >= $field || die "Invalid scp file line $_";
+ if ((!$exclude && $seen{$A[$field-1]}) || ($exclude && !defined $seen{$A[$field-1]})) {
+ print $_;
+ }
+ }
+}
+
+# tests:
+# the following should print "foo 1"
+# ( echo foo 1; echo bar 2 ) | scripts/filter_scp.pl <(echo foo)
+# the following should print "bar 2".
+# ( echo foo 1; echo bar 2 ) | scripts/filter_scp.pl -f 2 <(echo 2)
diff --git a/egs/aishell/paraformer/utils/fix_data.sh b/egs/aishell/paraformer/utils/fix_data.sh
new file mode 100755
index 0000000..32cdde5
--- /dev/null
+++ b/egs/aishell/paraformer/utils/fix_data.sh
@@ -0,0 +1,35 @@
+#!/usr/bin/env bash
+
+echo "$0 $@"
+data_dir=$1
+
+if [ ! -f ${data_dir}/wav.scp ]; then
+ echo "$0: wav.scp is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/text ]; then
+ echo "$0: text is not found"
+ exit 1;
+fi
+
+
+
+mkdir -p ${data_dir}/.backup
+
+awk '{print $1}' ${data_dir}/wav.scp > ${data_dir}/.backup/wav_id
+awk '{print $1}' ${data_dir}/text > ${data_dir}/.backup/text_id
+
+sort ${data_dir}/.backup/wav_id ${data_dir}/.backup/text_id | uniq -d > ${data_dir}/.backup/id
+
+cp ${data_dir}/wav.scp ${data_dir}/.backup/wav.scp
+cp ${data_dir}/text ${data_dir}/.backup/text
+
+mv ${data_dir}/wav.scp ${data_dir}/wav.scp.bak
+mv ${data_dir}/text ${data_dir}/text.bak
+
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/wav.scp.bak > ${data_dir}/wav.scp
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/text.bak > ${data_dir}/text
+
+rm ${data_dir}/wav.scp.bak
+rm ${data_dir}/text.bak
diff --git a/egs/aishell/paraformer/utils/fix_data_feat.sh b/egs/aishell/paraformer/utils/fix_data_feat.sh
new file mode 100755
index 0000000..2c92d7f
--- /dev/null
+++ b/egs/aishell/paraformer/utils/fix_data_feat.sh
@@ -0,0 +1,52 @@
+#!/usr/bin/env bash
+
+echo "$0 $@"
+data_dir=$1
+
+if [ ! -f ${data_dir}/feats.scp ]; then
+ echo "$0: feats.scp is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/text ]; then
+ echo "$0: text is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/speech_shape ]; then
+ echo "$0: feature lengths is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/text_shape ]; then
+ echo "$0: text lengths is not found"
+ exit 1;
+fi
+
+mkdir -p ${data_dir}/.backup
+
+awk '{print $1}' ${data_dir}/feats.scp > ${data_dir}/.backup/wav_id
+awk '{print $1}' ${data_dir}/text > ${data_dir}/.backup/text_id
+
+sort ${data_dir}/.backup/wav_id ${data_dir}/.backup/text_id | uniq -d > ${data_dir}/.backup/id
+
+cp ${data_dir}/feats.scp ${data_dir}/.backup/feats.scp
+cp ${data_dir}/text ${data_dir}/.backup/text
+cp ${data_dir}/speech_shape ${data_dir}/.backup/speech_shape
+cp ${data_dir}/text_shape ${data_dir}/.backup/text_shape
+
+mv ${data_dir}/feats.scp ${data_dir}/feats.scp.bak
+mv ${data_dir}/text ${data_dir}/text.bak
+mv ${data_dir}/speech_shape ${data_dir}/speech_shape.bak
+mv ${data_dir}/text_shape ${data_dir}/text_shape.bak
+
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/feats.scp.bak > ${data_dir}/feats.scp
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/text.bak > ${data_dir}/text
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/speech_shape.bak > ${data_dir}/speech_shape
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/text_shape.bak > ${data_dir}/text_shape
+
+rm ${data_dir}/feats.scp.bak
+rm ${data_dir}/text.bak
+rm ${data_dir}/speech_shape.bak
+rm ${data_dir}/text_shape.bak
+
diff --git a/egs/aishell/paraformer/utils/gen_ark_list.sh b/egs/aishell/paraformer/utils/gen_ark_list.sh
new file mode 100755
index 0000000..aebf356
--- /dev/null
+++ b/egs/aishell/paraformer/utils/gen_ark_list.sh
@@ -0,0 +1,22 @@
+#!/usr/bin/env bash
+
+
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+ark_dir=$1
+txt_dir=$2
+output_dir=$3
+
+[ ! -d ${ark_dir}/ark ] && echo "$0: ark data is required" && exit 1;
+[ ! -d ${txt_dir}/txt ] && echo "$0: txt data is required" && exit 1;
+
+for n in $(seq $nj); do
+ echo "${ark_dir}/ark/feats.$n.ark ${txt_dir}/txt/text.$n.txt" || exit 1
+done > ${output_dir}/ark_txt.scp || exit 1
+
diff --git a/egs/aishell/paraformer/utils/parse_options.sh b/egs/aishell/paraformer/utils/parse_options.sh
new file mode 100755
index 0000000..71fb9e5
--- /dev/null
+++ b/egs/aishell/paraformer/utils/parse_options.sh
@@ -0,0 +1,97 @@
+#!/usr/bin/env bash
+
+# Copyright 2012 Johns Hopkins University (Author: Daniel Povey);
+# Arnab Ghoshal, Karel Vesely
+
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+# Parse command-line options.
+# To be sourced by another script (as in ". parse_options.sh").
+# Option format is: --option-name arg
+# and shell variable "option_name" gets set to value "arg."
+# The exception is --help, which takes no arguments, but prints the
+# $help_message variable (if defined).
+
+
+###
+### The --config file options have lower priority to command line
+### options, so we need to import them first...
+###
+
+# Now import all the configs specified by command-line, in left-to-right order
+for ((argpos=1; argpos<$#; argpos++)); do
+ if [ "${!argpos}" == "--config" ]; then
+ argpos_plus1=$((argpos+1))
+ config=${!argpos_plus1}
+ [ ! -r $config ] && echo "$0: missing config '$config'" && exit 1
+ . $config # source the config file.
+ fi
+done
+
+
+###
+### Now we process the command line options
+###
+while true; do
+ [ -z "${1:-}" ] && break; # break if there are no arguments
+ case "$1" in
+ # If the enclosing script is called with --help option, print the help
+ # message and exit. Scripts should put help messages in $help_message
+ --help|-h) if [ -z "$help_message" ]; then echo "No help found." 1>&2;
+ else printf "$help_message\n" 1>&2 ; fi;
+ exit 0 ;;
+ --*=*) echo "$0: options to scripts must be of the form --name value, got '$1'"
+ exit 1 ;;
+ # If the first command-line argument begins with "--" (e.g. --foo-bar),
+ # then work out the variable name as $name, which will equal "foo_bar".
+ --*) name=`echo "$1" | sed s/^--// | sed s/-/_/g`;
+ # Next we test whether the variable in question is undefned-- if so it's
+ # an invalid option and we die. Note: $0 evaluates to the name of the
+ # enclosing script.
+ # The test [ -z ${foo_bar+xxx} ] will return true if the variable foo_bar
+ # is undefined. We then have to wrap this test inside "eval" because
+ # foo_bar is itself inside a variable ($name).
+ eval '[ -z "${'$name'+xxx}" ]' && echo "$0: invalid option $1" 1>&2 && exit 1;
+
+ oldval="`eval echo \\$$name`";
+ # Work out whether we seem to be expecting a Boolean argument.
+ if [ "$oldval" == "true" ] || [ "$oldval" == "false" ]; then
+ was_bool=true;
+ else
+ was_bool=false;
+ fi
+
+ # Set the variable to the right value-- the escaped quotes make it work if
+ # the option had spaces, like --cmd "queue.pl -sync y"
+ eval $name=\"$2\";
+
+ # Check that Boolean-valued arguments are really Boolean.
+ if $was_bool && [[ "$2" != "true" && "$2" != "false" ]]; then
+ echo "$0: expected \"true\" or \"false\": $1 $2" 1>&2
+ exit 1;
+ fi
+ shift 2;
+ ;;
+ *) break;
+ esac
+done
+
+
+# Check for an empty argument to the --cmd option, which can easily occur as a
+# result of scripting errors.
+[ ! -z "${cmd+xxx}" ] && [ -z "$cmd" ] && echo "$0: empty argument to --cmd option" 1>&2 && exit 1;
+
+
+true; # so this script returns exit code 0.
diff --git a/egs/aishell/paraformer/utils/print_args.py b/egs/aishell/paraformer/utils/print_args.py
new file mode 100755
index 0000000..b0c61e5
--- /dev/null
+++ b/egs/aishell/paraformer/utils/print_args.py
@@ -0,0 +1,45 @@
+#!/usr/bin/env python
+import sys
+
+
+def get_commandline_args(no_executable=True):
+ extra_chars = [
+ " ",
+ ";",
+ "&",
+ "|",
+ "<",
+ ">",
+ "?",
+ "*",
+ "~",
+ "`",
+ '"',
+ "'",
+ "\\",
+ "{",
+ "}",
+ "(",
+ ")",
+ ]
+
+ # Escape the extra characters for shell
+ argv = [
+ arg.replace("'", "'\\''")
+ if all(char not in arg for char in extra_chars)
+ else "'" + arg.replace("'", "'\\''") + "'"
+ for arg in sys.argv
+ ]
+
+ if no_executable:
+ return " ".join(argv[1:])
+ else:
+ return sys.executable + " " + " ".join(argv)
+
+
+def main():
+ print(get_commandline_args())
+
+
+if __name__ == "__main__":
+ main()
diff --git a/egs/aishell/paraformer/utils/proc_conf_oss.py b/egs/aishell/paraformer/utils/proc_conf_oss.py
new file mode 100755
index 0000000..c4a90c5
--- /dev/null
+++ b/egs/aishell/paraformer/utils/proc_conf_oss.py
@@ -0,0 +1,35 @@
+from pathlib import Path
+
+import torch
+import yaml
+
+
+class NoAliasSafeDumper(yaml.SafeDumper):
+ # Disable anchor/alias in yaml because looks ugly
+ def ignore_aliases(self, data):
+ return True
+
+
+def yaml_no_alias_safe_dump(data, stream=None, **kwargs):
+ """Safe-dump in yaml with no anchor/alias"""
+ return yaml.dump(
+ data, stream, allow_unicode=True, Dumper=NoAliasSafeDumper, **kwargs
+ )
+
+
+def gen_conf(file, out_dir):
+ conf = torch.load(file)["config"]
+ conf["oss_bucket"] = "null"
+ print(conf)
+ output_dir = Path(out_dir)
+ output_dir.mkdir(parents=True, exist_ok=True)
+ with (output_dir / "config.yaml").open("w", encoding="utf-8") as f:
+ yaml_no_alias_safe_dump(conf, f, indent=4, sort_keys=False)
+
+
+if __name__ == "__main__":
+ import sys
+
+ in_f = sys.argv[1]
+ out_f = sys.argv[2]
+ gen_conf(in_f, out_f)
diff --git a/egs/aishell/paraformer/utils/proce_text.py b/egs/aishell/paraformer/utils/proce_text.py
new file mode 100755
index 0000000..9e517a4
--- /dev/null
+++ b/egs/aishell/paraformer/utils/proce_text.py
@@ -0,0 +1,31 @@
+
+import sys
+import re
+
+in_f = sys.argv[1]
+out_f = sys.argv[2]
+
+
+with open(in_f, "r", encoding="utf-8") as f:
+ lines = f.readlines()
+
+with open(out_f, "w", encoding="utf-8") as f:
+ for line in lines:
+ outs = line.strip().split(" ", 1)
+ if len(outs) == 2:
+ idx, text = outs
+ text = re.sub("</s>", "", text)
+ text = re.sub("<s>", "", text)
+ text = re.sub("@@", "", text)
+ text = re.sub("@", "", text)
+ text = re.sub("<unk>", "", text)
+ text = re.sub(" ", "", text)
+ text = text.lower()
+ else:
+ idx = outs[0]
+ text = " "
+
+ text = [x for x in text]
+ text = " ".join(text)
+ out = "{} {}\n".format(idx, text)
+ f.write(out)
diff --git a/egs/aishell/paraformer/utils/run.pl b/egs/aishell/paraformer/utils/run.pl
new file mode 100755
index 0000000..483f95b
--- /dev/null
+++ b/egs/aishell/paraformer/utils/run.pl
@@ -0,0 +1,356 @@
+#!/usr/bin/env perl
+use warnings; #sed replacement for -w perl parameter
+# In general, doing
+# run.pl some.log a b c is like running the command a b c in
+# the bash shell, and putting the standard error and output into some.log.
+# To run parallel jobs (backgrounded on the host machine), you can do (e.g.)
+# run.pl JOB=1:4 some.JOB.log a b c JOB is like running the command a b c JOB
+# and putting it in some.JOB.log, for each one. [Note: JOB can be any identifier].
+# If any of the jobs fails, this script will fail.
+
+# A typical example is:
+# run.pl some.log my-prog "--opt=foo bar" foo \| other-prog baz
+# and run.pl will run something like:
+# ( my-prog '--opt=foo bar' foo | other-prog baz ) >& some.log
+#
+# Basically it takes the command-line arguments, quotes them
+# as necessary to preserve spaces, and evaluates them with bash.
+# In addition it puts the command line at the top of the log, and
+# the start and end times of the command at the beginning and end.
+# The reason why this is useful is so that we can create a different
+# version of this program that uses a queueing system instead.
+
+#use Data::Dumper;
+
+@ARGV < 2 && die "usage: run.pl log-file command-line arguments...";
+
+#print STDERR "COMMAND-LINE: " . Dumper(\@ARGV) . "\n";
+$job_pick = 'all';
+$max_jobs_run = -1;
+$jobstart = 1;
+$jobend = 1;
+$ignored_opts = ""; # These will be ignored.
+
+# First parse an option like JOB=1:4, and any
+# options that would normally be given to
+# queue.pl, which we will just discard.
+
+for (my $x = 1; $x <= 2; $x++) { # This for-loop is to
+ # allow the JOB=1:n option to be interleaved with the
+ # options to qsub.
+ while (@ARGV >= 2 && $ARGV[0] =~ m:^-:) {
+ # parse any options that would normally go to qsub, but which will be ignored here.
+ my $switch = shift @ARGV;
+ if ($switch eq "-V") {
+ $ignored_opts .= "-V ";
+ } elsif ($switch eq "--max-jobs-run" || $switch eq "-tc") {
+ # we do support the option --max-jobs-run n, and its GridEngine form -tc n.
+ # if the command appears multiple times uses the smallest option.
+ if ( $max_jobs_run <= 0 ) {
+ $max_jobs_run = shift @ARGV;
+ } else {
+ my $new_constraint = shift @ARGV;
+ if ( ($new_constraint < $max_jobs_run) ) {
+ $max_jobs_run = $new_constraint;
+ }
+ }
+
+ if (! ($max_jobs_run > 0)) {
+ die "run.pl: invalid option --max-jobs-run $max_jobs_run";
+ }
+ } else {
+ my $argument = shift @ARGV;
+ if ($argument =~ m/^--/) {
+ print STDERR "run.pl: WARNING: suspicious argument '$argument' to $switch; starts with '-'\n";
+ }
+ if ($switch eq "-sync" && $argument =~ m/^[yY]/) {
+ $ignored_opts .= "-sync "; # Note: in the
+ # corresponding code in queue.pl it says instead, just "$sync = 1;".
+ } elsif ($switch eq "-pe") { # e.g. -pe smp 5
+ my $argument2 = shift @ARGV;
+ $ignored_opts .= "$switch $argument $argument2 ";
+ } elsif ($switch eq "--gpu") {
+ $using_gpu = $argument;
+ } elsif ($switch eq "--pick") {
+ if($argument =~ m/^(all|failed|incomplete)$/) {
+ $job_pick = $argument;
+ } else {
+ print STDERR "run.pl: ERROR: --pick argument must be one of 'all', 'failed' or 'incomplete'"
+ }
+ } else {
+ # Ignore option.
+ $ignored_opts .= "$switch $argument ";
+ }
+ }
+ }
+ if ($ARGV[0] =~ m/^([\w_][\w\d_]*)+=(\d+):(\d+)$/) { # e.g. JOB=1:20
+ $jobname = $1;
+ $jobstart = $2;
+ $jobend = $3;
+ if ($jobstart > $jobend) {
+ die "run.pl: invalid job range $ARGV[0]";
+ }
+ if ($jobstart <= 0) {
+ die "run.pl: invalid job range $ARGV[0], start must be strictly positive (this is required for GridEngine compatibility).";
+ }
+ shift;
+ } elsif ($ARGV[0] =~ m/^([\w_][\w\d_]*)+=(\d+)$/) { # e.g. JOB=1.
+ $jobname = $1;
+ $jobstart = $2;
+ $jobend = $2;
+ shift;
+ } elsif ($ARGV[0] =~ m/.+\=.*\:.*$/) {
+ print STDERR "run.pl: Warning: suspicious first argument to run.pl: $ARGV[0]\n";
+ }
+}
+
+# Users found this message confusing so we are removing it.
+# if ($ignored_opts ne "") {
+# print STDERR "run.pl: Warning: ignoring options \"$ignored_opts\"\n";
+# }
+
+if ($max_jobs_run == -1) { # If --max-jobs-run option not set,
+ # then work out the number of processors if possible,
+ # and set it based on that.
+ $max_jobs_run = 0;
+ if ($using_gpu) {
+ if (open(P, "nvidia-smi -L |")) {
+ $max_jobs_run++ while (<P>);
+ close(P);
+ }
+ if ($max_jobs_run == 0) {
+ $max_jobs_run = 1;
+ print STDERR "run.pl: Warning: failed to detect number of GPUs from nvidia-smi, using ${max_jobs_run}\n";
+ }
+ } elsif (open(P, "</proc/cpuinfo")) { # Linux
+ while (<P>) { if (m/^processor/) { $max_jobs_run++; } }
+ if ($max_jobs_run == 0) {
+ print STDERR "run.pl: Warning: failed to detect any processors from /proc/cpuinfo\n";
+ $max_jobs_run = 10; # reasonable default.
+ }
+ close(P);
+ } elsif (open(P, "sysctl -a |")) { # BSD/Darwin
+ while (<P>) {
+ if (m/hw\.ncpu\s*[:=]\s*(\d+)/) { # hw.ncpu = 4, or hw.ncpu: 4
+ $max_jobs_run = $1;
+ last;
+ }
+ }
+ close(P);
+ if ($max_jobs_run == 0) {
+ print STDERR "run.pl: Warning: failed to detect any processors from sysctl -a\n";
+ $max_jobs_run = 10; # reasonable default.
+ }
+ } else {
+ # allow at most 32 jobs at once, on non-UNIX systems; change this code
+ # if you need to change this default.
+ $max_jobs_run = 32;
+ }
+ # The just-computed value of $max_jobs_run is just the number of processors
+ # (or our best guess); and if it happens that the number of jobs we need to
+ # run is just slightly above $max_jobs_run, it will make sense to increase
+ # $max_jobs_run to equal the number of jobs, so we don't have a small number
+ # of leftover jobs.
+ $num_jobs = $jobend - $jobstart + 1;
+ if (!$using_gpu &&
+ $num_jobs > $max_jobs_run && $num_jobs < 1.4 * $max_jobs_run) {
+ $max_jobs_run = $num_jobs;
+ }
+}
+
+sub pick_or_exit {
+ # pick_or_exit ( $logfile )
+ # Invoked before each job is started helps to run jobs selectively.
+ #
+ # Given the name of the output logfile decides whether the job must be
+ # executed (by returning from the subroutine) or not (by terminating the
+ # process calling exit)
+ #
+ # PRE: $job_pick is a global variable set by command line switch --pick
+ # and indicates which class of jobs must be executed.
+ #
+ # 1) If a failed job is not executed the process exit code will indicate
+ # failure, just as if the task was just executed and failed.
+ #
+ # 2) If a task is incomplete it will be executed. Incomplete may be either
+ # a job whose log file does not contain the accounting notes in the end,
+ # or a job whose log file does not exist.
+ #
+ # 3) If the $job_pick is set to 'all' (default behavior) a task will be
+ # executed regardless of the result of previous attempts.
+ #
+ # This logic could have been implemented in the main execution loop
+ # but a subroutine to preserve the current level of readability of
+ # that part of the code.
+ #
+ # Alexandre Felipe, (o.alexandre.felipe@gmail.com) 14th of August of 2020
+ #
+ if($job_pick eq 'all'){
+ return; # no need to bother with the previous log
+ }
+ open my $fh, "<", $_[0] or return; # job not executed yet
+ my $log_line;
+ my $cur_line;
+ while ($cur_line = <$fh>) {
+ if( $cur_line =~ m/# Ended \(code .*/ ) {
+ $log_line = $cur_line;
+ }
+ }
+ close $fh;
+ if (! defined($log_line)){
+ return; # incomplete
+ }
+ if ( $log_line =~ m/# Ended \(code 0\).*/ ) {
+ exit(0); # complete
+ } elsif ( $log_line =~ m/# Ended \(code \d+(; signal \d+)?\).*/ ){
+ if ($job_pick !~ m/^(failed|all)$/) {
+ exit(1); # failed but not going to run
+ } else {
+ return; # failed
+ }
+ } elsif ( $log_line =~ m/.*\S.*/ ) {
+ return; # incomplete jobs are always run
+ }
+}
+
+
+$logfile = shift @ARGV;
+
+if (defined $jobname && $logfile !~ m/$jobname/ &&
+ $jobend > $jobstart) {
+ print STDERR "run.pl: you are trying to run a parallel job but "
+ . "you are putting the output into just one log file ($logfile)\n";
+ exit(1);
+}
+
+$cmd = "";
+
+foreach $x (@ARGV) {
+ if ($x =~ m/^\S+$/) { $cmd .= $x . " "; }
+ elsif ($x =~ m:\":) { $cmd .= "'$x' "; }
+ else { $cmd .= "\"$x\" "; }
+}
+
+#$Data::Dumper::Indent=0;
+$ret = 0;
+$numfail = 0;
+%active_pids=();
+
+use POSIX ":sys_wait_h";
+for ($jobid = $jobstart; $jobid <= $jobend; $jobid++) {
+ if (scalar(keys %active_pids) >= $max_jobs_run) {
+
+ # Lets wait for a change in any child's status
+ # Then we have to work out which child finished
+ $r = waitpid(-1, 0);
+ $code = $?;
+ if ($r < 0 ) { die "run.pl: Error waiting for child process"; } # should never happen.
+ if ( defined $active_pids{$r} ) {
+ $jid=$active_pids{$r};
+ $fail[$jid]=$code;
+ if ($code !=0) { $numfail++;}
+ delete $active_pids{$r};
+ # print STDERR "Finished: $r/$jid " . Dumper(\%active_pids) . "\n";
+ } else {
+ die "run.pl: Cannot find the PID of the child process that just finished.";
+ }
+
+ # In theory we could do a non-blocking waitpid over all jobs running just
+ # to find out if only one or more jobs finished during the previous waitpid()
+ # However, we just omit this and will reap the next one in the next pass
+ # through the for(;;) cycle
+ }
+ $childpid = fork();
+ if (!defined $childpid) { die "run.pl: Error forking in run.pl (writing to $logfile)"; }
+ if ($childpid == 0) { # We're in the child... this branch
+ # executes the job and returns (possibly with an error status).
+ if (defined $jobname) {
+ $cmd =~ s/$jobname/$jobid/g;
+ $logfile =~ s/$jobname/$jobid/g;
+ }
+ # exit if the job does not need to be executed
+ pick_or_exit( $logfile );
+
+ system("mkdir -p `dirname $logfile` 2>/dev/null");
+ open(F, ">$logfile") || die "run.pl: Error opening log file $logfile";
+ print F "# " . $cmd . "\n";
+ print F "# Started at " . `date`;
+ $starttime = `date +'%s'`;
+ print F "#\n";
+ close(F);
+
+ # Pipe into bash.. make sure we're not using any other shell.
+ open(B, "|bash") || die "run.pl: Error opening shell command";
+ print B "( " . $cmd . ") 2>>$logfile >> $logfile";
+ close(B); # If there was an error, exit status is in $?
+ $ret = $?;
+
+ $lowbits = $ret & 127;
+ $highbits = $ret >> 8;
+ if ($lowbits != 0) { $return_str = "code $highbits; signal $lowbits" }
+ else { $return_str = "code $highbits"; }
+
+ $endtime = `date +'%s'`;
+ open(F, ">>$logfile") || die "run.pl: Error opening log file $logfile (again)";
+ $enddate = `date`;
+ chop $enddate;
+ print F "# Accounting: time=" . ($endtime - $starttime) . " threads=1\n";
+ print F "# Ended ($return_str) at " . $enddate . ", elapsed time " . ($endtime-$starttime) . " seconds\n";
+ close(F);
+ exit($ret == 0 ? 0 : 1);
+ } else {
+ $pid[$jobid] = $childpid;
+ $active_pids{$childpid} = $jobid;
+ # print STDERR "Queued: " . Dumper(\%active_pids) . "\n";
+ }
+}
+
+# Now we have submitted all the jobs, lets wait until all the jobs finish
+foreach $child (keys %active_pids) {
+ $jobid=$active_pids{$child};
+ $r = waitpid($pid[$jobid], 0);
+ $code = $?;
+ if ($r == -1) { die "run.pl: Error waiting for child process"; } # should never happen.
+ if ($r != 0) { $fail[$jobid]=$code; $numfail++ if $code!=0; } # Completed successfully
+}
+
+# Some sanity checks:
+# The $fail array should not contain undefined codes
+# The number of non-zeros in that array should be equal to $numfail
+# We cannot do foreach() here, as the JOB ids do not start at zero
+$failed_jids=0;
+for ($jobid = $jobstart; $jobid <= $jobend; $jobid++) {
+ $job_return = $fail[$jobid];
+ if (not defined $job_return ) {
+ # print Dumper(\@fail);
+
+ die "run.pl: Sanity check failed: we have indication that some jobs are running " .
+ "even after we waited for all jobs to finish" ;
+ }
+ if ($job_return != 0 ){ $failed_jids++;}
+}
+if ($failed_jids != $numfail) {
+ die "run.pl: Sanity check failed: cannot find out how many jobs failed ($failed_jids x $numfail)."
+}
+if ($numfail > 0) { $ret = 1; }
+
+if ($ret != 0) {
+ $njobs = $jobend - $jobstart + 1;
+ if ($njobs == 1) {
+ if (defined $jobname) {
+ $logfile =~ s/$jobname/$jobstart/; # only one numbered job, so replace name with
+ # that job.
+ }
+ print STDERR "run.pl: job failed, log is in $logfile\n";
+ if ($logfile =~ m/JOB/) {
+ print STDERR "run.pl: probably you forgot to put JOB=1:\$nj in your script.";
+ }
+ }
+ else {
+ $logfile =~ s/$jobname/*/g;
+ print STDERR "run.pl: $numfail / $njobs failed, log is in $logfile\n";
+ }
+}
+
+
+exit ($ret);
diff --git a/egs/aishell/paraformer/utils/shuffle_list.pl b/egs/aishell/paraformer/utils/shuffle_list.pl
new file mode 100755
index 0000000..a116200
--- /dev/null
+++ b/egs/aishell/paraformer/utils/shuffle_list.pl
@@ -0,0 +1,44 @@
+#!/usr/bin/env perl
+
+# Copyright 2013 Johns Hopkins University (author: Daniel Povey)
+
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+if ($ARGV[0] eq "--srand") {
+ $n = $ARGV[1];
+ $n =~ m/\d+/ || die "Bad argument to --srand option: \"$n\"";
+ srand($ARGV[1]);
+ shift;
+ shift;
+} else {
+ srand(0); # Gives inconsistent behavior if we don't seed.
+}
+
+if (@ARGV > 1 || $ARGV[0] =~ m/^-.+/) { # >1 args, or an option we
+ # don't understand.
+ print "Usage: shuffle_list.pl [--srand N] [input file] > output\n";
+ print "randomizes the order of lines of input.\n";
+ exit(1);
+}
+
+@lines;
+while (<>) {
+ push @lines, [ (rand(), $_)] ;
+}
+
+@lines = sort { $a->[0] cmp $b->[0] } @lines;
+foreach $l (@lines) {
+ print $l->[1];
+}
\ No newline at end of file
diff --git a/egs/aishell/paraformer/utils/split_data.py b/egs/aishell/paraformer/utils/split_data.py
new file mode 100755
index 0000000..060eae6
--- /dev/null
+++ b/egs/aishell/paraformer/utils/split_data.py
@@ -0,0 +1,60 @@
+import os
+import sys
+import random
+
+
+in_dir = sys.argv[1]
+out_dir = sys.argv[2]
+num_split = sys.argv[3]
+
+
+def split_scp(scp, num):
+ assert len(scp) >= num
+ avg = len(scp) // num
+ out = []
+ begin = 0
+
+ for i in range(num):
+ if i == num - 1:
+ out.append(scp[begin:])
+ else:
+ out.append(scp[begin:begin+avg])
+ begin += avg
+
+ return out
+
+
+os.path.exists("{}/wav.scp".format(in_dir))
+os.path.exists("{}/text".format(in_dir))
+
+with open("{}/wav.scp".format(in_dir), 'r') as infile:
+ wav_list = infile.readlines()
+
+with open("{}/text".format(in_dir), 'r') as infile:
+ text_list = infile.readlines()
+
+assert len(wav_list) == len(text_list)
+
+x = list(zip(wav_list, text_list))
+random.shuffle(x)
+wav_shuffle_list, text_shuffle_list = zip(*x)
+
+num_split = int(num_split)
+wav_split_list = split_scp(wav_shuffle_list, num_split)
+text_split_list = split_scp(text_shuffle_list, num_split)
+
+for idx, wav_list in enumerate(wav_split_list, 1):
+ path = out_dir + "/split" + str(num_split) + "/" + str(idx)
+ if not os.path.exists(path):
+ os.makedirs(path)
+ with open("{}/wav.scp".format(path), 'w') as wav_writer:
+ for line in wav_list:
+ wav_writer.write(line)
+
+for idx, text_list in enumerate(text_split_list, 1):
+ path = out_dir + "/split" + str(num_split) + "/" + str(idx)
+ if not os.path.exists(path):
+ os.makedirs(path)
+ with open("{}/text".format(path), 'w') as text_writer:
+ for line in text_list:
+ text_writer.write(line)
diff --git a/egs/aishell/paraformer/utils/split_scp.pl b/egs/aishell/paraformer/utils/split_scp.pl
new file mode 100755
index 0000000..0876dcb
--- /dev/null
+++ b/egs/aishell/paraformer/utils/split_scp.pl
@@ -0,0 +1,246 @@
+#!/usr/bin/env perl
+
+# Copyright 2010-2011 Microsoft Corporation
+
+# See ../../COPYING for clarification regarding multiple authors
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+# This program splits up any kind of .scp or archive-type file.
+# If there is no utt2spk option it will work on any text file and
+# will split it up with an approximately equal number of lines in
+# each but.
+# With the --utt2spk option it will work on anything that has the
+# utterance-id as the first entry on each line; the utt2spk file is
+# of the form "utterance speaker" (on each line).
+# It splits it into equal size chunks as far as it can. If you use the utt2spk
+# option it will make sure these chunks coincide with speaker boundaries. In
+# this case, if there are more chunks than speakers (and in some other
+# circumstances), some of the resulting chunks will be empty and it will print
+# an error message and exit with nonzero status.
+# You will normally call this like:
+# split_scp.pl scp scp.1 scp.2 scp.3 ...
+# or
+# split_scp.pl --utt2spk=utt2spk scp scp.1 scp.2 scp.3 ...
+# Note that you can use this script to split the utt2spk file itself,
+# e.g. split_scp.pl --utt2spk=utt2spk utt2spk utt2spk.1 utt2spk.2 ...
+
+# You can also call the scripts like:
+# split_scp.pl -j 3 0 scp scp.0
+# [note: with this option, it assumes zero-based indexing of the split parts,
+# i.e. the second number must be 0 <= n < num-jobs.]
+
+use warnings;
+
+$num_jobs = 0;
+$job_id = 0;
+$utt2spk_file = "";
+$one_based = 0;
+
+for ($x = 1; $x <= 3 && @ARGV > 0; $x++) {
+ if ($ARGV[0] eq "-j") {
+ shift @ARGV;
+ $num_jobs = shift @ARGV;
+ $job_id = shift @ARGV;
+ }
+ if ($ARGV[0] =~ /--utt2spk=(.+)/) {
+ $utt2spk_file=$1;
+ shift;
+ }
+ if ($ARGV[0] eq '--one-based') {
+ $one_based = 1;
+ shift @ARGV;
+ }
+}
+
+if ($num_jobs != 0 && ($num_jobs < 0 || $job_id - $one_based < 0 ||
+ $job_id - $one_based >= $num_jobs)) {
+ die "$0: Invalid job number/index values for '-j $num_jobs $job_id" .
+ ($one_based ? " --one-based" : "") . "'\n"
+}
+
+$one_based
+ and $job_id--;
+
+if(($num_jobs == 0 && @ARGV < 2) || ($num_jobs > 0 && (@ARGV < 1 || @ARGV > 2))) {
+ die
+"Usage: split_scp.pl [--utt2spk=<utt2spk_file>] in.scp out1.scp out2.scp ...
+ or: split_scp.pl -j num-jobs job-id [--one-based] [--utt2spk=<utt2spk_file>] in.scp [out.scp]
+ ... where 0 <= job-id < num-jobs, or 1 <= job-id <- num-jobs if --one-based.\n";
+}
+
+$error = 0;
+$inscp = shift @ARGV;
+if ($num_jobs == 0) { # without -j option
+ @OUTPUTS = @ARGV;
+} else {
+ for ($j = 0; $j < $num_jobs; $j++) {
+ if ($j == $job_id) {
+ if (@ARGV > 0) { push @OUTPUTS, $ARGV[0]; }
+ else { push @OUTPUTS, "-"; }
+ } else {
+ push @OUTPUTS, "/dev/null";
+ }
+ }
+}
+
+if ($utt2spk_file ne "") { # We have the --utt2spk option...
+ open($u_fh, '<', $utt2spk_file) || die "$0: Error opening utt2spk file $utt2spk_file: $!\n";
+ while(<$u_fh>) {
+ @A = split;
+ @A == 2 || die "$0: Bad line $_ in utt2spk file $utt2spk_file\n";
+ ($u,$s) = @A;
+ $utt2spk{$u} = $s;
+ }
+ close $u_fh;
+ open($i_fh, '<', $inscp) || die "$0: Error opening input scp file $inscp: $!\n";
+ @spkrs = ();
+ while(<$i_fh>) {
+ @A = split;
+ if(@A == 0) { die "$0: Empty or space-only line in scp file $inscp\n"; }
+ $u = $A[0];
+ $s = $utt2spk{$u};
+ defined $s || die "$0: No utterance $u in utt2spk file $utt2spk_file\n";
+ if(!defined $spk_count{$s}) {
+ push @spkrs, $s;
+ $spk_count{$s} = 0;
+ $spk_data{$s} = []; # ref to new empty array.
+ }
+ $spk_count{$s}++;
+ push @{$spk_data{$s}}, $_;
+ }
+ # Now split as equally as possible ..
+ # First allocate spks to files by allocating an approximately
+ # equal number of speakers.
+ $numspks = @spkrs; # number of speakers.
+ $numscps = @OUTPUTS; # number of output files.
+ if ($numspks < $numscps) {
+ die "$0: Refusing to split data because number of speakers $numspks " .
+ "is less than the number of output .scp files $numscps\n";
+ }
+ for($scpidx = 0; $scpidx < $numscps; $scpidx++) {
+ $scparray[$scpidx] = []; # [] is array reference.
+ }
+ for ($spkidx = 0; $spkidx < $numspks; $spkidx++) {
+ $scpidx = int(($spkidx*$numscps) / $numspks);
+ $spk = $spkrs[$spkidx];
+ push @{$scparray[$scpidx]}, $spk;
+ $scpcount[$scpidx] += $spk_count{$spk};
+ }
+
+ # Now will try to reassign beginning + ending speakers
+ # to different scp's and see if it gets more balanced.
+ # Suppose objf we're minimizing is sum_i (num utts in scp[i] - average)^2.
+ # We can show that if considering changing just 2 scp's, we minimize
+ # this by minimizing the squared difference in sizes. This is
+ # equivalent to minimizing the absolute difference in sizes. This
+ # shows this method is bound to converge.
+
+ $changed = 1;
+ while($changed) {
+ $changed = 0;
+ for($scpidx = 0; $scpidx < $numscps; $scpidx++) {
+ # First try to reassign ending spk of this scp.
+ if($scpidx < $numscps-1) {
+ $sz = @{$scparray[$scpidx]};
+ if($sz > 0) {
+ $spk = $scparray[$scpidx]->[$sz-1];
+ $count = $spk_count{$spk};
+ $nutt1 = $scpcount[$scpidx];
+ $nutt2 = $scpcount[$scpidx+1];
+ if( abs( ($nutt2+$count) - ($nutt1-$count))
+ < abs($nutt2 - $nutt1)) { # Would decrease
+ # size-diff by reassigning spk...
+ $scpcount[$scpidx+1] += $count;
+ $scpcount[$scpidx] -= $count;
+ pop @{$scparray[$scpidx]};
+ unshift @{$scparray[$scpidx+1]}, $spk;
+ $changed = 1;
+ }
+ }
+ }
+ if($scpidx > 0 && @{$scparray[$scpidx]} > 0) {
+ $spk = $scparray[$scpidx]->[0];
+ $count = $spk_count{$spk};
+ $nutt1 = $scpcount[$scpidx-1];
+ $nutt2 = $scpcount[$scpidx];
+ if( abs( ($nutt2-$count) - ($nutt1+$count))
+ < abs($nutt2 - $nutt1)) { # Would decrease
+ # size-diff by reassigning spk...
+ $scpcount[$scpidx-1] += $count;
+ $scpcount[$scpidx] -= $count;
+ shift @{$scparray[$scpidx]};
+ push @{$scparray[$scpidx-1]}, $spk;
+ $changed = 1;
+ }
+ }
+ }
+ }
+ # Now print out the files...
+ for($scpidx = 0; $scpidx < $numscps; $scpidx++) {
+ $scpfile = $OUTPUTS[$scpidx];
+ ($scpfile ne '-' ? open($f_fh, '>', $scpfile)
+ : open($f_fh, '>&', \*STDOUT)) ||
+ die "$0: Could not open scp file $scpfile for writing: $!\n";
+ $count = 0;
+ if(@{$scparray[$scpidx]} == 0) {
+ print STDERR "$0: eError: split_scp.pl producing empty .scp file " .
+ "$scpfile (too many splits and too few speakers?)\n";
+ $error = 1;
+ } else {
+ foreach $spk ( @{$scparray[$scpidx]} ) {
+ print $f_fh @{$spk_data{$spk}};
+ $count += $spk_count{$spk};
+ }
+ $count == $scpcount[$scpidx] || die "Count mismatch [code error]";
+ }
+ close($f_fh);
+ }
+} else {
+ # This block is the "normal" case where there is no --utt2spk
+ # option and we just break into equal size chunks.
+
+ open($i_fh, '<', $inscp) || die "$0: Error opening input scp file $inscp: $!\n";
+
+ $numscps = @OUTPUTS; # size of array.
+ @F = ();
+ while(<$i_fh>) {
+ push @F, $_;
+ }
+ $numlines = @F;
+ if($numlines == 0) {
+ print STDERR "$0: error: empty input scp file $inscp\n";
+ $error = 1;
+ }
+ $linesperscp = int( $numlines / $numscps); # the "whole part"..
+ $linesperscp >= 1 || die "$0: You are splitting into too many pieces! [reduce \$nj ($numscps) to be smaller than the number of lines ($numlines) in $inscp]\n";
+ $remainder = $numlines - ($linesperscp * $numscps);
+ ($remainder >= 0 && $remainder < $numlines) || die "bad remainder $remainder";
+ # [just doing int() rounds down].
+ $n = 0;
+ for($scpidx = 0; $scpidx < @OUTPUTS; $scpidx++) {
+ $scpfile = $OUTPUTS[$scpidx];
+ ($scpfile ne '-' ? open($o_fh, '>', $scpfile)
+ : open($o_fh, '>&', \*STDOUT)) ||
+ die "$0: Could not open scp file $scpfile for writing: $!\n";
+ for($k = 0; $k < $linesperscp + ($scpidx < $remainder ? 1 : 0); $k++) {
+ print $o_fh $F[$n++];
+ }
+ close($o_fh) || die "$0: Eror closing scp file $scpfile: $!\n";
+ }
+ $n == $numlines || die "$n != $numlines [code error]";
+}
+
+exit ($error);
diff --git a/egs/aishell/paraformer/utils/subset_data_dir_tr_cv.sh b/egs/aishell/paraformer/utils/subset_data_dir_tr_cv.sh
new file mode 100755
index 0000000..e16cebd
--- /dev/null
+++ b/egs/aishell/paraformer/utils/subset_data_dir_tr_cv.sh
@@ -0,0 +1,30 @@
+#!/usr/bin/env bash
+
+dev_num_utt=1000
+
+echo "$0 $@"
+. utils/parse_options.sh || exit 1;
+
+train_data=$1
+out_dir=$2
+
+[ ! -f ${train_data}/wav.scp ] && echo "$0: no such file ${train_data}/wav.scp" && exit 1;
+[ ! -f ${train_data}/text ] && echo "$0: no such file ${train_data}/text" && exit 1;
+
+mkdir -p ${out_dir}/train && mkdir -p ${out_dir}/dev
+
+cp ${train_data}/wav.scp ${out_dir}/train/wav.scp.bak
+cp ${train_data}/text ${out_dir}/train/text.bak
+
+num_utt=$(wc -l <${out_dir}/train/wav.scp.bak)
+
+utils/shuffle_list.pl --srand 1 ${out_dir}/train/wav.scp.bak > ${out_dir}/train/wav.scp.shuf
+head -n ${dev_num_utt} ${out_dir}/train/wav.scp.shuf > ${out_dir}/dev/wav.scp
+tail -n $((${num_utt}-${dev_num_utt})) ${out_dir}/train/wav.scp.shuf > ${out_dir}/train/wav.scp
+
+utils/shuffle_list.pl --srand 1 ${out_dir}/train/text.bak > ${out_dir}/train/text.shuf
+head -n ${dev_num_utt} ${out_dir}/train/text.shuf > ${out_dir}/dev/text
+tail -n $((${num_utt}-${dev_num_utt})) ${out_dir}/train/text.shuf > ${out_dir}/train/text
+
+rm ${out_dir}/train/wav.scp.bak ${out_dir}/train/text.bak
+rm ${out_dir}/train/wav.scp.shuf ${out_dir}/train/text.shuf
diff --git a/egs/aishell/paraformer/utils/text2token.py b/egs/aishell/paraformer/utils/text2token.py
new file mode 100755
index 0000000..56c3913
--- /dev/null
+++ b/egs/aishell/paraformer/utils/text2token.py
@@ -0,0 +1,135 @@
+#!/usr/bin/env python3
+
+# Copyright 2017 Johns Hopkins University (Shinji Watanabe)
+# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
+
+
+import argparse
+import codecs
+import re
+import sys
+
+is_python2 = sys.version_info[0] == 2
+
+
+def exist_or_not(i, match_pos):
+ start_pos = None
+ end_pos = None
+ for pos in match_pos:
+ if pos[0] <= i < pos[1]:
+ start_pos = pos[0]
+ end_pos = pos[1]
+ break
+
+ return start_pos, end_pos
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="convert raw text to tokenized text",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--nchar",
+ "-n",
+ default=1,
+ type=int,
+ help="number of characters to split, i.e., \
+ aabb -> a a b b with -n 1 and aa bb with -n 2",
+ )
+ parser.add_argument(
+ "--skip-ncols", "-s", default=0, type=int, help="skip first n columns"
+ )
+ parser.add_argument("--space", default="<space>", type=str, help="space symbol")
+ parser.add_argument(
+ "--non-lang-syms",
+ "-l",
+ default=None,
+ type=str,
+ help="list of non-linguistic symobles, e.g., <NOISE> etc.",
+ )
+ parser.add_argument("text", type=str, default=False, nargs="?", help="input text")
+ parser.add_argument(
+ "--trans_type",
+ "-t",
+ type=str,
+ default="char",
+ choices=["char", "phn"],
+ help="""Transcript type. char/phn. e.g., for TIMIT FADG0_SI1279 -
+ If trans_type is char,
+ read from SI1279.WRD file -> "bricks are an alternative"
+ Else if trans_type is phn,
+ read from SI1279.PHN file -> "sil b r ih sil k s aa r er n aa l
+ sil t er n ih sil t ih v sil" """,
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ rs = []
+ if args.non_lang_syms is not None:
+ with codecs.open(args.non_lang_syms, "r", encoding="utf-8") as f:
+ nls = [x.rstrip() for x in f.readlines()]
+ rs = [re.compile(re.escape(x)) for x in nls]
+
+ if args.text:
+ f = codecs.open(args.text, encoding="utf-8")
+ else:
+ f = codecs.getreader("utf-8")(sys.stdin if is_python2 else sys.stdin.buffer)
+
+ sys.stdout = codecs.getwriter("utf-8")(
+ sys.stdout if is_python2 else sys.stdout.buffer
+ )
+ line = f.readline()
+ n = args.nchar
+ while line:
+ x = line.split()
+ print(" ".join(x[: args.skip_ncols]), end=" ")
+ a = " ".join(x[args.skip_ncols :])
+
+ # get all matched positions
+ match_pos = []
+ for r in rs:
+ i = 0
+ while i >= 0:
+ m = r.search(a, i)
+ if m:
+ match_pos.append([m.start(), m.end()])
+ i = m.end()
+ else:
+ break
+
+ if args.trans_type == "phn":
+ a = a.split(" ")
+ else:
+ if len(match_pos) > 0:
+ chars = []
+ i = 0
+ while i < len(a):
+ start_pos, end_pos = exist_or_not(i, match_pos)
+ if start_pos is not None:
+ chars.append(a[start_pos:end_pos])
+ i = end_pos
+ else:
+ chars.append(a[i])
+ i += 1
+ a = chars
+
+ a = [a[j : j + n] for j in range(0, len(a), n)]
+
+ a_flat = []
+ for z in a:
+ a_flat.append("".join(z))
+
+ a_chars = [z.replace(" ", args.space) for z in a_flat]
+ if args.trans_type == "phn":
+ a_chars = [z.replace("sil", args.space) for z in a_chars]
+ print(" ".join(a_chars))
+ line = f.readline()
+
+
+if __name__ == "__main__":
+ main()
diff --git a/egs/aishell/paraformer/utils/text_tokenize.py b/egs/aishell/paraformer/utils/text_tokenize.py
new file mode 100755
index 0000000..962ea11
--- /dev/null
+++ b/egs/aishell/paraformer/utils/text_tokenize.py
@@ -0,0 +1,106 @@
+import re
+import argparse
+
+
+def load_dict(seg_file):
+ seg_dict = {}
+ with open(seg_file, 'r') as infile:
+ for line in infile:
+ s = line.strip().split()
+ key = s[0]
+ value = s[1:]
+ seg_dict[key] = " ".join(value)
+ return seg_dict
+
+
+def forward_segment(text, dic):
+ word_list = []
+ i = 0
+ while i < len(text):
+ longest_word = text[i]
+ for j in range(i + 1, len(text) + 1):
+ word = text[i:j]
+ if word in dic:
+ if len(word) > len(longest_word):
+ longest_word = word
+ word_list.append(longest_word)
+ i += len(longest_word)
+ return word_list
+
+
+def tokenize(txt,
+ seg_dict):
+ out_txt = ""
+ pattern = re.compile(r"([\u4E00-\u9FA5A-Za-z0-9])")
+ for word in txt:
+ if pattern.match(word):
+ if word in seg_dict:
+ out_txt += seg_dict[word] + " "
+ else:
+ out_txt += "<unk>" + " "
+ else:
+ continue
+ return out_txt.strip()
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="text tokenize",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--text-file",
+ "-t",
+ default=False,
+ required=True,
+ type=str,
+ help="input text",
+ )
+ parser.add_argument(
+ "--seg-file",
+ "-s",
+ default=False,
+ required=True,
+ type=str,
+ help="seg file",
+ )
+ parser.add_argument(
+ "--txt-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="txt index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ txt_writer = open("{}/text.{}.txt".format(args.output_dir, args.txt_index), 'w')
+ shape_writer = open("{}/len.{}".format(args.output_dir, args.txt_index), 'w')
+ seg_dict = load_dict(args.seg_file)
+ with open(args.text_file, 'r') as infile:
+ for line in infile:
+ s = line.strip().split()
+ text_id = s[0]
+ text_list = forward_segment("".join(s[1:]).lower(), seg_dict)
+ text = tokenize(text_list, seg_dict)
+ lens = len(text.strip().split())
+ txt_writer.write(text_id + " " + text + '\n')
+ shape_writer.write(text_id + " " + str(lens) + '\n')
+
+
+if __name__ == '__main__':
+ main()
+
diff --git a/egs/aishell/paraformer/utils/text_tokenize.sh b/egs/aishell/paraformer/utils/text_tokenize.sh
new file mode 100755
index 0000000..6b74fef
--- /dev/null
+++ b/egs/aishell/paraformer/utils/text_tokenize.sh
@@ -0,0 +1,35 @@
+#!/usr/bin/env bash
+
+
+# Begin configuration section.
+nj=32
+cmd=utils/run.pl
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+# tokenize configuration
+text_dir=$1
+seg_file=$2
+logdir=$3
+output_dir=$4
+
+txt_dir=${output_dir}/txt; mkdir -p ${output_dir}/txt
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/text_tokenize.JOB.log \
+ python utils/text_tokenize.py -t ${text_dir}/txt/text.JOB.txt \
+ -s ${seg_file} -i JOB -o ${txt_dir} \
+ || exit 1;
+
+# concatenate the text files together.
+for n in $(seq $nj); do
+ cat ${txt_dir}/text.$n.txt || exit 1
+done > ${output_dir}/text || exit 1
+
+for n in $(seq $nj); do
+ cat ${txt_dir}/len.$n || exit 1
+done > ${output_dir}/text_shape || exit 1
+
+echo "$0: Succeeded text tokenize"
diff --git a/egs/aishell/paraformer/utils/textnorm_zh.py b/egs/aishell/paraformer/utils/textnorm_zh.py
new file mode 100755
index 0000000..79feb83
--- /dev/null
+++ b/egs/aishell/paraformer/utils/textnorm_zh.py
@@ -0,0 +1,834 @@
+#!/usr/bin/env python3
+# coding=utf-8
+
+# Authors:
+# 2019.5 Zhiyang Zhou (https://github.com/Joee1995/chn_text_norm.git)
+# 2019.9 Jiayu DU
+#
+# requirements:
+# - python 3.X
+# notes: python 2.X WILL fail or produce misleading results
+
+import sys, os, argparse, codecs, string, re
+
+# ================================================================================ #
+# basic constant
+# ================================================================================ #
+CHINESE_DIGIS = u'闆朵竴浜屼笁鍥涗簲鍏竷鍏節'
+BIG_CHINESE_DIGIS_SIMPLIFIED = u'闆跺9璐板弫鑲嗕紞闄嗘煉鎹岀帠'
+BIG_CHINESE_DIGIS_TRADITIONAL = u'闆跺9璨冲弮鑲嗕紞闄告煉鎹岀帠'
+SMALLER_BIG_CHINESE_UNITS_SIMPLIFIED = u'鍗佺櫨鍗冧竾'
+SMALLER_BIG_CHINESE_UNITS_TRADITIONAL = u'鎷句桨浠熻惉'
+LARGER_CHINESE_NUMERING_UNITS_SIMPLIFIED = u'浜垮厗浜灀绉┌娌熸锭姝h浇'
+LARGER_CHINESE_NUMERING_UNITS_TRADITIONAL = u'鍎勫厗浜灀绉┌婧濇緱姝h級'
+SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED = u'鍗佺櫨鍗冧竾'
+SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL = u'鎷句桨浠熻惉'
+
+ZERO_ALT = u'銆�'
+ONE_ALT = u'骞�'
+TWO_ALTS = [u'涓�', u'鍏�']
+
+POSITIVE = [u'姝�', u'姝�']
+NEGATIVE = [u'璐�', u'璨�']
+POINT = [u'鐐�', u'榛�']
+# PLUS = [u'鍔�', u'鍔�']
+# SIL = [u'鏉�', u'妲�']
+
+FILLER_CHARS = ['鍛�', '鍟�']
+ER_WHITELIST = '(鍎垮コ|鍎垮瓙|鍎垮瓩|濂冲効|鍎垮|濡诲効|' \
+ '鑳庡効|濠村効|鏂扮敓鍎縷濠村辜鍎縷骞煎効|灏戝効|灏忓効|鍎挎瓕|鍎跨|鍎跨|鎵樺効鎵�|瀛ゅ効|' \
+ '鍎挎垙|鍎垮寲|鍙板効搴剕楣垮効宀泑姝e効鍏粡|鍚婂効閮庡綋|鐢熷効鑲插コ|鎵樺効甯﹀コ|鍏诲効闃茶�亅鐥村効鍛嗗コ|' \
+ '浣冲効浣冲|鍎挎�滃吔鎵皘鍎挎棤甯哥埗|鍎夸笉瀚屾瘝涓憒鍎胯鍗冮噷姣嶆媴蹇鍎垮ぇ涓嶇敱鐖穦鑻忎篂鍎�)'
+
+# 涓枃鏁板瓧绯荤粺绫诲瀷
+NUMBERING_TYPES = ['low', 'mid', 'high']
+
+CURRENCY_NAMES = '(浜烘皯甯亅缇庡厓|鏃ュ厓|鑻遍晳|娆у厓|椹厠|娉曢儙|鍔犳嬁澶у厓|婢冲厓|娓竵|鍏堜护|鑺叞椹厠|鐖卞皵鍏伴晳|' \
+ '閲屾媺|鑽峰叞鐩緗鍩冩柉搴撳|姣斿濉攟鍗板凹鐩緗鏋楀悏鐗箌鏂拌タ鍏板厓|姣旂储|鍗㈠竷|鏂板姞鍧″厓|闊╁厓|娉伴摙)'
+CURRENCY_UNITS = '((浜縷鍗冧竾|鐧句竾|涓噟鍗億鐧�)|(浜縷鍗冧竾|鐧句竾|涓噟鍗億鐧緗)鍏億(浜縷鍗冧竾|鐧句竾|涓噟鍗億鐧緗)鍧梶瑙抾姣泑鍒�)'
+COM_QUANTIFIERS = '(鍖箌寮爘搴鍥瀨鍦簗灏緗鏉涓獆棣東闃檤闃祙缃憒鐐畖椤秥涓榺妫祙鍙獆鏀瘄琚瓅杈唡鎸憒鎷厊棰梶澹硘绐爘鏇瞸澧檤缇鑵攟' \
+ '鐮搴瀹璐瘄鎵巪鎹唡鍒�|浠鎵搢鎵媩缃梶鍧灞眧宀瓅姹焲婧獆閽焲闃焲鍗晐鍙寍瀵箌鍑簗鍙澶磡鑴殀鏉縷璺硘鏋潀浠秥璐磡' \
+ '閽坾绾縷绠鍚峾浣峾韬珅鍫倈璇緗鏈瑋椤祙瀹秥鎴穦灞倈涓潀姣珅鍘榺鍒唡閽眧涓鏂鎷厊閾鐭硘閽閿眧蹇絴(鍗億姣珅寰�)鍏媩' \
+ '姣珅鍘榺鍒唡瀵竱灏簗涓坾閲寍瀵粅甯竱閾簗绋媩(鍗億鍒唡鍘榺姣珅寰�)绫硘鎾畖鍕簗鍚坾鍗噟鏂梶鐭硘鐩榺纰梶纰焲鍙爘妗秥绗紎鐩唡' \
+ '鐩抾鏉瘄閽焲鏂泑閿厊绨媩绡畖鐩榺妗秥缃恷鐡秥澹秥鍗畖鐩弢绠﹟绠眧鐓瞸鍟東琚媩閽祙骞磡鏈坾鏃瀛鍒粅鏃秥鍛▅澶﹟绉抾鍒唡鏃瑋' \
+ '绾獆宀亅涓東鏇磡澶渱鏄澶弢绉媩鍐瑋浠浼弢杈坾涓竱娉绮抾棰梶骞鍫唡鏉鏍箌鏀瘄閬搢闈鐗噟寮爘棰梶鍧�)'
+
+# punctuation information are based on Zhon project (https://github.com/tsroten/zhon.git)
+CHINESE_PUNC_STOP = '锛侊紵锝°��'
+CHINESE_PUNC_NON_STOP = '锛傦純锛勶紖锛嗭紘锛堬級锛婏紜锛岋紞锛忥細锛涳紲锛濓紴锛狅蓟锛硷冀锛撅伎锝�锝涳綔锝濓綖锝燂綘锝剑锝ゃ�併�冦�嬨�屻�嶃�庛�忋�愩�戙�斻�曘�栥�椼�樸�欍�氥�涖�溿�濄�炪�熴�般�俱�库�撯�斺�樷�欌�涒�溾�濃�炩�熲�︹�э箯'
+CHINESE_PUNC_LIST = CHINESE_PUNC_STOP + CHINESE_PUNC_NON_STOP
+
+# ================================================================================ #
+# basic class
+# ================================================================================ #
+class ChineseChar(object):
+ """
+ 涓枃瀛楃
+ 姣忎釜瀛楃瀵瑰簲绠�浣撳拰绻佷綋,
+ e.g. 绠�浣� = '璐�', 绻佷綋 = '璨�'
+ 杞崲鏃跺彲杞崲涓虹畝浣撴垨绻佷綋
+ """
+
+ def __init__(self, simplified, traditional):
+ self.simplified = simplified
+ self.traditional = traditional
+ #self.__repr__ = self.__str__
+
+ def __str__(self):
+ return self.simplified or self.traditional or None
+
+ def __repr__(self):
+ return self.__str__()
+
+
+class ChineseNumberUnit(ChineseChar):
+ """
+ 涓枃鏁板瓧/鏁颁綅瀛楃
+ 姣忎釜瀛楃闄ょ箒绠�浣撳杩樻湁涓�涓澶栫殑澶у啓瀛楃
+ e.g. '闄�' 鍜� '闄�'
+ """
+
+ def __init__(self, power, simplified, traditional, big_s, big_t):
+ super(ChineseNumberUnit, self).__init__(simplified, traditional)
+ self.power = power
+ self.big_s = big_s
+ self.big_t = big_t
+
+ def __str__(self):
+ return '10^{}'.format(self.power)
+
+ @classmethod
+ def create(cls, index, value, numbering_type=NUMBERING_TYPES[1], small_unit=False):
+
+ if small_unit:
+ return ChineseNumberUnit(power=index + 1,
+ simplified=value[0], traditional=value[1], big_s=value[1], big_t=value[1])
+ elif numbering_type == NUMBERING_TYPES[0]:
+ return ChineseNumberUnit(power=index + 8,
+ simplified=value[0], traditional=value[1], big_s=value[0], big_t=value[1])
+ elif numbering_type == NUMBERING_TYPES[1]:
+ return ChineseNumberUnit(power=(index + 2) * 4,
+ simplified=value[0], traditional=value[1], big_s=value[0], big_t=value[1])
+ elif numbering_type == NUMBERING_TYPES[2]:
+ return ChineseNumberUnit(power=pow(2, index + 3),
+ simplified=value[0], traditional=value[1], big_s=value[0], big_t=value[1])
+ else:
+ raise ValueError(
+ 'Counting type should be in {0} ({1} provided).'.format(NUMBERING_TYPES, numbering_type))
+
+
+class ChineseNumberDigit(ChineseChar):
+ """
+ 涓枃鏁板瓧瀛楃
+ """
+
+ def __init__(self, value, simplified, traditional, big_s, big_t, alt_s=None, alt_t=None):
+ super(ChineseNumberDigit, self).__init__(simplified, traditional)
+ self.value = value
+ self.big_s = big_s
+ self.big_t = big_t
+ self.alt_s = alt_s
+ self.alt_t = alt_t
+
+ def __str__(self):
+ return str(self.value)
+
+ @classmethod
+ def create(cls, i, v):
+ return ChineseNumberDigit(i, v[0], v[1], v[2], v[3])
+
+
+class ChineseMath(ChineseChar):
+ """
+ 涓枃鏁颁綅瀛楃
+ """
+
+ def __init__(self, simplified, traditional, symbol, expression=None):
+ super(ChineseMath, self).__init__(simplified, traditional)
+ self.symbol = symbol
+ self.expression = expression
+ self.big_s = simplified
+ self.big_t = traditional
+
+
+CC, CNU, CND, CM = ChineseChar, ChineseNumberUnit, ChineseNumberDigit, ChineseMath
+
+
+class NumberSystem(object):
+ """
+ 涓枃鏁板瓧绯荤粺
+ """
+ pass
+
+
+class MathSymbol(object):
+ """
+ 鐢ㄤ簬涓枃鏁板瓧绯荤粺鐨勬暟瀛︾鍙� (绻�/绠�浣�), e.g.
+ positive = ['姝�', '姝�']
+ negative = ['璐�', '璨�']
+ point = ['鐐�', '榛�']
+ """
+
+ def __init__(self, positive, negative, point):
+ self.positive = positive
+ self.negative = negative
+ self.point = point
+
+ def __iter__(self):
+ for v in self.__dict__.values():
+ yield v
+
+
+# class OtherSymbol(object):
+# """
+# 鍏朵粬绗﹀彿
+# """
+#
+# def __init__(self, sil):
+# self.sil = sil
+#
+# def __iter__(self):
+# for v in self.__dict__.values():
+# yield v
+
+
+# ================================================================================ #
+# basic utils
+# ================================================================================ #
+def create_system(numbering_type=NUMBERING_TYPES[1]):
+ """
+ 鏍规嵁鏁板瓧绯荤粺绫诲瀷杩斿洖鍒涘缓鐩稿簲鐨勬暟瀛楃郴缁燂紝榛樿涓� mid
+ NUMBERING_TYPES = ['low', 'mid', 'high']: 涓枃鏁板瓧绯荤粺绫诲瀷
+ low: '鍏�' = '浜�' * '鍗�' = $10^{9}$, '浜�' = '鍏�' * '鍗�', etc.
+ mid: '鍏�' = '浜�' * '涓�' = $10^{12}$, '浜�' = '鍏�' * '涓�', etc.
+ high: '鍏�' = '浜�' * '浜�' = $10^{16}$, '浜�' = '鍏�' * '鍏�', etc.
+ 杩斿洖瀵瑰簲鐨勬暟瀛楃郴缁�
+ """
+
+ # chinese number units of '浜�' and larger
+ all_larger_units = zip(
+ LARGER_CHINESE_NUMERING_UNITS_SIMPLIFIED, LARGER_CHINESE_NUMERING_UNITS_TRADITIONAL)
+ larger_units = [CNU.create(i, v, numbering_type, False)
+ for i, v in enumerate(all_larger_units)]
+ # chinese number units of '鍗�, 鐧�, 鍗�, 涓�'
+ all_smaller_units = zip(
+ SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED, SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL)
+ smaller_units = [CNU.create(i, v, small_unit=True)
+ for i, v in enumerate(all_smaller_units)]
+ # digis
+ chinese_digis = zip(CHINESE_DIGIS, CHINESE_DIGIS,
+ BIG_CHINESE_DIGIS_SIMPLIFIED, BIG_CHINESE_DIGIS_TRADITIONAL)
+ digits = [CND.create(i, v) for i, v in enumerate(chinese_digis)]
+ digits[0].alt_s, digits[0].alt_t = ZERO_ALT, ZERO_ALT
+ digits[1].alt_s, digits[1].alt_t = ONE_ALT, ONE_ALT
+ digits[2].alt_s, digits[2].alt_t = TWO_ALTS[0], TWO_ALTS[1]
+
+ # symbols
+ positive_cn = CM(POSITIVE[0], POSITIVE[1], '+', lambda x: x)
+ negative_cn = CM(NEGATIVE[0], NEGATIVE[1], '-', lambda x: -x)
+ point_cn = CM(POINT[0], POINT[1], '.', lambda x,
+ y: float(str(x) + '.' + str(y)))
+ # sil_cn = CM(SIL[0], SIL[1], '-', lambda x, y: float(str(x) + '-' + str(y)))
+ system = NumberSystem()
+ system.units = smaller_units + larger_units
+ system.digits = digits
+ system.math = MathSymbol(positive_cn, negative_cn, point_cn)
+ # system.symbols = OtherSymbol(sil_cn)
+ return system
+
+
+def chn2num(chinese_string, numbering_type=NUMBERING_TYPES[1]):
+
+ def get_symbol(char, system):
+ for u in system.units:
+ if char in [u.traditional, u.simplified, u.big_s, u.big_t]:
+ return u
+ for d in system.digits:
+ if char in [d.traditional, d.simplified, d.big_s, d.big_t, d.alt_s, d.alt_t]:
+ return d
+ for m in system.math:
+ if char in [m.traditional, m.simplified]:
+ return m
+
+ def string2symbols(chinese_string, system):
+ int_string, dec_string = chinese_string, ''
+ for p in [system.math.point.simplified, system.math.point.traditional]:
+ if p in chinese_string:
+ int_string, dec_string = chinese_string.split(p)
+ break
+ return [get_symbol(c, system) for c in int_string], \
+ [get_symbol(c, system) for c in dec_string]
+
+ def correct_symbols(integer_symbols, system):
+ """
+ 涓�鐧惧叓 to 涓�鐧惧叓鍗�
+ 涓�浜夸竴鍗冧笁鐧句竾 to 涓�浜� 涓�鍗冧竾 涓夌櫨涓�
+ """
+
+ if integer_symbols and isinstance(integer_symbols[0], CNU):
+ if integer_symbols[0].power == 1:
+ integer_symbols = [system.digits[1]] + integer_symbols
+
+ if len(integer_symbols) > 1:
+ if isinstance(integer_symbols[-1], CND) and isinstance(integer_symbols[-2], CNU):
+ integer_symbols.append(
+ CNU(integer_symbols[-2].power - 1, None, None, None, None))
+
+ result = []
+ unit_count = 0
+ for s in integer_symbols:
+ if isinstance(s, CND):
+ result.append(s)
+ unit_count = 0
+ elif isinstance(s, CNU):
+ current_unit = CNU(s.power, None, None, None, None)
+ unit_count += 1
+
+ if unit_count == 1:
+ result.append(current_unit)
+ elif unit_count > 1:
+ for i in range(len(result)):
+ if isinstance(result[-i - 1], CNU) and result[-i - 1].power < current_unit.power:
+ result[-i - 1] = CNU(result[-i - 1].power +
+ current_unit.power, None, None, None, None)
+ return result
+
+ def compute_value(integer_symbols):
+ """
+ Compute the value.
+ When current unit is larger than previous unit, current unit * all previous units will be used as all previous units.
+ e.g. '涓ゅ崈涓�' = 2000 * 10000 not 2000 + 10000
+ """
+ value = [0]
+ last_power = 0
+ for s in integer_symbols:
+ if isinstance(s, CND):
+ value[-1] = s.value
+ elif isinstance(s, CNU):
+ value[-1] *= pow(10, s.power)
+ if s.power > last_power:
+ value[:-1] = list(map(lambda v: v *
+ pow(10, s.power), value[:-1]))
+ last_power = s.power
+ value.append(0)
+ return sum(value)
+
+ system = create_system(numbering_type)
+ int_part, dec_part = string2symbols(chinese_string, system)
+ int_part = correct_symbols(int_part, system)
+ int_str = str(compute_value(int_part))
+ dec_str = ''.join([str(d.value) for d in dec_part])
+ if dec_part:
+ return '{0}.{1}'.format(int_str, dec_str)
+ else:
+ return int_str
+
+
+def num2chn(number_string, numbering_type=NUMBERING_TYPES[1], big=False,
+ traditional=False, alt_zero=False, alt_one=False, alt_two=True,
+ use_zeros=True, use_units=True):
+
+ def get_value(value_string, use_zeros=True):
+
+ striped_string = value_string.lstrip('0')
+
+ # record nothing if all zeros
+ if not striped_string:
+ return []
+
+ # record one digits
+ elif len(striped_string) == 1:
+ if use_zeros and len(value_string) != len(striped_string):
+ return [system.digits[0], system.digits[int(striped_string)]]
+ else:
+ return [system.digits[int(striped_string)]]
+
+ # recursively record multiple digits
+ else:
+ result_unit = next(u for u in reversed(
+ system.units) if u.power < len(striped_string))
+ result_string = value_string[:-result_unit.power]
+ return get_value(result_string) + [result_unit] + get_value(striped_string[-result_unit.power:])
+
+ system = create_system(numbering_type)
+
+ int_dec = number_string.split('.')
+ if len(int_dec) == 1:
+ int_string = int_dec[0]
+ dec_string = ""
+ elif len(int_dec) == 2:
+ int_string = int_dec[0]
+ dec_string = int_dec[1]
+ else:
+ raise ValueError(
+ "invalid input num string with more than one dot: {}".format(number_string))
+
+ if use_units and len(int_string) > 1:
+ result_symbols = get_value(int_string)
+ else:
+ result_symbols = [system.digits[int(c)] for c in int_string]
+ dec_symbols = [system.digits[int(c)] for c in dec_string]
+ if dec_string:
+ result_symbols += [system.math.point] + dec_symbols
+
+ if alt_two:
+ liang = CND(2, system.digits[2].alt_s, system.digits[2].alt_t,
+ system.digits[2].big_s, system.digits[2].big_t)
+ for i, v in enumerate(result_symbols):
+ if isinstance(v, CND) and v.value == 2:
+ next_symbol = result_symbols[i +
+ 1] if i < len(result_symbols) - 1 else None
+ previous_symbol = result_symbols[i - 1] if i > 0 else None
+ if isinstance(next_symbol, CNU) and isinstance(previous_symbol, (CNU, type(None))):
+ if next_symbol.power != 1 and ((previous_symbol is None) or (previous_symbol.power != 1)):
+ result_symbols[i] = liang
+
+ # if big is True, '涓�' will not be used and `alt_two` has no impact on output
+ if big:
+ attr_name = 'big_'
+ if traditional:
+ attr_name += 't'
+ else:
+ attr_name += 's'
+ else:
+ if traditional:
+ attr_name = 'traditional'
+ else:
+ attr_name = 'simplified'
+
+ result = ''.join([getattr(s, attr_name) for s in result_symbols])
+
+ # if not use_zeros:
+ # result = result.strip(getattr(system.digits[0], attr_name))
+
+ if alt_zero:
+ result = result.replace(
+ getattr(system.digits[0], attr_name), system.digits[0].alt_s)
+
+ if alt_one:
+ result = result.replace(
+ getattr(system.digits[1], attr_name), system.digits[1].alt_s)
+
+ for i, p in enumerate(POINT):
+ if result.startswith(p):
+ return CHINESE_DIGIS[0] + result
+
+ # ^10, 11, .., 19
+ if len(result) >= 2 and result[1] in [SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED[0],
+ SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL[0]] and \
+ result[0] in [CHINESE_DIGIS[1], BIG_CHINESE_DIGIS_SIMPLIFIED[1], BIG_CHINESE_DIGIS_TRADITIONAL[1]]:
+ result = result[1:]
+
+ return result
+
+
+# ================================================================================ #
+# different types of rewriters
+# ================================================================================ #
+class Cardinal:
+ """
+ CARDINAL绫�
+ """
+
+ def __init__(self, cardinal=None, chntext=None):
+ self.cardinal = cardinal
+ self.chntext = chntext
+
+ def chntext2cardinal(self):
+ return chn2num(self.chntext)
+
+ def cardinal2chntext(self):
+ return num2chn(self.cardinal)
+
+class Digit:
+ """
+ DIGIT绫�
+ """
+
+ def __init__(self, digit=None, chntext=None):
+ self.digit = digit
+ self.chntext = chntext
+
+ # def chntext2digit(self):
+ # return chn2num(self.chntext)
+
+ def digit2chntext(self):
+ return num2chn(self.digit, alt_two=False, use_units=False)
+
+
+class TelePhone:
+ """
+ TELEPHONE绫�
+ """
+
+ def __init__(self, telephone=None, raw_chntext=None, chntext=None):
+ self.telephone = telephone
+ self.raw_chntext = raw_chntext
+ self.chntext = chntext
+
+ # def chntext2telephone(self):
+ # sil_parts = self.raw_chntext.split('<SIL>')
+ # self.telephone = '-'.join([
+ # str(chn2num(p)) for p in sil_parts
+ # ])
+ # return self.telephone
+
+ def telephone2chntext(self, fixed=False):
+
+ if fixed:
+ sil_parts = self.telephone.split('-')
+ self.raw_chntext = '<SIL>'.join([
+ num2chn(part, alt_two=False, use_units=False) for part in sil_parts
+ ])
+ self.chntext = self.raw_chntext.replace('<SIL>', '')
+ else:
+ sp_parts = self.telephone.strip('+').split()
+ self.raw_chntext = '<SP>'.join([
+ num2chn(part, alt_two=False, use_units=False) for part in sp_parts
+ ])
+ self.chntext = self.raw_chntext.replace('<SP>', '')
+ return self.chntext
+
+
+class Fraction:
+ """
+ FRACTION绫�
+ """
+
+ def __init__(self, fraction=None, chntext=None):
+ self.fraction = fraction
+ self.chntext = chntext
+
+ def chntext2fraction(self):
+ denominator, numerator = self.chntext.split('鍒嗕箣')
+ return chn2num(numerator) + '/' + chn2num(denominator)
+
+ def fraction2chntext(self):
+ numerator, denominator = self.fraction.split('/')
+ return num2chn(denominator) + '鍒嗕箣' + num2chn(numerator)
+
+
+class Date:
+ """
+ DATE绫�
+ """
+
+ def __init__(self, date=None, chntext=None):
+ self.date = date
+ self.chntext = chntext
+
+ # def chntext2date(self):
+ # chntext = self.chntext
+ # try:
+ # year, other = chntext.strip().split('骞�', maxsplit=1)
+ # year = Digit(chntext=year).digit2chntext() + '骞�'
+ # except ValueError:
+ # other = chntext
+ # year = ''
+ # if other:
+ # try:
+ # month, day = other.strip().split('鏈�', maxsplit=1)
+ # month = Cardinal(chntext=month).chntext2cardinal() + '鏈�'
+ # except ValueError:
+ # day = chntext
+ # month = ''
+ # if day:
+ # day = Cardinal(chntext=day[:-1]).chntext2cardinal() + day[-1]
+ # else:
+ # month = ''
+ # day = ''
+ # date = year + month + day
+ # self.date = date
+ # return self.date
+
+ def date2chntext(self):
+ date = self.date
+ try:
+ year, other = date.strip().split('骞�', 1)
+ year = Digit(digit=year).digit2chntext() + '骞�'
+ except ValueError:
+ other = date
+ year = ''
+ if other:
+ try:
+ month, day = other.strip().split('鏈�', 1)
+ month = Cardinal(cardinal=month).cardinal2chntext() + '鏈�'
+ except ValueError:
+ day = date
+ month = ''
+ if day:
+ day = Cardinal(cardinal=day[:-1]).cardinal2chntext() + day[-1]
+ else:
+ month = ''
+ day = ''
+ chntext = year + month + day
+ self.chntext = chntext
+ return self.chntext
+
+
+class Money:
+ """
+ MONEY绫�
+ """
+
+ def __init__(self, money=None, chntext=None):
+ self.money = money
+ self.chntext = chntext
+
+ # def chntext2money(self):
+ # return self.money
+
+ def money2chntext(self):
+ money = self.money
+ pattern = re.compile(r'(\d+(\.\d+)?)')
+ matchers = pattern.findall(money)
+ if matchers:
+ for matcher in matchers:
+ money = money.replace(matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext())
+ self.chntext = money
+ return self.chntext
+
+
+class Percentage:
+ """
+ PERCENTAGE绫�
+ """
+
+ def __init__(self, percentage=None, chntext=None):
+ self.percentage = percentage
+ self.chntext = chntext
+
+ def chntext2percentage(self):
+ return chn2num(self.chntext.strip().strip('鐧惧垎涔�')) + '%'
+
+ def percentage2chntext(self):
+ return '鐧惧垎涔�' + num2chn(self.percentage.strip().strip('%'))
+
+
+def remove_erhua(text, er_whitelist):
+ """
+ 鍘婚櫎鍎垮寲闊宠瘝涓殑鍎�:
+ 浠栧コ鍎垮湪閭h竟鍎� -> 浠栧コ鍎垮湪閭h竟
+ """
+
+ er_pattern = re.compile(er_whitelist)
+ new_str=''
+ while re.search('鍎�',text):
+ a = re.search('鍎�',text).span()
+ remove_er_flag = 0
+
+ if er_pattern.search(text):
+ b = er_pattern.search(text).span()
+ if b[0] <= a[0]:
+ remove_er_flag = 1
+
+ if remove_er_flag == 0 :
+ new_str = new_str + text[0:a[0]]
+ text = text[a[1]:]
+ else:
+ new_str = new_str + text[0:b[1]]
+ text = text[b[1]:]
+
+ text = new_str + text
+ return text
+
+# ================================================================================ #
+# NSW Normalizer
+# ================================================================================ #
+class NSWNormalizer:
+ def __init__(self, raw_text):
+ self.raw_text = '^' + raw_text + '$'
+ self.norm_text = ''
+
+ def _particular(self):
+ text = self.norm_text
+ pattern = re.compile(r"(([a-zA-Z]+)浜�([a-zA-Z]+))")
+ matchers = pattern.findall(text)
+ if matchers:
+ # print('particular')
+ for matcher in matchers:
+ text = text.replace(matcher[0], matcher[1]+'2'+matcher[2], 1)
+ self.norm_text = text
+ return self.norm_text
+
+ def normalize(self):
+ text = self.raw_text
+
+ # 瑙勮寖鍖栨棩鏈�
+ pattern = re.compile(r"\D+((([089]\d|(19|20)\d{2})骞�)?(\d{1,2}鏈�(\d{1,2}[鏃ュ彿])?)?)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('date')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Date(date=matcher[0]).date2chntext(), 1)
+
+ # 瑙勮寖鍖栭噾閽�
+ pattern = re.compile(r"\D+((\d+(\.\d+)?)[澶氫綑鍑燷?" + CURRENCY_UNITS + r"(\d" + CURRENCY_UNITS + r"?)?)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('money')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Money(money=matcher[0]).money2chntext(), 1)
+
+ # 瑙勮寖鍖栧浐璇�/鎵嬫満鍙风爜
+ # 鎵嬫満
+ # http://www.jihaoba.com/news/show/13680
+ # 绉诲姩锛�139銆�138銆�137銆�136銆�135銆�134銆�159銆�158銆�157銆�150銆�151銆�152銆�188銆�187銆�182銆�183銆�184銆�178銆�198
+ # 鑱旈�氾細130銆�131銆�132銆�156銆�155銆�186銆�185銆�176
+ # 鐢典俊锛�133銆�153銆�189銆�180銆�181銆�177
+ pattern = re.compile(r"\D((\+?86 ?)?1([38]\d|5[0-35-9]|7[678]|9[89])\d{8})\D")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('telephone')
+ for matcher in matchers:
+ text = text.replace(matcher[0], TelePhone(telephone=matcher[0]).telephone2chntext(), 1)
+ # 鍥鸿瘽
+ pattern = re.compile(r"\D((0(10|2[1-3]|[3-9]\d{2})-?)?[1-9]\d{6,7})\D")
+ matchers = pattern.findall(text)
+ if matchers:
+ # print('fixed telephone')
+ for matcher in matchers:
+ text = text.replace(matcher[0], TelePhone(telephone=matcher[0]).telephone2chntext(fixed=True), 1)
+
+ # 瑙勮寖鍖栧垎鏁�
+ pattern = re.compile(r"(\d+/\d+)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('fraction')
+ for matcher in matchers:
+ text = text.replace(matcher, Fraction(fraction=matcher).fraction2chntext(), 1)
+
+ # 瑙勮寖鍖栫櫨鍒嗘暟
+ text = text.replace('锛�', '%')
+ pattern = re.compile(r"(\d+(\.\d+)?%)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('percentage')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Percentage(percentage=matcher[0]).percentage2chntext(), 1)
+
+ # 瑙勮寖鍖栫函鏁�+閲忚瘝
+ pattern = re.compile(r"(\d+(\.\d+)?)[澶氫綑鍑燷?" + COM_QUANTIFIERS)
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('cardinal+quantifier')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext(), 1)
+
+ # 瑙勮寖鍖栨暟瀛楃紪鍙�
+ pattern = re.compile(r"(\d{4,32})")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('digit')
+ for matcher in matchers:
+ text = text.replace(matcher, Digit(digit=matcher).digit2chntext(), 1)
+
+ # 瑙勮寖鍖栫函鏁�
+ pattern = re.compile(r"(\d+(\.\d+)?)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('cardinal')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext(), 1)
+
+ self.norm_text = text
+ self._particular()
+
+ return self.norm_text.lstrip('^').rstrip('$')
+
+
+def nsw_test_case(raw_text):
+ print('I:' + raw_text)
+ print('O:' + NSWNormalizer(raw_text).normalize())
+ print('')
+
+
+def nsw_test():
+ nsw_test_case('鍥鸿瘽锛�0595-23865596鎴�23880880銆�')
+ nsw_test_case('鍥鸿瘽锛�0595-23865596鎴�23880880銆�')
+ nsw_test_case('鎵嬫満锛�+86 19859213959鎴�15659451527銆�')
+ nsw_test_case('鍒嗘暟锛�32477/76391銆�')
+ nsw_test_case('鐧惧垎鏁帮細80.03%銆�')
+ nsw_test_case('缂栧彿锛�31520181154418銆�')
+ nsw_test_case('绾暟锛�2983.07鍏嬫垨12345.60绫炽��')
+ nsw_test_case('鏃ユ湡锛�1999骞�2鏈�20鏃ユ垨09骞�3鏈�15鍙枫��')
+ nsw_test_case('閲戦挶锛�12鍧�5锛�34.5鍏冿紝20.1涓�')
+ nsw_test_case('鐗规畩锛歄2O鎴朆2C銆�')
+ nsw_test_case('3456涓囧惃')
+ nsw_test_case('2938涓�')
+ nsw_test_case('938')
+ nsw_test_case('浠婂ぉ鍚冧簡115涓皬绗煎寘231涓澶�')
+ nsw_test_case('鏈�62锛呯殑姒傜巼')
+
+
+if __name__ == '__main__':
+ #nsw_test()
+
+ p = argparse.ArgumentParser()
+ p.add_argument('ifile', help='input filename, assume utf-8 encoding')
+ p.add_argument('ofile', help='output filename')
+ p.add_argument('--to_upper', action='store_true', help='convert to upper case')
+ p.add_argument('--to_lower', action='store_true', help='convert to lower case')
+ p.add_argument('--has_key', action='store_true', help="input text has Kaldi's key as first field.")
+ p.add_argument('--remove_fillers', type=bool, default=True, help='remove filler chars such as "鍛�, 鍟�"')
+ p.add_argument('--remove_erhua', type=bool, default=True, help='remove erhua chars such as "杩欏効"')
+ p.add_argument('--log_interval', type=int, default=10000, help='log interval in number of processed lines')
+ args = p.parse_args()
+
+ ifile = codecs.open(args.ifile, 'r', 'utf8')
+ ofile = codecs.open(args.ofile, 'w+', 'utf8')
+
+ n = 0
+ for l in ifile:
+ key = ''
+ text = ''
+ if args.has_key:
+ cols = l.split(maxsplit=1)
+ key = cols[0]
+ if len(cols) == 2:
+ text = cols[1].strip()
+ else:
+ text = ''
+ else:
+ text = l.strip()
+
+ # cases
+ if args.to_upper and args.to_lower:
+ sys.stderr.write('text norm: to_upper OR to_lower?')
+ exit(1)
+ if args.to_upper:
+ text = text.upper()
+ if args.to_lower:
+ text = text.lower()
+
+ # Filler chars removal
+ if args.remove_fillers:
+ for ch in FILLER_CHARS:
+ text = text.replace(ch, '')
+
+ if args.remove_erhua:
+ text = remove_erhua(text, ER_WHITELIST)
+
+ # NSW(Non-Standard-Word) normalization
+ text = NSWNormalizer(text).normalize()
+
+ # Punctuations removal
+ old_chars = CHINESE_PUNC_LIST + string.punctuation # includes all CN and EN punctuations
+ new_chars = ' ' * len(old_chars)
+ del_chars = ''
+ text = text.translate(str.maketrans(old_chars, new_chars, del_chars))
+
+ #
+ if args.has_key:
+ ofile.write(key + '\t' + text + '\n')
+ else:
+ ofile.write(text + '\n')
+
+ n += 1
+ if n % args.log_interval == 0:
+ sys.stderr.write("text norm: {} lines done.\n".format(n))
+
+ sys.stderr.write("text norm: {} lines done in total.\n".format(n))
+
+ ifile.close()
+ ofile.close()
diff --git a/egs/aishell/paraformerbert/conf/train_asr_paraformer_conformer_12e_6d_2048_256.yaml b/egs/aishell/paraformerbert/conf/train_asr_paraformer_conformer_12e_6d_2048_256.yaml
deleted file mode 100644
index 779c7a9..0000000
--- a/egs/aishell/paraformerbert/conf/train_asr_paraformer_conformer_12e_6d_2048_256.yaml
+++ /dev/null
@@ -1,92 +0,0 @@
-# network architecture
-# encoder related
-encoder: conformer
-encoder_conf:
- output_size: 256 # dimension of attention
- attention_heads: 4
- linear_units: 2048 # the number of units of position-wise feed forward
- num_blocks: 12 # the number of encoder blocks
- dropout_rate: 0.1
- positional_dropout_rate: 0.1
- attention_dropout_rate: 0.0
- input_layer: conv2d # encoder architecture type
- normalize_before: true
- pos_enc_layer_type: rel_pos
- selfattention_layer_type: rel_selfattn
- activation_type: swish
- macaron_style: true
- use_cnn_module: true
- cnn_module_kernel: 15
-
-# decoder related
-decoder: paraformer_decoder_san
-decoder_conf:
- attention_heads: 4
- linear_units: 2048
- num_blocks: 6
- dropout_rate: 0.1
- positional_dropout_rate: 0.1
- self_attention_dropout_rate: 0.0
- src_attention_dropout_rate: 0.0
-
-# hybrid CTC/attention
-model: paraformer
-model_conf:
- ctc_weight: 0.3
- lsm_weight: 0.1 # label smoothing option
- length_normalized_loss: false
- predictor_weight: 1.0
- sampling_ratio: 0.4
-
-# minibatch related
-batch_type: length
-batch_bins: 25000
-num_workers: 16
-
-# optimization related
-accum_grad: 1
-grad_clip: 5
-max_epoch: 50
-val_scheduler_criterion:
- - valid
- - acc
-best_model_criterion:
-- - valid
- - acc
- - max
-keep_nbest_models: 10
-
-optim: adam
-optim_conf:
- lr: 0.0005
-scheduler: warmuplr
-scheduler_conf:
- warmup_steps: 30000
-
-specaug: specaug
-specaug_conf:
- apply_time_warp: true
- time_warp_window: 5
- time_warp_mode: bicubic
- apply_freq_mask: true
- freq_mask_width_range:
- - 0
- - 30
- num_freq_mask: 2
- apply_time_mask: true
- time_mask_width_range:
- - 0
- - 40
- num_time_mask: 2
-
-predictor: cif_predictor
-predictor_conf:
- idim: 256
- threshold: 1.0
- l_order: 1
- r_order: 1
- tail_threshold: 0.45
-
-
-log_interval: 50
-normalize: None
\ No newline at end of file
diff --git a/egs/aishell/paraformerbert/conf/train_asr_paraformer_conformer_20e_6d_1280_320.yaml b/egs/aishell/paraformerbert/conf/train_asr_paraformer_conformer_20e_6d_1280_320.yaml
deleted file mode 100644
index 29b9ca6..0000000
--- a/egs/aishell/paraformerbert/conf/train_asr_paraformer_conformer_20e_6d_1280_320.yaml
+++ /dev/null
@@ -1,94 +0,0 @@
-# network architecture
-# encoder related
-encoder: conformer
-encoder_conf:
- output_size: 320 # dimension of attention
- attention_heads: 4
- linear_units: 1280 # the number of units of position-wise feed forward
- num_blocks: 20 # the number of encoder blocks
- dropout_rate: 0.1
- positional_dropout_rate: 0.1
- attention_dropout_rate: 0.0
- input_layer: conv2d # encoder architecture type
- normalize_before: true
- pos_enc_layer_type: rel_pos
- selfattention_layer_type: rel_selfattn
- activation_type: swish
- macaron_style: true
- use_cnn_module: true
- cnn_module_kernel: 15
-
-# decoder related
-decoder: paraformer_decoder_san
-decoder_conf:
- attention_heads: 4
- linear_units: 2048
- num_blocks: 6
- dropout_rate: 0.1
- positional_dropout_rate: 0.1
- self_attention_dropout_rate: 0.0
- src_attention_dropout_rate: 0.0
-
-# hybrid CTC/attention
-model: paraformer
-model_conf:
- ctc_weight: 0.3
- lsm_weight: 0.1 # label smoothing option
- length_normalized_loss: false
- predictor_weight: 1.0
- sampling_ratio: 0.4
-
-# minibatch related
-#batch_type: length
-#batch_bins: 40000
-batch_type: numel
-batch_bins: 2000000
-num_workers: 16
-
-# optimization related
-accum_grad: 4
-grad_clip: 5
-max_epoch: 50
-val_scheduler_criterion:
- - valid
- - acc
-best_model_criterion:
-- - valid
- - acc
- - max
-keep_nbest_models: 10
-
-optim: adam
-optim_conf:
- lr: 0.0005
-scheduler: warmuplr
-scheduler_conf:
- warmup_steps: 30000
-
-specaug: specaug
-specaug_conf:
- apply_time_warp: true
- time_warp_window: 5
- time_warp_mode: bicubic
- apply_freq_mask: true
- freq_mask_width_range:
- - 0
- - 30
- num_freq_mask: 2
- apply_time_mask: true
- time_mask_width_range:
- - 0
- - 40
- num_time_mask: 2
-
-predictor: cif_predictor
-predictor_conf:
- idim: 256
- threshold: 1.0
- l_order: 1
- r_order: 1
- tail_threshold: 0.45
-
-
-log_interval: 50
-normalize: None
\ No newline at end of file
diff --git a/egs/aishell/paraformerbert/run.sh b/egs/aishell/paraformerbert/run.sh
index 15d659c..6f331ec 100755
--- a/egs/aishell/paraformerbert/run.sh
+++ b/egs/aishell/paraformerbert/run.sh
@@ -10,9 +10,11 @@
# for gpu decoding, inference_nj=ngpu*njob; for cpu decoding, inference_nj=njob
njob=8
train_cmd=utils/run.pl
+infer_cmd=utils/run.pl
# general configuration
-feats_dir=".." #feature output dictionary, for large data
+feats_dir="../DATA" #feature output dictionary, for large data
+exp_dir="."
lang=zh
dumpdir=dump/fbank
feats_type=fbank
@@ -51,11 +53,9 @@
test_sets="dev test"
asr_config=conf/train_asr_paraformerbert_conformer_12e_6d_2048_256.yaml
-run_dir="exp"
model_dir="baseline_$(basename "${asr_config}" .yaml)_${feats_type}_${lang}_${token_type}_${tag}"
-exp_dir=$run_dir/$model_dir
-inference_config=conf/decode_asr_transformer.yaml
+inference_config=conf/decode_asr_transformer_noctc_1best.yaml
inference_asr_model=valid.acc.ave_10best.pth
# you can set gpu num for decoding here
@@ -64,20 +64,22 @@
if ${gpu_inference}; then
inference_nj=$[${ngpu}*${njob}]
+ _ngpu=1
else
inference_nj=$njob
+ _ngpu=0
fi
if [ ${stage} -le 0 ] && [ ${stop_stage} -ge 0 ]; then
echo "stage 0: Data preparation"
# Data preparation
- local/aishell_data_prep.sh ${data_aishell}/data_aishell/wav ${data_aishell}/data_aishell/transcript
+ local/aishell_data_prep.sh ${data_aishell}/data_aishell/wav ${data_aishell}/data_aishell/transcript ${feats_dir}
for x in train dev test; do
- cp data/${x}/text data/${x}/text.org
- paste -d " " <(cut -f 1 -d" " data/${x}/text.org) <(cut -f 2- -d" " data/${x}/text.org | tr -d " ") \
- > data/${x}/text
- utils/text2token.py -n 1 -s 1 data/${x}/text > data/${x}/text.org
- mv data/${x}/text.org data/${x}/text
+ cp ${feats_dir}/data/${x}/text ${feats_dir}/data/${x}/text.org
+ paste -d " " <(cut -f 1 -d" " ${feats_dir}/data/${x}/text.org) <(cut -f 2- -d" " ${feats_dir}/data/${x}/text.org | tr -d " ") \
+ > ${feats_dir}/data/${x}/text
+ utils/text2token.py -n 1 -s 1 ${feats_dir}/data/${x}/text > ${feats_dir}/data/${x}/text.org
+ mv ${feats_dir}/data/${x}/text.org ${feats_dir}/data/${x}/text
done
fi
@@ -88,27 +90,27 @@
echo "stage 1: Feature Generation"
# compute fbank features
fbankdir=${feats_dir}/fbank
- utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj --speed_perturb ${speed_perturb} \
- data/train exp/make_fbank/train ${fbankdir}/train
+ utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj --feats_dim ${feats_dim} --sample_frequency ${sample_frequency} --speed_perturb ${speed_perturb} \
+ ${feats_dir}/data/train ${exp_dir}/exp/make_fbank/train ${fbankdir}/train
utils/fix_data_feat.sh ${fbankdir}/train
- utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj \
- data/dev exp/make_fbank/dev ${fbankdir}/dev
+ utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj --feats_dim ${feats_dim} --sample_frequency ${sample_frequency} \
+ ${feats_dir}/data/dev ${exp_dir}/exp/make_fbank/dev ${fbankdir}/dev
utils/fix_data_feat.sh ${fbankdir}/dev
- utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj \
- data/test exp/make_fbank/test ${fbankdir}/test
+ utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj --feats_dim ${feats_dim} --sample_frequency ${sample_frequency} \
+ ${feats_dir}/data/test ${exp_dir}/exp/make_fbank/test ${fbankdir}/test
utils/fix_data_feat.sh ${fbankdir}/test
# compute global cmvn
- utils/compute_cmvn.sh --cmd "$train_cmd" --nj $nj \
- ${fbankdir}/train exp/make_fbank/train
+ utils/compute_cmvn.sh --cmd "$train_cmd" --nj $nj --feats_dim ${feats_dim} \
+ ${fbankdir}/train ${exp_dir}/exp/make_fbank/train
# apply cmvn
utils/apply_cmvn.sh --cmd "$train_cmd" --nj $nj \
- ${fbankdir}/train ${fbankdir}/train/cmvn.json exp/make_fbank/train ${feat_train_dir}
+ ${fbankdir}/train ${fbankdir}/train/cmvn.json ${exp_dir}/exp/make_fbank/train ${feat_train_dir}
utils/apply_cmvn.sh --cmd "$train_cmd" --nj $nj \
- ${fbankdir}/dev ${fbankdir}/train/cmvn.json exp/make_fbank/dev ${feat_dev_dir}
+ ${fbankdir}/dev ${fbankdir}/train/cmvn.json ${exp_dir}/exp/make_fbank/dev ${feat_dev_dir}
utils/apply_cmvn.sh --cmd "$train_cmd" --nj $nj \
- ${fbankdir}/test ${fbankdir}/train/cmvn.json exp/make_fbank/test ${feat_test_dir}
+ ${fbankdir}/test ${fbankdir}/train/cmvn.json ${exp_dir}/exp/make_fbank/test ${feat_test_dir}
cp ${fbankdir}/train/text ${fbankdir}/train/speech_shape ${fbankdir}/train/text_shape ${feat_train_dir}
cp ${fbankdir}/dev/text ${fbankdir}/dev/speech_shape ${fbankdir}/dev/text_shape ${feat_dev_dir}
@@ -117,29 +119,33 @@
utils/fix_data_feat.sh ${feat_train_dir}
utils/fix_data_feat.sh ${feat_dev_dir}
utils/fix_data_feat.sh ${feat_test_dir}
+
+ #generate ark list
+ utils/gen_ark_list.sh --cmd "$train_cmd" --nj $nj ${feat_train_dir} ${fbankdir}/train ${feat_train_dir}
+ utils/gen_ark_list.sh --cmd "$train_cmd" --nj $nj ${feat_dev_dir} ${fbankdir}/dev ${feat_dev_dir}
fi
token_list=${feats_dir}/data/${lang}_token_list/char/tokens.txt
echo "dictionary: ${token_list}"
if [ ${stage} -le 2 ] && [ ${stop_stage} -ge 2 ]; then
echo "stage 2: Dictionary Preparation"
- mkdir -p data/${lang}_token_list/char/
+ mkdir -p ${feats_dir}/data/${lang}_token_list/char/
echo "make a dictionary"
echo "<blank>" > ${token_list}
echo "<s>" >> ${token_list}
echo "</s>" >> ${token_list}
- utils/text2token.py -s 1 -n 1 --space "" data/train/text | cut -f 2- -d" " | tr " " "\n" \
+ utils/text2token.py -s 1 -n 1 --space "" ${feats_dir}/data/train/text | cut -f 2- -d" " | tr " " "\n" \
| sort | uniq | grep -a -v -e '^\s*$' | awk '{print $0}' >> ${token_list}
num_token=$(cat ${token_list} | wc -l)
echo "<unk>" >> ${token_list}
vocab_size=$(cat ${token_list} | wc -l)
awk -v v=,${vocab_size} '{print $0v}' ${feat_train_dir}/text_shape > ${feat_train_dir}/text_shape.char
awk -v v=,${vocab_size} '{print $0v}' ${feat_dev_dir}/text_shape > ${feat_dev_dir}/text_shape.char
- mkdir -p asr_stats_fbank_zh_char/train
- mkdir -p asr_stats_fbank_zh_char/dev
- cp ${feat_train_dir}/speech_shape ${feat_train_dir}/text_shape ${feat_train_dir}/text_shape.char asr_stats_fbank_zh_char/train
- cp ${feat_dev_dir}/speech_shape ${feat_dev_dir}/text_shape ${feat_dev_dir}/text_shape.char asr_stats_fbank_zh_char/dev
+ mkdir -p ${feats_dir}/asr_stats_fbank_zh_char/train
+ mkdir -p ${feats_dir}/asr_stats_fbank_zh_char/dev
+ cp ${feat_train_dir}/speech_shape ${feat_train_dir}/text_shape ${feat_train_dir}/text_shape.char ${feats_dir}/asr_stats_fbank_zh_char/train
+ cp ${feat_dev_dir}/speech_shape ${feat_dev_dir}/text_shape ${feat_dev_dir}/text_shape.char ${feats_dir}/asr_stats_fbank_zh_char/dev
fi
if ! "${skip_extract_embed}"; then
@@ -152,9 +158,10 @@
# Training Stage
world_size=$gpu_num # run on one machine
if [ ${stage} -le 3 ] && [ ${stop_stage} -ge 3 ]; then
- mkdir -p $exp_dir
- mkdir -p $exp_dir/log
- INIT_FILE=$exp_dir/ddp_init
+ echo "stage 3: Training"
+ mkdir -p ${exp_dir}/exp/${model_dir}
+ mkdir -p ${exp_dir}/exp/${model_dir}/log
+ INIT_FILE=${exp_dir}/exp/${model_dir}/ddp_init
if [ -f $INIT_FILE ];then
rm -f $INIT_FILE
fi
@@ -183,7 +190,7 @@
--valid_shape_file ${feats_dir}/asr_stats_fbank_zh_char/${valid_set}/text_shape.char \
--valid_shape_file ${feats_dir}/embeds/${bert_model_name}/${valid_set}/embeds.shape \
--resume true \
- --output_dir $exp_dir \
+ --output_dir ${exp_dir}/exp/${model_dir} \
--config $asr_config \
--input_size $feats_dim \
--ngpu $gpu_num \
@@ -201,26 +208,57 @@
# Testing Stage
if [ ${stage} -le 4 ] && [ ${stop_stage} -ge 4 ]; then
- utils/easy_asr_infer.sh \
- --lang zh \
- --datadir ${feats_dir} \
- --feats_type ${feats_type} \
- --feats_dim ${feats_dim} \
- --token_type ${token_type} \
- --gpu_inference ${gpu_inference} \
- --inference_config "${inference_config}" \
- --test_sets "${test_sets}" \
- --token_list $token_list \
- --asr_exp $exp_dir \
- --stage 12 \
- --stop_stage 12 \
- --scp $scp \
- --text text \
- --inference_nj $inference_nj \
- --njob $njob \
- --inference_asr_model $inference_asr_model \
- --gpuid_list $gpuid_list \
- --gpu_inference ${gpu_inference} \
- --mode paraformer
+ echo "stage 4: Inference"
+ for dset in ${test_sets}; do
+ asr_exp=${exp_dir}/exp/${model_dir}
+ inference_tag="$(basename "${inference_config}" .yaml)"
+ _dir="${asr_exp}/${inference_tag}/${inference_asr_model}/${dset}"
+ _logdir="${_dir}/logdir"
+ if [ -d ${_dir} ]; then
+ echo "${_dir} is already exists. if you want to decode again, please delete this dir first."
+ exit 0
+ fi
+ mkdir -p "${_logdir}"
+ _data="${feats_dir}/${dumpdir}/${dset}"
+ key_file=${_data}/${scp}
+ num_scp_file="$(<${key_file} wc -l)"
+ _nj=$([ $inference_nj -le $num_scp_file ] && echo "$inference_nj" || echo "$num_scp_file")
+ split_scps=
+ for n in $(seq "${_nj}"); do
+ split_scps+=" ${_logdir}/keys.${n}.scp"
+ done
+ # shellcheck disable=SC2086
+ utils/split_scp.pl "${key_file}" ${split_scps}
+ _opts=
+ if [ -n "${inference_config}" ]; then
+ _opts+="--config ${inference_config} "
+ fi
+ ${infer_cmd} --gpu "${_ngpu}" --max-jobs-run "${_nj}" JOB=1:"${_nj}" "${_logdir}"/asr_inference.JOB.log \
+ python -m funasr.bin.asr_inference_launch \
+ --batch_size 1 \
+ --ngpu "${_ngpu}" \
+ --njob ${njob} \
+ --gpuid_list ${gpuid_list} \
+ --data_path_and_name_and_type "${_data}/${scp},speech,${type}" \
+ --key_file "${_logdir}"/keys.JOB.scp \
+ --asr_train_config "${asr_exp}"/config.yaml \
+ --asr_model_file "${asr_exp}"/"${inference_asr_model}" \
+ --output_dir "${_logdir}"/output.JOB \
+ --mode paraformer \
+ ${_opts}
+
+ for f in token token_int score text; do
+ if [ -f "${_logdir}/output.1/1best_recog/${f}" ]; then
+ for i in $(seq "${_nj}"); do
+ cat "${_logdir}/output.${i}/1best_recog/${f}"
+ done | sort -k1 >"${_dir}/${f}"
+ fi
+ done
+ python utils/proce_text.py ${_dir}/text ${_dir}/text.proc
+ python utils/proce_text.py ${_data}/text ${_data}/text.proc
+ python utils/compute_wer.py ${_data}/text.proc ${_dir}/text.proc ${_dir}/text.cer
+ tail -n 3 ${_dir}/text.cer > ${_dir}/text.cer.txt
+ cat ${_dir}/text.cer.txt
+ done
fi
diff --git a/egs/aishell/paraformerbert/utils b/egs/aishell/paraformerbert/utils
deleted file mode 120000
index 40e14f5..0000000
--- a/egs/aishell/paraformerbert/utils
+++ /dev/null
@@ -1 +0,0 @@
-../tranformer/utils
\ No newline at end of file
diff --git a/egs/aishell/paraformerbert/utils/__init__.py b/egs/aishell/paraformerbert/utils/__init__.py
new file mode 100644
index 0000000..e69de29
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/__init__.py
diff --git a/egs/aishell/paraformerbert/utils/apply_cmvn.py b/egs/aishell/paraformerbert/utils/apply_cmvn.py
new file mode 100755
index 0000000..b5c5086
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/apply_cmvn.py
@@ -0,0 +1,79 @@
+from kaldiio import ReadHelper
+from kaldiio import WriteHelper
+
+import argparse
+import json
+import math
+import numpy as np
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="apply cmvn",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--ark-file",
+ "-a",
+ default=False,
+ required=True,
+ type=str,
+ help="fbank ark file",
+ )
+ parser.add_argument(
+ "--cmvn-file",
+ "-c",
+ default=False,
+ required=True,
+ type=str,
+ help="cmvn file",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ ark_file = args.output_dir + "/feats." + str(args.ark_index) + ".ark"
+ scp_file = args.output_dir + "/feats." + str(args.ark_index) + ".scp"
+ ark_writer = WriteHelper('ark,scp:{},{}'.format(ark_file, scp_file))
+
+ with open(args.cmvn_file) as f:
+ cmvn_stats = json.load(f)
+
+ means = cmvn_stats['mean_stats']
+ vars = cmvn_stats['var_stats']
+ total_frames = cmvn_stats['total_frames']
+
+ for i in range(len(means)):
+ means[i] /= total_frames
+ vars[i] = vars[i] / total_frames - means[i] * means[i]
+ if vars[i] < 1.0e-20:
+ vars[i] = 1.0e-20
+ vars[i] = 1.0 / math.sqrt(vars[i])
+
+ with ReadHelper('ark:{}'.format(args.ark_file)) as ark_reader:
+ for key, mat in ark_reader:
+ mat = (mat - means) * vars
+ ark_writer(key, mat)
+
+
+if __name__ == '__main__':
+ main()
diff --git a/egs/aishell/paraformerbert/utils/apply_cmvn.sh b/egs/aishell/paraformerbert/utils/apply_cmvn.sh
new file mode 100755
index 0000000..f8fd1d1
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/apply_cmvn.sh
@@ -0,0 +1,29 @@
+#!/usr/bin/env bash
+
+. ./path.sh || exit 1;
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+fbankdir=$1
+cmvn_file=$2
+logdir=$3
+output_dir=$4
+
+dump_dir=${output_dir}/ark; mkdir -p ${dump_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/apply_cmvn.JOB.log \
+ python utils/apply_cmvn.py -a $fbankdir/ark/feats.JOB.ark \
+ -c $cmvn_file -i JOB -o ${dump_dir} \
+ || exit 1;
+
+for n in $(seq $nj); do
+ cat ${dump_dir}/feats.$n.scp || exit 1
+done > ${output_dir}/feats.scp || exit 1
+
+echo "$0: Succeeded apply cmvn"
diff --git a/egs/aishell/paraformerbert/utils/apply_lfr_and_cmvn.py b/egs/aishell/paraformerbert/utils/apply_lfr_and_cmvn.py
new file mode 100755
index 0000000..50d18d1
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/apply_lfr_and_cmvn.py
@@ -0,0 +1,143 @@
+from kaldiio import ReadHelper, WriteHelper
+
+import argparse
+import numpy as np
+
+
+def build_LFR_features(inputs, m=7, n=6):
+ LFR_inputs = []
+ T = inputs.shape[0]
+ T_lfr = int(np.ceil(T / n))
+ left_padding = np.tile(inputs[0], ((m - 1) // 2, 1))
+ inputs = np.vstack((left_padding, inputs))
+ T = T + (m - 1) // 2
+ for i in range(T_lfr):
+ if m <= T - i * n:
+ LFR_inputs.append(np.hstack(inputs[i * n:i * n + m]))
+ else:
+ num_padding = m - (T - i * n)
+ frame = np.hstack(inputs[i * n:])
+ for _ in range(num_padding):
+ frame = np.hstack((frame, inputs[-1]))
+ LFR_inputs.append(frame)
+ return np.vstack(LFR_inputs)
+
+
+def build_CMVN_features(inputs, mvn_file): # noqa
+ with open(mvn_file, 'r', encoding='utf-8') as f:
+ lines = f.readlines()
+
+ add_shift_list = []
+ rescale_list = []
+ for i in range(len(lines)):
+ line_item = lines[i].split()
+ if line_item[0] == '<AddShift>':
+ line_item = lines[i + 1].split()
+ if line_item[0] == '<LearnRateCoef>':
+ add_shift_line = line_item[3:(len(line_item) - 1)]
+ add_shift_list = list(add_shift_line)
+ continue
+ elif line_item[0] == '<Rescale>':
+ line_item = lines[i + 1].split()
+ if line_item[0] == '<LearnRateCoef>':
+ rescale_line = line_item[3:(len(line_item) - 1)]
+ rescale_list = list(rescale_line)
+ continue
+
+ for j in range(inputs.shape[0]):
+ for k in range(inputs.shape[1]):
+ add_shift_value = add_shift_list[k]
+ rescale_value = rescale_list[k]
+ inputs[j, k] = float(inputs[j, k]) + float(add_shift_value)
+ inputs[j, k] = float(inputs[j, k]) * float(rescale_value)
+
+ return inputs
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="apply low_frame_rate and cmvn",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--ark-file",
+ "-a",
+ default=False,
+ required=True,
+ type=str,
+ help="fbank ark file",
+ )
+ parser.add_argument(
+ "--lfr",
+ "-f",
+ default=True,
+ type=str,
+ help="low frame rate",
+ )
+ parser.add_argument(
+ "--lfr-m",
+ "-m",
+ default=7,
+ type=int,
+ help="number of frames to stack",
+ )
+ parser.add_argument(
+ "--lfr-n",
+ "-n",
+ default=6,
+ type=int,
+ help="number of frames to skip",
+ )
+ parser.add_argument(
+ "--cmvn-file",
+ "-c",
+ default=False,
+ required=True,
+ type=str,
+ help="global cmvn file",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ dump_ark_file = args.output_dir + "/feats." + str(args.ark_index) + ".ark"
+ dump_scp_file = args.output_dir + "/feats." + str(args.ark_index) + ".scp"
+ shape_file = args.output_dir + "/len." + str(args.ark_index)
+ ark_writer = WriteHelper('ark,scp:{},{}'.format(dump_ark_file, dump_scp_file))
+
+ shape_writer = open(shape_file, 'w')
+ with ReadHelper('ark:{}'.format(args.ark_file)) as ark_reader:
+ for key, mat in ark_reader:
+ if args.lfr:
+ lfr = build_LFR_features(mat, args.lfr_m, args.lfr_n)
+ else:
+ lfr = mat
+ cmvn = build_CMVN_features(lfr, args.cmvn_file)
+ dims = cmvn.shape[1]
+ lens = cmvn.shape[0]
+ shape_writer.write(key + " " + str(lens) + "," + str(dims) + '\n')
+ ark_writer(key, cmvn)
+
+
+if __name__ == '__main__':
+ main()
+
diff --git a/egs/aishell/paraformerbert/utils/apply_lfr_and_cmvn.sh b/egs/aishell/paraformerbert/utils/apply_lfr_and_cmvn.sh
new file mode 100755
index 0000000..3119fdb
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/apply_lfr_and_cmvn.sh
@@ -0,0 +1,38 @@
+#!/usr/bin/env bash
+
+
+# Begin configuration section.
+nj=32
+cmd=utils/run.pl
+
+# feature configuration
+lfr=True
+lfr_m=7
+lfr_n=6
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+fbankdir=$1
+cmvn_file=$2
+logdir=$3
+output_dir=$4
+
+dump_dir=${output_dir}/ark; mkdir -p ${dump_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/apply_lfr_and_cmvn.JOB.log \
+ python utils/apply_lfr_and_cmvn.py -a $fbankdir/ark/feats.JOB.ark \
+ -f $lfr -m $lfr_m -n $lfr_n -c $cmvn_file -i JOB -o ${dump_dir} \
+ || exit 1;
+
+for n in $(seq $nj); do
+ cat ${dump_dir}/feats.$n.scp || exit 1
+done > ${output_dir}/feats.scp || exit 1
+
+for n in $(seq $nj); do
+ cat ${dump_dir}/len.$n || exit 1
+done > ${output_dir}/speech_shape || exit 1
+
+echo "$0: Succeeded apply low frame rate and cmvn"
diff --git a/egs/aishell/paraformerbert/utils/combine_cmvn_file.py b/egs/aishell/paraformerbert/utils/combine_cmvn_file.py
new file mode 100755
index 0000000..b2974a4
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/combine_cmvn_file.py
@@ -0,0 +1,73 @@
+import argparse
+import json
+import numpy as np
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="combine cmvn file",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--dims",
+ "-d",
+ default=80,
+ type=int,
+ help="feature dims",
+ )
+ parser.add_argument(
+ "--cmvn-dir",
+ "-c",
+ default=False,
+ required=True,
+ type=str,
+ help="cmvn dir",
+ )
+
+ parser.add_argument(
+ "--nj",
+ "-n",
+ default=1,
+ required=True,
+ type=int,
+ help="num of cmvn file",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ total_means = np.zeros(args.dims)
+ total_vars = np.zeros(args.dims)
+ total_frames = 0
+
+ cmvn_file = args.output_dir + "/cmvn.json"
+
+ for i in range(1, args.nj+1):
+ with open(args.cmvn_dir + "/cmvn." + str(i) + ".json", "r") as fin:
+ cmvn_stats = json.load(fin)
+
+ total_means += np.array(cmvn_stats["mean_stats"])
+ total_vars += np.array(cmvn_stats["var_stats"])
+ total_frames += cmvn_stats["total_frames"]
+
+ cmvn_info = {
+ 'mean_stats': list(total_means.tolist()),
+ 'var_stats': list(total_vars.tolist()),
+ 'total_frames': total_frames
+ }
+ with open(cmvn_file, 'w') as fout:
+ fout.write(json.dumps(cmvn_info))
+
+
+if __name__ == '__main__':
+ main()
diff --git a/egs/aishell/paraformerbert/utils/compute_cmvn.py b/egs/aishell/paraformerbert/utils/compute_cmvn.py
new file mode 100755
index 0000000..2b96e26
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/compute_cmvn.py
@@ -0,0 +1,74 @@
+from kaldiio import ReadHelper
+
+import argparse
+import numpy as np
+import json
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="computer global cmvn",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--dims",
+ "-d",
+ default=80,
+ type=int,
+ help="feature dims",
+ )
+ parser.add_argument(
+ "--ark-file",
+ "-a",
+ default=False,
+ required=True,
+ type=str,
+ help="fbank ark file",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ ark_file = args.ark_file + "/feats." + str(args.ark_index) + ".ark"
+ cmvn_file = args.output_dir + "/cmvn." + str(args.ark_index) + ".json"
+
+ mean_stats = np.zeros(args.dims)
+ var_stats = np.zeros(args.dims)
+ total_frames = 0
+
+ with ReadHelper('ark:{}'.format(ark_file)) as ark_reader:
+ for key, mat in ark_reader:
+ mean_stats += np.sum(mat, axis=0)
+ var_stats += np.sum(np.square(mat), axis=0)
+ total_frames += mat.shape[0]
+
+ cmvn_info = {
+ 'mean_stats': list(mean_stats.tolist()),
+ 'var_stats': list(var_stats.tolist()),
+ 'total_frames': total_frames
+ }
+ with open(cmvn_file, 'w') as fout:
+ fout.write(json.dumps(cmvn_info))
+
+
+if __name__ == '__main__':
+ main()
diff --git a/egs/aishell/paraformerbert/utils/compute_cmvn.sh b/egs/aishell/paraformerbert/utils/compute_cmvn.sh
new file mode 100755
index 0000000..12173ee
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/compute_cmvn.sh
@@ -0,0 +1,25 @@
+#!/usr/bin/env bash
+
+. ./path.sh || exit 1;
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+feats_dim=80
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+fbankdir=$1
+logdir=$2
+
+output_dir=${fbankdir}/cmvn; mkdir -p ${output_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/cmvn.JOB.log \
+ python utils/compute_cmvn.py -d ${feats_dim} -a $fbankdir/ark -i JOB -o ${output_dir} \
+ || exit 1;
+
+python utils/combine_cmvn_file.py -d ${feats_dim} -c ${output_dir} -n $nj -o $fbankdir
+
+echo "$0: Succeeded compute global cmvn"
diff --git a/egs/aishell/paraformerbert/utils/compute_fbank.py b/egs/aishell/paraformerbert/utils/compute_fbank.py
new file mode 100755
index 0000000..d03b5a8
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/compute_fbank.py
@@ -0,0 +1,153 @@
+from kaldiio import WriteHelper
+
+import argparse
+import numpy as np
+import json
+import torch
+import torchaudio
+import torchaudio.compliance.kaldi as kaldi
+
+
+def compute_fbank(wav_file,
+ num_mel_bins=80,
+ frame_length=25,
+ frame_shift=10,
+ dither=0.0,
+ resample_rate=16000,
+ speed=1.0):
+
+ waveform, sample_rate = torchaudio.load(wav_file)
+ if resample_rate != sample_rate:
+ waveform = torchaudio.transforms.Resample(orig_freq=sample_rate,
+ new_freq=resample_rate)(waveform)
+ if speed != 1.0:
+ waveform, _ = torchaudio.sox_effects.apply_effects_tensor(
+ waveform, resample_rate,
+ [['speed', str(speed)], ['rate', str(resample_rate)]]
+ )
+
+ waveform = waveform * (1 << 15)
+ mat = kaldi.fbank(waveform,
+ num_mel_bins=num_mel_bins,
+ frame_length=frame_length,
+ frame_shift=frame_shift,
+ dither=dither,
+ energy_floor=0.0,
+ window_type='hamming',
+ sample_frequency=resample_rate)
+
+ return mat.numpy()
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="computer features",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--wav-lists",
+ "-w",
+ default=False,
+ required=True,
+ type=str,
+ help="input wav lists",
+ )
+ parser.add_argument(
+ "--text-files",
+ "-t",
+ default=False,
+ required=True,
+ type=str,
+ help="input text files",
+ )
+ parser.add_argument(
+ "--dims",
+ "-d",
+ default=80,
+ type=int,
+ help="feature dims",
+ )
+ parser.add_argument(
+ "--sample-frequency",
+ "-s",
+ default=16000,
+ type=int,
+ help="sample frequency",
+ )
+ parser.add_argument(
+ "--speed-perturb",
+ "-p",
+ default="1.0",
+ type=str,
+ help="speed perturb",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-a",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ ark_file = args.output_dir + "/ark/feats." + str(args.ark_index) + ".ark"
+ scp_file = args.output_dir + "/ark/feats." + str(args.ark_index) + ".scp"
+ text_file = args.output_dir + "/txt/text." + str(args.ark_index) + ".txt"
+ feats_shape_file = args.output_dir + "/ark/len." + str(args.ark_index)
+ text_shape_file = args.output_dir + "/txt/len." + str(args.ark_index)
+
+ ark_writer = WriteHelper('ark,scp:{},{}'.format(ark_file, scp_file))
+ text_writer = open(text_file, 'w')
+ feats_shape_writer = open(feats_shape_file, 'w')
+ text_shape_writer = open(text_shape_file, 'w')
+
+ speed_perturb_list = args.speed_perturb.split(',')
+
+ for speed in speed_perturb_list:
+ with open(args.wav_lists, 'r', encoding='utf-8') as wavfile:
+ with open(args.text_files, 'r', encoding='utf-8') as textfile:
+ for wav, text in zip(wavfile, textfile):
+ s_w = wav.strip().split()
+ wav_id = s_w[0]
+ wav_file = s_w[1]
+
+ s_t = text.strip().split()
+ text_id = s_t[0]
+ txt = s_t[1:]
+ fbank = compute_fbank(wav_file,
+ num_mel_bins=args.dims,
+ resample_rate=args.sample_frequency,
+ speed=float(speed)
+ )
+ feats_dims = fbank.shape[1]
+ feats_lens = fbank.shape[0]
+ txt_lens = len(txt)
+ if speed == "1.0":
+ wav_id_sp = wav_id
+ else:
+ wav_id_sp = wav_id + "_sp" + speed
+
+ feats_shape_writer.write(wav_id_sp + " " + str(feats_lens) + "," + str(feats_dims) + '\n')
+ text_shape_writer.write(wav_id_sp + " " + str(txt_lens) + '\n')
+
+ text_writer.write(wav_id_sp + " " + " ".join(txt) + '\n')
+ ark_writer(wav_id_sp, fbank)
+
+
+if __name__ == '__main__':
+ main()
+
diff --git a/egs/aishell/paraformerbert/utils/compute_fbank.sh b/egs/aishell/paraformerbert/utils/compute_fbank.sh
new file mode 100755
index 0000000..92a4fe6
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/compute_fbank.sh
@@ -0,0 +1,51 @@
+#!/usr/bin/env bash
+
+. ./path.sh || exit 1;
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+
+# feature configuration
+feats_dim=80
+sample_frequency=16000
+speed_perturb="1.0"
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+data=$1
+logdir=$2
+fbankdir=$3
+
+[ ! -f $data/wav.scp ] && echo "$0: no such file $data/wav.scp" && exit 1;
+[ ! -f $data/text ] && echo "$0: no such file $data/text" && exit 1;
+
+python utils/split_data.py $data $data $nj
+
+ark_dir=${fbankdir}/ark; mkdir -p ${ark_dir}
+text_dir=${fbankdir}/txt; mkdir -p ${text_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/make_fbank.JOB.log \
+ python utils/compute_fbank.py -w $data/split${nj}/JOB/wav.scp -t $data/split${nj}/JOB/text \
+ -d $feats_dim -s $sample_frequency -p ${speed_perturb} -a JOB -o ${fbankdir} \
+ || exit 1;
+
+for n in $(seq $nj); do
+ cat ${ark_dir}/feats.$n.scp || exit 1
+done > $fbankdir/feats.scp || exit 1
+
+for n in $(seq $nj); do
+ cat ${text_dir}/text.$n.txt || exit 1
+done > $fbankdir/text || exit 1
+
+for n in $(seq $nj); do
+ cat ${ark_dir}/len.$n || exit 1
+done > $fbankdir/speech_shape || exit 1
+
+for n in $(seq $nj); do
+ cat ${text_dir}/len.$n || exit 1
+done > $fbankdir/text_shape || exit 1
+
+echo "$0: Succeeded compute FBANK features"
diff --git a/egs/aishell/paraformerbert/utils/compute_wer.py b/egs/aishell/paraformerbert/utils/compute_wer.py
new file mode 100755
index 0000000..349a3f6
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/compute_wer.py
@@ -0,0 +1,157 @@
+import os
+import numpy as np
+import sys
+
+def compute_wer(ref_file,
+ hyp_file,
+ cer_detail_file):
+ rst = {
+ 'Wrd': 0,
+ 'Corr': 0,
+ 'Ins': 0,
+ 'Del': 0,
+ 'Sub': 0,
+ 'Snt': 0,
+ 'Err': 0.0,
+ 'S.Err': 0.0,
+ 'wrong_words': 0,
+ 'wrong_sentences': 0
+ }
+
+ hyp_dict = {}
+ ref_dict = {}
+ with open(hyp_file, 'r') as hyp_reader:
+ for line in hyp_reader:
+ key = line.strip().split()[0]
+ value = line.strip().split()[1:]
+ hyp_dict[key] = value
+ with open(ref_file, 'r') as ref_reader:
+ for line in ref_reader:
+ key = line.strip().split()[0]
+ value = line.strip().split()[1:]
+ ref_dict[key] = value
+
+ cer_detail_writer = open(cer_detail_file, 'w')
+ for hyp_key in hyp_dict:
+ if hyp_key in ref_dict:
+ out_item = compute_wer_by_line(hyp_dict[hyp_key], ref_dict[hyp_key])
+ rst['Wrd'] += out_item['nwords']
+ rst['Corr'] += out_item['cor']
+ rst['wrong_words'] += out_item['wrong']
+ rst['Ins'] += out_item['ins']
+ rst['Del'] += out_item['del']
+ rst['Sub'] += out_item['sub']
+ rst['Snt'] += 1
+ if out_item['wrong'] > 0:
+ rst['wrong_sentences'] += 1
+ cer_detail_writer.write(hyp_key + print_cer_detail(out_item) + '\n')
+ cer_detail_writer.write("ref:" + '\t' + "".join(ref_dict[hyp_key]) + '\n')
+ cer_detail_writer.write("hyp:" + '\t' + "".join(hyp_dict[hyp_key]) + '\n')
+
+ if rst['Wrd'] > 0:
+ rst['Err'] = round(rst['wrong_words'] * 100 / rst['Wrd'], 2)
+ if rst['Snt'] > 0:
+ rst['S.Err'] = round(rst['wrong_sentences'] * 100 / rst['Snt'], 2)
+
+ cer_detail_writer.write('\n')
+ cer_detail_writer.write("%WER " + str(rst['Err']) + " [ " + str(rst['wrong_words'])+ " / " + str(rst['Wrd']) +
+ ", " + str(rst['Ins']) + " ins, " + str(rst['Del']) + " del, " + str(rst['Sub']) + " sub ]" + '\n')
+ cer_detail_writer.write("%SER " + str(rst['S.Err']) + " [ " + str(rst['wrong_sentences']) + " / " + str(rst['Snt']) + " ]" + '\n')
+ cer_detail_writer.write("Scored " + str(len(hyp_dict)) + " sentences, " + str(len(hyp_dict) - rst['Snt']) + " not present in hyp." + '\n')
+
+
+def compute_wer_by_line(hyp,
+ ref):
+ hyp = list(map(lambda x: x.lower(), hyp))
+ ref = list(map(lambda x: x.lower(), ref))
+
+ len_hyp = len(hyp)
+ len_ref = len(ref)
+
+ cost_matrix = np.zeros((len_hyp + 1, len_ref + 1), dtype=np.int16)
+
+ ops_matrix = np.zeros((len_hyp + 1, len_ref + 1), dtype=np.int8)
+
+ for i in range(len_hyp + 1):
+ cost_matrix[i][0] = i
+ for j in range(len_ref + 1):
+ cost_matrix[0][j] = j
+
+ for i in range(1, len_hyp + 1):
+ for j in range(1, len_ref + 1):
+ if hyp[i - 1] == ref[j - 1]:
+ cost_matrix[i][j] = cost_matrix[i - 1][j - 1]
+ else:
+ substitution = cost_matrix[i - 1][j - 1] + 1
+ insertion = cost_matrix[i - 1][j] + 1
+ deletion = cost_matrix[i][j - 1] + 1
+
+ compare_val = [substitution, insertion, deletion]
+
+ min_val = min(compare_val)
+ operation_idx = compare_val.index(min_val) + 1
+ cost_matrix[i][j] = min_val
+ ops_matrix[i][j] = operation_idx
+
+ match_idx = []
+ i = len_hyp
+ j = len_ref
+ rst = {
+ 'nwords': len_ref,
+ 'cor': 0,
+ 'wrong': 0,
+ 'ins': 0,
+ 'del': 0,
+ 'sub': 0
+ }
+ while i >= 0 or j >= 0:
+ i_idx = max(0, i)
+ j_idx = max(0, j)
+
+ if ops_matrix[i_idx][j_idx] == 0: # correct
+ if i - 1 >= 0 and j - 1 >= 0:
+ match_idx.append((j - 1, i - 1))
+ rst['cor'] += 1
+
+ i -= 1
+ j -= 1
+
+ elif ops_matrix[i_idx][j_idx] == 2: # insert
+ i -= 1
+ rst['ins'] += 1
+
+ elif ops_matrix[i_idx][j_idx] == 3: # delete
+ j -= 1
+ rst['del'] += 1
+
+ elif ops_matrix[i_idx][j_idx] == 1: # substitute
+ i -= 1
+ j -= 1
+ rst['sub'] += 1
+
+ if i < 0 and j >= 0:
+ rst['del'] += 1
+ elif j < 0 and i >= 0:
+ rst['ins'] += 1
+
+ match_idx.reverse()
+ wrong_cnt = cost_matrix[len_hyp][len_ref]
+ rst['wrong'] = wrong_cnt
+
+ return rst
+
+def print_cer_detail(rst):
+ return ("(" + "nwords=" + str(rst['nwords']) + ",cor=" + str(rst['cor'])
+ + ",ins=" + str(rst['ins']) + ",del=" + str(rst['del']) + ",sub="
+ + str(rst['sub']) + ") corr:" + '{:.2%}'.format(rst['cor']/rst['nwords'])
+ + ",cer:" + '{:.2%}'.format(rst['wrong']/rst['nwords']))
+
+if __name__ == '__main__':
+ if len(sys.argv) != 4:
+ print("usage : python compute-wer.py test.ref test.hyp test.wer")
+ sys.exit(0)
+
+ ref_file = sys.argv[1]
+ hyp_file = sys.argv[2]
+ cer_detail_file = sys.argv[3]
+ compute_wer(ref_file, hyp_file, cer_detail_file)
diff --git a/egs/aishell/paraformerbert/utils/error_rate_zh b/egs/aishell/paraformerbert/utils/error_rate_zh
new file mode 100755
index 0000000..6871a07
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/error_rate_zh
@@ -0,0 +1,370 @@
+#!/usr/bin/env python3
+# coding=utf8
+
+# Copyright 2021 Jiayu DU
+
+import sys
+import argparse
+import json
+import logging
+logging.basicConfig(stream=sys.stderr, level=logging.INFO, format='[%(levelname)s] %(message)s')
+
+DEBUG = None
+
+def GetEditType(ref_token, hyp_token):
+ if ref_token == None and hyp_token != None:
+ return 'I'
+ elif ref_token != None and hyp_token == None:
+ return 'D'
+ elif ref_token == hyp_token:
+ return 'C'
+ elif ref_token != hyp_token:
+ return 'S'
+ else:
+ raise RuntimeError
+
+class AlignmentArc:
+ def __init__(self, src, dst, ref, hyp):
+ self.src = src
+ self.dst = dst
+ self.ref = ref
+ self.hyp = hyp
+ self.edit_type = GetEditType(ref, hyp)
+
+def similarity_score_function(ref_token, hyp_token):
+ return 0 if (ref_token == hyp_token) else -1.0
+
+def insertion_score_function(token):
+ return -1.0
+
+def deletion_score_function(token):
+ return -1.0
+
+def EditDistance(
+ ref,
+ hyp,
+ similarity_score_function = similarity_score_function,
+ insertion_score_function = insertion_score_function,
+ deletion_score_function = deletion_score_function):
+ assert(len(ref) != 0)
+ class DPState:
+ def __init__(self):
+ self.score = -float('inf')
+ # backpointer
+ self.prev_r = None
+ self.prev_h = None
+
+ def print_search_grid(S, R, H, fstream):
+ print(file=fstream)
+ for r in range(R):
+ for h in range(H):
+ print(F'[{r},{h}]:{S[r][h].score:4.3f}:({S[r][h].prev_r},{S[r][h].prev_h}) ', end='', file=fstream)
+ print(file=fstream)
+
+ R = len(ref) + 1
+ H = len(hyp) + 1
+
+ # Construct DP search space, a (R x H) grid
+ S = [ [] for r in range(R) ]
+ for r in range(R):
+ S[r] = [ DPState() for x in range(H) ]
+
+ # initialize DP search grid origin, S(r = 0, h = 0)
+ S[0][0].score = 0.0
+ S[0][0].prev_r = None
+ S[0][0].prev_h = None
+
+ # initialize REF axis
+ for r in range(1, R):
+ S[r][0].score = S[r-1][0].score + deletion_score_function(ref[r-1])
+ S[r][0].prev_r = r-1
+ S[r][0].prev_h = 0
+
+ # initialize HYP axis
+ for h in range(1, H):
+ S[0][h].score = S[0][h-1].score + insertion_score_function(hyp[h-1])
+ S[0][h].prev_r = 0
+ S[0][h].prev_h = h-1
+
+ best_score = S[0][0].score
+ best_state = (0, 0)
+
+ for r in range(1, R):
+ for h in range(1, H):
+ sub_or_cor_score = similarity_score_function(ref[r-1], hyp[h-1])
+ new_score = S[r-1][h-1].score + sub_or_cor_score
+ if new_score >= S[r][h].score:
+ S[r][h].score = new_score
+ S[r][h].prev_r = r-1
+ S[r][h].prev_h = h-1
+
+ del_score = deletion_score_function(ref[r-1])
+ new_score = S[r-1][h].score + del_score
+ if new_score >= S[r][h].score:
+ S[r][h].score = new_score
+ S[r][h].prev_r = r - 1
+ S[r][h].prev_h = h
+
+ ins_score = insertion_score_function(hyp[h-1])
+ new_score = S[r][h-1].score + ins_score
+ if new_score >= S[r][h].score:
+ S[r][h].score = new_score
+ S[r][h].prev_r = r
+ S[r][h].prev_h = h-1
+
+ best_score = S[R-1][H-1].score
+ best_state = (R-1, H-1)
+
+ if DEBUG:
+ print_search_grid(S, R, H, sys.stderr)
+
+ # Backtracing best alignment path, i.e. a list of arcs
+ # arc = (src, dst, ref, hyp, edit_type)
+ # src/dst = (r, h), where r/h refers to search grid state-id along Ref/Hyp axis
+ best_path = []
+ r, h = best_state[0], best_state[1]
+ prev_r, prev_h = S[r][h].prev_r, S[r][h].prev_h
+ score = S[r][h].score
+ # loop invariant:
+ # 1. (prev_r, prev_h) -> (r, h) is a "forward arc" on best alignment path
+ # 2. score is the value of point(r, h) on DP search grid
+ while prev_r != None or prev_h != None:
+ src = (prev_r, prev_h)
+ dst = (r, h)
+ if (r == prev_r + 1 and h == prev_h + 1): # Substitution or correct
+ arc = AlignmentArc(src, dst, ref[prev_r], hyp[prev_h])
+ elif (r == prev_r + 1 and h == prev_h): # Deletion
+ arc = AlignmentArc(src, dst, ref[prev_r], None)
+ elif (r == prev_r and h == prev_h + 1): # Insertion
+ arc = AlignmentArc(src, dst, None, hyp[prev_h])
+ else:
+ raise RuntimeError
+ best_path.append(arc)
+ r, h = prev_r, prev_h
+ prev_r, prev_h = S[r][h].prev_r, S[r][h].prev_h
+ score = S[r][h].score
+
+ best_path.reverse()
+ return (best_path, best_score)
+
+def PrettyPrintAlignment(alignment, stream = sys.stderr):
+ def get_token_str(token):
+ if token == None:
+ return "*"
+ return token
+
+ def is_double_width_char(ch):
+ if (ch >= '\u4e00') and (ch <= '\u9fa5'): # codepoint ranges for Chinese chars
+ return True
+ # TODO: support other double-width-char language such as Japanese, Korean
+ else:
+ return False
+
+ def display_width(token_str):
+ m = 0
+ for c in token_str:
+ if is_double_width_char(c):
+ m += 2
+ else:
+ m += 1
+ return m
+
+ R = ' REF : '
+ H = ' HYP : '
+ E = ' EDIT : '
+ for arc in alignment:
+ r = get_token_str(arc.ref)
+ h = get_token_str(arc.hyp)
+ e = arc.edit_type if arc.edit_type != 'C' else ''
+
+ nr, nh, ne = display_width(r), display_width(h), display_width(e)
+ n = max(nr, nh, ne) + 1
+
+ R += r + ' ' * (n-nr)
+ H += h + ' ' * (n-nh)
+ E += e + ' ' * (n-ne)
+
+ print(R, file=stream)
+ print(H, file=stream)
+ print(E, file=stream)
+
+def CountEdits(alignment):
+ c, s, i, d = 0, 0, 0, 0
+ for arc in alignment:
+ if arc.edit_type == 'C':
+ c += 1
+ elif arc.edit_type == 'S':
+ s += 1
+ elif arc.edit_type == 'I':
+ i += 1
+ elif arc.edit_type == 'D':
+ d += 1
+ else:
+ raise RuntimeError
+ return (c, s, i, d)
+
+def ComputeTokenErrorRate(c, s, i, d):
+ return 100.0 * (s + d + i) / (s + d + c)
+
+def ComputeSentenceErrorRate(num_err_utts, num_utts):
+ assert(num_utts != 0)
+ return 100.0 * num_err_utts / num_utts
+
+
+class EvaluationResult:
+ def __init__(self):
+ self.num_ref_utts = 0
+ self.num_hyp_utts = 0
+ self.num_eval_utts = 0 # seen in both ref & hyp
+ self.num_hyp_without_ref = 0
+
+ self.C = 0
+ self.S = 0
+ self.I = 0
+ self.D = 0
+ self.token_error_rate = 0.0
+
+ self.num_utts_with_error = 0
+ self.sentence_error_rate = 0.0
+
+ def to_json(self):
+ return json.dumps(self.__dict__)
+
+ def to_kaldi(self):
+ info = (
+ F'%WER {self.token_error_rate:.2f} [ {self.S + self.D + self.I} / {self.C + self.S + self.D}, {self.I} ins, {self.D} del, {self.S} sub ]\n'
+ F'%SER {self.sentence_error_rate:.2f} [ {self.num_utts_with_error} / {self.num_eval_utts} ]\n'
+ )
+ return info
+
+ def to_sclite(self):
+ return "TODO"
+
+ def to_espnet(self):
+ return "TODO"
+
+ def to_summary(self):
+ #return json.dumps(self.__dict__, indent=4)
+ summary = (
+ '==================== Overall Statistics ====================\n'
+ F'num_ref_utts: {self.num_ref_utts}\n'
+ F'num_hyp_utts: {self.num_hyp_utts}\n'
+ F'num_hyp_without_ref: {self.num_hyp_without_ref}\n'
+ F'num_eval_utts: {self.num_eval_utts}\n'
+ F'sentence_error_rate: {self.sentence_error_rate:.2f}%\n'
+ F'token_error_rate: {self.token_error_rate:.2f}%\n'
+ F'token_stats:\n'
+ F' - tokens:{self.C + self.S + self.D:>7}\n'
+ F' - edits: {self.S + self.I + self.D:>7}\n'
+ F' - cor: {self.C:>7}\n'
+ F' - sub: {self.S:>7}\n'
+ F' - ins: {self.I:>7}\n'
+ F' - del: {self.D:>7}\n'
+ '============================================================\n'
+ )
+ return summary
+
+
+class Utterance:
+ def __init__(self, uid, text):
+ self.uid = uid
+ self.text = text
+
+
+def LoadUtterances(filepath, format):
+ utts = {}
+ if format == 'text': # utt_id word1 word2 ...
+ with open(filepath, 'r', encoding='utf8') as f:
+ for line in f:
+ line = line.strip()
+ if line:
+ cols = line.split(maxsplit=1)
+ assert(len(cols) == 2 or len(cols) == 1)
+ uid = cols[0]
+ text = cols[1] if len(cols) == 2 else ''
+ if utts.get(uid) != None:
+ raise RuntimeError(F'Found duplicated utterence id {uid}')
+ utts[uid] = Utterance(uid, text)
+ else:
+ raise RuntimeError(F'Unsupported text format {format}')
+ return utts
+
+
+def tokenize_text(text, tokenizer):
+ if tokenizer == 'whitespace':
+ return text.split()
+ elif tokenizer == 'char':
+ return [ ch for ch in ''.join(text.split()) ]
+ else:
+ raise RuntimeError(F'ERROR: Unsupported tokenizer {tokenizer}')
+
+
+if __name__ == '__main__':
+ parser = argparse.ArgumentParser()
+ # optional
+ parser.add_argument('--tokenizer', choices=['whitespace', 'char'], default='whitespace', help='whitespace for WER, char for CER')
+ parser.add_argument('--ref-format', choices=['text'], default='text', help='reference format, first col is utt_id, the rest is text')
+ parser.add_argument('--hyp-format', choices=['text'], default='text', help='hypothesis format, first col is utt_id, the rest is text')
+ # required
+ parser.add_argument('--ref', type=str, required=True, help='input reference file')
+ parser.add_argument('--hyp', type=str, required=True, help='input hypothesis file')
+
+ parser.add_argument('result_file', type=str)
+ args = parser.parse_args()
+ logging.info(args)
+
+ ref_utts = LoadUtterances(args.ref, args.ref_format)
+ hyp_utts = LoadUtterances(args.hyp, args.hyp_format)
+
+ r = EvaluationResult()
+
+ # check valid utterances in hyp that have matched non-empty reference
+ eval_utts = []
+ r.num_hyp_without_ref = 0
+ for uid in sorted(hyp_utts.keys()):
+ if uid in ref_utts.keys(): # TODO: efficiency
+ if ref_utts[uid].text.strip(): # non-empty reference
+ eval_utts.append(uid)
+ else:
+ logging.warn(F'Found {uid} with empty reference, skipping...')
+ else:
+ logging.warn(F'Found {uid} without reference, skipping...')
+ r.num_hyp_without_ref += 1
+
+ r.num_hyp_utts = len(hyp_utts)
+ r.num_ref_utts = len(ref_utts)
+ r.num_eval_utts = len(eval_utts)
+
+ with open(args.result_file, 'w+', encoding='utf8') as fo:
+ for uid in eval_utts:
+ ref = ref_utts[uid]
+ hyp = hyp_utts[uid]
+
+ alignment, score = EditDistance(
+ tokenize_text(ref.text, args.tokenizer),
+ tokenize_text(hyp.text, args.tokenizer)
+ )
+
+ c, s, i, d = CountEdits(alignment)
+ utt_ter = ComputeTokenErrorRate(c, s, i, d)
+
+ # utt-level evaluation result
+ print(F'{{"uid":{uid}, "score":{score}, "ter":{utt_ter:.2f}, "cor":{c}, "sub":{s}, "ins":{i}, "del":{d}}}', file=fo)
+ PrettyPrintAlignment(alignment, fo)
+
+ r.C += c
+ r.S += s
+ r.I += i
+ r.D += d
+
+ if utt_ter > 0:
+ r.num_utts_with_error += 1
+
+ # corpus level evaluation result
+ r.sentence_error_rate = ComputeSentenceErrorRate(r.num_utts_with_error, r.num_eval_utts)
+ r.token_error_rate = ComputeTokenErrorRate(r.C, r.S, r.I, r.D)
+
+ print(r.to_summary(), file=fo)
+
+ print(r.to_json())
+ print(r.to_kaldi())
diff --git a/egs/aishell/paraformerbert/utils/extract_embeds.py b/egs/aishell/paraformerbert/utils/extract_embeds.py
new file mode 100755
index 0000000..7b817d8
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/extract_embeds.py
@@ -0,0 +1,47 @@
+from transformers import AutoTokenizer, AutoModel, pipeline
+import numpy as np
+import sys
+import os
+import torch
+from kaldiio import WriteHelper
+import re
+text_file_json = sys.argv[1]
+out_ark = sys.argv[2]
+out_scp = sys.argv[3]
+out_shape = sys.argv[4]
+device = int(sys.argv[5])
+model_path = sys.argv[6]
+
+model = AutoModel.from_pretrained(model_path)
+tokenizer = AutoTokenizer.from_pretrained(model_path)
+extractor = pipeline(task="feature-extraction", model=model, tokenizer=tokenizer, device=device)
+
+with open(text_file_json, 'r') as f:
+ js = f.readlines()
+
+
+f_shape = open(out_shape, "w")
+with WriteHelper('ark,scp:{},{}'.format(out_ark, out_scp)) as writer:
+ with torch.no_grad():
+ for idx, line in enumerate(js):
+ id, tokens = line.strip().split(" ", 1)
+ tokens = re.sub(" ", "", tokens.strip())
+ tokens = ' '.join([j for j in tokens])
+ token_num = len(tokens.split(" "))
+ outputs = extractor(tokens)
+ outputs = np.array(outputs)
+ embeds = outputs[0, 1:-1, :]
+
+ token_num_embeds, dim = embeds.shape
+ if token_num == token_num_embeds:
+ writer(id, embeds)
+ shape_line = "{} {},{}\n".format(id, token_num_embeds, dim)
+ f_shape.write(shape_line)
+ else:
+ print("{}, size has changed, {}, {}, {}".format(id, token_num, token_num_embeds, tokens))
+
+
+
+f_shape.close()
+
+
diff --git a/egs/aishell/paraformerbert/utils/filter_scp.pl b/egs/aishell/paraformerbert/utils/filter_scp.pl
new file mode 100755
index 0000000..003530d
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/filter_scp.pl
@@ -0,0 +1,87 @@
+#!/usr/bin/env perl
+# Copyright 2010-2012 Microsoft Corporation
+# Johns Hopkins University (author: Daniel Povey)
+
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+# This script takes a list of utterance-ids or any file whose first field
+# of each line is an utterance-id, and filters an scp
+# file (or any file whose "n-th" field is an utterance id), printing
+# out only those lines whose "n-th" field is in id_list. The index of
+# the "n-th" field is 1, by default, but can be changed by using
+# the -f <n> switch
+
+$exclude = 0;
+$field = 1;
+$shifted = 0;
+
+do {
+ $shifted=0;
+ if ($ARGV[0] eq "--exclude") {
+ $exclude = 1;
+ shift @ARGV;
+ $shifted=1;
+ }
+ if ($ARGV[0] eq "-f") {
+ $field = $ARGV[1];
+ shift @ARGV; shift @ARGV;
+ $shifted=1
+ }
+} while ($shifted);
+
+if(@ARGV < 1 || @ARGV > 2) {
+ die "Usage: filter_scp.pl [--exclude] [-f <field-to-filter-on>] id_list [in.scp] > out.scp \n" .
+ "Prints only the input lines whose f'th field (default: first) is in 'id_list'.\n" .
+ "Note: only the first field of each line in id_list matters. With --exclude, prints\n" .
+ "only the lines that were *not* in id_list.\n" .
+ "Caution: previously, the -f option was interpreted as a zero-based field index.\n" .
+ "If your older scripts (written before Oct 2014) stopped working and you used the\n" .
+ "-f option, add 1 to the argument.\n" .
+ "See also: scripts/filter_scp.pl .\n";
+}
+
+
+$idlist = shift @ARGV;
+open(F, "<$idlist") || die "Could not open id-list file $idlist";
+while(<F>) {
+ @A = split;
+ @A>=1 || die "Invalid id-list file line $_";
+ $seen{$A[0]} = 1;
+}
+
+if ($field == 1) { # Treat this as special case, since it is common.
+ while(<>) {
+ $_ =~ m/\s*(\S+)\s*/ || die "Bad line $_, could not get first field.";
+ # $1 is what we filter on.
+ if ((!$exclude && $seen{$1}) || ($exclude && !defined $seen{$1})) {
+ print $_;
+ }
+ }
+} else {
+ while(<>) {
+ @A = split;
+ @A > 0 || die "Invalid scp file line $_";
+ @A >= $field || die "Invalid scp file line $_";
+ if ((!$exclude && $seen{$A[$field-1]}) || ($exclude && !defined $seen{$A[$field-1]})) {
+ print $_;
+ }
+ }
+}
+
+# tests:
+# the following should print "foo 1"
+# ( echo foo 1; echo bar 2 ) | scripts/filter_scp.pl <(echo foo)
+# the following should print "bar 2".
+# ( echo foo 1; echo bar 2 ) | scripts/filter_scp.pl -f 2 <(echo 2)
diff --git a/egs/aishell/paraformerbert/utils/fix_data.sh b/egs/aishell/paraformerbert/utils/fix_data.sh
new file mode 100755
index 0000000..32cdde5
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/fix_data.sh
@@ -0,0 +1,35 @@
+#!/usr/bin/env bash
+
+echo "$0 $@"
+data_dir=$1
+
+if [ ! -f ${data_dir}/wav.scp ]; then
+ echo "$0: wav.scp is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/text ]; then
+ echo "$0: text is not found"
+ exit 1;
+fi
+
+
+
+mkdir -p ${data_dir}/.backup
+
+awk '{print $1}' ${data_dir}/wav.scp > ${data_dir}/.backup/wav_id
+awk '{print $1}' ${data_dir}/text > ${data_dir}/.backup/text_id
+
+sort ${data_dir}/.backup/wav_id ${data_dir}/.backup/text_id | uniq -d > ${data_dir}/.backup/id
+
+cp ${data_dir}/wav.scp ${data_dir}/.backup/wav.scp
+cp ${data_dir}/text ${data_dir}/.backup/text
+
+mv ${data_dir}/wav.scp ${data_dir}/wav.scp.bak
+mv ${data_dir}/text ${data_dir}/text.bak
+
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/wav.scp.bak > ${data_dir}/wav.scp
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/text.bak > ${data_dir}/text
+
+rm ${data_dir}/wav.scp.bak
+rm ${data_dir}/text.bak
diff --git a/egs/aishell/paraformerbert/utils/fix_data_feat.sh b/egs/aishell/paraformerbert/utils/fix_data_feat.sh
new file mode 100755
index 0000000..2c92d7f
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/fix_data_feat.sh
@@ -0,0 +1,52 @@
+#!/usr/bin/env bash
+
+echo "$0 $@"
+data_dir=$1
+
+if [ ! -f ${data_dir}/feats.scp ]; then
+ echo "$0: feats.scp is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/text ]; then
+ echo "$0: text is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/speech_shape ]; then
+ echo "$0: feature lengths is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/text_shape ]; then
+ echo "$0: text lengths is not found"
+ exit 1;
+fi
+
+mkdir -p ${data_dir}/.backup
+
+awk '{print $1}' ${data_dir}/feats.scp > ${data_dir}/.backup/wav_id
+awk '{print $1}' ${data_dir}/text > ${data_dir}/.backup/text_id
+
+sort ${data_dir}/.backup/wav_id ${data_dir}/.backup/text_id | uniq -d > ${data_dir}/.backup/id
+
+cp ${data_dir}/feats.scp ${data_dir}/.backup/feats.scp
+cp ${data_dir}/text ${data_dir}/.backup/text
+cp ${data_dir}/speech_shape ${data_dir}/.backup/speech_shape
+cp ${data_dir}/text_shape ${data_dir}/.backup/text_shape
+
+mv ${data_dir}/feats.scp ${data_dir}/feats.scp.bak
+mv ${data_dir}/text ${data_dir}/text.bak
+mv ${data_dir}/speech_shape ${data_dir}/speech_shape.bak
+mv ${data_dir}/text_shape ${data_dir}/text_shape.bak
+
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/feats.scp.bak > ${data_dir}/feats.scp
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/text.bak > ${data_dir}/text
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/speech_shape.bak > ${data_dir}/speech_shape
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/text_shape.bak > ${data_dir}/text_shape
+
+rm ${data_dir}/feats.scp.bak
+rm ${data_dir}/text.bak
+rm ${data_dir}/speech_shape.bak
+rm ${data_dir}/text_shape.bak
+
diff --git a/egs/aishell/paraformerbert/utils/gen_ark_list.sh b/egs/aishell/paraformerbert/utils/gen_ark_list.sh
new file mode 100755
index 0000000..aebf356
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/gen_ark_list.sh
@@ -0,0 +1,22 @@
+#!/usr/bin/env bash
+
+
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+ark_dir=$1
+txt_dir=$2
+output_dir=$3
+
+[ ! -d ${ark_dir}/ark ] && echo "$0: ark data is required" && exit 1;
+[ ! -d ${txt_dir}/txt ] && echo "$0: txt data is required" && exit 1;
+
+for n in $(seq $nj); do
+ echo "${ark_dir}/ark/feats.$n.ark ${txt_dir}/txt/text.$n.txt" || exit 1
+done > ${output_dir}/ark_txt.scp || exit 1
+
diff --git a/egs/aishell/paraformerbert/utils/parse_options.sh b/egs/aishell/paraformerbert/utils/parse_options.sh
new file mode 100755
index 0000000..71fb9e5
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/parse_options.sh
@@ -0,0 +1,97 @@
+#!/usr/bin/env bash
+
+# Copyright 2012 Johns Hopkins University (Author: Daniel Povey);
+# Arnab Ghoshal, Karel Vesely
+
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+# Parse command-line options.
+# To be sourced by another script (as in ". parse_options.sh").
+# Option format is: --option-name arg
+# and shell variable "option_name" gets set to value "arg."
+# The exception is --help, which takes no arguments, but prints the
+# $help_message variable (if defined).
+
+
+###
+### The --config file options have lower priority to command line
+### options, so we need to import them first...
+###
+
+# Now import all the configs specified by command-line, in left-to-right order
+for ((argpos=1; argpos<$#; argpos++)); do
+ if [ "${!argpos}" == "--config" ]; then
+ argpos_plus1=$((argpos+1))
+ config=${!argpos_plus1}
+ [ ! -r $config ] && echo "$0: missing config '$config'" && exit 1
+ . $config # source the config file.
+ fi
+done
+
+
+###
+### Now we process the command line options
+###
+while true; do
+ [ -z "${1:-}" ] && break; # break if there are no arguments
+ case "$1" in
+ # If the enclosing script is called with --help option, print the help
+ # message and exit. Scripts should put help messages in $help_message
+ --help|-h) if [ -z "$help_message" ]; then echo "No help found." 1>&2;
+ else printf "$help_message\n" 1>&2 ; fi;
+ exit 0 ;;
+ --*=*) echo "$0: options to scripts must be of the form --name value, got '$1'"
+ exit 1 ;;
+ # If the first command-line argument begins with "--" (e.g. --foo-bar),
+ # then work out the variable name as $name, which will equal "foo_bar".
+ --*) name=`echo "$1" | sed s/^--// | sed s/-/_/g`;
+ # Next we test whether the variable in question is undefned-- if so it's
+ # an invalid option and we die. Note: $0 evaluates to the name of the
+ # enclosing script.
+ # The test [ -z ${foo_bar+xxx} ] will return true if the variable foo_bar
+ # is undefined. We then have to wrap this test inside "eval" because
+ # foo_bar is itself inside a variable ($name).
+ eval '[ -z "${'$name'+xxx}" ]' && echo "$0: invalid option $1" 1>&2 && exit 1;
+
+ oldval="`eval echo \\$$name`";
+ # Work out whether we seem to be expecting a Boolean argument.
+ if [ "$oldval" == "true" ] || [ "$oldval" == "false" ]; then
+ was_bool=true;
+ else
+ was_bool=false;
+ fi
+
+ # Set the variable to the right value-- the escaped quotes make it work if
+ # the option had spaces, like --cmd "queue.pl -sync y"
+ eval $name=\"$2\";
+
+ # Check that Boolean-valued arguments are really Boolean.
+ if $was_bool && [[ "$2" != "true" && "$2" != "false" ]]; then
+ echo "$0: expected \"true\" or \"false\": $1 $2" 1>&2
+ exit 1;
+ fi
+ shift 2;
+ ;;
+ *) break;
+ esac
+done
+
+
+# Check for an empty argument to the --cmd option, which can easily occur as a
+# result of scripting errors.
+[ ! -z "${cmd+xxx}" ] && [ -z "$cmd" ] && echo "$0: empty argument to --cmd option" 1>&2 && exit 1;
+
+
+true; # so this script returns exit code 0.
diff --git a/egs/aishell/paraformerbert/utils/print_args.py b/egs/aishell/paraformerbert/utils/print_args.py
new file mode 100755
index 0000000..b0c61e5
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/print_args.py
@@ -0,0 +1,45 @@
+#!/usr/bin/env python
+import sys
+
+
+def get_commandline_args(no_executable=True):
+ extra_chars = [
+ " ",
+ ";",
+ "&",
+ "|",
+ "<",
+ ">",
+ "?",
+ "*",
+ "~",
+ "`",
+ '"',
+ "'",
+ "\\",
+ "{",
+ "}",
+ "(",
+ ")",
+ ]
+
+ # Escape the extra characters for shell
+ argv = [
+ arg.replace("'", "'\\''")
+ if all(char not in arg for char in extra_chars)
+ else "'" + arg.replace("'", "'\\''") + "'"
+ for arg in sys.argv
+ ]
+
+ if no_executable:
+ return " ".join(argv[1:])
+ else:
+ return sys.executable + " " + " ".join(argv)
+
+
+def main():
+ print(get_commandline_args())
+
+
+if __name__ == "__main__":
+ main()
diff --git a/egs/aishell/paraformerbert/utils/proc_conf_oss.py b/egs/aishell/paraformerbert/utils/proc_conf_oss.py
new file mode 100755
index 0000000..c4a90c5
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/proc_conf_oss.py
@@ -0,0 +1,35 @@
+from pathlib import Path
+
+import torch
+import yaml
+
+
+class NoAliasSafeDumper(yaml.SafeDumper):
+ # Disable anchor/alias in yaml because looks ugly
+ def ignore_aliases(self, data):
+ return True
+
+
+def yaml_no_alias_safe_dump(data, stream=None, **kwargs):
+ """Safe-dump in yaml with no anchor/alias"""
+ return yaml.dump(
+ data, stream, allow_unicode=True, Dumper=NoAliasSafeDumper, **kwargs
+ )
+
+
+def gen_conf(file, out_dir):
+ conf = torch.load(file)["config"]
+ conf["oss_bucket"] = "null"
+ print(conf)
+ output_dir = Path(out_dir)
+ output_dir.mkdir(parents=True, exist_ok=True)
+ with (output_dir / "config.yaml").open("w", encoding="utf-8") as f:
+ yaml_no_alias_safe_dump(conf, f, indent=4, sort_keys=False)
+
+
+if __name__ == "__main__":
+ import sys
+
+ in_f = sys.argv[1]
+ out_f = sys.argv[2]
+ gen_conf(in_f, out_f)
diff --git a/egs/aishell/paraformerbert/utils/proce_text.py b/egs/aishell/paraformerbert/utils/proce_text.py
new file mode 100755
index 0000000..9e517a4
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/proce_text.py
@@ -0,0 +1,31 @@
+
+import sys
+import re
+
+in_f = sys.argv[1]
+out_f = sys.argv[2]
+
+
+with open(in_f, "r", encoding="utf-8") as f:
+ lines = f.readlines()
+
+with open(out_f, "w", encoding="utf-8") as f:
+ for line in lines:
+ outs = line.strip().split(" ", 1)
+ if len(outs) == 2:
+ idx, text = outs
+ text = re.sub("</s>", "", text)
+ text = re.sub("<s>", "", text)
+ text = re.sub("@@", "", text)
+ text = re.sub("@", "", text)
+ text = re.sub("<unk>", "", text)
+ text = re.sub(" ", "", text)
+ text = text.lower()
+ else:
+ idx = outs[0]
+ text = " "
+
+ text = [x for x in text]
+ text = " ".join(text)
+ out = "{} {}\n".format(idx, text)
+ f.write(out)
diff --git a/egs/aishell/paraformerbert/utils/run.pl b/egs/aishell/paraformerbert/utils/run.pl
new file mode 100755
index 0000000..483f95b
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/run.pl
@@ -0,0 +1,356 @@
+#!/usr/bin/env perl
+use warnings; #sed replacement for -w perl parameter
+# In general, doing
+# run.pl some.log a b c is like running the command a b c in
+# the bash shell, and putting the standard error and output into some.log.
+# To run parallel jobs (backgrounded on the host machine), you can do (e.g.)
+# run.pl JOB=1:4 some.JOB.log a b c JOB is like running the command a b c JOB
+# and putting it in some.JOB.log, for each one. [Note: JOB can be any identifier].
+# If any of the jobs fails, this script will fail.
+
+# A typical example is:
+# run.pl some.log my-prog "--opt=foo bar" foo \| other-prog baz
+# and run.pl will run something like:
+# ( my-prog '--opt=foo bar' foo | other-prog baz ) >& some.log
+#
+# Basically it takes the command-line arguments, quotes them
+# as necessary to preserve spaces, and evaluates them with bash.
+# In addition it puts the command line at the top of the log, and
+# the start and end times of the command at the beginning and end.
+# The reason why this is useful is so that we can create a different
+# version of this program that uses a queueing system instead.
+
+#use Data::Dumper;
+
+@ARGV < 2 && die "usage: run.pl log-file command-line arguments...";
+
+#print STDERR "COMMAND-LINE: " . Dumper(\@ARGV) . "\n";
+$job_pick = 'all';
+$max_jobs_run = -1;
+$jobstart = 1;
+$jobend = 1;
+$ignored_opts = ""; # These will be ignored.
+
+# First parse an option like JOB=1:4, and any
+# options that would normally be given to
+# queue.pl, which we will just discard.
+
+for (my $x = 1; $x <= 2; $x++) { # This for-loop is to
+ # allow the JOB=1:n option to be interleaved with the
+ # options to qsub.
+ while (@ARGV >= 2 && $ARGV[0] =~ m:^-:) {
+ # parse any options that would normally go to qsub, but which will be ignored here.
+ my $switch = shift @ARGV;
+ if ($switch eq "-V") {
+ $ignored_opts .= "-V ";
+ } elsif ($switch eq "--max-jobs-run" || $switch eq "-tc") {
+ # we do support the option --max-jobs-run n, and its GridEngine form -tc n.
+ # if the command appears multiple times uses the smallest option.
+ if ( $max_jobs_run <= 0 ) {
+ $max_jobs_run = shift @ARGV;
+ } else {
+ my $new_constraint = shift @ARGV;
+ if ( ($new_constraint < $max_jobs_run) ) {
+ $max_jobs_run = $new_constraint;
+ }
+ }
+
+ if (! ($max_jobs_run > 0)) {
+ die "run.pl: invalid option --max-jobs-run $max_jobs_run";
+ }
+ } else {
+ my $argument = shift @ARGV;
+ if ($argument =~ m/^--/) {
+ print STDERR "run.pl: WARNING: suspicious argument '$argument' to $switch; starts with '-'\n";
+ }
+ if ($switch eq "-sync" && $argument =~ m/^[yY]/) {
+ $ignored_opts .= "-sync "; # Note: in the
+ # corresponding code in queue.pl it says instead, just "$sync = 1;".
+ } elsif ($switch eq "-pe") { # e.g. -pe smp 5
+ my $argument2 = shift @ARGV;
+ $ignored_opts .= "$switch $argument $argument2 ";
+ } elsif ($switch eq "--gpu") {
+ $using_gpu = $argument;
+ } elsif ($switch eq "--pick") {
+ if($argument =~ m/^(all|failed|incomplete)$/) {
+ $job_pick = $argument;
+ } else {
+ print STDERR "run.pl: ERROR: --pick argument must be one of 'all', 'failed' or 'incomplete'"
+ }
+ } else {
+ # Ignore option.
+ $ignored_opts .= "$switch $argument ";
+ }
+ }
+ }
+ if ($ARGV[0] =~ m/^([\w_][\w\d_]*)+=(\d+):(\d+)$/) { # e.g. JOB=1:20
+ $jobname = $1;
+ $jobstart = $2;
+ $jobend = $3;
+ if ($jobstart > $jobend) {
+ die "run.pl: invalid job range $ARGV[0]";
+ }
+ if ($jobstart <= 0) {
+ die "run.pl: invalid job range $ARGV[0], start must be strictly positive (this is required for GridEngine compatibility).";
+ }
+ shift;
+ } elsif ($ARGV[0] =~ m/^([\w_][\w\d_]*)+=(\d+)$/) { # e.g. JOB=1.
+ $jobname = $1;
+ $jobstart = $2;
+ $jobend = $2;
+ shift;
+ } elsif ($ARGV[0] =~ m/.+\=.*\:.*$/) {
+ print STDERR "run.pl: Warning: suspicious first argument to run.pl: $ARGV[0]\n";
+ }
+}
+
+# Users found this message confusing so we are removing it.
+# if ($ignored_opts ne "") {
+# print STDERR "run.pl: Warning: ignoring options \"$ignored_opts\"\n";
+# }
+
+if ($max_jobs_run == -1) { # If --max-jobs-run option not set,
+ # then work out the number of processors if possible,
+ # and set it based on that.
+ $max_jobs_run = 0;
+ if ($using_gpu) {
+ if (open(P, "nvidia-smi -L |")) {
+ $max_jobs_run++ while (<P>);
+ close(P);
+ }
+ if ($max_jobs_run == 0) {
+ $max_jobs_run = 1;
+ print STDERR "run.pl: Warning: failed to detect number of GPUs from nvidia-smi, using ${max_jobs_run}\n";
+ }
+ } elsif (open(P, "</proc/cpuinfo")) { # Linux
+ while (<P>) { if (m/^processor/) { $max_jobs_run++; } }
+ if ($max_jobs_run == 0) {
+ print STDERR "run.pl: Warning: failed to detect any processors from /proc/cpuinfo\n";
+ $max_jobs_run = 10; # reasonable default.
+ }
+ close(P);
+ } elsif (open(P, "sysctl -a |")) { # BSD/Darwin
+ while (<P>) {
+ if (m/hw\.ncpu\s*[:=]\s*(\d+)/) { # hw.ncpu = 4, or hw.ncpu: 4
+ $max_jobs_run = $1;
+ last;
+ }
+ }
+ close(P);
+ if ($max_jobs_run == 0) {
+ print STDERR "run.pl: Warning: failed to detect any processors from sysctl -a\n";
+ $max_jobs_run = 10; # reasonable default.
+ }
+ } else {
+ # allow at most 32 jobs at once, on non-UNIX systems; change this code
+ # if you need to change this default.
+ $max_jobs_run = 32;
+ }
+ # The just-computed value of $max_jobs_run is just the number of processors
+ # (or our best guess); and if it happens that the number of jobs we need to
+ # run is just slightly above $max_jobs_run, it will make sense to increase
+ # $max_jobs_run to equal the number of jobs, so we don't have a small number
+ # of leftover jobs.
+ $num_jobs = $jobend - $jobstart + 1;
+ if (!$using_gpu &&
+ $num_jobs > $max_jobs_run && $num_jobs < 1.4 * $max_jobs_run) {
+ $max_jobs_run = $num_jobs;
+ }
+}
+
+sub pick_or_exit {
+ # pick_or_exit ( $logfile )
+ # Invoked before each job is started helps to run jobs selectively.
+ #
+ # Given the name of the output logfile decides whether the job must be
+ # executed (by returning from the subroutine) or not (by terminating the
+ # process calling exit)
+ #
+ # PRE: $job_pick is a global variable set by command line switch --pick
+ # and indicates which class of jobs must be executed.
+ #
+ # 1) If a failed job is not executed the process exit code will indicate
+ # failure, just as if the task was just executed and failed.
+ #
+ # 2) If a task is incomplete it will be executed. Incomplete may be either
+ # a job whose log file does not contain the accounting notes in the end,
+ # or a job whose log file does not exist.
+ #
+ # 3) If the $job_pick is set to 'all' (default behavior) a task will be
+ # executed regardless of the result of previous attempts.
+ #
+ # This logic could have been implemented in the main execution loop
+ # but a subroutine to preserve the current level of readability of
+ # that part of the code.
+ #
+ # Alexandre Felipe, (o.alexandre.felipe@gmail.com) 14th of August of 2020
+ #
+ if($job_pick eq 'all'){
+ return; # no need to bother with the previous log
+ }
+ open my $fh, "<", $_[0] or return; # job not executed yet
+ my $log_line;
+ my $cur_line;
+ while ($cur_line = <$fh>) {
+ if( $cur_line =~ m/# Ended \(code .*/ ) {
+ $log_line = $cur_line;
+ }
+ }
+ close $fh;
+ if (! defined($log_line)){
+ return; # incomplete
+ }
+ if ( $log_line =~ m/# Ended \(code 0\).*/ ) {
+ exit(0); # complete
+ } elsif ( $log_line =~ m/# Ended \(code \d+(; signal \d+)?\).*/ ){
+ if ($job_pick !~ m/^(failed|all)$/) {
+ exit(1); # failed but not going to run
+ } else {
+ return; # failed
+ }
+ } elsif ( $log_line =~ m/.*\S.*/ ) {
+ return; # incomplete jobs are always run
+ }
+}
+
+
+$logfile = shift @ARGV;
+
+if (defined $jobname && $logfile !~ m/$jobname/ &&
+ $jobend > $jobstart) {
+ print STDERR "run.pl: you are trying to run a parallel job but "
+ . "you are putting the output into just one log file ($logfile)\n";
+ exit(1);
+}
+
+$cmd = "";
+
+foreach $x (@ARGV) {
+ if ($x =~ m/^\S+$/) { $cmd .= $x . " "; }
+ elsif ($x =~ m:\":) { $cmd .= "'$x' "; }
+ else { $cmd .= "\"$x\" "; }
+}
+
+#$Data::Dumper::Indent=0;
+$ret = 0;
+$numfail = 0;
+%active_pids=();
+
+use POSIX ":sys_wait_h";
+for ($jobid = $jobstart; $jobid <= $jobend; $jobid++) {
+ if (scalar(keys %active_pids) >= $max_jobs_run) {
+
+ # Lets wait for a change in any child's status
+ # Then we have to work out which child finished
+ $r = waitpid(-1, 0);
+ $code = $?;
+ if ($r < 0 ) { die "run.pl: Error waiting for child process"; } # should never happen.
+ if ( defined $active_pids{$r} ) {
+ $jid=$active_pids{$r};
+ $fail[$jid]=$code;
+ if ($code !=0) { $numfail++;}
+ delete $active_pids{$r};
+ # print STDERR "Finished: $r/$jid " . Dumper(\%active_pids) . "\n";
+ } else {
+ die "run.pl: Cannot find the PID of the child process that just finished.";
+ }
+
+ # In theory we could do a non-blocking waitpid over all jobs running just
+ # to find out if only one or more jobs finished during the previous waitpid()
+ # However, we just omit this and will reap the next one in the next pass
+ # through the for(;;) cycle
+ }
+ $childpid = fork();
+ if (!defined $childpid) { die "run.pl: Error forking in run.pl (writing to $logfile)"; }
+ if ($childpid == 0) { # We're in the child... this branch
+ # executes the job and returns (possibly with an error status).
+ if (defined $jobname) {
+ $cmd =~ s/$jobname/$jobid/g;
+ $logfile =~ s/$jobname/$jobid/g;
+ }
+ # exit if the job does not need to be executed
+ pick_or_exit( $logfile );
+
+ system("mkdir -p `dirname $logfile` 2>/dev/null");
+ open(F, ">$logfile") || die "run.pl: Error opening log file $logfile";
+ print F "# " . $cmd . "\n";
+ print F "# Started at " . `date`;
+ $starttime = `date +'%s'`;
+ print F "#\n";
+ close(F);
+
+ # Pipe into bash.. make sure we're not using any other shell.
+ open(B, "|bash") || die "run.pl: Error opening shell command";
+ print B "( " . $cmd . ") 2>>$logfile >> $logfile";
+ close(B); # If there was an error, exit status is in $?
+ $ret = $?;
+
+ $lowbits = $ret & 127;
+ $highbits = $ret >> 8;
+ if ($lowbits != 0) { $return_str = "code $highbits; signal $lowbits" }
+ else { $return_str = "code $highbits"; }
+
+ $endtime = `date +'%s'`;
+ open(F, ">>$logfile") || die "run.pl: Error opening log file $logfile (again)";
+ $enddate = `date`;
+ chop $enddate;
+ print F "# Accounting: time=" . ($endtime - $starttime) . " threads=1\n";
+ print F "# Ended ($return_str) at " . $enddate . ", elapsed time " . ($endtime-$starttime) . " seconds\n";
+ close(F);
+ exit($ret == 0 ? 0 : 1);
+ } else {
+ $pid[$jobid] = $childpid;
+ $active_pids{$childpid} = $jobid;
+ # print STDERR "Queued: " . Dumper(\%active_pids) . "\n";
+ }
+}
+
+# Now we have submitted all the jobs, lets wait until all the jobs finish
+foreach $child (keys %active_pids) {
+ $jobid=$active_pids{$child};
+ $r = waitpid($pid[$jobid], 0);
+ $code = $?;
+ if ($r == -1) { die "run.pl: Error waiting for child process"; } # should never happen.
+ if ($r != 0) { $fail[$jobid]=$code; $numfail++ if $code!=0; } # Completed successfully
+}
+
+# Some sanity checks:
+# The $fail array should not contain undefined codes
+# The number of non-zeros in that array should be equal to $numfail
+# We cannot do foreach() here, as the JOB ids do not start at zero
+$failed_jids=0;
+for ($jobid = $jobstart; $jobid <= $jobend; $jobid++) {
+ $job_return = $fail[$jobid];
+ if (not defined $job_return ) {
+ # print Dumper(\@fail);
+
+ die "run.pl: Sanity check failed: we have indication that some jobs are running " .
+ "even after we waited for all jobs to finish" ;
+ }
+ if ($job_return != 0 ){ $failed_jids++;}
+}
+if ($failed_jids != $numfail) {
+ die "run.pl: Sanity check failed: cannot find out how many jobs failed ($failed_jids x $numfail)."
+}
+if ($numfail > 0) { $ret = 1; }
+
+if ($ret != 0) {
+ $njobs = $jobend - $jobstart + 1;
+ if ($njobs == 1) {
+ if (defined $jobname) {
+ $logfile =~ s/$jobname/$jobstart/; # only one numbered job, so replace name with
+ # that job.
+ }
+ print STDERR "run.pl: job failed, log is in $logfile\n";
+ if ($logfile =~ m/JOB/) {
+ print STDERR "run.pl: probably you forgot to put JOB=1:\$nj in your script.";
+ }
+ }
+ else {
+ $logfile =~ s/$jobname/*/g;
+ print STDERR "run.pl: $numfail / $njobs failed, log is in $logfile\n";
+ }
+}
+
+
+exit ($ret);
diff --git a/egs/aishell/paraformerbert/utils/shuffle_list.pl b/egs/aishell/paraformerbert/utils/shuffle_list.pl
new file mode 100755
index 0000000..a116200
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/shuffle_list.pl
@@ -0,0 +1,44 @@
+#!/usr/bin/env perl
+
+# Copyright 2013 Johns Hopkins University (author: Daniel Povey)
+
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+if ($ARGV[0] eq "--srand") {
+ $n = $ARGV[1];
+ $n =~ m/\d+/ || die "Bad argument to --srand option: \"$n\"";
+ srand($ARGV[1]);
+ shift;
+ shift;
+} else {
+ srand(0); # Gives inconsistent behavior if we don't seed.
+}
+
+if (@ARGV > 1 || $ARGV[0] =~ m/^-.+/) { # >1 args, or an option we
+ # don't understand.
+ print "Usage: shuffle_list.pl [--srand N] [input file] > output\n";
+ print "randomizes the order of lines of input.\n";
+ exit(1);
+}
+
+@lines;
+while (<>) {
+ push @lines, [ (rand(), $_)] ;
+}
+
+@lines = sort { $a->[0] cmp $b->[0] } @lines;
+foreach $l (@lines) {
+ print $l->[1];
+}
\ No newline at end of file
diff --git a/egs/aishell/paraformerbert/utils/split_data.py b/egs/aishell/paraformerbert/utils/split_data.py
new file mode 100755
index 0000000..060eae6
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/split_data.py
@@ -0,0 +1,60 @@
+import os
+import sys
+import random
+
+
+in_dir = sys.argv[1]
+out_dir = sys.argv[2]
+num_split = sys.argv[3]
+
+
+def split_scp(scp, num):
+ assert len(scp) >= num
+ avg = len(scp) // num
+ out = []
+ begin = 0
+
+ for i in range(num):
+ if i == num - 1:
+ out.append(scp[begin:])
+ else:
+ out.append(scp[begin:begin+avg])
+ begin += avg
+
+ return out
+
+
+os.path.exists("{}/wav.scp".format(in_dir))
+os.path.exists("{}/text".format(in_dir))
+
+with open("{}/wav.scp".format(in_dir), 'r') as infile:
+ wav_list = infile.readlines()
+
+with open("{}/text".format(in_dir), 'r') as infile:
+ text_list = infile.readlines()
+
+assert len(wav_list) == len(text_list)
+
+x = list(zip(wav_list, text_list))
+random.shuffle(x)
+wav_shuffle_list, text_shuffle_list = zip(*x)
+
+num_split = int(num_split)
+wav_split_list = split_scp(wav_shuffle_list, num_split)
+text_split_list = split_scp(text_shuffle_list, num_split)
+
+for idx, wav_list in enumerate(wav_split_list, 1):
+ path = out_dir + "/split" + str(num_split) + "/" + str(idx)
+ if not os.path.exists(path):
+ os.makedirs(path)
+ with open("{}/wav.scp".format(path), 'w') as wav_writer:
+ for line in wav_list:
+ wav_writer.write(line)
+
+for idx, text_list in enumerate(text_split_list, 1):
+ path = out_dir + "/split" + str(num_split) + "/" + str(idx)
+ if not os.path.exists(path):
+ os.makedirs(path)
+ with open("{}/text".format(path), 'w') as text_writer:
+ for line in text_list:
+ text_writer.write(line)
diff --git a/egs/aishell/paraformerbert/utils/split_scp.pl b/egs/aishell/paraformerbert/utils/split_scp.pl
new file mode 100755
index 0000000..0876dcb
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/split_scp.pl
@@ -0,0 +1,246 @@
+#!/usr/bin/env perl
+
+# Copyright 2010-2011 Microsoft Corporation
+
+# See ../../COPYING for clarification regarding multiple authors
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+# This program splits up any kind of .scp or archive-type file.
+# If there is no utt2spk option it will work on any text file and
+# will split it up with an approximately equal number of lines in
+# each but.
+# With the --utt2spk option it will work on anything that has the
+# utterance-id as the first entry on each line; the utt2spk file is
+# of the form "utterance speaker" (on each line).
+# It splits it into equal size chunks as far as it can. If you use the utt2spk
+# option it will make sure these chunks coincide with speaker boundaries. In
+# this case, if there are more chunks than speakers (and in some other
+# circumstances), some of the resulting chunks will be empty and it will print
+# an error message and exit with nonzero status.
+# You will normally call this like:
+# split_scp.pl scp scp.1 scp.2 scp.3 ...
+# or
+# split_scp.pl --utt2spk=utt2spk scp scp.1 scp.2 scp.3 ...
+# Note that you can use this script to split the utt2spk file itself,
+# e.g. split_scp.pl --utt2spk=utt2spk utt2spk utt2spk.1 utt2spk.2 ...
+
+# You can also call the scripts like:
+# split_scp.pl -j 3 0 scp scp.0
+# [note: with this option, it assumes zero-based indexing of the split parts,
+# i.e. the second number must be 0 <= n < num-jobs.]
+
+use warnings;
+
+$num_jobs = 0;
+$job_id = 0;
+$utt2spk_file = "";
+$one_based = 0;
+
+for ($x = 1; $x <= 3 && @ARGV > 0; $x++) {
+ if ($ARGV[0] eq "-j") {
+ shift @ARGV;
+ $num_jobs = shift @ARGV;
+ $job_id = shift @ARGV;
+ }
+ if ($ARGV[0] =~ /--utt2spk=(.+)/) {
+ $utt2spk_file=$1;
+ shift;
+ }
+ if ($ARGV[0] eq '--one-based') {
+ $one_based = 1;
+ shift @ARGV;
+ }
+}
+
+if ($num_jobs != 0 && ($num_jobs < 0 || $job_id - $one_based < 0 ||
+ $job_id - $one_based >= $num_jobs)) {
+ die "$0: Invalid job number/index values for '-j $num_jobs $job_id" .
+ ($one_based ? " --one-based" : "") . "'\n"
+}
+
+$one_based
+ and $job_id--;
+
+if(($num_jobs == 0 && @ARGV < 2) || ($num_jobs > 0 && (@ARGV < 1 || @ARGV > 2))) {
+ die
+"Usage: split_scp.pl [--utt2spk=<utt2spk_file>] in.scp out1.scp out2.scp ...
+ or: split_scp.pl -j num-jobs job-id [--one-based] [--utt2spk=<utt2spk_file>] in.scp [out.scp]
+ ... where 0 <= job-id < num-jobs, or 1 <= job-id <- num-jobs if --one-based.\n";
+}
+
+$error = 0;
+$inscp = shift @ARGV;
+if ($num_jobs == 0) { # without -j option
+ @OUTPUTS = @ARGV;
+} else {
+ for ($j = 0; $j < $num_jobs; $j++) {
+ if ($j == $job_id) {
+ if (@ARGV > 0) { push @OUTPUTS, $ARGV[0]; }
+ else { push @OUTPUTS, "-"; }
+ } else {
+ push @OUTPUTS, "/dev/null";
+ }
+ }
+}
+
+if ($utt2spk_file ne "") { # We have the --utt2spk option...
+ open($u_fh, '<', $utt2spk_file) || die "$0: Error opening utt2spk file $utt2spk_file: $!\n";
+ while(<$u_fh>) {
+ @A = split;
+ @A == 2 || die "$0: Bad line $_ in utt2spk file $utt2spk_file\n";
+ ($u,$s) = @A;
+ $utt2spk{$u} = $s;
+ }
+ close $u_fh;
+ open($i_fh, '<', $inscp) || die "$0: Error opening input scp file $inscp: $!\n";
+ @spkrs = ();
+ while(<$i_fh>) {
+ @A = split;
+ if(@A == 0) { die "$0: Empty or space-only line in scp file $inscp\n"; }
+ $u = $A[0];
+ $s = $utt2spk{$u};
+ defined $s || die "$0: No utterance $u in utt2spk file $utt2spk_file\n";
+ if(!defined $spk_count{$s}) {
+ push @spkrs, $s;
+ $spk_count{$s} = 0;
+ $spk_data{$s} = []; # ref to new empty array.
+ }
+ $spk_count{$s}++;
+ push @{$spk_data{$s}}, $_;
+ }
+ # Now split as equally as possible ..
+ # First allocate spks to files by allocating an approximately
+ # equal number of speakers.
+ $numspks = @spkrs; # number of speakers.
+ $numscps = @OUTPUTS; # number of output files.
+ if ($numspks < $numscps) {
+ die "$0: Refusing to split data because number of speakers $numspks " .
+ "is less than the number of output .scp files $numscps\n";
+ }
+ for($scpidx = 0; $scpidx < $numscps; $scpidx++) {
+ $scparray[$scpidx] = []; # [] is array reference.
+ }
+ for ($spkidx = 0; $spkidx < $numspks; $spkidx++) {
+ $scpidx = int(($spkidx*$numscps) / $numspks);
+ $spk = $spkrs[$spkidx];
+ push @{$scparray[$scpidx]}, $spk;
+ $scpcount[$scpidx] += $spk_count{$spk};
+ }
+
+ # Now will try to reassign beginning + ending speakers
+ # to different scp's and see if it gets more balanced.
+ # Suppose objf we're minimizing is sum_i (num utts in scp[i] - average)^2.
+ # We can show that if considering changing just 2 scp's, we minimize
+ # this by minimizing the squared difference in sizes. This is
+ # equivalent to minimizing the absolute difference in sizes. This
+ # shows this method is bound to converge.
+
+ $changed = 1;
+ while($changed) {
+ $changed = 0;
+ for($scpidx = 0; $scpidx < $numscps; $scpidx++) {
+ # First try to reassign ending spk of this scp.
+ if($scpidx < $numscps-1) {
+ $sz = @{$scparray[$scpidx]};
+ if($sz > 0) {
+ $spk = $scparray[$scpidx]->[$sz-1];
+ $count = $spk_count{$spk};
+ $nutt1 = $scpcount[$scpidx];
+ $nutt2 = $scpcount[$scpidx+1];
+ if( abs( ($nutt2+$count) - ($nutt1-$count))
+ < abs($nutt2 - $nutt1)) { # Would decrease
+ # size-diff by reassigning spk...
+ $scpcount[$scpidx+1] += $count;
+ $scpcount[$scpidx] -= $count;
+ pop @{$scparray[$scpidx]};
+ unshift @{$scparray[$scpidx+1]}, $spk;
+ $changed = 1;
+ }
+ }
+ }
+ if($scpidx > 0 && @{$scparray[$scpidx]} > 0) {
+ $spk = $scparray[$scpidx]->[0];
+ $count = $spk_count{$spk};
+ $nutt1 = $scpcount[$scpidx-1];
+ $nutt2 = $scpcount[$scpidx];
+ if( abs( ($nutt2-$count) - ($nutt1+$count))
+ < abs($nutt2 - $nutt1)) { # Would decrease
+ # size-diff by reassigning spk...
+ $scpcount[$scpidx-1] += $count;
+ $scpcount[$scpidx] -= $count;
+ shift @{$scparray[$scpidx]};
+ push @{$scparray[$scpidx-1]}, $spk;
+ $changed = 1;
+ }
+ }
+ }
+ }
+ # Now print out the files...
+ for($scpidx = 0; $scpidx < $numscps; $scpidx++) {
+ $scpfile = $OUTPUTS[$scpidx];
+ ($scpfile ne '-' ? open($f_fh, '>', $scpfile)
+ : open($f_fh, '>&', \*STDOUT)) ||
+ die "$0: Could not open scp file $scpfile for writing: $!\n";
+ $count = 0;
+ if(@{$scparray[$scpidx]} == 0) {
+ print STDERR "$0: eError: split_scp.pl producing empty .scp file " .
+ "$scpfile (too many splits and too few speakers?)\n";
+ $error = 1;
+ } else {
+ foreach $spk ( @{$scparray[$scpidx]} ) {
+ print $f_fh @{$spk_data{$spk}};
+ $count += $spk_count{$spk};
+ }
+ $count == $scpcount[$scpidx] || die "Count mismatch [code error]";
+ }
+ close($f_fh);
+ }
+} else {
+ # This block is the "normal" case where there is no --utt2spk
+ # option and we just break into equal size chunks.
+
+ open($i_fh, '<', $inscp) || die "$0: Error opening input scp file $inscp: $!\n";
+
+ $numscps = @OUTPUTS; # size of array.
+ @F = ();
+ while(<$i_fh>) {
+ push @F, $_;
+ }
+ $numlines = @F;
+ if($numlines == 0) {
+ print STDERR "$0: error: empty input scp file $inscp\n";
+ $error = 1;
+ }
+ $linesperscp = int( $numlines / $numscps); # the "whole part"..
+ $linesperscp >= 1 || die "$0: You are splitting into too many pieces! [reduce \$nj ($numscps) to be smaller than the number of lines ($numlines) in $inscp]\n";
+ $remainder = $numlines - ($linesperscp * $numscps);
+ ($remainder >= 0 && $remainder < $numlines) || die "bad remainder $remainder";
+ # [just doing int() rounds down].
+ $n = 0;
+ for($scpidx = 0; $scpidx < @OUTPUTS; $scpidx++) {
+ $scpfile = $OUTPUTS[$scpidx];
+ ($scpfile ne '-' ? open($o_fh, '>', $scpfile)
+ : open($o_fh, '>&', \*STDOUT)) ||
+ die "$0: Could not open scp file $scpfile for writing: $!\n";
+ for($k = 0; $k < $linesperscp + ($scpidx < $remainder ? 1 : 0); $k++) {
+ print $o_fh $F[$n++];
+ }
+ close($o_fh) || die "$0: Eror closing scp file $scpfile: $!\n";
+ }
+ $n == $numlines || die "$n != $numlines [code error]";
+}
+
+exit ($error);
diff --git a/egs/aishell/paraformerbert/utils/subset_data_dir_tr_cv.sh b/egs/aishell/paraformerbert/utils/subset_data_dir_tr_cv.sh
new file mode 100755
index 0000000..e16cebd
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/subset_data_dir_tr_cv.sh
@@ -0,0 +1,30 @@
+#!/usr/bin/env bash
+
+dev_num_utt=1000
+
+echo "$0 $@"
+. utils/parse_options.sh || exit 1;
+
+train_data=$1
+out_dir=$2
+
+[ ! -f ${train_data}/wav.scp ] && echo "$0: no such file ${train_data}/wav.scp" && exit 1;
+[ ! -f ${train_data}/text ] && echo "$0: no such file ${train_data}/text" && exit 1;
+
+mkdir -p ${out_dir}/train && mkdir -p ${out_dir}/dev
+
+cp ${train_data}/wav.scp ${out_dir}/train/wav.scp.bak
+cp ${train_data}/text ${out_dir}/train/text.bak
+
+num_utt=$(wc -l <${out_dir}/train/wav.scp.bak)
+
+utils/shuffle_list.pl --srand 1 ${out_dir}/train/wav.scp.bak > ${out_dir}/train/wav.scp.shuf
+head -n ${dev_num_utt} ${out_dir}/train/wav.scp.shuf > ${out_dir}/dev/wav.scp
+tail -n $((${num_utt}-${dev_num_utt})) ${out_dir}/train/wav.scp.shuf > ${out_dir}/train/wav.scp
+
+utils/shuffle_list.pl --srand 1 ${out_dir}/train/text.bak > ${out_dir}/train/text.shuf
+head -n ${dev_num_utt} ${out_dir}/train/text.shuf > ${out_dir}/dev/text
+tail -n $((${num_utt}-${dev_num_utt})) ${out_dir}/train/text.shuf > ${out_dir}/train/text
+
+rm ${out_dir}/train/wav.scp.bak ${out_dir}/train/text.bak
+rm ${out_dir}/train/wav.scp.shuf ${out_dir}/train/text.shuf
diff --git a/egs/aishell/paraformerbert/utils/text2token.py b/egs/aishell/paraformerbert/utils/text2token.py
new file mode 100755
index 0000000..56c3913
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/text2token.py
@@ -0,0 +1,135 @@
+#!/usr/bin/env python3
+
+# Copyright 2017 Johns Hopkins University (Shinji Watanabe)
+# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
+
+
+import argparse
+import codecs
+import re
+import sys
+
+is_python2 = sys.version_info[0] == 2
+
+
+def exist_or_not(i, match_pos):
+ start_pos = None
+ end_pos = None
+ for pos in match_pos:
+ if pos[0] <= i < pos[1]:
+ start_pos = pos[0]
+ end_pos = pos[1]
+ break
+
+ return start_pos, end_pos
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="convert raw text to tokenized text",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--nchar",
+ "-n",
+ default=1,
+ type=int,
+ help="number of characters to split, i.e., \
+ aabb -> a a b b with -n 1 and aa bb with -n 2",
+ )
+ parser.add_argument(
+ "--skip-ncols", "-s", default=0, type=int, help="skip first n columns"
+ )
+ parser.add_argument("--space", default="<space>", type=str, help="space symbol")
+ parser.add_argument(
+ "--non-lang-syms",
+ "-l",
+ default=None,
+ type=str,
+ help="list of non-linguistic symobles, e.g., <NOISE> etc.",
+ )
+ parser.add_argument("text", type=str, default=False, nargs="?", help="input text")
+ parser.add_argument(
+ "--trans_type",
+ "-t",
+ type=str,
+ default="char",
+ choices=["char", "phn"],
+ help="""Transcript type. char/phn. e.g., for TIMIT FADG0_SI1279 -
+ If trans_type is char,
+ read from SI1279.WRD file -> "bricks are an alternative"
+ Else if trans_type is phn,
+ read from SI1279.PHN file -> "sil b r ih sil k s aa r er n aa l
+ sil t er n ih sil t ih v sil" """,
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ rs = []
+ if args.non_lang_syms is not None:
+ with codecs.open(args.non_lang_syms, "r", encoding="utf-8") as f:
+ nls = [x.rstrip() for x in f.readlines()]
+ rs = [re.compile(re.escape(x)) for x in nls]
+
+ if args.text:
+ f = codecs.open(args.text, encoding="utf-8")
+ else:
+ f = codecs.getreader("utf-8")(sys.stdin if is_python2 else sys.stdin.buffer)
+
+ sys.stdout = codecs.getwriter("utf-8")(
+ sys.stdout if is_python2 else sys.stdout.buffer
+ )
+ line = f.readline()
+ n = args.nchar
+ while line:
+ x = line.split()
+ print(" ".join(x[: args.skip_ncols]), end=" ")
+ a = " ".join(x[args.skip_ncols :])
+
+ # get all matched positions
+ match_pos = []
+ for r in rs:
+ i = 0
+ while i >= 0:
+ m = r.search(a, i)
+ if m:
+ match_pos.append([m.start(), m.end()])
+ i = m.end()
+ else:
+ break
+
+ if args.trans_type == "phn":
+ a = a.split(" ")
+ else:
+ if len(match_pos) > 0:
+ chars = []
+ i = 0
+ while i < len(a):
+ start_pos, end_pos = exist_or_not(i, match_pos)
+ if start_pos is not None:
+ chars.append(a[start_pos:end_pos])
+ i = end_pos
+ else:
+ chars.append(a[i])
+ i += 1
+ a = chars
+
+ a = [a[j : j + n] for j in range(0, len(a), n)]
+
+ a_flat = []
+ for z in a:
+ a_flat.append("".join(z))
+
+ a_chars = [z.replace(" ", args.space) for z in a_flat]
+ if args.trans_type == "phn":
+ a_chars = [z.replace("sil", args.space) for z in a_chars]
+ print(" ".join(a_chars))
+ line = f.readline()
+
+
+if __name__ == "__main__":
+ main()
diff --git a/egs/aishell/paraformerbert/utils/text_tokenize.py b/egs/aishell/paraformerbert/utils/text_tokenize.py
new file mode 100755
index 0000000..962ea11
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/text_tokenize.py
@@ -0,0 +1,106 @@
+import re
+import argparse
+
+
+def load_dict(seg_file):
+ seg_dict = {}
+ with open(seg_file, 'r') as infile:
+ for line in infile:
+ s = line.strip().split()
+ key = s[0]
+ value = s[1:]
+ seg_dict[key] = " ".join(value)
+ return seg_dict
+
+
+def forward_segment(text, dic):
+ word_list = []
+ i = 0
+ while i < len(text):
+ longest_word = text[i]
+ for j in range(i + 1, len(text) + 1):
+ word = text[i:j]
+ if word in dic:
+ if len(word) > len(longest_word):
+ longest_word = word
+ word_list.append(longest_word)
+ i += len(longest_word)
+ return word_list
+
+
+def tokenize(txt,
+ seg_dict):
+ out_txt = ""
+ pattern = re.compile(r"([\u4E00-\u9FA5A-Za-z0-9])")
+ for word in txt:
+ if pattern.match(word):
+ if word in seg_dict:
+ out_txt += seg_dict[word] + " "
+ else:
+ out_txt += "<unk>" + " "
+ else:
+ continue
+ return out_txt.strip()
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="text tokenize",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--text-file",
+ "-t",
+ default=False,
+ required=True,
+ type=str,
+ help="input text",
+ )
+ parser.add_argument(
+ "--seg-file",
+ "-s",
+ default=False,
+ required=True,
+ type=str,
+ help="seg file",
+ )
+ parser.add_argument(
+ "--txt-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="txt index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ txt_writer = open("{}/text.{}.txt".format(args.output_dir, args.txt_index), 'w')
+ shape_writer = open("{}/len.{}".format(args.output_dir, args.txt_index), 'w')
+ seg_dict = load_dict(args.seg_file)
+ with open(args.text_file, 'r') as infile:
+ for line in infile:
+ s = line.strip().split()
+ text_id = s[0]
+ text_list = forward_segment("".join(s[1:]).lower(), seg_dict)
+ text = tokenize(text_list, seg_dict)
+ lens = len(text.strip().split())
+ txt_writer.write(text_id + " " + text + '\n')
+ shape_writer.write(text_id + " " + str(lens) + '\n')
+
+
+if __name__ == '__main__':
+ main()
+
diff --git a/egs/aishell/paraformerbert/utils/text_tokenize.sh b/egs/aishell/paraformerbert/utils/text_tokenize.sh
new file mode 100755
index 0000000..6b74fef
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/text_tokenize.sh
@@ -0,0 +1,35 @@
+#!/usr/bin/env bash
+
+
+# Begin configuration section.
+nj=32
+cmd=utils/run.pl
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+# tokenize configuration
+text_dir=$1
+seg_file=$2
+logdir=$3
+output_dir=$4
+
+txt_dir=${output_dir}/txt; mkdir -p ${output_dir}/txt
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/text_tokenize.JOB.log \
+ python utils/text_tokenize.py -t ${text_dir}/txt/text.JOB.txt \
+ -s ${seg_file} -i JOB -o ${txt_dir} \
+ || exit 1;
+
+# concatenate the text files together.
+for n in $(seq $nj); do
+ cat ${txt_dir}/text.$n.txt || exit 1
+done > ${output_dir}/text || exit 1
+
+for n in $(seq $nj); do
+ cat ${txt_dir}/len.$n || exit 1
+done > ${output_dir}/text_shape || exit 1
+
+echo "$0: Succeeded text tokenize"
diff --git a/egs/aishell/paraformerbert/utils/textnorm_zh.py b/egs/aishell/paraformerbert/utils/textnorm_zh.py
new file mode 100755
index 0000000..79feb83
--- /dev/null
+++ b/egs/aishell/paraformerbert/utils/textnorm_zh.py
@@ -0,0 +1,834 @@
+#!/usr/bin/env python3
+# coding=utf-8
+
+# Authors:
+# 2019.5 Zhiyang Zhou (https://github.com/Joee1995/chn_text_norm.git)
+# 2019.9 Jiayu DU
+#
+# requirements:
+# - python 3.X
+# notes: python 2.X WILL fail or produce misleading results
+
+import sys, os, argparse, codecs, string, re
+
+# ================================================================================ #
+# basic constant
+# ================================================================================ #
+CHINESE_DIGIS = u'闆朵竴浜屼笁鍥涗簲鍏竷鍏節'
+BIG_CHINESE_DIGIS_SIMPLIFIED = u'闆跺9璐板弫鑲嗕紞闄嗘煉鎹岀帠'
+BIG_CHINESE_DIGIS_TRADITIONAL = u'闆跺9璨冲弮鑲嗕紞闄告煉鎹岀帠'
+SMALLER_BIG_CHINESE_UNITS_SIMPLIFIED = u'鍗佺櫨鍗冧竾'
+SMALLER_BIG_CHINESE_UNITS_TRADITIONAL = u'鎷句桨浠熻惉'
+LARGER_CHINESE_NUMERING_UNITS_SIMPLIFIED = u'浜垮厗浜灀绉┌娌熸锭姝h浇'
+LARGER_CHINESE_NUMERING_UNITS_TRADITIONAL = u'鍎勫厗浜灀绉┌婧濇緱姝h級'
+SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED = u'鍗佺櫨鍗冧竾'
+SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL = u'鎷句桨浠熻惉'
+
+ZERO_ALT = u'銆�'
+ONE_ALT = u'骞�'
+TWO_ALTS = [u'涓�', u'鍏�']
+
+POSITIVE = [u'姝�', u'姝�']
+NEGATIVE = [u'璐�', u'璨�']
+POINT = [u'鐐�', u'榛�']
+# PLUS = [u'鍔�', u'鍔�']
+# SIL = [u'鏉�', u'妲�']
+
+FILLER_CHARS = ['鍛�', '鍟�']
+ER_WHITELIST = '(鍎垮コ|鍎垮瓙|鍎垮瓩|濂冲効|鍎垮|濡诲効|' \
+ '鑳庡効|濠村効|鏂扮敓鍎縷濠村辜鍎縷骞煎効|灏戝効|灏忓効|鍎挎瓕|鍎跨|鍎跨|鎵樺効鎵�|瀛ゅ効|' \
+ '鍎挎垙|鍎垮寲|鍙板効搴剕楣垮効宀泑姝e効鍏粡|鍚婂効閮庡綋|鐢熷効鑲插コ|鎵樺効甯﹀コ|鍏诲効闃茶�亅鐥村効鍛嗗コ|' \
+ '浣冲効浣冲|鍎挎�滃吔鎵皘鍎挎棤甯哥埗|鍎夸笉瀚屾瘝涓憒鍎胯鍗冮噷姣嶆媴蹇鍎垮ぇ涓嶇敱鐖穦鑻忎篂鍎�)'
+
+# 涓枃鏁板瓧绯荤粺绫诲瀷
+NUMBERING_TYPES = ['low', 'mid', 'high']
+
+CURRENCY_NAMES = '(浜烘皯甯亅缇庡厓|鏃ュ厓|鑻遍晳|娆у厓|椹厠|娉曢儙|鍔犳嬁澶у厓|婢冲厓|娓竵|鍏堜护|鑺叞椹厠|鐖卞皵鍏伴晳|' \
+ '閲屾媺|鑽峰叞鐩緗鍩冩柉搴撳|姣斿濉攟鍗板凹鐩緗鏋楀悏鐗箌鏂拌タ鍏板厓|姣旂储|鍗㈠竷|鏂板姞鍧″厓|闊╁厓|娉伴摙)'
+CURRENCY_UNITS = '((浜縷鍗冧竾|鐧句竾|涓噟鍗億鐧�)|(浜縷鍗冧竾|鐧句竾|涓噟鍗億鐧緗)鍏億(浜縷鍗冧竾|鐧句竾|涓噟鍗億鐧緗)鍧梶瑙抾姣泑鍒�)'
+COM_QUANTIFIERS = '(鍖箌寮爘搴鍥瀨鍦簗灏緗鏉涓獆棣東闃檤闃祙缃憒鐐畖椤秥涓榺妫祙鍙獆鏀瘄琚瓅杈唡鎸憒鎷厊棰梶澹硘绐爘鏇瞸澧檤缇鑵攟' \
+ '鐮搴瀹璐瘄鎵巪鎹唡鍒�|浠鎵搢鎵媩缃梶鍧灞眧宀瓅姹焲婧獆閽焲闃焲鍗晐鍙寍瀵箌鍑簗鍙澶磡鑴殀鏉縷璺硘鏋潀浠秥璐磡' \
+ '閽坾绾縷绠鍚峾浣峾韬珅鍫倈璇緗鏈瑋椤祙瀹秥鎴穦灞倈涓潀姣珅鍘榺鍒唡閽眧涓鏂鎷厊閾鐭硘閽閿眧蹇絴(鍗億姣珅寰�)鍏媩' \
+ '姣珅鍘榺鍒唡瀵竱灏簗涓坾閲寍瀵粅甯竱閾簗绋媩(鍗億鍒唡鍘榺姣珅寰�)绫硘鎾畖鍕簗鍚坾鍗噟鏂梶鐭硘鐩榺纰梶纰焲鍙爘妗秥绗紎鐩唡' \
+ '鐩抾鏉瘄閽焲鏂泑閿厊绨媩绡畖鐩榺妗秥缃恷鐡秥澹秥鍗畖鐩弢绠﹟绠眧鐓瞸鍟東琚媩閽祙骞磡鏈坾鏃瀛鍒粅鏃秥鍛▅澶﹟绉抾鍒唡鏃瑋' \
+ '绾獆宀亅涓東鏇磡澶渱鏄澶弢绉媩鍐瑋浠浼弢杈坾涓竱娉绮抾棰梶骞鍫唡鏉鏍箌鏀瘄閬搢闈鐗噟寮爘棰梶鍧�)'
+
+# punctuation information are based on Zhon project (https://github.com/tsroten/zhon.git)
+CHINESE_PUNC_STOP = '锛侊紵锝°��'
+CHINESE_PUNC_NON_STOP = '锛傦純锛勶紖锛嗭紘锛堬級锛婏紜锛岋紞锛忥細锛涳紲锛濓紴锛狅蓟锛硷冀锛撅伎锝�锝涳綔锝濓綖锝燂綘锝剑锝ゃ�併�冦�嬨�屻�嶃�庛�忋�愩�戙�斻�曘�栥�椼�樸�欍�氥�涖�溿�濄�炪�熴�般�俱�库�撯�斺�樷�欌�涒�溾�濃�炩�熲�︹�э箯'
+CHINESE_PUNC_LIST = CHINESE_PUNC_STOP + CHINESE_PUNC_NON_STOP
+
+# ================================================================================ #
+# basic class
+# ================================================================================ #
+class ChineseChar(object):
+ """
+ 涓枃瀛楃
+ 姣忎釜瀛楃瀵瑰簲绠�浣撳拰绻佷綋,
+ e.g. 绠�浣� = '璐�', 绻佷綋 = '璨�'
+ 杞崲鏃跺彲杞崲涓虹畝浣撴垨绻佷綋
+ """
+
+ def __init__(self, simplified, traditional):
+ self.simplified = simplified
+ self.traditional = traditional
+ #self.__repr__ = self.__str__
+
+ def __str__(self):
+ return self.simplified or self.traditional or None
+
+ def __repr__(self):
+ return self.__str__()
+
+
+class ChineseNumberUnit(ChineseChar):
+ """
+ 涓枃鏁板瓧/鏁颁綅瀛楃
+ 姣忎釜瀛楃闄ょ箒绠�浣撳杩樻湁涓�涓澶栫殑澶у啓瀛楃
+ e.g. '闄�' 鍜� '闄�'
+ """
+
+ def __init__(self, power, simplified, traditional, big_s, big_t):
+ super(ChineseNumberUnit, self).__init__(simplified, traditional)
+ self.power = power
+ self.big_s = big_s
+ self.big_t = big_t
+
+ def __str__(self):
+ return '10^{}'.format(self.power)
+
+ @classmethod
+ def create(cls, index, value, numbering_type=NUMBERING_TYPES[1], small_unit=False):
+
+ if small_unit:
+ return ChineseNumberUnit(power=index + 1,
+ simplified=value[0], traditional=value[1], big_s=value[1], big_t=value[1])
+ elif numbering_type == NUMBERING_TYPES[0]:
+ return ChineseNumberUnit(power=index + 8,
+ simplified=value[0], traditional=value[1], big_s=value[0], big_t=value[1])
+ elif numbering_type == NUMBERING_TYPES[1]:
+ return ChineseNumberUnit(power=(index + 2) * 4,
+ simplified=value[0], traditional=value[1], big_s=value[0], big_t=value[1])
+ elif numbering_type == NUMBERING_TYPES[2]:
+ return ChineseNumberUnit(power=pow(2, index + 3),
+ simplified=value[0], traditional=value[1], big_s=value[0], big_t=value[1])
+ else:
+ raise ValueError(
+ 'Counting type should be in {0} ({1} provided).'.format(NUMBERING_TYPES, numbering_type))
+
+
+class ChineseNumberDigit(ChineseChar):
+ """
+ 涓枃鏁板瓧瀛楃
+ """
+
+ def __init__(self, value, simplified, traditional, big_s, big_t, alt_s=None, alt_t=None):
+ super(ChineseNumberDigit, self).__init__(simplified, traditional)
+ self.value = value
+ self.big_s = big_s
+ self.big_t = big_t
+ self.alt_s = alt_s
+ self.alt_t = alt_t
+
+ def __str__(self):
+ return str(self.value)
+
+ @classmethod
+ def create(cls, i, v):
+ return ChineseNumberDigit(i, v[0], v[1], v[2], v[3])
+
+
+class ChineseMath(ChineseChar):
+ """
+ 涓枃鏁颁綅瀛楃
+ """
+
+ def __init__(self, simplified, traditional, symbol, expression=None):
+ super(ChineseMath, self).__init__(simplified, traditional)
+ self.symbol = symbol
+ self.expression = expression
+ self.big_s = simplified
+ self.big_t = traditional
+
+
+CC, CNU, CND, CM = ChineseChar, ChineseNumberUnit, ChineseNumberDigit, ChineseMath
+
+
+class NumberSystem(object):
+ """
+ 涓枃鏁板瓧绯荤粺
+ """
+ pass
+
+
+class MathSymbol(object):
+ """
+ 鐢ㄤ簬涓枃鏁板瓧绯荤粺鐨勬暟瀛︾鍙� (绻�/绠�浣�), e.g.
+ positive = ['姝�', '姝�']
+ negative = ['璐�', '璨�']
+ point = ['鐐�', '榛�']
+ """
+
+ def __init__(self, positive, negative, point):
+ self.positive = positive
+ self.negative = negative
+ self.point = point
+
+ def __iter__(self):
+ for v in self.__dict__.values():
+ yield v
+
+
+# class OtherSymbol(object):
+# """
+# 鍏朵粬绗﹀彿
+# """
+#
+# def __init__(self, sil):
+# self.sil = sil
+#
+# def __iter__(self):
+# for v in self.__dict__.values():
+# yield v
+
+
+# ================================================================================ #
+# basic utils
+# ================================================================================ #
+def create_system(numbering_type=NUMBERING_TYPES[1]):
+ """
+ 鏍规嵁鏁板瓧绯荤粺绫诲瀷杩斿洖鍒涘缓鐩稿簲鐨勬暟瀛楃郴缁燂紝榛樿涓� mid
+ NUMBERING_TYPES = ['low', 'mid', 'high']: 涓枃鏁板瓧绯荤粺绫诲瀷
+ low: '鍏�' = '浜�' * '鍗�' = $10^{9}$, '浜�' = '鍏�' * '鍗�', etc.
+ mid: '鍏�' = '浜�' * '涓�' = $10^{12}$, '浜�' = '鍏�' * '涓�', etc.
+ high: '鍏�' = '浜�' * '浜�' = $10^{16}$, '浜�' = '鍏�' * '鍏�', etc.
+ 杩斿洖瀵瑰簲鐨勬暟瀛楃郴缁�
+ """
+
+ # chinese number units of '浜�' and larger
+ all_larger_units = zip(
+ LARGER_CHINESE_NUMERING_UNITS_SIMPLIFIED, LARGER_CHINESE_NUMERING_UNITS_TRADITIONAL)
+ larger_units = [CNU.create(i, v, numbering_type, False)
+ for i, v in enumerate(all_larger_units)]
+ # chinese number units of '鍗�, 鐧�, 鍗�, 涓�'
+ all_smaller_units = zip(
+ SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED, SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL)
+ smaller_units = [CNU.create(i, v, small_unit=True)
+ for i, v in enumerate(all_smaller_units)]
+ # digis
+ chinese_digis = zip(CHINESE_DIGIS, CHINESE_DIGIS,
+ BIG_CHINESE_DIGIS_SIMPLIFIED, BIG_CHINESE_DIGIS_TRADITIONAL)
+ digits = [CND.create(i, v) for i, v in enumerate(chinese_digis)]
+ digits[0].alt_s, digits[0].alt_t = ZERO_ALT, ZERO_ALT
+ digits[1].alt_s, digits[1].alt_t = ONE_ALT, ONE_ALT
+ digits[2].alt_s, digits[2].alt_t = TWO_ALTS[0], TWO_ALTS[1]
+
+ # symbols
+ positive_cn = CM(POSITIVE[0], POSITIVE[1], '+', lambda x: x)
+ negative_cn = CM(NEGATIVE[0], NEGATIVE[1], '-', lambda x: -x)
+ point_cn = CM(POINT[0], POINT[1], '.', lambda x,
+ y: float(str(x) + '.' + str(y)))
+ # sil_cn = CM(SIL[0], SIL[1], '-', lambda x, y: float(str(x) + '-' + str(y)))
+ system = NumberSystem()
+ system.units = smaller_units + larger_units
+ system.digits = digits
+ system.math = MathSymbol(positive_cn, negative_cn, point_cn)
+ # system.symbols = OtherSymbol(sil_cn)
+ return system
+
+
+def chn2num(chinese_string, numbering_type=NUMBERING_TYPES[1]):
+
+ def get_symbol(char, system):
+ for u in system.units:
+ if char in [u.traditional, u.simplified, u.big_s, u.big_t]:
+ return u
+ for d in system.digits:
+ if char in [d.traditional, d.simplified, d.big_s, d.big_t, d.alt_s, d.alt_t]:
+ return d
+ for m in system.math:
+ if char in [m.traditional, m.simplified]:
+ return m
+
+ def string2symbols(chinese_string, system):
+ int_string, dec_string = chinese_string, ''
+ for p in [system.math.point.simplified, system.math.point.traditional]:
+ if p in chinese_string:
+ int_string, dec_string = chinese_string.split(p)
+ break
+ return [get_symbol(c, system) for c in int_string], \
+ [get_symbol(c, system) for c in dec_string]
+
+ def correct_symbols(integer_symbols, system):
+ """
+ 涓�鐧惧叓 to 涓�鐧惧叓鍗�
+ 涓�浜夸竴鍗冧笁鐧句竾 to 涓�浜� 涓�鍗冧竾 涓夌櫨涓�
+ """
+
+ if integer_symbols and isinstance(integer_symbols[0], CNU):
+ if integer_symbols[0].power == 1:
+ integer_symbols = [system.digits[1]] + integer_symbols
+
+ if len(integer_symbols) > 1:
+ if isinstance(integer_symbols[-1], CND) and isinstance(integer_symbols[-2], CNU):
+ integer_symbols.append(
+ CNU(integer_symbols[-2].power - 1, None, None, None, None))
+
+ result = []
+ unit_count = 0
+ for s in integer_symbols:
+ if isinstance(s, CND):
+ result.append(s)
+ unit_count = 0
+ elif isinstance(s, CNU):
+ current_unit = CNU(s.power, None, None, None, None)
+ unit_count += 1
+
+ if unit_count == 1:
+ result.append(current_unit)
+ elif unit_count > 1:
+ for i in range(len(result)):
+ if isinstance(result[-i - 1], CNU) and result[-i - 1].power < current_unit.power:
+ result[-i - 1] = CNU(result[-i - 1].power +
+ current_unit.power, None, None, None, None)
+ return result
+
+ def compute_value(integer_symbols):
+ """
+ Compute the value.
+ When current unit is larger than previous unit, current unit * all previous units will be used as all previous units.
+ e.g. '涓ゅ崈涓�' = 2000 * 10000 not 2000 + 10000
+ """
+ value = [0]
+ last_power = 0
+ for s in integer_symbols:
+ if isinstance(s, CND):
+ value[-1] = s.value
+ elif isinstance(s, CNU):
+ value[-1] *= pow(10, s.power)
+ if s.power > last_power:
+ value[:-1] = list(map(lambda v: v *
+ pow(10, s.power), value[:-1]))
+ last_power = s.power
+ value.append(0)
+ return sum(value)
+
+ system = create_system(numbering_type)
+ int_part, dec_part = string2symbols(chinese_string, system)
+ int_part = correct_symbols(int_part, system)
+ int_str = str(compute_value(int_part))
+ dec_str = ''.join([str(d.value) for d in dec_part])
+ if dec_part:
+ return '{0}.{1}'.format(int_str, dec_str)
+ else:
+ return int_str
+
+
+def num2chn(number_string, numbering_type=NUMBERING_TYPES[1], big=False,
+ traditional=False, alt_zero=False, alt_one=False, alt_two=True,
+ use_zeros=True, use_units=True):
+
+ def get_value(value_string, use_zeros=True):
+
+ striped_string = value_string.lstrip('0')
+
+ # record nothing if all zeros
+ if not striped_string:
+ return []
+
+ # record one digits
+ elif len(striped_string) == 1:
+ if use_zeros and len(value_string) != len(striped_string):
+ return [system.digits[0], system.digits[int(striped_string)]]
+ else:
+ return [system.digits[int(striped_string)]]
+
+ # recursively record multiple digits
+ else:
+ result_unit = next(u for u in reversed(
+ system.units) if u.power < len(striped_string))
+ result_string = value_string[:-result_unit.power]
+ return get_value(result_string) + [result_unit] + get_value(striped_string[-result_unit.power:])
+
+ system = create_system(numbering_type)
+
+ int_dec = number_string.split('.')
+ if len(int_dec) == 1:
+ int_string = int_dec[0]
+ dec_string = ""
+ elif len(int_dec) == 2:
+ int_string = int_dec[0]
+ dec_string = int_dec[1]
+ else:
+ raise ValueError(
+ "invalid input num string with more than one dot: {}".format(number_string))
+
+ if use_units and len(int_string) > 1:
+ result_symbols = get_value(int_string)
+ else:
+ result_symbols = [system.digits[int(c)] for c in int_string]
+ dec_symbols = [system.digits[int(c)] for c in dec_string]
+ if dec_string:
+ result_symbols += [system.math.point] + dec_symbols
+
+ if alt_two:
+ liang = CND(2, system.digits[2].alt_s, system.digits[2].alt_t,
+ system.digits[2].big_s, system.digits[2].big_t)
+ for i, v in enumerate(result_symbols):
+ if isinstance(v, CND) and v.value == 2:
+ next_symbol = result_symbols[i +
+ 1] if i < len(result_symbols) - 1 else None
+ previous_symbol = result_symbols[i - 1] if i > 0 else None
+ if isinstance(next_symbol, CNU) and isinstance(previous_symbol, (CNU, type(None))):
+ if next_symbol.power != 1 and ((previous_symbol is None) or (previous_symbol.power != 1)):
+ result_symbols[i] = liang
+
+ # if big is True, '涓�' will not be used and `alt_two` has no impact on output
+ if big:
+ attr_name = 'big_'
+ if traditional:
+ attr_name += 't'
+ else:
+ attr_name += 's'
+ else:
+ if traditional:
+ attr_name = 'traditional'
+ else:
+ attr_name = 'simplified'
+
+ result = ''.join([getattr(s, attr_name) for s in result_symbols])
+
+ # if not use_zeros:
+ # result = result.strip(getattr(system.digits[0], attr_name))
+
+ if alt_zero:
+ result = result.replace(
+ getattr(system.digits[0], attr_name), system.digits[0].alt_s)
+
+ if alt_one:
+ result = result.replace(
+ getattr(system.digits[1], attr_name), system.digits[1].alt_s)
+
+ for i, p in enumerate(POINT):
+ if result.startswith(p):
+ return CHINESE_DIGIS[0] + result
+
+ # ^10, 11, .., 19
+ if len(result) >= 2 and result[1] in [SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED[0],
+ SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL[0]] and \
+ result[0] in [CHINESE_DIGIS[1], BIG_CHINESE_DIGIS_SIMPLIFIED[1], BIG_CHINESE_DIGIS_TRADITIONAL[1]]:
+ result = result[1:]
+
+ return result
+
+
+# ================================================================================ #
+# different types of rewriters
+# ================================================================================ #
+class Cardinal:
+ """
+ CARDINAL绫�
+ """
+
+ def __init__(self, cardinal=None, chntext=None):
+ self.cardinal = cardinal
+ self.chntext = chntext
+
+ def chntext2cardinal(self):
+ return chn2num(self.chntext)
+
+ def cardinal2chntext(self):
+ return num2chn(self.cardinal)
+
+class Digit:
+ """
+ DIGIT绫�
+ """
+
+ def __init__(self, digit=None, chntext=None):
+ self.digit = digit
+ self.chntext = chntext
+
+ # def chntext2digit(self):
+ # return chn2num(self.chntext)
+
+ def digit2chntext(self):
+ return num2chn(self.digit, alt_two=False, use_units=False)
+
+
+class TelePhone:
+ """
+ TELEPHONE绫�
+ """
+
+ def __init__(self, telephone=None, raw_chntext=None, chntext=None):
+ self.telephone = telephone
+ self.raw_chntext = raw_chntext
+ self.chntext = chntext
+
+ # def chntext2telephone(self):
+ # sil_parts = self.raw_chntext.split('<SIL>')
+ # self.telephone = '-'.join([
+ # str(chn2num(p)) for p in sil_parts
+ # ])
+ # return self.telephone
+
+ def telephone2chntext(self, fixed=False):
+
+ if fixed:
+ sil_parts = self.telephone.split('-')
+ self.raw_chntext = '<SIL>'.join([
+ num2chn(part, alt_two=False, use_units=False) for part in sil_parts
+ ])
+ self.chntext = self.raw_chntext.replace('<SIL>', '')
+ else:
+ sp_parts = self.telephone.strip('+').split()
+ self.raw_chntext = '<SP>'.join([
+ num2chn(part, alt_two=False, use_units=False) for part in sp_parts
+ ])
+ self.chntext = self.raw_chntext.replace('<SP>', '')
+ return self.chntext
+
+
+class Fraction:
+ """
+ FRACTION绫�
+ """
+
+ def __init__(self, fraction=None, chntext=None):
+ self.fraction = fraction
+ self.chntext = chntext
+
+ def chntext2fraction(self):
+ denominator, numerator = self.chntext.split('鍒嗕箣')
+ return chn2num(numerator) + '/' + chn2num(denominator)
+
+ def fraction2chntext(self):
+ numerator, denominator = self.fraction.split('/')
+ return num2chn(denominator) + '鍒嗕箣' + num2chn(numerator)
+
+
+class Date:
+ """
+ DATE绫�
+ """
+
+ def __init__(self, date=None, chntext=None):
+ self.date = date
+ self.chntext = chntext
+
+ # def chntext2date(self):
+ # chntext = self.chntext
+ # try:
+ # year, other = chntext.strip().split('骞�', maxsplit=1)
+ # year = Digit(chntext=year).digit2chntext() + '骞�'
+ # except ValueError:
+ # other = chntext
+ # year = ''
+ # if other:
+ # try:
+ # month, day = other.strip().split('鏈�', maxsplit=1)
+ # month = Cardinal(chntext=month).chntext2cardinal() + '鏈�'
+ # except ValueError:
+ # day = chntext
+ # month = ''
+ # if day:
+ # day = Cardinal(chntext=day[:-1]).chntext2cardinal() + day[-1]
+ # else:
+ # month = ''
+ # day = ''
+ # date = year + month + day
+ # self.date = date
+ # return self.date
+
+ def date2chntext(self):
+ date = self.date
+ try:
+ year, other = date.strip().split('骞�', 1)
+ year = Digit(digit=year).digit2chntext() + '骞�'
+ except ValueError:
+ other = date
+ year = ''
+ if other:
+ try:
+ month, day = other.strip().split('鏈�', 1)
+ month = Cardinal(cardinal=month).cardinal2chntext() + '鏈�'
+ except ValueError:
+ day = date
+ month = ''
+ if day:
+ day = Cardinal(cardinal=day[:-1]).cardinal2chntext() + day[-1]
+ else:
+ month = ''
+ day = ''
+ chntext = year + month + day
+ self.chntext = chntext
+ return self.chntext
+
+
+class Money:
+ """
+ MONEY绫�
+ """
+
+ def __init__(self, money=None, chntext=None):
+ self.money = money
+ self.chntext = chntext
+
+ # def chntext2money(self):
+ # return self.money
+
+ def money2chntext(self):
+ money = self.money
+ pattern = re.compile(r'(\d+(\.\d+)?)')
+ matchers = pattern.findall(money)
+ if matchers:
+ for matcher in matchers:
+ money = money.replace(matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext())
+ self.chntext = money
+ return self.chntext
+
+
+class Percentage:
+ """
+ PERCENTAGE绫�
+ """
+
+ def __init__(self, percentage=None, chntext=None):
+ self.percentage = percentage
+ self.chntext = chntext
+
+ def chntext2percentage(self):
+ return chn2num(self.chntext.strip().strip('鐧惧垎涔�')) + '%'
+
+ def percentage2chntext(self):
+ return '鐧惧垎涔�' + num2chn(self.percentage.strip().strip('%'))
+
+
+def remove_erhua(text, er_whitelist):
+ """
+ 鍘婚櫎鍎垮寲闊宠瘝涓殑鍎�:
+ 浠栧コ鍎垮湪閭h竟鍎� -> 浠栧コ鍎垮湪閭h竟
+ """
+
+ er_pattern = re.compile(er_whitelist)
+ new_str=''
+ while re.search('鍎�',text):
+ a = re.search('鍎�',text).span()
+ remove_er_flag = 0
+
+ if er_pattern.search(text):
+ b = er_pattern.search(text).span()
+ if b[0] <= a[0]:
+ remove_er_flag = 1
+
+ if remove_er_flag == 0 :
+ new_str = new_str + text[0:a[0]]
+ text = text[a[1]:]
+ else:
+ new_str = new_str + text[0:b[1]]
+ text = text[b[1]:]
+
+ text = new_str + text
+ return text
+
+# ================================================================================ #
+# NSW Normalizer
+# ================================================================================ #
+class NSWNormalizer:
+ def __init__(self, raw_text):
+ self.raw_text = '^' + raw_text + '$'
+ self.norm_text = ''
+
+ def _particular(self):
+ text = self.norm_text
+ pattern = re.compile(r"(([a-zA-Z]+)浜�([a-zA-Z]+))")
+ matchers = pattern.findall(text)
+ if matchers:
+ # print('particular')
+ for matcher in matchers:
+ text = text.replace(matcher[0], matcher[1]+'2'+matcher[2], 1)
+ self.norm_text = text
+ return self.norm_text
+
+ def normalize(self):
+ text = self.raw_text
+
+ # 瑙勮寖鍖栨棩鏈�
+ pattern = re.compile(r"\D+((([089]\d|(19|20)\d{2})骞�)?(\d{1,2}鏈�(\d{1,2}[鏃ュ彿])?)?)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('date')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Date(date=matcher[0]).date2chntext(), 1)
+
+ # 瑙勮寖鍖栭噾閽�
+ pattern = re.compile(r"\D+((\d+(\.\d+)?)[澶氫綑鍑燷?" + CURRENCY_UNITS + r"(\d" + CURRENCY_UNITS + r"?)?)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('money')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Money(money=matcher[0]).money2chntext(), 1)
+
+ # 瑙勮寖鍖栧浐璇�/鎵嬫満鍙风爜
+ # 鎵嬫満
+ # http://www.jihaoba.com/news/show/13680
+ # 绉诲姩锛�139銆�138銆�137銆�136銆�135銆�134銆�159銆�158銆�157銆�150銆�151銆�152銆�188銆�187銆�182銆�183銆�184銆�178銆�198
+ # 鑱旈�氾細130銆�131銆�132銆�156銆�155銆�186銆�185銆�176
+ # 鐢典俊锛�133銆�153銆�189銆�180銆�181銆�177
+ pattern = re.compile(r"\D((\+?86 ?)?1([38]\d|5[0-35-9]|7[678]|9[89])\d{8})\D")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('telephone')
+ for matcher in matchers:
+ text = text.replace(matcher[0], TelePhone(telephone=matcher[0]).telephone2chntext(), 1)
+ # 鍥鸿瘽
+ pattern = re.compile(r"\D((0(10|2[1-3]|[3-9]\d{2})-?)?[1-9]\d{6,7})\D")
+ matchers = pattern.findall(text)
+ if matchers:
+ # print('fixed telephone')
+ for matcher in matchers:
+ text = text.replace(matcher[0], TelePhone(telephone=matcher[0]).telephone2chntext(fixed=True), 1)
+
+ # 瑙勮寖鍖栧垎鏁�
+ pattern = re.compile(r"(\d+/\d+)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('fraction')
+ for matcher in matchers:
+ text = text.replace(matcher, Fraction(fraction=matcher).fraction2chntext(), 1)
+
+ # 瑙勮寖鍖栫櫨鍒嗘暟
+ text = text.replace('锛�', '%')
+ pattern = re.compile(r"(\d+(\.\d+)?%)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('percentage')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Percentage(percentage=matcher[0]).percentage2chntext(), 1)
+
+ # 瑙勮寖鍖栫函鏁�+閲忚瘝
+ pattern = re.compile(r"(\d+(\.\d+)?)[澶氫綑鍑燷?" + COM_QUANTIFIERS)
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('cardinal+quantifier')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext(), 1)
+
+ # 瑙勮寖鍖栨暟瀛楃紪鍙�
+ pattern = re.compile(r"(\d{4,32})")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('digit')
+ for matcher in matchers:
+ text = text.replace(matcher, Digit(digit=matcher).digit2chntext(), 1)
+
+ # 瑙勮寖鍖栫函鏁�
+ pattern = re.compile(r"(\d+(\.\d+)?)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('cardinal')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext(), 1)
+
+ self.norm_text = text
+ self._particular()
+
+ return self.norm_text.lstrip('^').rstrip('$')
+
+
+def nsw_test_case(raw_text):
+ print('I:' + raw_text)
+ print('O:' + NSWNormalizer(raw_text).normalize())
+ print('')
+
+
+def nsw_test():
+ nsw_test_case('鍥鸿瘽锛�0595-23865596鎴�23880880銆�')
+ nsw_test_case('鍥鸿瘽锛�0595-23865596鎴�23880880銆�')
+ nsw_test_case('鎵嬫満锛�+86 19859213959鎴�15659451527銆�')
+ nsw_test_case('鍒嗘暟锛�32477/76391銆�')
+ nsw_test_case('鐧惧垎鏁帮細80.03%銆�')
+ nsw_test_case('缂栧彿锛�31520181154418銆�')
+ nsw_test_case('绾暟锛�2983.07鍏嬫垨12345.60绫炽��')
+ nsw_test_case('鏃ユ湡锛�1999骞�2鏈�20鏃ユ垨09骞�3鏈�15鍙枫��')
+ nsw_test_case('閲戦挶锛�12鍧�5锛�34.5鍏冿紝20.1涓�')
+ nsw_test_case('鐗规畩锛歄2O鎴朆2C銆�')
+ nsw_test_case('3456涓囧惃')
+ nsw_test_case('2938涓�')
+ nsw_test_case('938')
+ nsw_test_case('浠婂ぉ鍚冧簡115涓皬绗煎寘231涓澶�')
+ nsw_test_case('鏈�62锛呯殑姒傜巼')
+
+
+if __name__ == '__main__':
+ #nsw_test()
+
+ p = argparse.ArgumentParser()
+ p.add_argument('ifile', help='input filename, assume utf-8 encoding')
+ p.add_argument('ofile', help='output filename')
+ p.add_argument('--to_upper', action='store_true', help='convert to upper case')
+ p.add_argument('--to_lower', action='store_true', help='convert to lower case')
+ p.add_argument('--has_key', action='store_true', help="input text has Kaldi's key as first field.")
+ p.add_argument('--remove_fillers', type=bool, default=True, help='remove filler chars such as "鍛�, 鍟�"')
+ p.add_argument('--remove_erhua', type=bool, default=True, help='remove erhua chars such as "杩欏効"')
+ p.add_argument('--log_interval', type=int, default=10000, help='log interval in number of processed lines')
+ args = p.parse_args()
+
+ ifile = codecs.open(args.ifile, 'r', 'utf8')
+ ofile = codecs.open(args.ofile, 'w+', 'utf8')
+
+ n = 0
+ for l in ifile:
+ key = ''
+ text = ''
+ if args.has_key:
+ cols = l.split(maxsplit=1)
+ key = cols[0]
+ if len(cols) == 2:
+ text = cols[1].strip()
+ else:
+ text = ''
+ else:
+ text = l.strip()
+
+ # cases
+ if args.to_upper and args.to_lower:
+ sys.stderr.write('text norm: to_upper OR to_lower?')
+ exit(1)
+ if args.to_upper:
+ text = text.upper()
+ if args.to_lower:
+ text = text.lower()
+
+ # Filler chars removal
+ if args.remove_fillers:
+ for ch in FILLER_CHARS:
+ text = text.replace(ch, '')
+
+ if args.remove_erhua:
+ text = remove_erhua(text, ER_WHITELIST)
+
+ # NSW(Non-Standard-Word) normalization
+ text = NSWNormalizer(text).normalize()
+
+ # Punctuations removal
+ old_chars = CHINESE_PUNC_LIST + string.punctuation # includes all CN and EN punctuations
+ new_chars = ' ' * len(old_chars)
+ del_chars = ''
+ text = text.translate(str.maketrans(old_chars, new_chars, del_chars))
+
+ #
+ if args.has_key:
+ ofile.write(key + '\t' + text + '\n')
+ else:
+ ofile.write(text + '\n')
+
+ n += 1
+ if n % args.log_interval == 0:
+ sys.stderr.write("text norm: {} lines done.\n".format(n))
+
+ sys.stderr.write("text norm: {} lines done in total.\n".format(n))
+
+ ifile.close()
+ ofile.close()
diff --git a/egs/aishell/tranformer/conf/train_asr_conformer.yaml b/egs/aishell/tranformer/conf/train_asr_conformer.yaml
deleted file mode 100644
index ddf217e..0000000
--- a/egs/aishell/tranformer/conf/train_asr_conformer.yaml
+++ /dev/null
@@ -1,80 +0,0 @@
-# network architecture
-# encoder related
-encoder: conformer
-encoder_conf:
- output_size: 256 # dimension of attention
- attention_heads: 4
- linear_units: 2048 # the number of units of position-wise feed forward
- num_blocks: 12 # the number of encoder blocks
- dropout_rate: 0.1
- positional_dropout_rate: 0.1
- attention_dropout_rate: 0.0
- input_layer: conv2d # encoder architecture type
- normalize_before: true
- pos_enc_layer_type: rel_pos
- selfattention_layer_type: rel_selfattn
- activation_type: swish
- macaron_style: true
- use_cnn_module: true
- cnn_module_kernel: 15
-
-# decoder related
-decoder: transformer
-decoder_conf:
- attention_heads: 4
- linear_units: 2048
- num_blocks: 6
- dropout_rate: 0.1
- positional_dropout_rate: 0.1
- self_attention_dropout_rate: 0.0
- src_attention_dropout_rate: 0.0
-
-# hybrid CTC/attention
-model_conf:
- ctc_weight: 0.3
- lsm_weight: 0.1 # label smoothing option
- length_normalized_loss: false
-
-# minibatch related
-batch_type: length
-batch_bins: 25000
-num_workers: 16
-
-# optimization related
-accum_grad: 1
-grad_clip: 5
-max_epoch: 50
-val_scheduler_criterion:
- - valid
- - acc
-best_model_criterion:
-- - valid
- - acc
- - max
-keep_nbest_models: 10
-
-optim: adam
-optim_conf:
- lr: 0.0005
-scheduler: warmuplr
-scheduler_conf:
- warmup_steps: 30000
-
-specaug: specaug
-specaug_conf:
- apply_time_warp: true
- time_warp_window: 5
- time_warp_mode: bicubic
- apply_freq_mask: true
- freq_mask_width_range:
- - 0
- - 30
- num_freq_mask: 2
- apply_time_mask: true
- time_mask_width_range:
- - 0
- - 40
- num_time_mask: 2
-
-log_interval: 50
-normalize: None
diff --git a/egs/aishell/tranformer/run.sh b/egs/aishell/tranformer/run.sh
index 16ebc67..4c307b0 100755
--- a/egs/aishell/tranformer/run.sh
+++ b/egs/aishell/tranformer/run.sh
@@ -10,9 +10,10 @@
# for gpu decoding, inference_nj=ngpu*njob; for cpu decoding, inference_nj=njob
njob=8
train_cmd=utils/run.pl
+infer_cmd=utils/run.pl
# general configuration
-feats_dir=".." #feature output dictionary, for large data
+feats_dir="../DATA" #feature output dictionary, for large data
exp_dir="."
lang=zh
dumpdir=dump/fbank
@@ -59,8 +60,10 @@
if ${gpu_inference}; then
inference_nj=$[${ngpu}*${njob}]
+ _ngpu=1
else
inference_nj=$njob
+ _ngpu=0
fi
if [ ${stage} -le 0 ] && [ ${stop_stage} -ge 0 ]; then
@@ -83,18 +86,18 @@
echo "stage 1: Feature Generation"
# compute fbank features
fbankdir=${feats_dir}/fbank
- utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj --speed_perturb ${speed_perturb} \
+ utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj --feats_dim ${feats_dim} --sample_frequency ${sample_frequency} --speed_perturb ${speed_perturb} \
${feats_dir}/data/train ${exp_dir}/exp/make_fbank/train ${fbankdir}/train
utils/fix_data_feat.sh ${fbankdir}/train
- utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj \
+ utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj --feats_dim ${feats_dim} --sample_frequency ${sample_frequency} \
${feats_dir}/data/dev ${exp_dir}/exp/make_fbank/dev ${fbankdir}/dev
utils/fix_data_feat.sh ${fbankdir}/dev
- utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj \
+ utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj --feats_dim ${feats_dim} --sample_frequency ${sample_frequency} \
${feats_dir}/data/test ${exp_dir}/exp/make_fbank/test ${fbankdir}/test
utils/fix_data_feat.sh ${fbankdir}/test
# compute global cmvn
- utils/compute_cmvn.sh --cmd "$train_cmd" --nj $nj \
+ utils/compute_cmvn.sh --cmd "$train_cmd" --nj $nj --feats_dim ${feats_dim} \
${fbankdir}/train ${exp_dir}/exp/make_fbank/train
# apply cmvn
@@ -112,6 +115,10 @@
utils/fix_data_feat.sh ${feat_train_dir}
utils/fix_data_feat.sh ${feat_dev_dir}
utils/fix_data_feat.sh ${feat_test_dir}
+
+ #generate ark list
+ utils/gen_ark_list.sh --cmd "$train_cmd" --nj $nj ${feat_train_dir} ${fbankdir}/train ${feat_train_dir}
+ utils/gen_ark_list.sh --cmd "$train_cmd" --nj $nj ${feat_dev_dir} ${fbankdir}/dev ${feat_dev_dir}
fi
token_list=${feats_dir}/data/${lang}_token_list/char/tokens.txt
@@ -140,9 +147,10 @@
# Training Stage
world_size=$gpu_num # run on one machine
if [ ${stage} -le 3 ] && [ ${stop_stage} -ge 3 ]; then
+ echo "stage 3: Training"
mkdir -p ${exp_dir}/exp/${model_dir}
mkdir -p ${exp_dir}/exp/${model_dir}/log
- INIT_FILE=$exp_dir/ddp_init
+ INIT_FILE=${exp_dir}/exp/${model_dir}/ddp_init
if [ -f $INIT_FILE ];then
rm -f $INIT_FILE
fi
@@ -184,25 +192,56 @@
# Testing Stage
if [ ${stage} -le 4 ] && [ ${stop_stage} -ge 4 ]; then
- utils/easy_asr_infer.sh \
- --lang zh \
- --datadir ${feats_dir} \
- --feats_type ${feats_type} \
- --feats_dim ${feats_dim} \
- --token_type ${token_type} \
- --gpu_inference ${gpu_inference} \
- --inference_config "${inference_config}" \
- --test_sets "${test_sets}" \
- --token_list $token_list \
- --asr_exp ${exp_dir}/${model_dir} \
- --stage 12 \
- --stop_stage 12 \
- --scp $scp \
- --text text \
- --inference_nj $inference_nj \
- --njob $njob \
- --inference_asr_model $inference_asr_model \
- --gpuid_list $gpuid_list \
- --mode asr
-fi
+ echo "stage 4: Inference"
+ for dset in ${test_sets}; do
+ asr_exp=${exp_dir}/exp/${model_dir}
+ inference_tag="$(basename "${inference_config}" .yaml)"
+ _dir="${asr_exp}/${inference_tag}/${inference_asr_model}/${dset}"
+ _logdir="${_dir}/logdir"
+ if [ -d ${_dir} ]; then
+ echo "${_dir} is already exists. if you want to decode again, please delete this dir first."
+ exit 0
+ fi
+ mkdir -p "${_logdir}"
+ _data="${feats_dir}/${dumpdir}/${dset}"
+ key_file=${_data}/${scp}
+ num_scp_file="$(<${key_file} wc -l)"
+ _nj=$([ $inference_nj -le $num_scp_file ] && echo "$inference_nj" || echo "$num_scp_file")
+ split_scps=
+ for n in $(seq "${_nj}"); do
+ split_scps+=" ${_logdir}/keys.${n}.scp"
+ done
+ # shellcheck disable=SC2086
+ utils/split_scp.pl "${key_file}" ${split_scps}
+ _opts=
+ if [ -n "${inference_config}" ]; then
+ _opts+="--config ${inference_config} "
+ fi
+ ${infer_cmd} --gpu "${_ngpu}" --max-jobs-run "${_nj}" JOB=1:"${_nj}" "${_logdir}"/asr_inference.JOB.log \
+ python -m funasr.bin.asr_inference_launch \
+ --batch_size 1 \
+ --ngpu "${_ngpu}" \
+ --njob ${njob} \
+ --gpuid_list ${gpuid_list} \
+ --data_path_and_name_and_type "${_data}/${scp},speech,${type}" \
+ --key_file "${_logdir}"/keys.JOB.scp \
+ --asr_train_config "${asr_exp}"/config.yaml \
+ --asr_model_file "${asr_exp}"/"${inference_asr_model}" \
+ --output_dir "${_logdir}"/output.JOB \
+ --mode asr \
+ ${_opts}
+ for f in token token_int score text; do
+ if [ -f "${_logdir}/output.1/1best_recog/${f}" ]; then
+ for i in $(seq "${_nj}"); do
+ cat "${_logdir}/output.${i}/1best_recog/${f}"
+ done | sort -k1 >"${_dir}/${f}"
+ fi
+ done
+ python utils/proce_text.py ${_dir}/text ${_dir}/text.proc
+ python utils/proce_text.py ${_data}/text ${_data}/text.proc
+ python utils/compute_wer.py ${_data}/text.proc ${_dir}/text.proc ${_dir}/text.cer
+ tail -n 3 ${_dir}/text.cer > ${_dir}/text.cer.txt
+ cat ${_dir}/text.cer.txt
+ done
+fi
diff --git a/egs/aishell/tranformer/utils/combine_cmvn_file.py b/egs/aishell/tranformer/utils/combine_cmvn_file.py
index e16174c..b2974a4 100755
--- a/egs/aishell/tranformer/utils/combine_cmvn_file.py
+++ b/egs/aishell/tranformer/utils/combine_cmvn_file.py
@@ -8,6 +8,13 @@
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
)
parser.add_argument(
+ "--dims",
+ "-d",
+ default=80,
+ type=int,
+ help="feature dims",
+ )
+ parser.add_argument(
"--cmvn-dir",
"-c",
default=False,
@@ -39,8 +46,8 @@
parser = get_parser()
args = parser.parse_args()
- total_means = 0.0
- total_vars = 0.0
+ total_means = np.zeros(args.dims)
+ total_vars = np.zeros(args.dims)
total_frames = 0
cmvn_file = args.output_dir + "/cmvn.json"
diff --git a/egs/aishell/tranformer/utils/compute_cmvn.py b/egs/aishell/tranformer/utils/compute_cmvn.py
index 988d6dc..2b96e26 100755
--- a/egs/aishell/tranformer/utils/compute_cmvn.py
+++ b/egs/aishell/tranformer/utils/compute_cmvn.py
@@ -11,6 +11,13 @@
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
)
parser.add_argument(
+ "--dims",
+ "-d",
+ default=80,
+ type=int,
+ help="feature dims",
+ )
+ parser.add_argument(
"--ark-file",
"-a",
default=False,
@@ -44,8 +51,8 @@
ark_file = args.ark_file + "/feats." + str(args.ark_index) + ".ark"
cmvn_file = args.output_dir + "/cmvn." + str(args.ark_index) + ".json"
- mean_stats = 0.0
- var_stats = 0.0
+ mean_stats = np.zeros(args.dims)
+ var_stats = np.zeros(args.dims)
total_frames = 0
with ReadHelper('ark:{}'.format(ark_file)) as ark_reader:
diff --git a/egs/aishell/tranformer/utils/compute_cmvn.sh b/egs/aishell/tranformer/utils/compute_cmvn.sh
index 3a30190..12173ee 100755
--- a/egs/aishell/tranformer/utils/compute_cmvn.sh
+++ b/egs/aishell/tranformer/utils/compute_cmvn.sh
@@ -4,6 +4,7 @@
# Begin configuration section.
nj=32
cmd=./utils/run.pl
+feats_dim=80
echo "$0 $@"
@@ -16,9 +17,9 @@
mkdir -p ${logdir}
$cmd JOB=1:$nj $logdir/cmvn.JOB.log \
- python utils/compute_cmvn.py -a $fbankdir/ark -i JOB -o ${output_dir} \
+ python utils/compute_cmvn.py -d ${feats_dim} -a $fbankdir/ark -i JOB -o ${output_dir} \
|| exit 1;
-python utils/combine_cmvn_file.py -c ${output_dir} -n $nj -o $fbankdir
+python utils/combine_cmvn_file.py -d ${feats_dim} -c ${output_dir} -n $nj -o $fbankdir
echo "$0: Succeeded compute global cmvn"
diff --git a/egs/aishell/tranformer/utils/compute_fbank.sh b/egs/aishell/tranformer/utils/compute_fbank.sh
index b456b4d..92a4fe6 100755
--- a/egs/aishell/tranformer/utils/compute_fbank.sh
+++ b/egs/aishell/tranformer/utils/compute_fbank.sh
@@ -6,7 +6,7 @@
cmd=./utils/run.pl
# feature configuration
-feat_dims=80
+feats_dim=80
sample_frequency=16000
speed_perturb="1.0"
@@ -29,7 +29,7 @@
$cmd JOB=1:$nj $logdir/make_fbank.JOB.log \
python utils/compute_fbank.py -w $data/split${nj}/JOB/wav.scp -t $data/split${nj}/JOB/text \
- -d $feat_dims -s $sample_frequency -p ${speed_perturb} -a JOB -o ${fbankdir} \
+ -d $feats_dim -s $sample_frequency -p ${speed_perturb} -a JOB -o ${fbankdir} \
|| exit 1;
for n in $(seq $nj); do
diff --git a/egs/aishell/tranformer/utils/easy_asr_infer.sh b/egs/aishell/tranformer/utils/easy_asr_infer.sh
deleted file mode 100755
index 1b8db34..0000000
--- a/egs/aishell/tranformer/utils/easy_asr_infer.sh
+++ /dev/null
@@ -1,407 +0,0 @@
-#!/usr/bin/env bash
-
-# Set bash to 'debug' mode, it will exit on :
-# -e 'error', -u 'undefined variable', -o ... 'error in pipeline', -x 'print commands',
-set -e
-set -u
-set -o pipefail
-
-log() {
- local fname=${BASH_SOURCE[1]##*/}
- echo -e "$(date '+%Y-%m-%dT%H:%M:%S') (${fname}:${BASH_LINENO[0]}:${FUNCNAME[1]}) $*"
-}
-min() {
- local a b
- a=$1
- for b in "$@"; do
- if [ "${b}" -le "${a}" ]; then
- a="${b}"
- fi
- done
- echo "${a}"
-}
-SECONDS=0
-
-# General configuration
-stage=1 # Processes starts from the specified stage.
-stop_stage=10000 # Processes is stopped at the specified stage.
-skip_data_prep=true # Skip data preparation stages.
-skip_train=false # Skip training stages.
-skip_eval=false # Skip decoding and evaluation stages.
-skip_upload=true # Skip packing and uploading stages.
-skip_upload_hf=true # Skip uploading to hugging face stages.
-cuda_cmd=utils/run.pl
-decode_cmd=utils/run.pl
-ngpu=1 # The number of gpus ("0" uses cpu, otherwise use gpu).
-njob=1 # the number of jobs for each gpu
-gpuid_list=
-num_nodes=1 # The number of nodes.
-nj=32 # The number of parallel jobs.
-inference_nj=32 # The number of parallel jobs in decoding.
-gpu_inference=true # Whether to perform gpu decoding, set false for cpu decoding
-datadir="./"
-dumpdir=dump # Directory to dump features.
-expdir=exp # Directory to save experiments.
-python=python # Specify python to execute funasr commands.
-
-# Data preparation related
-local_data_opts= # The options given to local/data.sh.
-
-# Speed perturbation related
-speed_perturb_factors= # perturbation factors, e.g. "0.9 1.0 1.1" (separated by space).
-
-# Feature extraction related
-feats_type=fbank # Feature type (raw or fbank_pitch).
-feats_dim=
-audio_format=flac # Audio format: wav, flac, wav.ark, flac.ark (only in feats_type=raw).
-fs=16k # Sampling rate.
-min_wav_duration=0.1 # Minimum duration in second.
-max_wav_duration=20 # Maximum duration in second.
-
-# Tokenization related
-token_type=bpe # Tokenization type (char or bpe).
-nbpe=30 # The number of BPE vocabulary.
-bpemode=unigram # Mode of BPE (unigram or bpe).
-oov="<unk>" # Out of vocabulary symbol.
-blank="<blank>" # CTC blank symbol
-sos_eos="<sos/eos>" # sos and eos symbole
-bpe_input_sentence_size=100000000 # Size of input sentence for BPE.
-bpe_nlsyms= # non-linguistic symbols list, separated by a comma, for BPE
-bpe_char_cover=1.0 # character coverage when modeling BPE
-
-# Ngram model related
-use_ngram=false
-ngram_exp=
-ngram_num=3
-
-# Language model related
-use_lm=false # Use language model for ASR decoding.
-lm_tag= # Suffix to the result dir for language model training.
-lm_exp= # Specify the directory path for LM experiment.
- # If this option is specified, lm_tag is ignored.
-lm_stats_dir= # Specify the directory path for LM statistics.
-lm_config= # Config for language model training.
-lm_args= # Arguments for language model training, e.g., "--max_epoch 10".
- # Note that it will overwrite args in lm config.
-use_word_lm=false # Whether to use word language model.
-num_splits_lm=1 # Number of splitting for lm corpus.
-# shellcheck disable=SC2034
-word_vocab_size=10000 # Size of word vocabulary.
-
-# ASR model related
-asr_tag= # Suffix to the result dir for asr model training.
-asr_exp= # Specify the directory path for ASR experiment.
- # If this option is specified, asr_tag is ignored.
-asr_stats_dir= # Specify the directory path for ASR statistics.
-asr_config= # Config for asr model training.
-asr_args= # Arguments for asr model training, e.g., "--max_epoch 10".
- # Note that it will overwrite args in asr config.
-pretrained_model= # Pretrained model to load
-ignore_init_mismatch=false # Ignore initial mismatch
-feats_normalize=global_mvn # Normalizaton layer type.
-num_splits_asr=1 # Number of splitting for lm corpus.
-
-# Upload model related
-hf_repo=
-
-# Decoding related
-use_k2=false # Whether to use k2 based decoder
-k2_ctc_decoding=true
-use_nbest_rescoring=true # use transformer-decoder
- # and transformer language model for nbest rescoring
-num_paths=1000 # The 3rd argument of k2.random_paths.
-nll_batch_size=100 # Affect GPU memory usage when computing nll
- # during nbest rescoring
-k2_config=./conf/decode_asr_transformer_with_k2.yaml
-
-use_streaming=false # Whether to use streaming decoding
-
-use_maskctc=false # Whether to use maskctc decoding
-
-batch_size=1
-inference_tag= # Suffix to the result dir for decoding.
-inference_config= # Config for decoding.
-inference_args= # Arguments for decoding, e.g., "--lm_weight 0.1".
- # Note that it will overwrite args in inference config.
-inference_lm=valid.loss.ave.pth # Language model path for decoding.
-inference_ngram=${ngram_num}gram.bin
-inference_asr_model=valid.acc.ave.pth # ASR model path for decoding.
- # e.g.
- # inference_asr_model=train.loss.best.pth
- # inference_asr_model=3epoch.pth
- # inference_asr_model=valid.acc.best.pth
- # inference_asr_model=valid.loss.ave.pth
-download_model= # Download a model from Model Zoo and use it for decoding.
-
-# [Task dependent] Set the datadir name created by local/data.sh
-train_set= # Name of training set.
-valid_set= # Name of validation set used for monitoring/tuning network training.
-test_sets= # Names of test sets. Multiple items (e.g., both dev and eval sets) can be specified.
-bpe_train_text= # Text file path of bpe training set.
-lm_train_text= # Text file path of language model training set.
-lm_dev_text= # Text file path of language model development set.
-lm_test_text= # Text file path of language model evaluation set.
-nlsyms_txt=none # Non-linguistic symbol list if existing.
-cleaner=none # Text cleaner.
-g2p=none # g2p method (needed if token_type=phn).
-lang=noinfo # The language type of corpus.
-score_opts= # The options given to sclite scoring
-local_score_opts= # The options given to local/score.sh.
-asr_speech_fold_length=800 # fold_length for speech data during ASR training.
-asr_text_fold_length=150 # fold_length for text data during ASR training.
-lm_fold_length=150 # fold_length for LM training.
-
-oss_path=
-token_list=
-scp=
-text=
-
-mode=
-
-help_message=$(cat << EOF
-Usage: $0 --train-set "<train_set_name>" --valid-set "<valid_set_name>" --test_sets "<test_set_names>"
-
-Options:
- # General configuration
- --stage # Processes starts from the specified stage (default="${stage}").
- --stop_stage # Processes is stopped at the specified stage (default="${stop_stage}").
- --skip_data_prep # Skip data preparation stages (default="${skip_data_prep}").
- --skip_train # Skip training stages (default="${skip_train}").
- --skip_eval # Skip decoding and evaluation stages (default="${skip_eval}").
- --skip_upload # Skip packing and uploading stages (default="${skip_upload}").
- --ngpu # The number of gpus ("0" uses cpu, otherwise use gpu, default="${ngpu}").
- --num_nodes # The number of nodes (default="${num_nodes}").
- --nj # The number of parallel jobs (default="${nj}").
- --inference_nj # The number of parallel jobs in decoding (default="${inference_nj}").
- --gpu_inference # Whether to perform gpu decoding (default="${gpu_inference}").
- --dumpdir # Directory to dump features (default="${dumpdir}").
- --expdir # Directory to save experiments (default="${expdir}").
- --python # Specify python to execute espnet commands (default="${python}").
-
- # Data preparation related
- --local_data_opts # The options given to local/data.sh (default="${local_data_opts}").
-
- # Speed perturbation related
- --speed_perturb_factors # speed perturbation factors, e.g. "0.9 1.0 1.1" (separated by space, default="${speed_perturb_factors}").
-
- # Feature extraction related
- --feats_type # Feature type (raw, fbank_pitch or extracted, default="${feats_type}").
- --audio_format # Audio format: wav, flac, wav.ark, flac.ark (only in feats_type=raw, default="${audio_format}").
- --fs # Sampling rate (default="${fs}").
- --min_wav_duration # Minimum duration in second (default="${min_wav_duration}").
- --max_wav_duration # Maximum duration in second (default="${max_wav_duration}").
-
- # Tokenization related
- --token_type # Tokenization type (char or bpe, default="${token_type}").
- --nbpe # The number of BPE vocabulary (default="${nbpe}").
- --bpemode # Mode of BPE (unigram or bpe, default="${bpemode}").
- --oov # Out of vocabulary symbol (default="${oov}").
- --blank # CTC blank symbol (default="${blank}").
- --sos_eos # sos and eos symbole (default="${sos_eos}").
- --bpe_input_sentence_size # Size of input sentence for BPE (default="${bpe_input_sentence_size}").
- --bpe_nlsyms # Non-linguistic symbol list for sentencepiece, separated by a comma. (default="${bpe_nlsyms}").
- --bpe_char_cover # Character coverage when modeling BPE (default="${bpe_char_cover}").
-
- # Language model related
- --lm_tag # Suffix to the result dir for language model training (default="${lm_tag}").
- --lm_exp # Specify the directory path for LM experiment.
- # If this option is specified, lm_tag is ignored (default="${lm_exp}").
- --lm_stats_dir # Specify the directory path for LM statistics (default="${lm_stats_dir}").
- --lm_config # Config for language model training (default="${lm_config}").
- --lm_args # Arguments for language model training (default="${lm_args}").
- # e.g., --lm_args "--max_epoch 10"
- # Note that it will overwrite args in lm config.
- --use_word_lm # Whether to use word language model (default="${use_word_lm}").
- --word_vocab_size # Size of word vocabulary (default="${word_vocab_size}").
- --num_splits_lm # Number of splitting for lm corpus (default="${num_splits_lm}").
-
- # ASR model related
- --asr_tag # Suffix to the result dir for asr model training (default="${asr_tag}").
- --asr_exp # Specify the directory path for ASR experiment.
- # If this option is specified, asr_tag is ignored (default="${asr_exp}").
- --asr_stats_dir # Specify the directory path for ASR statistics (default="${asr_stats_dir}").
- --asr_config # Config for asr model training (default="${asr_config}").
- --asr_args # Arguments for asr model training (default="${asr_args}").
- # e.g., --asr_args "--max_epoch 10"
- # Note that it will overwrite args in asr config.
- --pretrained_model= # Pretrained model to load (default="${pretrained_model}").
- --ignore_init_mismatch= # Ignore mismatch parameter init with pretrained model (default="${ignore_init_mismatch}").
- --feats_normalize # Normalizaton layer type (default="${feats_normalize}").
- --num_splits_asr # Number of splitting for lm corpus (default="${num_splits_asr}").
-
- # Decoding related
- --inference_tag # Suffix to the result dir for decoding (default="${inference_tag}").
- --inference_config # Config for decoding (default="${inference_config}").
- --inference_args # Arguments for decoding (default="${inference_args}").
- # e.g., --inference_args "--lm_weight 0.1"
- # Note that it will overwrite args in inference config.
- --inference_lm # Language model path for decoding (default="${inference_lm}").
- --inference_asr_model # ASR model path for decoding (default="${inference_asr_model}").
- --download_model # Download a model from Model Zoo and use it for decoding (default="${download_model}").
- --use_streaming # Whether to use streaming decoding (default="${use_streaming}").
- --use_maskctc # Whether to use maskctc decoding (default="${use_streaming}").
-
- # [Task dependent] Set the datadir name created by local/data.sh
- --train_set # Name of training set (required).
- --valid_set # Name of validation set used for monitoring/tuning network training (required).
- --test_sets # Names of test sets.
- # Multiple items (e.g., both dev and eval sets) can be specified (required).
- --bpe_train_text # Text file path of bpe training set.
- --lm_train_text # Text file path of language model training set.
- --lm_dev_text # Text file path of language model development set (default="${lm_dev_text}").
- --lm_test_text # Text file path of language model evaluation set (default="${lm_test_text}").
- --nlsyms_txt # Non-linguistic symbol list if existing (default="${nlsyms_txt}").
- --cleaner # Text cleaner (default="${cleaner}").
- --g2p # g2p method (default="${g2p}").
- --lang # The language type of corpus (default=${lang}).
- --score_opts # The options given to sclite scoring (default="{score_opts}").
- --local_score_opts # The options given to local/score.sh (default="{local_score_opts}").
- --asr_speech_fold_length # fold_length for speech data during ASR training (default="${asr_speech_fold_length}").
- --asr_text_fold_length # fold_length for text data during ASR training (default="${asr_text_fold_length}").
- --lm_fold_length # fold_length for LM training (default="${lm_fold_length}").
-EOF
-)
-
-log "$0 $*"
-# Save command line args for logging (they will be lost after utils/parse_options.sh)
-run_args=$(utils/print_args.py $0 "$@")
-. utils/parse_options.sh
-
-if [ $# -ne 0 ]; then
- log "${help_message}"
- log "Error: No positional arguments are required."
- exit 2
-fi
-
-# set absolute dump dir path
-dumpdir=${datadir}/${dumpdir}
-
-if [ -z "${inference_tag}" ]; then
- if [ -n "${inference_config}" ]; then
- inference_tag="$(basename "${inference_config}" .yaml)"
- else
- inference_tag=inference
- fi
-
- if "${use_k2}"; then
- inference_tag+="_use_k2"
- inference_tag+="_k2_ctc_decoding_${k2_ctc_decoding}"
- inference_tag+="_use_nbest_rescoring_${use_nbest_rescoring}"
- fi
-fi
-
-# ========================== Main stages start from here. ==========================
-
-if [ ${stage} -le 12 ] && [ ${stop_stage} -ge 12 ]; then
- log "Stage 12: Decoding: training_dir=${asr_exp}"
-
- if ${gpu_inference}; then
- _cmd="${cuda_cmd}"
- _ngpu=1
- else
- _cmd="${decode_cmd}"
- _ngpu=0
- fi
-
- _opts=
- if [ -n "${inference_config}" ]; then
- _opts+="--config ${inference_config} "
- fi
-
- if "${use_lm}"; then
- if "${use_word_lm}"; then
- _opts+="--word_lm_train_config ${lm_exp}/config.yaml "
- _opts+="--word_lm_file ${lm_exp}/${inference_lm} "
- else
- _opts+="--lm_train_config ${lm_exp}/config.yaml "
- _opts+="--lm_file ${lm_exp}/${inference_lm} "
- fi
- fi
-
- if "${use_ngram}"; then
- _opts+="--ngram_file ${ngram_exp}/${inference_ngram}"
- inference_tag=${inference_tag}.${inference_ngram}
- fi
-
- # 2. Generate run.sh
- log "Generate '${asr_exp}/${inference_tag}/run.sh'. You can resume the process from stage 12 using this script"
- mkdir -p "${asr_exp}/${inference_tag}"; echo "${run_args} --stage 12 \"\$@\"; exit \$?" > "${asr_exp}/${inference_tag}/run.sh"; chmod +x "${asr_exp}/${inference_tag}/run.sh"
-
- if "${use_streaming}"; then
- asr_inference_tool="funasr.bin.asr_inference_streaming"
- elif "${use_maskctc}"; then
- asr_inference_tool="funasr.bin.asr_inference_maskctc"
- else
- asr_inference_tool="funasr.bin.asr_inference_launch"
- fi
-
- for dset in ${test_sets}; do
- if [ $feats_type == "ark_wav" ]; then
- _data="${dumpdir}/wav/${dset}"
- else
- _data="${dumpdir}/$feats_type/${dset}"
- fi
- _dir="${asr_exp}/${inference_tag}/${inference_asr_model}/${dset}"
- _logdir="${_dir}/logdir"
-
- if [ -d ${_dir} ]; then
- #echo "${_dir} is already exists. if you want to decode again, please delete this dir first."
- rm -r ${_dir}
- fi
- mkdir -p "${_logdir}"
-
- _scp=$scp
- _type=kaldi_ark
-
-
- # 1. Split the key file
- key_file=${_data}/${_scp}
- split_scps=""
- if "${use_k2}"; then
- # Now only _nj=1 is verified if using k2
- _nj=1
- else
- _nj=$(min "${inference_nj}" "$(<${key_file} wc -l)")
- fi
-
- for n in $(seq "${_nj}"); do
- split_scps+=" ${_logdir}/keys.${n}.scp"
- done
- # shellcheck disable=SC2086
- utils/split_scp.pl "${key_file}" ${split_scps}
-
- # 2. Submit decoding jobs
- log "Decoding started... log: '${_logdir}/asr_inference.*.log'"
- # shellcheck disable=SC2086
- ${_cmd} --gpu "${_ngpu}" --max-jobs-run "${_nj}" JOB=1:"${_nj}" "${_logdir}"/asr_inference.JOB.log \
- ${python} -m ${asr_inference_tool} \
- --batch_size ${batch_size} \
- --ngpu "${_ngpu}" \
- --njob ${njob} \
- --gpuid_list ${gpuid_list} \
- --data_path_and_name_and_type "${_data}/${_scp},speech,${_type}" \
- --key_file "${_logdir}"/keys.JOB.scp \
- --asr_train_config "${asr_exp}"/config.yaml \
- --asr_model_file "${asr_exp}"/"${inference_asr_model}" \
- --output_dir "${_logdir}"/output.JOB \
- --mode $mode \
- ${_opts} ${inference_args}
-
- # 3. Concatenates the output files from each jobs
- for f in token token_int score text; do
- if [ -f "${_logdir}/output.1/1best_recog/${f}" ]; then
- for i in $(seq "${_nj}"); do
- cat "${_logdir}/output.${i}/1best_recog/${f}"
- done | sort -k1 >"${_dir}/${f}"
- fi
- done
- python utils/proce_text.py ${_dir}/text ${_dir}/${text}.proc
- python utils/proce_text.py ${_data}/text ${_data}/${text}.proc
- python utils/compute_wer.py ${_data}/text.proc ${_dir}/text.proc ${_dir}/text.cer
- tail -n 3 ${_dir}/text.cer > ${_dir}/text.cer.txt
- cat ${_dir}/text.cer.txt
- done
-fi
-
-log "Successfully finished. [elapsed=${SECONDS}s]"
-
diff --git a/egs/aishell/tranformer/utils/gen_ark_list.sh b/egs/aishell/tranformer/utils/gen_ark_list.sh
index be60f7b..aebf356 100755
--- a/egs/aishell/tranformer/utils/gen_ark_list.sh
+++ b/egs/aishell/tranformer/utils/gen_ark_list.sh
@@ -2,19 +2,21 @@
# Begin configuration section.
-nj=4
+nj=32
cmd=./utils/run.pl
echo "$0 $@"
. utils/parse_options.sh || exit 1;
-data=$1
+ark_dir=$1
+txt_dir=$2
+output_dir=$3
-[ ! -d ${data}/ark ] && echo "$0: ark data is required" && exit 1;
-[ ! -d ${data}/txt ] && echo "$0: txt data is required" && exit 1;
+[ ! -d ${ark_dir}/ark ] && echo "$0: ark data is required" && exit 1;
+[ ! -d ${txt_dir}/txt ] && echo "$0: txt data is required" && exit 1;
for n in $(seq $nj); do
- echo "$data/ark/feats.$n.ark $data/txt/text.$n" || exit 1
-done > $data/ark_txt.scp || exit 1
+ echo "${ark_dir}/ark/feats.$n.ark ${txt_dir}/txt/text.$n.txt" || exit 1
+done > ${output_dir}/ark_txt.scp || exit 1
diff --git a/egs_modelscope/aishell/paraformer/conf/train_asr_paraformer_sanm_50e_16d_2048_512_lfr6.yaml b/egs_modelscope/aishell/paraformer/conf/train_asr_paraformer_sanm_50e_16d_2048_512_lfr6.yaml
index e9210f3..cb8b0af 100644
--- a/egs_modelscope/aishell/paraformer/conf/train_asr_paraformer_sanm_50e_16d_2048_512_lfr6.yaml
+++ b/egs_modelscope/aishell/paraformer/conf/train_asr_paraformer_sanm_50e_16d_2048_512_lfr6.yaml
@@ -66,7 +66,7 @@
lr: 0.0005
scheduler: warmuplr
scheduler_conf:
- warmup_steps: 30000
+ warmup_steps: 50000
specaug: specaug_lfr
specaug_conf:
diff --git a/egs_modelscope/aishell/paraformer/modelscope_utils b/egs_modelscope/aishell/paraformer/modelscope_utils
deleted file mode 120000
index fc97768..0000000
--- a/egs_modelscope/aishell/paraformer/modelscope_utils
+++ /dev/null
@@ -1 +0,0 @@
-../../common/modelscope_utils
\ No newline at end of file
diff --git a/egs_modelscope/aishell/paraformer/modelscope_utils/download_model.py b/egs_modelscope/aishell/paraformer/modelscope_utils/download_model.py
new file mode 100755
index 0000000..51ba6b8
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/modelscope_utils/download_model.py
@@ -0,0 +1,25 @@
+#!/usr/bin/env python3
+import argparse
+
+from modelscope.pipelines import pipeline
+from modelscope.utils.constant import Tasks
+
+if __name__ == '__main__':
+ parser = argparse.ArgumentParser(
+ description="download model configs",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument("--model_name",
+ type=str,
+ default="speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-pytorch",
+ help="model name in modelscope")
+ parser.add_argument("--model_revision",
+ type=str,
+ default="v1.0.3",
+ help="model revision in modelscope")
+ args = parser.parse_args()
+
+ inference_pipeline = pipeline(
+ task=Tasks.auto_speech_recognition,
+ model='damo/{}'.format(args.model_name),
+ model_revision=args.model_revision)
diff --git a/egs_modelscope/aishell/paraformer/modelscope_utils/modelscope_infer.sh b/egs_modelscope/aishell/paraformer/modelscope_utils/modelscope_infer.sh
new file mode 100755
index 0000000..a0c606f
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/modelscope_utils/modelscope_infer.sh
@@ -0,0 +1,88 @@
+#!/usr/bin/env bash
+
+set -e
+set -u
+set -o pipefail
+
+data_dir=
+exp_dir=
+model_name=
+model_revision=
+inference_nj=32
+gpuid_list="0,1,2,3"
+njob=32
+gpu_inference=true
+
+test_sets="dev test"
+decode_cmd=utils/run.pl
+
+# LM configs
+use_lm=false
+beam_size=1
+lm_weight=0.0
+
+. utils/parse_options.sh
+
+if ${gpu_inference}; then
+ _ngpu=1
+else
+ _ngpu=0
+fi
+
+# download model from modelscope
+python modelscope_utils/download_model.py \
+ --model_name ${model_name} --model_revision ${model_revision}
+
+modelscope_dir=${HOME}/.cache/modelscope/hub/damo/${model_name}
+
+
+for dset in ${test_sets}; do
+ _dir=${exp_dir}/${model_name}/decode_asr/${dset}
+ _logdir=${_dir}/logdir
+ _data=${data_dir}/${dset}
+ if [ -d ${_dir} ]; then
+ echo "${_dir} is already exists. if you want to decode again, please delete ${_dir} first."
+ exit 1
+ else
+ mkdir -p "${_dir}"
+ mkdir -p "${_logdir}"
+ fi
+
+ if "${use_lm}"; then
+ cp ${modelscope_dir}/decoding.yaml ${modelscope_dir}/decoding.yaml.back
+ sed -i "s#beam_size: [0-9]*#beam_size: `echo $beam_size`#g" ${modelscope_dir}/decoding.yaml
+ sed -i "s#lm_weight: 0.[0-9]*#lm_weight: `echo $lm_weight`#g" ${modelscope_dir}/decoding.yaml
+ fi
+
+ for n in $(seq "${inference_nj}"); do
+ split_scps+=" ${_logdir}/keys.${n}.scp"
+ done
+ # shellcheck disable=SC2086
+ utils/split_scp.pl "${data_dir}/${dset}/wav.scp" ${split_scps}
+
+ echo "Decoding started... log: '${_logdir}/asr_inference.*.log'"
+ # shellcheck disable=SC2086
+ ${decode_cmd} --max-jobs-run "${inference_nj}" JOB=1:"${inference_nj}" "${_logdir}"/asr_inference.JOB.log \
+ python -m funasr.bin.modelscope_infer \
+ --model_name ${model_name} \
+ --model_revision ${model_revision} \
+ --wav_list ${_logdir}/keys.JOB.scp \
+ --output_file ${_logdir}/text.JOB \
+ --gpuid_list ${gpuid_list} \
+ --njob ${njob} \
+ --ngpu ${_ngpu} \
+
+ for i in $(seq ${inference_nj}); do
+ cat ${_logdir}/text.${i}
+ done | sort -k1 >${_dir}/text
+
+ python utils/proce_text.py ${_dir}/text ${_dir}/text.proc
+ python utils/proce_text.py ${_data}/text ${_data}/text.proc
+ python utils/compute_wer.py ${_data}/text.proc ${_dir}/text.proc ${_dir}/text.cer
+ tail -n 3 ${_dir}/text.cer > ${_dir}/text.cer.txt
+ cat ${_dir}/text.cer.txt
+done
+
+if "${use_lm}"; then
+ mv ${modelscope_dir}/decoding.yaml.back ${modelscope_dir}/decoding.yaml
+fi
diff --git a/egs_modelscope/aishell/paraformer/modelscope_utils/update_config.py b/egs_modelscope/aishell/paraformer/modelscope_utils/update_config.py
new file mode 100644
index 0000000..88466ed
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/modelscope_utils/update_config.py
@@ -0,0 +1,41 @@
+import yaml
+import argparse
+
+def update_dct(fin_configs, root):
+ if root == {}:
+ return {}
+ for root_key, root_value in root.items():
+ if not isinstance(root[root_key],dict):
+ fin_configs[root_key] = root[root_key]
+ else:
+ result = update_dct(fin_configs[root_key], root[root_key])
+ fin_configs[root_key] = result
+ return fin_configs
+
+if __name__ == '__main__':
+ parser = argparse.ArgumentParser(
+ description="update configs",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument("--modelscope_config",
+ type=str,
+ help="modelscope config file")
+ parser.add_argument("--finetune_config",
+ type=str,
+ help="finetune config file")
+ parser.add_argument("--output_config",
+ type=str,
+ help="output config file")
+ args = parser.parse_args()
+
+ with open(args.modelscope_config) as f:
+ modelscope_configs = yaml.safe_load(f)
+
+ with open(args.finetune_config) as f:
+ finetune_configs = yaml.safe_load(f)
+
+ # update configs, e.g., lr, batch_size, ...
+ modelscope_configs = update_dct(modelscope_configs, finetune_configs)
+
+ with open(args.output_config, "w") as f:
+ yaml.dump(modelscope_configs, f, indent=4)
diff --git a/egs_modelscope/aishell/paraformer/paraformer_large_finetune.sh b/egs_modelscope/aishell/paraformer/paraformer_large_finetune.sh
index 3c75774..453b762 100755
--- a/egs_modelscope/aishell/paraformer/paraformer_large_finetune.sh
+++ b/egs_modelscope/aishell/paraformer/paraformer_large_finetune.sh
@@ -9,6 +9,7 @@
gpu_inference=true # Whether to perform gpu decoding, set false for cpu decoding
njob=4 # the number of jobs for each gpu
train_cmd=utils/run.pl
+infer_cmd=utils/run.pl
# general configuration
feats_dir="../DATA" #feature output dictionary, for large data
@@ -32,7 +33,7 @@
lfr_n=6
init_model_name=speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-pytorch # pre-trained model, download from modelscope during fine-tuning
-model_revision="v1.0.3" # please do not modify the model revision
+model_revision="v1.0.4" # please do not modify the model revision
cmvn_file=init_model/${init_model_name}/am.mvn
seg_file=init_model/${init_model_name}/seg_dict
vocab=init_model/${init_model_name}/tokens.txt
@@ -81,7 +82,14 @@
# you can set gpu num for decoding here
gpuid_list=$CUDA_VISIBLE_DEVICES # set gpus for decoding, the same as training stage by default
ngpu=$(echo $gpuid_list | awk -F "," '{print NF}')
-inference_nj=$[${ngpu}*${njob}]
+
+if ${gpu_inference}; then
+ inference_nj=$[${ngpu}*${njob}]
+ _ngpu=1
+else
+ inference_nj=$njob
+ _ngpu=0
+fi
if [ ${stage} -le 0 ] && [ ${stop_stage} -ge 0 ]; then
echo "stage 0: Data preparation"
@@ -99,7 +107,7 @@
feat_dev_dir=${feats_dir}/${dumpdir}/dev; mkdir -p ${feat_dev_dir}
feat_test_dir=${feats_dir}/${dumpdir}/test; mkdir -p ${feat_test_dir}
if [ ${stage} -le 1 ] && [ ${stop_stage} -ge 1 ]; then
- echo "Feature Generation"
+ echo "stage 1: Feature Generation"
# compute fbank features
fbankdir=${feats_dir}/fbank
utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj --speed_perturb ${speed_perturb} \
@@ -152,6 +160,7 @@
# Training Stage
world_size=$gpu_num # run on one machine
if [ ${stage} -le 3 ] && [ ${stop_stage} -ge 3 ]; then
+ echo "stage 3: Training"
# update asr train config.yaml
python modelscope_utils/update_config.py --modelscope_config init_model/${init_model_name}/finetune.yaml --finetune_config ${asr_config} --output_config init_model/${init_model_name}/asr_finetune_config.yaml
finetune_config=init_model/${init_model_name}/asr_finetune_config.yaml
@@ -201,25 +210,58 @@
# Testing Stage
if [ ${stage} -le 4 ] && [ ${stop_stage} -ge 4 ]; then
- ./utils/easy_asr_infer.sh \
- --lang zh \
- --datadir ${feats_dir} \
- --feats_type ${feats_type} \
- --feats_dim ${feats_dim} \
- --token_type ${token_type} \
- --gpu_inference ${gpu_inference} \
- --inference_config "${inference_config}" \
- --test_sets "${test_sets}" \
- --token_list $token_list \
- --asr_exp ${exp_dir}/exp/${model_dir} \
- --stage 12 \
- --stop_stage 12 \
- --scp $scp \
- --text text \
- --inference_nj $inference_nj \
- --njob $njob \
- --inference_asr_model $inference_asr_model \
- --gpuid_list $gpuid_list \
- --mode paraformer
+ echo "stage 4: Inference"
+ for dset in ${test_sets}; do
+ asr_exp=${exp_dir}/exp/${model_dir}
+ inference_tag="$(basename "${inference_config}" .yaml)"
+ _dir="${asr_exp}/${inference_tag}/${inference_asr_model}/${dset}"
+ _logdir="${_dir}/logdir"
+ if [ -d ${_dir} ]; then
+ echo "${_dir} is already exists. if you want to decode again, please delete this dir first."
+ exit 0
+ fi
+ mkdir -p "${_logdir}"
+ _data="${feats_dir}/${dumpdir}/${dset}"
+ key_file=${_data}/${scp}
+ num_scp_file="$(<${key_file} wc -l)"
+ _nj=$([ $inference_nj -le $num_scp_file ] && echo "$inference_nj" || echo "$num_scp_file")
+ split_scps=
+ for n in $(seq "${_nj}"); do
+ split_scps+=" ${_logdir}/keys.${n}.scp"
+ done
+ # shellcheck disable=SC2086
+ utils/split_scp.pl "${key_file}" ${split_scps}
+ _opts=
+ if [ -n "${inference_config}" ]; then
+ _opts+="--config ${inference_config} "
+ fi
+ ${infer_cmd} --gpu "${_ngpu}" --max-jobs-run "${_nj}" JOB=1:"${_nj}" "${_logdir}"/asr_inference.JOB.log \
+ python -m funasr.bin.asr_inference_launch \
+ --batch_size 1 \
+ --ngpu "${_ngpu}" \
+ --njob ${njob} \
+ --gpuid_list ${gpuid_list} \
+ --data_path_and_name_and_type "${_data}/${scp},speech,${type}" \
+ --key_file "${_logdir}"/keys.JOB.scp \
+ --asr_train_config "${asr_exp}"/config.yaml \
+ --asr_model_file "${asr_exp}"/"${inference_asr_model}" \
+ --output_dir "${_logdir}"/output.JOB \
+ --mode paraformer \
+ ${_opts}
+
+ for f in token token_int score text; do
+ if [ -f "${_logdir}/output.1/1best_recog/${f}" ]; then
+ for i in $(seq "${_nj}"); do
+ cat "${_logdir}/output.${i}/1best_recog/${f}"
+ done | sort -k1 >"${_dir}/${f}"
+ fi
+ done
+ python utils/proce_text.py ${_dir}/text ${_dir}/text.proc
+ python utils/proce_text.py ${_data}/text ${_data}/text.proc
+ python utils/compute_wer.py ${_data}/text.proc ${_dir}/text.proc ${_dir}/text.cer
+ tail -n 3 ${_dir}/text.cer > ${_dir}/text.cer.txt
+ cat ${_dir}/text.cer.txt
+ done
fi
+
diff --git a/egs_modelscope/aishell/paraformer/paraformer_large_infer.sh b/egs_modelscope/aishell/paraformer/paraformer_large_infer.sh
index ab65cde..ac70211 100755
--- a/egs_modelscope/aishell/paraformer/paraformer_large_infer.sh
+++ b/egs_modelscope/aishell/paraformer/paraformer_large_infer.sh
@@ -8,7 +8,7 @@
data_dir=
exp_dir=
model_name=speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-pytorch
-model_revision="v1.0.3" # please do not modify the model revision
+model_revision="v1.0.4" # please do not modify the model revision
inference_nj=32
gpuid_list="0,1" # set gpus, e.g., gpuid_list="0,1"
ngpu=$(echo $gpuid_list | awk -F "," '{print NF}')
diff --git a/egs_modelscope/aishell/paraformer/utils b/egs_modelscope/aishell/paraformer/utils
deleted file mode 120000
index 37d9761..0000000
--- a/egs_modelscope/aishell/paraformer/utils
+++ /dev/null
@@ -1 +0,0 @@
-../../../egs/aishell/tranformer/utils/
\ No newline at end of file
diff --git a/egs_modelscope/aishell/paraformer/utils/__init__.py b/egs_modelscope/aishell/paraformer/utils/__init__.py
new file mode 100644
index 0000000..e69de29
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/__init__.py
diff --git a/egs_modelscope/aishell/paraformer/utils/apply_cmvn.py b/egs_modelscope/aishell/paraformer/utils/apply_cmvn.py
new file mode 100755
index 0000000..b5c5086
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/apply_cmvn.py
@@ -0,0 +1,79 @@
+from kaldiio import ReadHelper
+from kaldiio import WriteHelper
+
+import argparse
+import json
+import math
+import numpy as np
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="apply cmvn",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--ark-file",
+ "-a",
+ default=False,
+ required=True,
+ type=str,
+ help="fbank ark file",
+ )
+ parser.add_argument(
+ "--cmvn-file",
+ "-c",
+ default=False,
+ required=True,
+ type=str,
+ help="cmvn file",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ ark_file = args.output_dir + "/feats." + str(args.ark_index) + ".ark"
+ scp_file = args.output_dir + "/feats." + str(args.ark_index) + ".scp"
+ ark_writer = WriteHelper('ark,scp:{},{}'.format(ark_file, scp_file))
+
+ with open(args.cmvn_file) as f:
+ cmvn_stats = json.load(f)
+
+ means = cmvn_stats['mean_stats']
+ vars = cmvn_stats['var_stats']
+ total_frames = cmvn_stats['total_frames']
+
+ for i in range(len(means)):
+ means[i] /= total_frames
+ vars[i] = vars[i] / total_frames - means[i] * means[i]
+ if vars[i] < 1.0e-20:
+ vars[i] = 1.0e-20
+ vars[i] = 1.0 / math.sqrt(vars[i])
+
+ with ReadHelper('ark:{}'.format(args.ark_file)) as ark_reader:
+ for key, mat in ark_reader:
+ mat = (mat - means) * vars
+ ark_writer(key, mat)
+
+
+if __name__ == '__main__':
+ main()
diff --git a/egs_modelscope/aishell/paraformer/utils/apply_cmvn.sh b/egs_modelscope/aishell/paraformer/utils/apply_cmvn.sh
new file mode 100755
index 0000000..f8fd1d1
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/apply_cmvn.sh
@@ -0,0 +1,29 @@
+#!/usr/bin/env bash
+
+. ./path.sh || exit 1;
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+fbankdir=$1
+cmvn_file=$2
+logdir=$3
+output_dir=$4
+
+dump_dir=${output_dir}/ark; mkdir -p ${dump_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/apply_cmvn.JOB.log \
+ python utils/apply_cmvn.py -a $fbankdir/ark/feats.JOB.ark \
+ -c $cmvn_file -i JOB -o ${dump_dir} \
+ || exit 1;
+
+for n in $(seq $nj); do
+ cat ${dump_dir}/feats.$n.scp || exit 1
+done > ${output_dir}/feats.scp || exit 1
+
+echo "$0: Succeeded apply cmvn"
diff --git a/egs_modelscope/aishell/paraformer/utils/apply_lfr_and_cmvn.py b/egs_modelscope/aishell/paraformer/utils/apply_lfr_and_cmvn.py
new file mode 100755
index 0000000..50d18d1
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/apply_lfr_and_cmvn.py
@@ -0,0 +1,143 @@
+from kaldiio import ReadHelper, WriteHelper
+
+import argparse
+import numpy as np
+
+
+def build_LFR_features(inputs, m=7, n=6):
+ LFR_inputs = []
+ T = inputs.shape[0]
+ T_lfr = int(np.ceil(T / n))
+ left_padding = np.tile(inputs[0], ((m - 1) // 2, 1))
+ inputs = np.vstack((left_padding, inputs))
+ T = T + (m - 1) // 2
+ for i in range(T_lfr):
+ if m <= T - i * n:
+ LFR_inputs.append(np.hstack(inputs[i * n:i * n + m]))
+ else:
+ num_padding = m - (T - i * n)
+ frame = np.hstack(inputs[i * n:])
+ for _ in range(num_padding):
+ frame = np.hstack((frame, inputs[-1]))
+ LFR_inputs.append(frame)
+ return np.vstack(LFR_inputs)
+
+
+def build_CMVN_features(inputs, mvn_file): # noqa
+ with open(mvn_file, 'r', encoding='utf-8') as f:
+ lines = f.readlines()
+
+ add_shift_list = []
+ rescale_list = []
+ for i in range(len(lines)):
+ line_item = lines[i].split()
+ if line_item[0] == '<AddShift>':
+ line_item = lines[i + 1].split()
+ if line_item[0] == '<LearnRateCoef>':
+ add_shift_line = line_item[3:(len(line_item) - 1)]
+ add_shift_list = list(add_shift_line)
+ continue
+ elif line_item[0] == '<Rescale>':
+ line_item = lines[i + 1].split()
+ if line_item[0] == '<LearnRateCoef>':
+ rescale_line = line_item[3:(len(line_item) - 1)]
+ rescale_list = list(rescale_line)
+ continue
+
+ for j in range(inputs.shape[0]):
+ for k in range(inputs.shape[1]):
+ add_shift_value = add_shift_list[k]
+ rescale_value = rescale_list[k]
+ inputs[j, k] = float(inputs[j, k]) + float(add_shift_value)
+ inputs[j, k] = float(inputs[j, k]) * float(rescale_value)
+
+ return inputs
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="apply low_frame_rate and cmvn",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--ark-file",
+ "-a",
+ default=False,
+ required=True,
+ type=str,
+ help="fbank ark file",
+ )
+ parser.add_argument(
+ "--lfr",
+ "-f",
+ default=True,
+ type=str,
+ help="low frame rate",
+ )
+ parser.add_argument(
+ "--lfr-m",
+ "-m",
+ default=7,
+ type=int,
+ help="number of frames to stack",
+ )
+ parser.add_argument(
+ "--lfr-n",
+ "-n",
+ default=6,
+ type=int,
+ help="number of frames to skip",
+ )
+ parser.add_argument(
+ "--cmvn-file",
+ "-c",
+ default=False,
+ required=True,
+ type=str,
+ help="global cmvn file",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ dump_ark_file = args.output_dir + "/feats." + str(args.ark_index) + ".ark"
+ dump_scp_file = args.output_dir + "/feats." + str(args.ark_index) + ".scp"
+ shape_file = args.output_dir + "/len." + str(args.ark_index)
+ ark_writer = WriteHelper('ark,scp:{},{}'.format(dump_ark_file, dump_scp_file))
+
+ shape_writer = open(shape_file, 'w')
+ with ReadHelper('ark:{}'.format(args.ark_file)) as ark_reader:
+ for key, mat in ark_reader:
+ if args.lfr:
+ lfr = build_LFR_features(mat, args.lfr_m, args.lfr_n)
+ else:
+ lfr = mat
+ cmvn = build_CMVN_features(lfr, args.cmvn_file)
+ dims = cmvn.shape[1]
+ lens = cmvn.shape[0]
+ shape_writer.write(key + " " + str(lens) + "," + str(dims) + '\n')
+ ark_writer(key, cmvn)
+
+
+if __name__ == '__main__':
+ main()
+
diff --git a/egs_modelscope/aishell/paraformer/utils/apply_lfr_and_cmvn.sh b/egs_modelscope/aishell/paraformer/utils/apply_lfr_and_cmvn.sh
new file mode 100755
index 0000000..3119fdb
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/apply_lfr_and_cmvn.sh
@@ -0,0 +1,38 @@
+#!/usr/bin/env bash
+
+
+# Begin configuration section.
+nj=32
+cmd=utils/run.pl
+
+# feature configuration
+lfr=True
+lfr_m=7
+lfr_n=6
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+fbankdir=$1
+cmvn_file=$2
+logdir=$3
+output_dir=$4
+
+dump_dir=${output_dir}/ark; mkdir -p ${dump_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/apply_lfr_and_cmvn.JOB.log \
+ python utils/apply_lfr_and_cmvn.py -a $fbankdir/ark/feats.JOB.ark \
+ -f $lfr -m $lfr_m -n $lfr_n -c $cmvn_file -i JOB -o ${dump_dir} \
+ || exit 1;
+
+for n in $(seq $nj); do
+ cat ${dump_dir}/feats.$n.scp || exit 1
+done > ${output_dir}/feats.scp || exit 1
+
+for n in $(seq $nj); do
+ cat ${dump_dir}/len.$n || exit 1
+done > ${output_dir}/speech_shape || exit 1
+
+echo "$0: Succeeded apply low frame rate and cmvn"
diff --git a/egs_modelscope/aishell/paraformer/utils/combine_cmvn_file.py b/egs_modelscope/aishell/paraformer/utils/combine_cmvn_file.py
new file mode 100755
index 0000000..b2974a4
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/combine_cmvn_file.py
@@ -0,0 +1,73 @@
+import argparse
+import json
+import numpy as np
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="combine cmvn file",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--dims",
+ "-d",
+ default=80,
+ type=int,
+ help="feature dims",
+ )
+ parser.add_argument(
+ "--cmvn-dir",
+ "-c",
+ default=False,
+ required=True,
+ type=str,
+ help="cmvn dir",
+ )
+
+ parser.add_argument(
+ "--nj",
+ "-n",
+ default=1,
+ required=True,
+ type=int,
+ help="num of cmvn file",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ total_means = np.zeros(args.dims)
+ total_vars = np.zeros(args.dims)
+ total_frames = 0
+
+ cmvn_file = args.output_dir + "/cmvn.json"
+
+ for i in range(1, args.nj+1):
+ with open(args.cmvn_dir + "/cmvn." + str(i) + ".json", "r") as fin:
+ cmvn_stats = json.load(fin)
+
+ total_means += np.array(cmvn_stats["mean_stats"])
+ total_vars += np.array(cmvn_stats["var_stats"])
+ total_frames += cmvn_stats["total_frames"]
+
+ cmvn_info = {
+ 'mean_stats': list(total_means.tolist()),
+ 'var_stats': list(total_vars.tolist()),
+ 'total_frames': total_frames
+ }
+ with open(cmvn_file, 'w') as fout:
+ fout.write(json.dumps(cmvn_info))
+
+
+if __name__ == '__main__':
+ main()
diff --git a/egs_modelscope/aishell/paraformer/utils/compute_cmvn.py b/egs_modelscope/aishell/paraformer/utils/compute_cmvn.py
new file mode 100755
index 0000000..2b96e26
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/compute_cmvn.py
@@ -0,0 +1,74 @@
+from kaldiio import ReadHelper
+
+import argparse
+import numpy as np
+import json
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="computer global cmvn",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--dims",
+ "-d",
+ default=80,
+ type=int,
+ help="feature dims",
+ )
+ parser.add_argument(
+ "--ark-file",
+ "-a",
+ default=False,
+ required=True,
+ type=str,
+ help="fbank ark file",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ ark_file = args.ark_file + "/feats." + str(args.ark_index) + ".ark"
+ cmvn_file = args.output_dir + "/cmvn." + str(args.ark_index) + ".json"
+
+ mean_stats = np.zeros(args.dims)
+ var_stats = np.zeros(args.dims)
+ total_frames = 0
+
+ with ReadHelper('ark:{}'.format(ark_file)) as ark_reader:
+ for key, mat in ark_reader:
+ mean_stats += np.sum(mat, axis=0)
+ var_stats += np.sum(np.square(mat), axis=0)
+ total_frames += mat.shape[0]
+
+ cmvn_info = {
+ 'mean_stats': list(mean_stats.tolist()),
+ 'var_stats': list(var_stats.tolist()),
+ 'total_frames': total_frames
+ }
+ with open(cmvn_file, 'w') as fout:
+ fout.write(json.dumps(cmvn_info))
+
+
+if __name__ == '__main__':
+ main()
diff --git a/egs_modelscope/aishell/paraformer/utils/compute_cmvn.sh b/egs_modelscope/aishell/paraformer/utils/compute_cmvn.sh
new file mode 100755
index 0000000..12173ee
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/compute_cmvn.sh
@@ -0,0 +1,25 @@
+#!/usr/bin/env bash
+
+. ./path.sh || exit 1;
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+feats_dim=80
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+fbankdir=$1
+logdir=$2
+
+output_dir=${fbankdir}/cmvn; mkdir -p ${output_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/cmvn.JOB.log \
+ python utils/compute_cmvn.py -d ${feats_dim} -a $fbankdir/ark -i JOB -o ${output_dir} \
+ || exit 1;
+
+python utils/combine_cmvn_file.py -d ${feats_dim} -c ${output_dir} -n $nj -o $fbankdir
+
+echo "$0: Succeeded compute global cmvn"
diff --git a/egs_modelscope/aishell/paraformer/utils/compute_fbank.py b/egs_modelscope/aishell/paraformer/utils/compute_fbank.py
new file mode 100755
index 0000000..d03b5a8
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/compute_fbank.py
@@ -0,0 +1,153 @@
+from kaldiio import WriteHelper
+
+import argparse
+import numpy as np
+import json
+import torch
+import torchaudio
+import torchaudio.compliance.kaldi as kaldi
+
+
+def compute_fbank(wav_file,
+ num_mel_bins=80,
+ frame_length=25,
+ frame_shift=10,
+ dither=0.0,
+ resample_rate=16000,
+ speed=1.0):
+
+ waveform, sample_rate = torchaudio.load(wav_file)
+ if resample_rate != sample_rate:
+ waveform = torchaudio.transforms.Resample(orig_freq=sample_rate,
+ new_freq=resample_rate)(waveform)
+ if speed != 1.0:
+ waveform, _ = torchaudio.sox_effects.apply_effects_tensor(
+ waveform, resample_rate,
+ [['speed', str(speed)], ['rate', str(resample_rate)]]
+ )
+
+ waveform = waveform * (1 << 15)
+ mat = kaldi.fbank(waveform,
+ num_mel_bins=num_mel_bins,
+ frame_length=frame_length,
+ frame_shift=frame_shift,
+ dither=dither,
+ energy_floor=0.0,
+ window_type='hamming',
+ sample_frequency=resample_rate)
+
+ return mat.numpy()
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="computer features",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--wav-lists",
+ "-w",
+ default=False,
+ required=True,
+ type=str,
+ help="input wav lists",
+ )
+ parser.add_argument(
+ "--text-files",
+ "-t",
+ default=False,
+ required=True,
+ type=str,
+ help="input text files",
+ )
+ parser.add_argument(
+ "--dims",
+ "-d",
+ default=80,
+ type=int,
+ help="feature dims",
+ )
+ parser.add_argument(
+ "--sample-frequency",
+ "-s",
+ default=16000,
+ type=int,
+ help="sample frequency",
+ )
+ parser.add_argument(
+ "--speed-perturb",
+ "-p",
+ default="1.0",
+ type=str,
+ help="speed perturb",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-a",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ ark_file = args.output_dir + "/ark/feats." + str(args.ark_index) + ".ark"
+ scp_file = args.output_dir + "/ark/feats." + str(args.ark_index) + ".scp"
+ text_file = args.output_dir + "/txt/text." + str(args.ark_index) + ".txt"
+ feats_shape_file = args.output_dir + "/ark/len." + str(args.ark_index)
+ text_shape_file = args.output_dir + "/txt/len." + str(args.ark_index)
+
+ ark_writer = WriteHelper('ark,scp:{},{}'.format(ark_file, scp_file))
+ text_writer = open(text_file, 'w')
+ feats_shape_writer = open(feats_shape_file, 'w')
+ text_shape_writer = open(text_shape_file, 'w')
+
+ speed_perturb_list = args.speed_perturb.split(',')
+
+ for speed in speed_perturb_list:
+ with open(args.wav_lists, 'r', encoding='utf-8') as wavfile:
+ with open(args.text_files, 'r', encoding='utf-8') as textfile:
+ for wav, text in zip(wavfile, textfile):
+ s_w = wav.strip().split()
+ wav_id = s_w[0]
+ wav_file = s_w[1]
+
+ s_t = text.strip().split()
+ text_id = s_t[0]
+ txt = s_t[1:]
+ fbank = compute_fbank(wav_file,
+ num_mel_bins=args.dims,
+ resample_rate=args.sample_frequency,
+ speed=float(speed)
+ )
+ feats_dims = fbank.shape[1]
+ feats_lens = fbank.shape[0]
+ txt_lens = len(txt)
+ if speed == "1.0":
+ wav_id_sp = wav_id
+ else:
+ wav_id_sp = wav_id + "_sp" + speed
+
+ feats_shape_writer.write(wav_id_sp + " " + str(feats_lens) + "," + str(feats_dims) + '\n')
+ text_shape_writer.write(wav_id_sp + " " + str(txt_lens) + '\n')
+
+ text_writer.write(wav_id_sp + " " + " ".join(txt) + '\n')
+ ark_writer(wav_id_sp, fbank)
+
+
+if __name__ == '__main__':
+ main()
+
diff --git a/egs_modelscope/aishell/paraformer/utils/compute_fbank.sh b/egs_modelscope/aishell/paraformer/utils/compute_fbank.sh
new file mode 100755
index 0000000..92a4fe6
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/compute_fbank.sh
@@ -0,0 +1,51 @@
+#!/usr/bin/env bash
+
+. ./path.sh || exit 1;
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+
+# feature configuration
+feats_dim=80
+sample_frequency=16000
+speed_perturb="1.0"
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+data=$1
+logdir=$2
+fbankdir=$3
+
+[ ! -f $data/wav.scp ] && echo "$0: no such file $data/wav.scp" && exit 1;
+[ ! -f $data/text ] && echo "$0: no such file $data/text" && exit 1;
+
+python utils/split_data.py $data $data $nj
+
+ark_dir=${fbankdir}/ark; mkdir -p ${ark_dir}
+text_dir=${fbankdir}/txt; mkdir -p ${text_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/make_fbank.JOB.log \
+ python utils/compute_fbank.py -w $data/split${nj}/JOB/wav.scp -t $data/split${nj}/JOB/text \
+ -d $feats_dim -s $sample_frequency -p ${speed_perturb} -a JOB -o ${fbankdir} \
+ || exit 1;
+
+for n in $(seq $nj); do
+ cat ${ark_dir}/feats.$n.scp || exit 1
+done > $fbankdir/feats.scp || exit 1
+
+for n in $(seq $nj); do
+ cat ${text_dir}/text.$n.txt || exit 1
+done > $fbankdir/text || exit 1
+
+for n in $(seq $nj); do
+ cat ${ark_dir}/len.$n || exit 1
+done > $fbankdir/speech_shape || exit 1
+
+for n in $(seq $nj); do
+ cat ${text_dir}/len.$n || exit 1
+done > $fbankdir/text_shape || exit 1
+
+echo "$0: Succeeded compute FBANK features"
diff --git a/egs_modelscope/aishell/paraformer/utils/compute_wer.py b/egs_modelscope/aishell/paraformer/utils/compute_wer.py
new file mode 100755
index 0000000..349a3f6
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/compute_wer.py
@@ -0,0 +1,157 @@
+import os
+import numpy as np
+import sys
+
+def compute_wer(ref_file,
+ hyp_file,
+ cer_detail_file):
+ rst = {
+ 'Wrd': 0,
+ 'Corr': 0,
+ 'Ins': 0,
+ 'Del': 0,
+ 'Sub': 0,
+ 'Snt': 0,
+ 'Err': 0.0,
+ 'S.Err': 0.0,
+ 'wrong_words': 0,
+ 'wrong_sentences': 0
+ }
+
+ hyp_dict = {}
+ ref_dict = {}
+ with open(hyp_file, 'r') as hyp_reader:
+ for line in hyp_reader:
+ key = line.strip().split()[0]
+ value = line.strip().split()[1:]
+ hyp_dict[key] = value
+ with open(ref_file, 'r') as ref_reader:
+ for line in ref_reader:
+ key = line.strip().split()[0]
+ value = line.strip().split()[1:]
+ ref_dict[key] = value
+
+ cer_detail_writer = open(cer_detail_file, 'w')
+ for hyp_key in hyp_dict:
+ if hyp_key in ref_dict:
+ out_item = compute_wer_by_line(hyp_dict[hyp_key], ref_dict[hyp_key])
+ rst['Wrd'] += out_item['nwords']
+ rst['Corr'] += out_item['cor']
+ rst['wrong_words'] += out_item['wrong']
+ rst['Ins'] += out_item['ins']
+ rst['Del'] += out_item['del']
+ rst['Sub'] += out_item['sub']
+ rst['Snt'] += 1
+ if out_item['wrong'] > 0:
+ rst['wrong_sentences'] += 1
+ cer_detail_writer.write(hyp_key + print_cer_detail(out_item) + '\n')
+ cer_detail_writer.write("ref:" + '\t' + "".join(ref_dict[hyp_key]) + '\n')
+ cer_detail_writer.write("hyp:" + '\t' + "".join(hyp_dict[hyp_key]) + '\n')
+
+ if rst['Wrd'] > 0:
+ rst['Err'] = round(rst['wrong_words'] * 100 / rst['Wrd'], 2)
+ if rst['Snt'] > 0:
+ rst['S.Err'] = round(rst['wrong_sentences'] * 100 / rst['Snt'], 2)
+
+ cer_detail_writer.write('\n')
+ cer_detail_writer.write("%WER " + str(rst['Err']) + " [ " + str(rst['wrong_words'])+ " / " + str(rst['Wrd']) +
+ ", " + str(rst['Ins']) + " ins, " + str(rst['Del']) + " del, " + str(rst['Sub']) + " sub ]" + '\n')
+ cer_detail_writer.write("%SER " + str(rst['S.Err']) + " [ " + str(rst['wrong_sentences']) + " / " + str(rst['Snt']) + " ]" + '\n')
+ cer_detail_writer.write("Scored " + str(len(hyp_dict)) + " sentences, " + str(len(hyp_dict) - rst['Snt']) + " not present in hyp." + '\n')
+
+
+def compute_wer_by_line(hyp,
+ ref):
+ hyp = list(map(lambda x: x.lower(), hyp))
+ ref = list(map(lambda x: x.lower(), ref))
+
+ len_hyp = len(hyp)
+ len_ref = len(ref)
+
+ cost_matrix = np.zeros((len_hyp + 1, len_ref + 1), dtype=np.int16)
+
+ ops_matrix = np.zeros((len_hyp + 1, len_ref + 1), dtype=np.int8)
+
+ for i in range(len_hyp + 1):
+ cost_matrix[i][0] = i
+ for j in range(len_ref + 1):
+ cost_matrix[0][j] = j
+
+ for i in range(1, len_hyp + 1):
+ for j in range(1, len_ref + 1):
+ if hyp[i - 1] == ref[j - 1]:
+ cost_matrix[i][j] = cost_matrix[i - 1][j - 1]
+ else:
+ substitution = cost_matrix[i - 1][j - 1] + 1
+ insertion = cost_matrix[i - 1][j] + 1
+ deletion = cost_matrix[i][j - 1] + 1
+
+ compare_val = [substitution, insertion, deletion]
+
+ min_val = min(compare_val)
+ operation_idx = compare_val.index(min_val) + 1
+ cost_matrix[i][j] = min_val
+ ops_matrix[i][j] = operation_idx
+
+ match_idx = []
+ i = len_hyp
+ j = len_ref
+ rst = {
+ 'nwords': len_ref,
+ 'cor': 0,
+ 'wrong': 0,
+ 'ins': 0,
+ 'del': 0,
+ 'sub': 0
+ }
+ while i >= 0 or j >= 0:
+ i_idx = max(0, i)
+ j_idx = max(0, j)
+
+ if ops_matrix[i_idx][j_idx] == 0: # correct
+ if i - 1 >= 0 and j - 1 >= 0:
+ match_idx.append((j - 1, i - 1))
+ rst['cor'] += 1
+
+ i -= 1
+ j -= 1
+
+ elif ops_matrix[i_idx][j_idx] == 2: # insert
+ i -= 1
+ rst['ins'] += 1
+
+ elif ops_matrix[i_idx][j_idx] == 3: # delete
+ j -= 1
+ rst['del'] += 1
+
+ elif ops_matrix[i_idx][j_idx] == 1: # substitute
+ i -= 1
+ j -= 1
+ rst['sub'] += 1
+
+ if i < 0 and j >= 0:
+ rst['del'] += 1
+ elif j < 0 and i >= 0:
+ rst['ins'] += 1
+
+ match_idx.reverse()
+ wrong_cnt = cost_matrix[len_hyp][len_ref]
+ rst['wrong'] = wrong_cnt
+
+ return rst
+
+def print_cer_detail(rst):
+ return ("(" + "nwords=" + str(rst['nwords']) + ",cor=" + str(rst['cor'])
+ + ",ins=" + str(rst['ins']) + ",del=" + str(rst['del']) + ",sub="
+ + str(rst['sub']) + ") corr:" + '{:.2%}'.format(rst['cor']/rst['nwords'])
+ + ",cer:" + '{:.2%}'.format(rst['wrong']/rst['nwords']))
+
+if __name__ == '__main__':
+ if len(sys.argv) != 4:
+ print("usage : python compute-wer.py test.ref test.hyp test.wer")
+ sys.exit(0)
+
+ ref_file = sys.argv[1]
+ hyp_file = sys.argv[2]
+ cer_detail_file = sys.argv[3]
+ compute_wer(ref_file, hyp_file, cer_detail_file)
diff --git a/egs_modelscope/aishell/paraformer/utils/error_rate_zh b/egs_modelscope/aishell/paraformer/utils/error_rate_zh
new file mode 100755
index 0000000..6871a07
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/error_rate_zh
@@ -0,0 +1,370 @@
+#!/usr/bin/env python3
+# coding=utf8
+
+# Copyright 2021 Jiayu DU
+
+import sys
+import argparse
+import json
+import logging
+logging.basicConfig(stream=sys.stderr, level=logging.INFO, format='[%(levelname)s] %(message)s')
+
+DEBUG = None
+
+def GetEditType(ref_token, hyp_token):
+ if ref_token == None and hyp_token != None:
+ return 'I'
+ elif ref_token != None and hyp_token == None:
+ return 'D'
+ elif ref_token == hyp_token:
+ return 'C'
+ elif ref_token != hyp_token:
+ return 'S'
+ else:
+ raise RuntimeError
+
+class AlignmentArc:
+ def __init__(self, src, dst, ref, hyp):
+ self.src = src
+ self.dst = dst
+ self.ref = ref
+ self.hyp = hyp
+ self.edit_type = GetEditType(ref, hyp)
+
+def similarity_score_function(ref_token, hyp_token):
+ return 0 if (ref_token == hyp_token) else -1.0
+
+def insertion_score_function(token):
+ return -1.0
+
+def deletion_score_function(token):
+ return -1.0
+
+def EditDistance(
+ ref,
+ hyp,
+ similarity_score_function = similarity_score_function,
+ insertion_score_function = insertion_score_function,
+ deletion_score_function = deletion_score_function):
+ assert(len(ref) != 0)
+ class DPState:
+ def __init__(self):
+ self.score = -float('inf')
+ # backpointer
+ self.prev_r = None
+ self.prev_h = None
+
+ def print_search_grid(S, R, H, fstream):
+ print(file=fstream)
+ for r in range(R):
+ for h in range(H):
+ print(F'[{r},{h}]:{S[r][h].score:4.3f}:({S[r][h].prev_r},{S[r][h].prev_h}) ', end='', file=fstream)
+ print(file=fstream)
+
+ R = len(ref) + 1
+ H = len(hyp) + 1
+
+ # Construct DP search space, a (R x H) grid
+ S = [ [] for r in range(R) ]
+ for r in range(R):
+ S[r] = [ DPState() for x in range(H) ]
+
+ # initialize DP search grid origin, S(r = 0, h = 0)
+ S[0][0].score = 0.0
+ S[0][0].prev_r = None
+ S[0][0].prev_h = None
+
+ # initialize REF axis
+ for r in range(1, R):
+ S[r][0].score = S[r-1][0].score + deletion_score_function(ref[r-1])
+ S[r][0].prev_r = r-1
+ S[r][0].prev_h = 0
+
+ # initialize HYP axis
+ for h in range(1, H):
+ S[0][h].score = S[0][h-1].score + insertion_score_function(hyp[h-1])
+ S[0][h].prev_r = 0
+ S[0][h].prev_h = h-1
+
+ best_score = S[0][0].score
+ best_state = (0, 0)
+
+ for r in range(1, R):
+ for h in range(1, H):
+ sub_or_cor_score = similarity_score_function(ref[r-1], hyp[h-1])
+ new_score = S[r-1][h-1].score + sub_or_cor_score
+ if new_score >= S[r][h].score:
+ S[r][h].score = new_score
+ S[r][h].prev_r = r-1
+ S[r][h].prev_h = h-1
+
+ del_score = deletion_score_function(ref[r-1])
+ new_score = S[r-1][h].score + del_score
+ if new_score >= S[r][h].score:
+ S[r][h].score = new_score
+ S[r][h].prev_r = r - 1
+ S[r][h].prev_h = h
+
+ ins_score = insertion_score_function(hyp[h-1])
+ new_score = S[r][h-1].score + ins_score
+ if new_score >= S[r][h].score:
+ S[r][h].score = new_score
+ S[r][h].prev_r = r
+ S[r][h].prev_h = h-1
+
+ best_score = S[R-1][H-1].score
+ best_state = (R-1, H-1)
+
+ if DEBUG:
+ print_search_grid(S, R, H, sys.stderr)
+
+ # Backtracing best alignment path, i.e. a list of arcs
+ # arc = (src, dst, ref, hyp, edit_type)
+ # src/dst = (r, h), where r/h refers to search grid state-id along Ref/Hyp axis
+ best_path = []
+ r, h = best_state[0], best_state[1]
+ prev_r, prev_h = S[r][h].prev_r, S[r][h].prev_h
+ score = S[r][h].score
+ # loop invariant:
+ # 1. (prev_r, prev_h) -> (r, h) is a "forward arc" on best alignment path
+ # 2. score is the value of point(r, h) on DP search grid
+ while prev_r != None or prev_h != None:
+ src = (prev_r, prev_h)
+ dst = (r, h)
+ if (r == prev_r + 1 and h == prev_h + 1): # Substitution or correct
+ arc = AlignmentArc(src, dst, ref[prev_r], hyp[prev_h])
+ elif (r == prev_r + 1 and h == prev_h): # Deletion
+ arc = AlignmentArc(src, dst, ref[prev_r], None)
+ elif (r == prev_r and h == prev_h + 1): # Insertion
+ arc = AlignmentArc(src, dst, None, hyp[prev_h])
+ else:
+ raise RuntimeError
+ best_path.append(arc)
+ r, h = prev_r, prev_h
+ prev_r, prev_h = S[r][h].prev_r, S[r][h].prev_h
+ score = S[r][h].score
+
+ best_path.reverse()
+ return (best_path, best_score)
+
+def PrettyPrintAlignment(alignment, stream = sys.stderr):
+ def get_token_str(token):
+ if token == None:
+ return "*"
+ return token
+
+ def is_double_width_char(ch):
+ if (ch >= '\u4e00') and (ch <= '\u9fa5'): # codepoint ranges for Chinese chars
+ return True
+ # TODO: support other double-width-char language such as Japanese, Korean
+ else:
+ return False
+
+ def display_width(token_str):
+ m = 0
+ for c in token_str:
+ if is_double_width_char(c):
+ m += 2
+ else:
+ m += 1
+ return m
+
+ R = ' REF : '
+ H = ' HYP : '
+ E = ' EDIT : '
+ for arc in alignment:
+ r = get_token_str(arc.ref)
+ h = get_token_str(arc.hyp)
+ e = arc.edit_type if arc.edit_type != 'C' else ''
+
+ nr, nh, ne = display_width(r), display_width(h), display_width(e)
+ n = max(nr, nh, ne) + 1
+
+ R += r + ' ' * (n-nr)
+ H += h + ' ' * (n-nh)
+ E += e + ' ' * (n-ne)
+
+ print(R, file=stream)
+ print(H, file=stream)
+ print(E, file=stream)
+
+def CountEdits(alignment):
+ c, s, i, d = 0, 0, 0, 0
+ for arc in alignment:
+ if arc.edit_type == 'C':
+ c += 1
+ elif arc.edit_type == 'S':
+ s += 1
+ elif arc.edit_type == 'I':
+ i += 1
+ elif arc.edit_type == 'D':
+ d += 1
+ else:
+ raise RuntimeError
+ return (c, s, i, d)
+
+def ComputeTokenErrorRate(c, s, i, d):
+ return 100.0 * (s + d + i) / (s + d + c)
+
+def ComputeSentenceErrorRate(num_err_utts, num_utts):
+ assert(num_utts != 0)
+ return 100.0 * num_err_utts / num_utts
+
+
+class EvaluationResult:
+ def __init__(self):
+ self.num_ref_utts = 0
+ self.num_hyp_utts = 0
+ self.num_eval_utts = 0 # seen in both ref & hyp
+ self.num_hyp_without_ref = 0
+
+ self.C = 0
+ self.S = 0
+ self.I = 0
+ self.D = 0
+ self.token_error_rate = 0.0
+
+ self.num_utts_with_error = 0
+ self.sentence_error_rate = 0.0
+
+ def to_json(self):
+ return json.dumps(self.__dict__)
+
+ def to_kaldi(self):
+ info = (
+ F'%WER {self.token_error_rate:.2f} [ {self.S + self.D + self.I} / {self.C + self.S + self.D}, {self.I} ins, {self.D} del, {self.S} sub ]\n'
+ F'%SER {self.sentence_error_rate:.2f} [ {self.num_utts_with_error} / {self.num_eval_utts} ]\n'
+ )
+ return info
+
+ def to_sclite(self):
+ return "TODO"
+
+ def to_espnet(self):
+ return "TODO"
+
+ def to_summary(self):
+ #return json.dumps(self.__dict__, indent=4)
+ summary = (
+ '==================== Overall Statistics ====================\n'
+ F'num_ref_utts: {self.num_ref_utts}\n'
+ F'num_hyp_utts: {self.num_hyp_utts}\n'
+ F'num_hyp_without_ref: {self.num_hyp_without_ref}\n'
+ F'num_eval_utts: {self.num_eval_utts}\n'
+ F'sentence_error_rate: {self.sentence_error_rate:.2f}%\n'
+ F'token_error_rate: {self.token_error_rate:.2f}%\n'
+ F'token_stats:\n'
+ F' - tokens:{self.C + self.S + self.D:>7}\n'
+ F' - edits: {self.S + self.I + self.D:>7}\n'
+ F' - cor: {self.C:>7}\n'
+ F' - sub: {self.S:>7}\n'
+ F' - ins: {self.I:>7}\n'
+ F' - del: {self.D:>7}\n'
+ '============================================================\n'
+ )
+ return summary
+
+
+class Utterance:
+ def __init__(self, uid, text):
+ self.uid = uid
+ self.text = text
+
+
+def LoadUtterances(filepath, format):
+ utts = {}
+ if format == 'text': # utt_id word1 word2 ...
+ with open(filepath, 'r', encoding='utf8') as f:
+ for line in f:
+ line = line.strip()
+ if line:
+ cols = line.split(maxsplit=1)
+ assert(len(cols) == 2 or len(cols) == 1)
+ uid = cols[0]
+ text = cols[1] if len(cols) == 2 else ''
+ if utts.get(uid) != None:
+ raise RuntimeError(F'Found duplicated utterence id {uid}')
+ utts[uid] = Utterance(uid, text)
+ else:
+ raise RuntimeError(F'Unsupported text format {format}')
+ return utts
+
+
+def tokenize_text(text, tokenizer):
+ if tokenizer == 'whitespace':
+ return text.split()
+ elif tokenizer == 'char':
+ return [ ch for ch in ''.join(text.split()) ]
+ else:
+ raise RuntimeError(F'ERROR: Unsupported tokenizer {tokenizer}')
+
+
+if __name__ == '__main__':
+ parser = argparse.ArgumentParser()
+ # optional
+ parser.add_argument('--tokenizer', choices=['whitespace', 'char'], default='whitespace', help='whitespace for WER, char for CER')
+ parser.add_argument('--ref-format', choices=['text'], default='text', help='reference format, first col is utt_id, the rest is text')
+ parser.add_argument('--hyp-format', choices=['text'], default='text', help='hypothesis format, first col is utt_id, the rest is text')
+ # required
+ parser.add_argument('--ref', type=str, required=True, help='input reference file')
+ parser.add_argument('--hyp', type=str, required=True, help='input hypothesis file')
+
+ parser.add_argument('result_file', type=str)
+ args = parser.parse_args()
+ logging.info(args)
+
+ ref_utts = LoadUtterances(args.ref, args.ref_format)
+ hyp_utts = LoadUtterances(args.hyp, args.hyp_format)
+
+ r = EvaluationResult()
+
+ # check valid utterances in hyp that have matched non-empty reference
+ eval_utts = []
+ r.num_hyp_without_ref = 0
+ for uid in sorted(hyp_utts.keys()):
+ if uid in ref_utts.keys(): # TODO: efficiency
+ if ref_utts[uid].text.strip(): # non-empty reference
+ eval_utts.append(uid)
+ else:
+ logging.warn(F'Found {uid} with empty reference, skipping...')
+ else:
+ logging.warn(F'Found {uid} without reference, skipping...')
+ r.num_hyp_without_ref += 1
+
+ r.num_hyp_utts = len(hyp_utts)
+ r.num_ref_utts = len(ref_utts)
+ r.num_eval_utts = len(eval_utts)
+
+ with open(args.result_file, 'w+', encoding='utf8') as fo:
+ for uid in eval_utts:
+ ref = ref_utts[uid]
+ hyp = hyp_utts[uid]
+
+ alignment, score = EditDistance(
+ tokenize_text(ref.text, args.tokenizer),
+ tokenize_text(hyp.text, args.tokenizer)
+ )
+
+ c, s, i, d = CountEdits(alignment)
+ utt_ter = ComputeTokenErrorRate(c, s, i, d)
+
+ # utt-level evaluation result
+ print(F'{{"uid":{uid}, "score":{score}, "ter":{utt_ter:.2f}, "cor":{c}, "sub":{s}, "ins":{i}, "del":{d}}}', file=fo)
+ PrettyPrintAlignment(alignment, fo)
+
+ r.C += c
+ r.S += s
+ r.I += i
+ r.D += d
+
+ if utt_ter > 0:
+ r.num_utts_with_error += 1
+
+ # corpus level evaluation result
+ r.sentence_error_rate = ComputeSentenceErrorRate(r.num_utts_with_error, r.num_eval_utts)
+ r.token_error_rate = ComputeTokenErrorRate(r.C, r.S, r.I, r.D)
+
+ print(r.to_summary(), file=fo)
+
+ print(r.to_json())
+ print(r.to_kaldi())
diff --git a/egs_modelscope/aishell/paraformer/utils/extract_embeds.py b/egs_modelscope/aishell/paraformer/utils/extract_embeds.py
new file mode 100755
index 0000000..7b817d8
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/extract_embeds.py
@@ -0,0 +1,47 @@
+from transformers import AutoTokenizer, AutoModel, pipeline
+import numpy as np
+import sys
+import os
+import torch
+from kaldiio import WriteHelper
+import re
+text_file_json = sys.argv[1]
+out_ark = sys.argv[2]
+out_scp = sys.argv[3]
+out_shape = sys.argv[4]
+device = int(sys.argv[5])
+model_path = sys.argv[6]
+
+model = AutoModel.from_pretrained(model_path)
+tokenizer = AutoTokenizer.from_pretrained(model_path)
+extractor = pipeline(task="feature-extraction", model=model, tokenizer=tokenizer, device=device)
+
+with open(text_file_json, 'r') as f:
+ js = f.readlines()
+
+
+f_shape = open(out_shape, "w")
+with WriteHelper('ark,scp:{},{}'.format(out_ark, out_scp)) as writer:
+ with torch.no_grad():
+ for idx, line in enumerate(js):
+ id, tokens = line.strip().split(" ", 1)
+ tokens = re.sub(" ", "", tokens.strip())
+ tokens = ' '.join([j for j in tokens])
+ token_num = len(tokens.split(" "))
+ outputs = extractor(tokens)
+ outputs = np.array(outputs)
+ embeds = outputs[0, 1:-1, :]
+
+ token_num_embeds, dim = embeds.shape
+ if token_num == token_num_embeds:
+ writer(id, embeds)
+ shape_line = "{} {},{}\n".format(id, token_num_embeds, dim)
+ f_shape.write(shape_line)
+ else:
+ print("{}, size has changed, {}, {}, {}".format(id, token_num, token_num_embeds, tokens))
+
+
+
+f_shape.close()
+
+
diff --git a/egs_modelscope/aishell/paraformer/utils/filter_scp.pl b/egs_modelscope/aishell/paraformer/utils/filter_scp.pl
new file mode 100755
index 0000000..003530d
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/filter_scp.pl
@@ -0,0 +1,87 @@
+#!/usr/bin/env perl
+# Copyright 2010-2012 Microsoft Corporation
+# Johns Hopkins University (author: Daniel Povey)
+
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+# This script takes a list of utterance-ids or any file whose first field
+# of each line is an utterance-id, and filters an scp
+# file (or any file whose "n-th" field is an utterance id), printing
+# out only those lines whose "n-th" field is in id_list. The index of
+# the "n-th" field is 1, by default, but can be changed by using
+# the -f <n> switch
+
+$exclude = 0;
+$field = 1;
+$shifted = 0;
+
+do {
+ $shifted=0;
+ if ($ARGV[0] eq "--exclude") {
+ $exclude = 1;
+ shift @ARGV;
+ $shifted=1;
+ }
+ if ($ARGV[0] eq "-f") {
+ $field = $ARGV[1];
+ shift @ARGV; shift @ARGV;
+ $shifted=1
+ }
+} while ($shifted);
+
+if(@ARGV < 1 || @ARGV > 2) {
+ die "Usage: filter_scp.pl [--exclude] [-f <field-to-filter-on>] id_list [in.scp] > out.scp \n" .
+ "Prints only the input lines whose f'th field (default: first) is in 'id_list'.\n" .
+ "Note: only the first field of each line in id_list matters. With --exclude, prints\n" .
+ "only the lines that were *not* in id_list.\n" .
+ "Caution: previously, the -f option was interpreted as a zero-based field index.\n" .
+ "If your older scripts (written before Oct 2014) stopped working and you used the\n" .
+ "-f option, add 1 to the argument.\n" .
+ "See also: scripts/filter_scp.pl .\n";
+}
+
+
+$idlist = shift @ARGV;
+open(F, "<$idlist") || die "Could not open id-list file $idlist";
+while(<F>) {
+ @A = split;
+ @A>=1 || die "Invalid id-list file line $_";
+ $seen{$A[0]} = 1;
+}
+
+if ($field == 1) { # Treat this as special case, since it is common.
+ while(<>) {
+ $_ =~ m/\s*(\S+)\s*/ || die "Bad line $_, could not get first field.";
+ # $1 is what we filter on.
+ if ((!$exclude && $seen{$1}) || ($exclude && !defined $seen{$1})) {
+ print $_;
+ }
+ }
+} else {
+ while(<>) {
+ @A = split;
+ @A > 0 || die "Invalid scp file line $_";
+ @A >= $field || die "Invalid scp file line $_";
+ if ((!$exclude && $seen{$A[$field-1]}) || ($exclude && !defined $seen{$A[$field-1]})) {
+ print $_;
+ }
+ }
+}
+
+# tests:
+# the following should print "foo 1"
+# ( echo foo 1; echo bar 2 ) | scripts/filter_scp.pl <(echo foo)
+# the following should print "bar 2".
+# ( echo foo 1; echo bar 2 ) | scripts/filter_scp.pl -f 2 <(echo 2)
diff --git a/egs_modelscope/aishell/paraformer/utils/fix_data.sh b/egs_modelscope/aishell/paraformer/utils/fix_data.sh
new file mode 100755
index 0000000..32cdde5
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/fix_data.sh
@@ -0,0 +1,35 @@
+#!/usr/bin/env bash
+
+echo "$0 $@"
+data_dir=$1
+
+if [ ! -f ${data_dir}/wav.scp ]; then
+ echo "$0: wav.scp is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/text ]; then
+ echo "$0: text is not found"
+ exit 1;
+fi
+
+
+
+mkdir -p ${data_dir}/.backup
+
+awk '{print $1}' ${data_dir}/wav.scp > ${data_dir}/.backup/wav_id
+awk '{print $1}' ${data_dir}/text > ${data_dir}/.backup/text_id
+
+sort ${data_dir}/.backup/wav_id ${data_dir}/.backup/text_id | uniq -d > ${data_dir}/.backup/id
+
+cp ${data_dir}/wav.scp ${data_dir}/.backup/wav.scp
+cp ${data_dir}/text ${data_dir}/.backup/text
+
+mv ${data_dir}/wav.scp ${data_dir}/wav.scp.bak
+mv ${data_dir}/text ${data_dir}/text.bak
+
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/wav.scp.bak > ${data_dir}/wav.scp
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/text.bak > ${data_dir}/text
+
+rm ${data_dir}/wav.scp.bak
+rm ${data_dir}/text.bak
diff --git a/egs_modelscope/aishell/paraformer/utils/fix_data_feat.sh b/egs_modelscope/aishell/paraformer/utils/fix_data_feat.sh
new file mode 100755
index 0000000..2c92d7f
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/fix_data_feat.sh
@@ -0,0 +1,52 @@
+#!/usr/bin/env bash
+
+echo "$0 $@"
+data_dir=$1
+
+if [ ! -f ${data_dir}/feats.scp ]; then
+ echo "$0: feats.scp is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/text ]; then
+ echo "$0: text is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/speech_shape ]; then
+ echo "$0: feature lengths is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/text_shape ]; then
+ echo "$0: text lengths is not found"
+ exit 1;
+fi
+
+mkdir -p ${data_dir}/.backup
+
+awk '{print $1}' ${data_dir}/feats.scp > ${data_dir}/.backup/wav_id
+awk '{print $1}' ${data_dir}/text > ${data_dir}/.backup/text_id
+
+sort ${data_dir}/.backup/wav_id ${data_dir}/.backup/text_id | uniq -d > ${data_dir}/.backup/id
+
+cp ${data_dir}/feats.scp ${data_dir}/.backup/feats.scp
+cp ${data_dir}/text ${data_dir}/.backup/text
+cp ${data_dir}/speech_shape ${data_dir}/.backup/speech_shape
+cp ${data_dir}/text_shape ${data_dir}/.backup/text_shape
+
+mv ${data_dir}/feats.scp ${data_dir}/feats.scp.bak
+mv ${data_dir}/text ${data_dir}/text.bak
+mv ${data_dir}/speech_shape ${data_dir}/speech_shape.bak
+mv ${data_dir}/text_shape ${data_dir}/text_shape.bak
+
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/feats.scp.bak > ${data_dir}/feats.scp
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/text.bak > ${data_dir}/text
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/speech_shape.bak > ${data_dir}/speech_shape
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/text_shape.bak > ${data_dir}/text_shape
+
+rm ${data_dir}/feats.scp.bak
+rm ${data_dir}/text.bak
+rm ${data_dir}/speech_shape.bak
+rm ${data_dir}/text_shape.bak
+
diff --git a/egs_modelscope/aishell/paraformer/utils/gen_ark_list.sh b/egs_modelscope/aishell/paraformer/utils/gen_ark_list.sh
new file mode 100755
index 0000000..aebf356
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/gen_ark_list.sh
@@ -0,0 +1,22 @@
+#!/usr/bin/env bash
+
+
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+ark_dir=$1
+txt_dir=$2
+output_dir=$3
+
+[ ! -d ${ark_dir}/ark ] && echo "$0: ark data is required" && exit 1;
+[ ! -d ${txt_dir}/txt ] && echo "$0: txt data is required" && exit 1;
+
+for n in $(seq $nj); do
+ echo "${ark_dir}/ark/feats.$n.ark ${txt_dir}/txt/text.$n.txt" || exit 1
+done > ${output_dir}/ark_txt.scp || exit 1
+
diff --git a/egs_modelscope/aishell/paraformer/utils/parse_options.sh b/egs_modelscope/aishell/paraformer/utils/parse_options.sh
new file mode 100755
index 0000000..71fb9e5
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/parse_options.sh
@@ -0,0 +1,97 @@
+#!/usr/bin/env bash
+
+# Copyright 2012 Johns Hopkins University (Author: Daniel Povey);
+# Arnab Ghoshal, Karel Vesely
+
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+# Parse command-line options.
+# To be sourced by another script (as in ". parse_options.sh").
+# Option format is: --option-name arg
+# and shell variable "option_name" gets set to value "arg."
+# The exception is --help, which takes no arguments, but prints the
+# $help_message variable (if defined).
+
+
+###
+### The --config file options have lower priority to command line
+### options, so we need to import them first...
+###
+
+# Now import all the configs specified by command-line, in left-to-right order
+for ((argpos=1; argpos<$#; argpos++)); do
+ if [ "${!argpos}" == "--config" ]; then
+ argpos_plus1=$((argpos+1))
+ config=${!argpos_plus1}
+ [ ! -r $config ] && echo "$0: missing config '$config'" && exit 1
+ . $config # source the config file.
+ fi
+done
+
+
+###
+### Now we process the command line options
+###
+while true; do
+ [ -z "${1:-}" ] && break; # break if there are no arguments
+ case "$1" in
+ # If the enclosing script is called with --help option, print the help
+ # message and exit. Scripts should put help messages in $help_message
+ --help|-h) if [ -z "$help_message" ]; then echo "No help found." 1>&2;
+ else printf "$help_message\n" 1>&2 ; fi;
+ exit 0 ;;
+ --*=*) echo "$0: options to scripts must be of the form --name value, got '$1'"
+ exit 1 ;;
+ # If the first command-line argument begins with "--" (e.g. --foo-bar),
+ # then work out the variable name as $name, which will equal "foo_bar".
+ --*) name=`echo "$1" | sed s/^--// | sed s/-/_/g`;
+ # Next we test whether the variable in question is undefned-- if so it's
+ # an invalid option and we die. Note: $0 evaluates to the name of the
+ # enclosing script.
+ # The test [ -z ${foo_bar+xxx} ] will return true if the variable foo_bar
+ # is undefined. We then have to wrap this test inside "eval" because
+ # foo_bar is itself inside a variable ($name).
+ eval '[ -z "${'$name'+xxx}" ]' && echo "$0: invalid option $1" 1>&2 && exit 1;
+
+ oldval="`eval echo \\$$name`";
+ # Work out whether we seem to be expecting a Boolean argument.
+ if [ "$oldval" == "true" ] || [ "$oldval" == "false" ]; then
+ was_bool=true;
+ else
+ was_bool=false;
+ fi
+
+ # Set the variable to the right value-- the escaped quotes make it work if
+ # the option had spaces, like --cmd "queue.pl -sync y"
+ eval $name=\"$2\";
+
+ # Check that Boolean-valued arguments are really Boolean.
+ if $was_bool && [[ "$2" != "true" && "$2" != "false" ]]; then
+ echo "$0: expected \"true\" or \"false\": $1 $2" 1>&2
+ exit 1;
+ fi
+ shift 2;
+ ;;
+ *) break;
+ esac
+done
+
+
+# Check for an empty argument to the --cmd option, which can easily occur as a
+# result of scripting errors.
+[ ! -z "${cmd+xxx}" ] && [ -z "$cmd" ] && echo "$0: empty argument to --cmd option" 1>&2 && exit 1;
+
+
+true; # so this script returns exit code 0.
diff --git a/egs_modelscope/aishell/paraformer/utils/print_args.py b/egs_modelscope/aishell/paraformer/utils/print_args.py
new file mode 100755
index 0000000..b0c61e5
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/print_args.py
@@ -0,0 +1,45 @@
+#!/usr/bin/env python
+import sys
+
+
+def get_commandline_args(no_executable=True):
+ extra_chars = [
+ " ",
+ ";",
+ "&",
+ "|",
+ "<",
+ ">",
+ "?",
+ "*",
+ "~",
+ "`",
+ '"',
+ "'",
+ "\\",
+ "{",
+ "}",
+ "(",
+ ")",
+ ]
+
+ # Escape the extra characters for shell
+ argv = [
+ arg.replace("'", "'\\''")
+ if all(char not in arg for char in extra_chars)
+ else "'" + arg.replace("'", "'\\''") + "'"
+ for arg in sys.argv
+ ]
+
+ if no_executable:
+ return " ".join(argv[1:])
+ else:
+ return sys.executable + " " + " ".join(argv)
+
+
+def main():
+ print(get_commandline_args())
+
+
+if __name__ == "__main__":
+ main()
diff --git a/egs_modelscope/aishell/paraformer/utils/proc_conf_oss.py b/egs_modelscope/aishell/paraformer/utils/proc_conf_oss.py
new file mode 100755
index 0000000..c4a90c5
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/proc_conf_oss.py
@@ -0,0 +1,35 @@
+from pathlib import Path
+
+import torch
+import yaml
+
+
+class NoAliasSafeDumper(yaml.SafeDumper):
+ # Disable anchor/alias in yaml because looks ugly
+ def ignore_aliases(self, data):
+ return True
+
+
+def yaml_no_alias_safe_dump(data, stream=None, **kwargs):
+ """Safe-dump in yaml with no anchor/alias"""
+ return yaml.dump(
+ data, stream, allow_unicode=True, Dumper=NoAliasSafeDumper, **kwargs
+ )
+
+
+def gen_conf(file, out_dir):
+ conf = torch.load(file)["config"]
+ conf["oss_bucket"] = "null"
+ print(conf)
+ output_dir = Path(out_dir)
+ output_dir.mkdir(parents=True, exist_ok=True)
+ with (output_dir / "config.yaml").open("w", encoding="utf-8") as f:
+ yaml_no_alias_safe_dump(conf, f, indent=4, sort_keys=False)
+
+
+if __name__ == "__main__":
+ import sys
+
+ in_f = sys.argv[1]
+ out_f = sys.argv[2]
+ gen_conf(in_f, out_f)
diff --git a/egs_modelscope/aishell/paraformer/utils/proce_text.py b/egs_modelscope/aishell/paraformer/utils/proce_text.py
new file mode 100755
index 0000000..9e517a4
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/proce_text.py
@@ -0,0 +1,31 @@
+
+import sys
+import re
+
+in_f = sys.argv[1]
+out_f = sys.argv[2]
+
+
+with open(in_f, "r", encoding="utf-8") as f:
+ lines = f.readlines()
+
+with open(out_f, "w", encoding="utf-8") as f:
+ for line in lines:
+ outs = line.strip().split(" ", 1)
+ if len(outs) == 2:
+ idx, text = outs
+ text = re.sub("</s>", "", text)
+ text = re.sub("<s>", "", text)
+ text = re.sub("@@", "", text)
+ text = re.sub("@", "", text)
+ text = re.sub("<unk>", "", text)
+ text = re.sub(" ", "", text)
+ text = text.lower()
+ else:
+ idx = outs[0]
+ text = " "
+
+ text = [x for x in text]
+ text = " ".join(text)
+ out = "{} {}\n".format(idx, text)
+ f.write(out)
diff --git a/egs_modelscope/aishell/paraformer/utils/run.pl b/egs_modelscope/aishell/paraformer/utils/run.pl
new file mode 100755
index 0000000..483f95b
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/run.pl
@@ -0,0 +1,356 @@
+#!/usr/bin/env perl
+use warnings; #sed replacement for -w perl parameter
+# In general, doing
+# run.pl some.log a b c is like running the command a b c in
+# the bash shell, and putting the standard error and output into some.log.
+# To run parallel jobs (backgrounded on the host machine), you can do (e.g.)
+# run.pl JOB=1:4 some.JOB.log a b c JOB is like running the command a b c JOB
+# and putting it in some.JOB.log, for each one. [Note: JOB can be any identifier].
+# If any of the jobs fails, this script will fail.
+
+# A typical example is:
+# run.pl some.log my-prog "--opt=foo bar" foo \| other-prog baz
+# and run.pl will run something like:
+# ( my-prog '--opt=foo bar' foo | other-prog baz ) >& some.log
+#
+# Basically it takes the command-line arguments, quotes them
+# as necessary to preserve spaces, and evaluates them with bash.
+# In addition it puts the command line at the top of the log, and
+# the start and end times of the command at the beginning and end.
+# The reason why this is useful is so that we can create a different
+# version of this program that uses a queueing system instead.
+
+#use Data::Dumper;
+
+@ARGV < 2 && die "usage: run.pl log-file command-line arguments...";
+
+#print STDERR "COMMAND-LINE: " . Dumper(\@ARGV) . "\n";
+$job_pick = 'all';
+$max_jobs_run = -1;
+$jobstart = 1;
+$jobend = 1;
+$ignored_opts = ""; # These will be ignored.
+
+# First parse an option like JOB=1:4, and any
+# options that would normally be given to
+# queue.pl, which we will just discard.
+
+for (my $x = 1; $x <= 2; $x++) { # This for-loop is to
+ # allow the JOB=1:n option to be interleaved with the
+ # options to qsub.
+ while (@ARGV >= 2 && $ARGV[0] =~ m:^-:) {
+ # parse any options that would normally go to qsub, but which will be ignored here.
+ my $switch = shift @ARGV;
+ if ($switch eq "-V") {
+ $ignored_opts .= "-V ";
+ } elsif ($switch eq "--max-jobs-run" || $switch eq "-tc") {
+ # we do support the option --max-jobs-run n, and its GridEngine form -tc n.
+ # if the command appears multiple times uses the smallest option.
+ if ( $max_jobs_run <= 0 ) {
+ $max_jobs_run = shift @ARGV;
+ } else {
+ my $new_constraint = shift @ARGV;
+ if ( ($new_constraint < $max_jobs_run) ) {
+ $max_jobs_run = $new_constraint;
+ }
+ }
+
+ if (! ($max_jobs_run > 0)) {
+ die "run.pl: invalid option --max-jobs-run $max_jobs_run";
+ }
+ } else {
+ my $argument = shift @ARGV;
+ if ($argument =~ m/^--/) {
+ print STDERR "run.pl: WARNING: suspicious argument '$argument' to $switch; starts with '-'\n";
+ }
+ if ($switch eq "-sync" && $argument =~ m/^[yY]/) {
+ $ignored_opts .= "-sync "; # Note: in the
+ # corresponding code in queue.pl it says instead, just "$sync = 1;".
+ } elsif ($switch eq "-pe") { # e.g. -pe smp 5
+ my $argument2 = shift @ARGV;
+ $ignored_opts .= "$switch $argument $argument2 ";
+ } elsif ($switch eq "--gpu") {
+ $using_gpu = $argument;
+ } elsif ($switch eq "--pick") {
+ if($argument =~ m/^(all|failed|incomplete)$/) {
+ $job_pick = $argument;
+ } else {
+ print STDERR "run.pl: ERROR: --pick argument must be one of 'all', 'failed' or 'incomplete'"
+ }
+ } else {
+ # Ignore option.
+ $ignored_opts .= "$switch $argument ";
+ }
+ }
+ }
+ if ($ARGV[0] =~ m/^([\w_][\w\d_]*)+=(\d+):(\d+)$/) { # e.g. JOB=1:20
+ $jobname = $1;
+ $jobstart = $2;
+ $jobend = $3;
+ if ($jobstart > $jobend) {
+ die "run.pl: invalid job range $ARGV[0]";
+ }
+ if ($jobstart <= 0) {
+ die "run.pl: invalid job range $ARGV[0], start must be strictly positive (this is required for GridEngine compatibility).";
+ }
+ shift;
+ } elsif ($ARGV[0] =~ m/^([\w_][\w\d_]*)+=(\d+)$/) { # e.g. JOB=1.
+ $jobname = $1;
+ $jobstart = $2;
+ $jobend = $2;
+ shift;
+ } elsif ($ARGV[0] =~ m/.+\=.*\:.*$/) {
+ print STDERR "run.pl: Warning: suspicious first argument to run.pl: $ARGV[0]\n";
+ }
+}
+
+# Users found this message confusing so we are removing it.
+# if ($ignored_opts ne "") {
+# print STDERR "run.pl: Warning: ignoring options \"$ignored_opts\"\n";
+# }
+
+if ($max_jobs_run == -1) { # If --max-jobs-run option not set,
+ # then work out the number of processors if possible,
+ # and set it based on that.
+ $max_jobs_run = 0;
+ if ($using_gpu) {
+ if (open(P, "nvidia-smi -L |")) {
+ $max_jobs_run++ while (<P>);
+ close(P);
+ }
+ if ($max_jobs_run == 0) {
+ $max_jobs_run = 1;
+ print STDERR "run.pl: Warning: failed to detect number of GPUs from nvidia-smi, using ${max_jobs_run}\n";
+ }
+ } elsif (open(P, "</proc/cpuinfo")) { # Linux
+ while (<P>) { if (m/^processor/) { $max_jobs_run++; } }
+ if ($max_jobs_run == 0) {
+ print STDERR "run.pl: Warning: failed to detect any processors from /proc/cpuinfo\n";
+ $max_jobs_run = 10; # reasonable default.
+ }
+ close(P);
+ } elsif (open(P, "sysctl -a |")) { # BSD/Darwin
+ while (<P>) {
+ if (m/hw\.ncpu\s*[:=]\s*(\d+)/) { # hw.ncpu = 4, or hw.ncpu: 4
+ $max_jobs_run = $1;
+ last;
+ }
+ }
+ close(P);
+ if ($max_jobs_run == 0) {
+ print STDERR "run.pl: Warning: failed to detect any processors from sysctl -a\n";
+ $max_jobs_run = 10; # reasonable default.
+ }
+ } else {
+ # allow at most 32 jobs at once, on non-UNIX systems; change this code
+ # if you need to change this default.
+ $max_jobs_run = 32;
+ }
+ # The just-computed value of $max_jobs_run is just the number of processors
+ # (or our best guess); and if it happens that the number of jobs we need to
+ # run is just slightly above $max_jobs_run, it will make sense to increase
+ # $max_jobs_run to equal the number of jobs, so we don't have a small number
+ # of leftover jobs.
+ $num_jobs = $jobend - $jobstart + 1;
+ if (!$using_gpu &&
+ $num_jobs > $max_jobs_run && $num_jobs < 1.4 * $max_jobs_run) {
+ $max_jobs_run = $num_jobs;
+ }
+}
+
+sub pick_or_exit {
+ # pick_or_exit ( $logfile )
+ # Invoked before each job is started helps to run jobs selectively.
+ #
+ # Given the name of the output logfile decides whether the job must be
+ # executed (by returning from the subroutine) or not (by terminating the
+ # process calling exit)
+ #
+ # PRE: $job_pick is a global variable set by command line switch --pick
+ # and indicates which class of jobs must be executed.
+ #
+ # 1) If a failed job is not executed the process exit code will indicate
+ # failure, just as if the task was just executed and failed.
+ #
+ # 2) If a task is incomplete it will be executed. Incomplete may be either
+ # a job whose log file does not contain the accounting notes in the end,
+ # or a job whose log file does not exist.
+ #
+ # 3) If the $job_pick is set to 'all' (default behavior) a task will be
+ # executed regardless of the result of previous attempts.
+ #
+ # This logic could have been implemented in the main execution loop
+ # but a subroutine to preserve the current level of readability of
+ # that part of the code.
+ #
+ # Alexandre Felipe, (o.alexandre.felipe@gmail.com) 14th of August of 2020
+ #
+ if($job_pick eq 'all'){
+ return; # no need to bother with the previous log
+ }
+ open my $fh, "<", $_[0] or return; # job not executed yet
+ my $log_line;
+ my $cur_line;
+ while ($cur_line = <$fh>) {
+ if( $cur_line =~ m/# Ended \(code .*/ ) {
+ $log_line = $cur_line;
+ }
+ }
+ close $fh;
+ if (! defined($log_line)){
+ return; # incomplete
+ }
+ if ( $log_line =~ m/# Ended \(code 0\).*/ ) {
+ exit(0); # complete
+ } elsif ( $log_line =~ m/# Ended \(code \d+(; signal \d+)?\).*/ ){
+ if ($job_pick !~ m/^(failed|all)$/) {
+ exit(1); # failed but not going to run
+ } else {
+ return; # failed
+ }
+ } elsif ( $log_line =~ m/.*\S.*/ ) {
+ return; # incomplete jobs are always run
+ }
+}
+
+
+$logfile = shift @ARGV;
+
+if (defined $jobname && $logfile !~ m/$jobname/ &&
+ $jobend > $jobstart) {
+ print STDERR "run.pl: you are trying to run a parallel job but "
+ . "you are putting the output into just one log file ($logfile)\n";
+ exit(1);
+}
+
+$cmd = "";
+
+foreach $x (@ARGV) {
+ if ($x =~ m/^\S+$/) { $cmd .= $x . " "; }
+ elsif ($x =~ m:\":) { $cmd .= "'$x' "; }
+ else { $cmd .= "\"$x\" "; }
+}
+
+#$Data::Dumper::Indent=0;
+$ret = 0;
+$numfail = 0;
+%active_pids=();
+
+use POSIX ":sys_wait_h";
+for ($jobid = $jobstart; $jobid <= $jobend; $jobid++) {
+ if (scalar(keys %active_pids) >= $max_jobs_run) {
+
+ # Lets wait for a change in any child's status
+ # Then we have to work out which child finished
+ $r = waitpid(-1, 0);
+ $code = $?;
+ if ($r < 0 ) { die "run.pl: Error waiting for child process"; } # should never happen.
+ if ( defined $active_pids{$r} ) {
+ $jid=$active_pids{$r};
+ $fail[$jid]=$code;
+ if ($code !=0) { $numfail++;}
+ delete $active_pids{$r};
+ # print STDERR "Finished: $r/$jid " . Dumper(\%active_pids) . "\n";
+ } else {
+ die "run.pl: Cannot find the PID of the child process that just finished.";
+ }
+
+ # In theory we could do a non-blocking waitpid over all jobs running just
+ # to find out if only one or more jobs finished during the previous waitpid()
+ # However, we just omit this and will reap the next one in the next pass
+ # through the for(;;) cycle
+ }
+ $childpid = fork();
+ if (!defined $childpid) { die "run.pl: Error forking in run.pl (writing to $logfile)"; }
+ if ($childpid == 0) { # We're in the child... this branch
+ # executes the job and returns (possibly with an error status).
+ if (defined $jobname) {
+ $cmd =~ s/$jobname/$jobid/g;
+ $logfile =~ s/$jobname/$jobid/g;
+ }
+ # exit if the job does not need to be executed
+ pick_or_exit( $logfile );
+
+ system("mkdir -p `dirname $logfile` 2>/dev/null");
+ open(F, ">$logfile") || die "run.pl: Error opening log file $logfile";
+ print F "# " . $cmd . "\n";
+ print F "# Started at " . `date`;
+ $starttime = `date +'%s'`;
+ print F "#\n";
+ close(F);
+
+ # Pipe into bash.. make sure we're not using any other shell.
+ open(B, "|bash") || die "run.pl: Error opening shell command";
+ print B "( " . $cmd . ") 2>>$logfile >> $logfile";
+ close(B); # If there was an error, exit status is in $?
+ $ret = $?;
+
+ $lowbits = $ret & 127;
+ $highbits = $ret >> 8;
+ if ($lowbits != 0) { $return_str = "code $highbits; signal $lowbits" }
+ else { $return_str = "code $highbits"; }
+
+ $endtime = `date +'%s'`;
+ open(F, ">>$logfile") || die "run.pl: Error opening log file $logfile (again)";
+ $enddate = `date`;
+ chop $enddate;
+ print F "# Accounting: time=" . ($endtime - $starttime) . " threads=1\n";
+ print F "# Ended ($return_str) at " . $enddate . ", elapsed time " . ($endtime-$starttime) . " seconds\n";
+ close(F);
+ exit($ret == 0 ? 0 : 1);
+ } else {
+ $pid[$jobid] = $childpid;
+ $active_pids{$childpid} = $jobid;
+ # print STDERR "Queued: " . Dumper(\%active_pids) . "\n";
+ }
+}
+
+# Now we have submitted all the jobs, lets wait until all the jobs finish
+foreach $child (keys %active_pids) {
+ $jobid=$active_pids{$child};
+ $r = waitpid($pid[$jobid], 0);
+ $code = $?;
+ if ($r == -1) { die "run.pl: Error waiting for child process"; } # should never happen.
+ if ($r != 0) { $fail[$jobid]=$code; $numfail++ if $code!=0; } # Completed successfully
+}
+
+# Some sanity checks:
+# The $fail array should not contain undefined codes
+# The number of non-zeros in that array should be equal to $numfail
+# We cannot do foreach() here, as the JOB ids do not start at zero
+$failed_jids=0;
+for ($jobid = $jobstart; $jobid <= $jobend; $jobid++) {
+ $job_return = $fail[$jobid];
+ if (not defined $job_return ) {
+ # print Dumper(\@fail);
+
+ die "run.pl: Sanity check failed: we have indication that some jobs are running " .
+ "even after we waited for all jobs to finish" ;
+ }
+ if ($job_return != 0 ){ $failed_jids++;}
+}
+if ($failed_jids != $numfail) {
+ die "run.pl: Sanity check failed: cannot find out how many jobs failed ($failed_jids x $numfail)."
+}
+if ($numfail > 0) { $ret = 1; }
+
+if ($ret != 0) {
+ $njobs = $jobend - $jobstart + 1;
+ if ($njobs == 1) {
+ if (defined $jobname) {
+ $logfile =~ s/$jobname/$jobstart/; # only one numbered job, so replace name with
+ # that job.
+ }
+ print STDERR "run.pl: job failed, log is in $logfile\n";
+ if ($logfile =~ m/JOB/) {
+ print STDERR "run.pl: probably you forgot to put JOB=1:\$nj in your script.";
+ }
+ }
+ else {
+ $logfile =~ s/$jobname/*/g;
+ print STDERR "run.pl: $numfail / $njobs failed, log is in $logfile\n";
+ }
+}
+
+
+exit ($ret);
diff --git a/egs_modelscope/aishell/paraformer/utils/shuffle_list.pl b/egs_modelscope/aishell/paraformer/utils/shuffle_list.pl
new file mode 100755
index 0000000..a116200
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/shuffle_list.pl
@@ -0,0 +1,44 @@
+#!/usr/bin/env perl
+
+# Copyright 2013 Johns Hopkins University (author: Daniel Povey)
+
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+if ($ARGV[0] eq "--srand") {
+ $n = $ARGV[1];
+ $n =~ m/\d+/ || die "Bad argument to --srand option: \"$n\"";
+ srand($ARGV[1]);
+ shift;
+ shift;
+} else {
+ srand(0); # Gives inconsistent behavior if we don't seed.
+}
+
+if (@ARGV > 1 || $ARGV[0] =~ m/^-.+/) { # >1 args, or an option we
+ # don't understand.
+ print "Usage: shuffle_list.pl [--srand N] [input file] > output\n";
+ print "randomizes the order of lines of input.\n";
+ exit(1);
+}
+
+@lines;
+while (<>) {
+ push @lines, [ (rand(), $_)] ;
+}
+
+@lines = sort { $a->[0] cmp $b->[0] } @lines;
+foreach $l (@lines) {
+ print $l->[1];
+}
\ No newline at end of file
diff --git a/egs_modelscope/aishell/paraformer/utils/split_data.py b/egs_modelscope/aishell/paraformer/utils/split_data.py
new file mode 100755
index 0000000..060eae6
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/split_data.py
@@ -0,0 +1,60 @@
+import os
+import sys
+import random
+
+
+in_dir = sys.argv[1]
+out_dir = sys.argv[2]
+num_split = sys.argv[3]
+
+
+def split_scp(scp, num):
+ assert len(scp) >= num
+ avg = len(scp) // num
+ out = []
+ begin = 0
+
+ for i in range(num):
+ if i == num - 1:
+ out.append(scp[begin:])
+ else:
+ out.append(scp[begin:begin+avg])
+ begin += avg
+
+ return out
+
+
+os.path.exists("{}/wav.scp".format(in_dir))
+os.path.exists("{}/text".format(in_dir))
+
+with open("{}/wav.scp".format(in_dir), 'r') as infile:
+ wav_list = infile.readlines()
+
+with open("{}/text".format(in_dir), 'r') as infile:
+ text_list = infile.readlines()
+
+assert len(wav_list) == len(text_list)
+
+x = list(zip(wav_list, text_list))
+random.shuffle(x)
+wav_shuffle_list, text_shuffle_list = zip(*x)
+
+num_split = int(num_split)
+wav_split_list = split_scp(wav_shuffle_list, num_split)
+text_split_list = split_scp(text_shuffle_list, num_split)
+
+for idx, wav_list in enumerate(wav_split_list, 1):
+ path = out_dir + "/split" + str(num_split) + "/" + str(idx)
+ if not os.path.exists(path):
+ os.makedirs(path)
+ with open("{}/wav.scp".format(path), 'w') as wav_writer:
+ for line in wav_list:
+ wav_writer.write(line)
+
+for idx, text_list in enumerate(text_split_list, 1):
+ path = out_dir + "/split" + str(num_split) + "/" + str(idx)
+ if not os.path.exists(path):
+ os.makedirs(path)
+ with open("{}/text".format(path), 'w') as text_writer:
+ for line in text_list:
+ text_writer.write(line)
diff --git a/egs_modelscope/aishell/paraformer/utils/split_scp.pl b/egs_modelscope/aishell/paraformer/utils/split_scp.pl
new file mode 100755
index 0000000..0876dcb
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/split_scp.pl
@@ -0,0 +1,246 @@
+#!/usr/bin/env perl
+
+# Copyright 2010-2011 Microsoft Corporation
+
+# See ../../COPYING for clarification regarding multiple authors
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+# This program splits up any kind of .scp or archive-type file.
+# If there is no utt2spk option it will work on any text file and
+# will split it up with an approximately equal number of lines in
+# each but.
+# With the --utt2spk option it will work on anything that has the
+# utterance-id as the first entry on each line; the utt2spk file is
+# of the form "utterance speaker" (on each line).
+# It splits it into equal size chunks as far as it can. If you use the utt2spk
+# option it will make sure these chunks coincide with speaker boundaries. In
+# this case, if there are more chunks than speakers (and in some other
+# circumstances), some of the resulting chunks will be empty and it will print
+# an error message and exit with nonzero status.
+# You will normally call this like:
+# split_scp.pl scp scp.1 scp.2 scp.3 ...
+# or
+# split_scp.pl --utt2spk=utt2spk scp scp.1 scp.2 scp.3 ...
+# Note that you can use this script to split the utt2spk file itself,
+# e.g. split_scp.pl --utt2spk=utt2spk utt2spk utt2spk.1 utt2spk.2 ...
+
+# You can also call the scripts like:
+# split_scp.pl -j 3 0 scp scp.0
+# [note: with this option, it assumes zero-based indexing of the split parts,
+# i.e. the second number must be 0 <= n < num-jobs.]
+
+use warnings;
+
+$num_jobs = 0;
+$job_id = 0;
+$utt2spk_file = "";
+$one_based = 0;
+
+for ($x = 1; $x <= 3 && @ARGV > 0; $x++) {
+ if ($ARGV[0] eq "-j") {
+ shift @ARGV;
+ $num_jobs = shift @ARGV;
+ $job_id = shift @ARGV;
+ }
+ if ($ARGV[0] =~ /--utt2spk=(.+)/) {
+ $utt2spk_file=$1;
+ shift;
+ }
+ if ($ARGV[0] eq '--one-based') {
+ $one_based = 1;
+ shift @ARGV;
+ }
+}
+
+if ($num_jobs != 0 && ($num_jobs < 0 || $job_id - $one_based < 0 ||
+ $job_id - $one_based >= $num_jobs)) {
+ die "$0: Invalid job number/index values for '-j $num_jobs $job_id" .
+ ($one_based ? " --one-based" : "") . "'\n"
+}
+
+$one_based
+ and $job_id--;
+
+if(($num_jobs == 0 && @ARGV < 2) || ($num_jobs > 0 && (@ARGV < 1 || @ARGV > 2))) {
+ die
+"Usage: split_scp.pl [--utt2spk=<utt2spk_file>] in.scp out1.scp out2.scp ...
+ or: split_scp.pl -j num-jobs job-id [--one-based] [--utt2spk=<utt2spk_file>] in.scp [out.scp]
+ ... where 0 <= job-id < num-jobs, or 1 <= job-id <- num-jobs if --one-based.\n";
+}
+
+$error = 0;
+$inscp = shift @ARGV;
+if ($num_jobs == 0) { # without -j option
+ @OUTPUTS = @ARGV;
+} else {
+ for ($j = 0; $j < $num_jobs; $j++) {
+ if ($j == $job_id) {
+ if (@ARGV > 0) { push @OUTPUTS, $ARGV[0]; }
+ else { push @OUTPUTS, "-"; }
+ } else {
+ push @OUTPUTS, "/dev/null";
+ }
+ }
+}
+
+if ($utt2spk_file ne "") { # We have the --utt2spk option...
+ open($u_fh, '<', $utt2spk_file) || die "$0: Error opening utt2spk file $utt2spk_file: $!\n";
+ while(<$u_fh>) {
+ @A = split;
+ @A == 2 || die "$0: Bad line $_ in utt2spk file $utt2spk_file\n";
+ ($u,$s) = @A;
+ $utt2spk{$u} = $s;
+ }
+ close $u_fh;
+ open($i_fh, '<', $inscp) || die "$0: Error opening input scp file $inscp: $!\n";
+ @spkrs = ();
+ while(<$i_fh>) {
+ @A = split;
+ if(@A == 0) { die "$0: Empty or space-only line in scp file $inscp\n"; }
+ $u = $A[0];
+ $s = $utt2spk{$u};
+ defined $s || die "$0: No utterance $u in utt2spk file $utt2spk_file\n";
+ if(!defined $spk_count{$s}) {
+ push @spkrs, $s;
+ $spk_count{$s} = 0;
+ $spk_data{$s} = []; # ref to new empty array.
+ }
+ $spk_count{$s}++;
+ push @{$spk_data{$s}}, $_;
+ }
+ # Now split as equally as possible ..
+ # First allocate spks to files by allocating an approximately
+ # equal number of speakers.
+ $numspks = @spkrs; # number of speakers.
+ $numscps = @OUTPUTS; # number of output files.
+ if ($numspks < $numscps) {
+ die "$0: Refusing to split data because number of speakers $numspks " .
+ "is less than the number of output .scp files $numscps\n";
+ }
+ for($scpidx = 0; $scpidx < $numscps; $scpidx++) {
+ $scparray[$scpidx] = []; # [] is array reference.
+ }
+ for ($spkidx = 0; $spkidx < $numspks; $spkidx++) {
+ $scpidx = int(($spkidx*$numscps) / $numspks);
+ $spk = $spkrs[$spkidx];
+ push @{$scparray[$scpidx]}, $spk;
+ $scpcount[$scpidx] += $spk_count{$spk};
+ }
+
+ # Now will try to reassign beginning + ending speakers
+ # to different scp's and see if it gets more balanced.
+ # Suppose objf we're minimizing is sum_i (num utts in scp[i] - average)^2.
+ # We can show that if considering changing just 2 scp's, we minimize
+ # this by minimizing the squared difference in sizes. This is
+ # equivalent to minimizing the absolute difference in sizes. This
+ # shows this method is bound to converge.
+
+ $changed = 1;
+ while($changed) {
+ $changed = 0;
+ for($scpidx = 0; $scpidx < $numscps; $scpidx++) {
+ # First try to reassign ending spk of this scp.
+ if($scpidx < $numscps-1) {
+ $sz = @{$scparray[$scpidx]};
+ if($sz > 0) {
+ $spk = $scparray[$scpidx]->[$sz-1];
+ $count = $spk_count{$spk};
+ $nutt1 = $scpcount[$scpidx];
+ $nutt2 = $scpcount[$scpidx+1];
+ if( abs( ($nutt2+$count) - ($nutt1-$count))
+ < abs($nutt2 - $nutt1)) { # Would decrease
+ # size-diff by reassigning spk...
+ $scpcount[$scpidx+1] += $count;
+ $scpcount[$scpidx] -= $count;
+ pop @{$scparray[$scpidx]};
+ unshift @{$scparray[$scpidx+1]}, $spk;
+ $changed = 1;
+ }
+ }
+ }
+ if($scpidx > 0 && @{$scparray[$scpidx]} > 0) {
+ $spk = $scparray[$scpidx]->[0];
+ $count = $spk_count{$spk};
+ $nutt1 = $scpcount[$scpidx-1];
+ $nutt2 = $scpcount[$scpidx];
+ if( abs( ($nutt2-$count) - ($nutt1+$count))
+ < abs($nutt2 - $nutt1)) { # Would decrease
+ # size-diff by reassigning spk...
+ $scpcount[$scpidx-1] += $count;
+ $scpcount[$scpidx] -= $count;
+ shift @{$scparray[$scpidx]};
+ push @{$scparray[$scpidx-1]}, $spk;
+ $changed = 1;
+ }
+ }
+ }
+ }
+ # Now print out the files...
+ for($scpidx = 0; $scpidx < $numscps; $scpidx++) {
+ $scpfile = $OUTPUTS[$scpidx];
+ ($scpfile ne '-' ? open($f_fh, '>', $scpfile)
+ : open($f_fh, '>&', \*STDOUT)) ||
+ die "$0: Could not open scp file $scpfile for writing: $!\n";
+ $count = 0;
+ if(@{$scparray[$scpidx]} == 0) {
+ print STDERR "$0: eError: split_scp.pl producing empty .scp file " .
+ "$scpfile (too many splits and too few speakers?)\n";
+ $error = 1;
+ } else {
+ foreach $spk ( @{$scparray[$scpidx]} ) {
+ print $f_fh @{$spk_data{$spk}};
+ $count += $spk_count{$spk};
+ }
+ $count == $scpcount[$scpidx] || die "Count mismatch [code error]";
+ }
+ close($f_fh);
+ }
+} else {
+ # This block is the "normal" case where there is no --utt2spk
+ # option and we just break into equal size chunks.
+
+ open($i_fh, '<', $inscp) || die "$0: Error opening input scp file $inscp: $!\n";
+
+ $numscps = @OUTPUTS; # size of array.
+ @F = ();
+ while(<$i_fh>) {
+ push @F, $_;
+ }
+ $numlines = @F;
+ if($numlines == 0) {
+ print STDERR "$0: error: empty input scp file $inscp\n";
+ $error = 1;
+ }
+ $linesperscp = int( $numlines / $numscps); # the "whole part"..
+ $linesperscp >= 1 || die "$0: You are splitting into too many pieces! [reduce \$nj ($numscps) to be smaller than the number of lines ($numlines) in $inscp]\n";
+ $remainder = $numlines - ($linesperscp * $numscps);
+ ($remainder >= 0 && $remainder < $numlines) || die "bad remainder $remainder";
+ # [just doing int() rounds down].
+ $n = 0;
+ for($scpidx = 0; $scpidx < @OUTPUTS; $scpidx++) {
+ $scpfile = $OUTPUTS[$scpidx];
+ ($scpfile ne '-' ? open($o_fh, '>', $scpfile)
+ : open($o_fh, '>&', \*STDOUT)) ||
+ die "$0: Could not open scp file $scpfile for writing: $!\n";
+ for($k = 0; $k < $linesperscp + ($scpidx < $remainder ? 1 : 0); $k++) {
+ print $o_fh $F[$n++];
+ }
+ close($o_fh) || die "$0: Eror closing scp file $scpfile: $!\n";
+ }
+ $n == $numlines || die "$n != $numlines [code error]";
+}
+
+exit ($error);
diff --git a/egs_modelscope/aishell/paraformer/utils/subset_data_dir_tr_cv.sh b/egs_modelscope/aishell/paraformer/utils/subset_data_dir_tr_cv.sh
new file mode 100755
index 0000000..e16cebd
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/subset_data_dir_tr_cv.sh
@@ -0,0 +1,30 @@
+#!/usr/bin/env bash
+
+dev_num_utt=1000
+
+echo "$0 $@"
+. utils/parse_options.sh || exit 1;
+
+train_data=$1
+out_dir=$2
+
+[ ! -f ${train_data}/wav.scp ] && echo "$0: no such file ${train_data}/wav.scp" && exit 1;
+[ ! -f ${train_data}/text ] && echo "$0: no such file ${train_data}/text" && exit 1;
+
+mkdir -p ${out_dir}/train && mkdir -p ${out_dir}/dev
+
+cp ${train_data}/wav.scp ${out_dir}/train/wav.scp.bak
+cp ${train_data}/text ${out_dir}/train/text.bak
+
+num_utt=$(wc -l <${out_dir}/train/wav.scp.bak)
+
+utils/shuffle_list.pl --srand 1 ${out_dir}/train/wav.scp.bak > ${out_dir}/train/wav.scp.shuf
+head -n ${dev_num_utt} ${out_dir}/train/wav.scp.shuf > ${out_dir}/dev/wav.scp
+tail -n $((${num_utt}-${dev_num_utt})) ${out_dir}/train/wav.scp.shuf > ${out_dir}/train/wav.scp
+
+utils/shuffle_list.pl --srand 1 ${out_dir}/train/text.bak > ${out_dir}/train/text.shuf
+head -n ${dev_num_utt} ${out_dir}/train/text.shuf > ${out_dir}/dev/text
+tail -n $((${num_utt}-${dev_num_utt})) ${out_dir}/train/text.shuf > ${out_dir}/train/text
+
+rm ${out_dir}/train/wav.scp.bak ${out_dir}/train/text.bak
+rm ${out_dir}/train/wav.scp.shuf ${out_dir}/train/text.shuf
diff --git a/egs_modelscope/aishell/paraformer/utils/text2token.py b/egs_modelscope/aishell/paraformer/utils/text2token.py
new file mode 100755
index 0000000..56c3913
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/text2token.py
@@ -0,0 +1,135 @@
+#!/usr/bin/env python3
+
+# Copyright 2017 Johns Hopkins University (Shinji Watanabe)
+# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
+
+
+import argparse
+import codecs
+import re
+import sys
+
+is_python2 = sys.version_info[0] == 2
+
+
+def exist_or_not(i, match_pos):
+ start_pos = None
+ end_pos = None
+ for pos in match_pos:
+ if pos[0] <= i < pos[1]:
+ start_pos = pos[0]
+ end_pos = pos[1]
+ break
+
+ return start_pos, end_pos
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="convert raw text to tokenized text",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--nchar",
+ "-n",
+ default=1,
+ type=int,
+ help="number of characters to split, i.e., \
+ aabb -> a a b b with -n 1 and aa bb with -n 2",
+ )
+ parser.add_argument(
+ "--skip-ncols", "-s", default=0, type=int, help="skip first n columns"
+ )
+ parser.add_argument("--space", default="<space>", type=str, help="space symbol")
+ parser.add_argument(
+ "--non-lang-syms",
+ "-l",
+ default=None,
+ type=str,
+ help="list of non-linguistic symobles, e.g., <NOISE> etc.",
+ )
+ parser.add_argument("text", type=str, default=False, nargs="?", help="input text")
+ parser.add_argument(
+ "--trans_type",
+ "-t",
+ type=str,
+ default="char",
+ choices=["char", "phn"],
+ help="""Transcript type. char/phn. e.g., for TIMIT FADG0_SI1279 -
+ If trans_type is char,
+ read from SI1279.WRD file -> "bricks are an alternative"
+ Else if trans_type is phn,
+ read from SI1279.PHN file -> "sil b r ih sil k s aa r er n aa l
+ sil t er n ih sil t ih v sil" """,
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ rs = []
+ if args.non_lang_syms is not None:
+ with codecs.open(args.non_lang_syms, "r", encoding="utf-8") as f:
+ nls = [x.rstrip() for x in f.readlines()]
+ rs = [re.compile(re.escape(x)) for x in nls]
+
+ if args.text:
+ f = codecs.open(args.text, encoding="utf-8")
+ else:
+ f = codecs.getreader("utf-8")(sys.stdin if is_python2 else sys.stdin.buffer)
+
+ sys.stdout = codecs.getwriter("utf-8")(
+ sys.stdout if is_python2 else sys.stdout.buffer
+ )
+ line = f.readline()
+ n = args.nchar
+ while line:
+ x = line.split()
+ print(" ".join(x[: args.skip_ncols]), end=" ")
+ a = " ".join(x[args.skip_ncols :])
+
+ # get all matched positions
+ match_pos = []
+ for r in rs:
+ i = 0
+ while i >= 0:
+ m = r.search(a, i)
+ if m:
+ match_pos.append([m.start(), m.end()])
+ i = m.end()
+ else:
+ break
+
+ if args.trans_type == "phn":
+ a = a.split(" ")
+ else:
+ if len(match_pos) > 0:
+ chars = []
+ i = 0
+ while i < len(a):
+ start_pos, end_pos = exist_or_not(i, match_pos)
+ if start_pos is not None:
+ chars.append(a[start_pos:end_pos])
+ i = end_pos
+ else:
+ chars.append(a[i])
+ i += 1
+ a = chars
+
+ a = [a[j : j + n] for j in range(0, len(a), n)]
+
+ a_flat = []
+ for z in a:
+ a_flat.append("".join(z))
+
+ a_chars = [z.replace(" ", args.space) for z in a_flat]
+ if args.trans_type == "phn":
+ a_chars = [z.replace("sil", args.space) for z in a_chars]
+ print(" ".join(a_chars))
+ line = f.readline()
+
+
+if __name__ == "__main__":
+ main()
diff --git a/egs_modelscope/aishell/paraformer/utils/text_tokenize.py b/egs_modelscope/aishell/paraformer/utils/text_tokenize.py
new file mode 100755
index 0000000..962ea11
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/text_tokenize.py
@@ -0,0 +1,106 @@
+import re
+import argparse
+
+
+def load_dict(seg_file):
+ seg_dict = {}
+ with open(seg_file, 'r') as infile:
+ for line in infile:
+ s = line.strip().split()
+ key = s[0]
+ value = s[1:]
+ seg_dict[key] = " ".join(value)
+ return seg_dict
+
+
+def forward_segment(text, dic):
+ word_list = []
+ i = 0
+ while i < len(text):
+ longest_word = text[i]
+ for j in range(i + 1, len(text) + 1):
+ word = text[i:j]
+ if word in dic:
+ if len(word) > len(longest_word):
+ longest_word = word
+ word_list.append(longest_word)
+ i += len(longest_word)
+ return word_list
+
+
+def tokenize(txt,
+ seg_dict):
+ out_txt = ""
+ pattern = re.compile(r"([\u4E00-\u9FA5A-Za-z0-9])")
+ for word in txt:
+ if pattern.match(word):
+ if word in seg_dict:
+ out_txt += seg_dict[word] + " "
+ else:
+ out_txt += "<unk>" + " "
+ else:
+ continue
+ return out_txt.strip()
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="text tokenize",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--text-file",
+ "-t",
+ default=False,
+ required=True,
+ type=str,
+ help="input text",
+ )
+ parser.add_argument(
+ "--seg-file",
+ "-s",
+ default=False,
+ required=True,
+ type=str,
+ help="seg file",
+ )
+ parser.add_argument(
+ "--txt-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="txt index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ txt_writer = open("{}/text.{}.txt".format(args.output_dir, args.txt_index), 'w')
+ shape_writer = open("{}/len.{}".format(args.output_dir, args.txt_index), 'w')
+ seg_dict = load_dict(args.seg_file)
+ with open(args.text_file, 'r') as infile:
+ for line in infile:
+ s = line.strip().split()
+ text_id = s[0]
+ text_list = forward_segment("".join(s[1:]).lower(), seg_dict)
+ text = tokenize(text_list, seg_dict)
+ lens = len(text.strip().split())
+ txt_writer.write(text_id + " " + text + '\n')
+ shape_writer.write(text_id + " " + str(lens) + '\n')
+
+
+if __name__ == '__main__':
+ main()
+
diff --git a/egs_modelscope/aishell/paraformer/utils/text_tokenize.sh b/egs_modelscope/aishell/paraformer/utils/text_tokenize.sh
new file mode 100755
index 0000000..6b74fef
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/text_tokenize.sh
@@ -0,0 +1,35 @@
+#!/usr/bin/env bash
+
+
+# Begin configuration section.
+nj=32
+cmd=utils/run.pl
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+# tokenize configuration
+text_dir=$1
+seg_file=$2
+logdir=$3
+output_dir=$4
+
+txt_dir=${output_dir}/txt; mkdir -p ${output_dir}/txt
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/text_tokenize.JOB.log \
+ python utils/text_tokenize.py -t ${text_dir}/txt/text.JOB.txt \
+ -s ${seg_file} -i JOB -o ${txt_dir} \
+ || exit 1;
+
+# concatenate the text files together.
+for n in $(seq $nj); do
+ cat ${txt_dir}/text.$n.txt || exit 1
+done > ${output_dir}/text || exit 1
+
+for n in $(seq $nj); do
+ cat ${txt_dir}/len.$n || exit 1
+done > ${output_dir}/text_shape || exit 1
+
+echo "$0: Succeeded text tokenize"
diff --git a/egs_modelscope/aishell/paraformer/utils/textnorm_zh.py b/egs_modelscope/aishell/paraformer/utils/textnorm_zh.py
new file mode 100755
index 0000000..79feb83
--- /dev/null
+++ b/egs_modelscope/aishell/paraformer/utils/textnorm_zh.py
@@ -0,0 +1,834 @@
+#!/usr/bin/env python3
+# coding=utf-8
+
+# Authors:
+# 2019.5 Zhiyang Zhou (https://github.com/Joee1995/chn_text_norm.git)
+# 2019.9 Jiayu DU
+#
+# requirements:
+# - python 3.X
+# notes: python 2.X WILL fail or produce misleading results
+
+import sys, os, argparse, codecs, string, re
+
+# ================================================================================ #
+# basic constant
+# ================================================================================ #
+CHINESE_DIGIS = u'闆朵竴浜屼笁鍥涗簲鍏竷鍏節'
+BIG_CHINESE_DIGIS_SIMPLIFIED = u'闆跺9璐板弫鑲嗕紞闄嗘煉鎹岀帠'
+BIG_CHINESE_DIGIS_TRADITIONAL = u'闆跺9璨冲弮鑲嗕紞闄告煉鎹岀帠'
+SMALLER_BIG_CHINESE_UNITS_SIMPLIFIED = u'鍗佺櫨鍗冧竾'
+SMALLER_BIG_CHINESE_UNITS_TRADITIONAL = u'鎷句桨浠熻惉'
+LARGER_CHINESE_NUMERING_UNITS_SIMPLIFIED = u'浜垮厗浜灀绉┌娌熸锭姝h浇'
+LARGER_CHINESE_NUMERING_UNITS_TRADITIONAL = u'鍎勫厗浜灀绉┌婧濇緱姝h級'
+SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED = u'鍗佺櫨鍗冧竾'
+SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL = u'鎷句桨浠熻惉'
+
+ZERO_ALT = u'銆�'
+ONE_ALT = u'骞�'
+TWO_ALTS = [u'涓�', u'鍏�']
+
+POSITIVE = [u'姝�', u'姝�']
+NEGATIVE = [u'璐�', u'璨�']
+POINT = [u'鐐�', u'榛�']
+# PLUS = [u'鍔�', u'鍔�']
+# SIL = [u'鏉�', u'妲�']
+
+FILLER_CHARS = ['鍛�', '鍟�']
+ER_WHITELIST = '(鍎垮コ|鍎垮瓙|鍎垮瓩|濂冲効|鍎垮|濡诲効|' \
+ '鑳庡効|濠村効|鏂扮敓鍎縷濠村辜鍎縷骞煎効|灏戝効|灏忓効|鍎挎瓕|鍎跨|鍎跨|鎵樺効鎵�|瀛ゅ効|' \
+ '鍎挎垙|鍎垮寲|鍙板効搴剕楣垮効宀泑姝e効鍏粡|鍚婂効閮庡綋|鐢熷効鑲插コ|鎵樺効甯﹀コ|鍏诲効闃茶�亅鐥村効鍛嗗コ|' \
+ '浣冲効浣冲|鍎挎�滃吔鎵皘鍎挎棤甯哥埗|鍎夸笉瀚屾瘝涓憒鍎胯鍗冮噷姣嶆媴蹇鍎垮ぇ涓嶇敱鐖穦鑻忎篂鍎�)'
+
+# 涓枃鏁板瓧绯荤粺绫诲瀷
+NUMBERING_TYPES = ['low', 'mid', 'high']
+
+CURRENCY_NAMES = '(浜烘皯甯亅缇庡厓|鏃ュ厓|鑻遍晳|娆у厓|椹厠|娉曢儙|鍔犳嬁澶у厓|婢冲厓|娓竵|鍏堜护|鑺叞椹厠|鐖卞皵鍏伴晳|' \
+ '閲屾媺|鑽峰叞鐩緗鍩冩柉搴撳|姣斿濉攟鍗板凹鐩緗鏋楀悏鐗箌鏂拌タ鍏板厓|姣旂储|鍗㈠竷|鏂板姞鍧″厓|闊╁厓|娉伴摙)'
+CURRENCY_UNITS = '((浜縷鍗冧竾|鐧句竾|涓噟鍗億鐧�)|(浜縷鍗冧竾|鐧句竾|涓噟鍗億鐧緗)鍏億(浜縷鍗冧竾|鐧句竾|涓噟鍗億鐧緗)鍧梶瑙抾姣泑鍒�)'
+COM_QUANTIFIERS = '(鍖箌寮爘搴鍥瀨鍦簗灏緗鏉涓獆棣東闃檤闃祙缃憒鐐畖椤秥涓榺妫祙鍙獆鏀瘄琚瓅杈唡鎸憒鎷厊棰梶澹硘绐爘鏇瞸澧檤缇鑵攟' \
+ '鐮搴瀹璐瘄鎵巪鎹唡鍒�|浠鎵搢鎵媩缃梶鍧灞眧宀瓅姹焲婧獆閽焲闃焲鍗晐鍙寍瀵箌鍑簗鍙澶磡鑴殀鏉縷璺硘鏋潀浠秥璐磡' \
+ '閽坾绾縷绠鍚峾浣峾韬珅鍫倈璇緗鏈瑋椤祙瀹秥鎴穦灞倈涓潀姣珅鍘榺鍒唡閽眧涓鏂鎷厊閾鐭硘閽閿眧蹇絴(鍗億姣珅寰�)鍏媩' \
+ '姣珅鍘榺鍒唡瀵竱灏簗涓坾閲寍瀵粅甯竱閾簗绋媩(鍗億鍒唡鍘榺姣珅寰�)绫硘鎾畖鍕簗鍚坾鍗噟鏂梶鐭硘鐩榺纰梶纰焲鍙爘妗秥绗紎鐩唡' \
+ '鐩抾鏉瘄閽焲鏂泑閿厊绨媩绡畖鐩榺妗秥缃恷鐡秥澹秥鍗畖鐩弢绠﹟绠眧鐓瞸鍟東琚媩閽祙骞磡鏈坾鏃瀛鍒粅鏃秥鍛▅澶﹟绉抾鍒唡鏃瑋' \
+ '绾獆宀亅涓東鏇磡澶渱鏄澶弢绉媩鍐瑋浠浼弢杈坾涓竱娉绮抾棰梶骞鍫唡鏉鏍箌鏀瘄閬搢闈鐗噟寮爘棰梶鍧�)'
+
+# punctuation information are based on Zhon project (https://github.com/tsroten/zhon.git)
+CHINESE_PUNC_STOP = '锛侊紵锝°��'
+CHINESE_PUNC_NON_STOP = '锛傦純锛勶紖锛嗭紘锛堬級锛婏紜锛岋紞锛忥細锛涳紲锛濓紴锛狅蓟锛硷冀锛撅伎锝�锝涳綔锝濓綖锝燂綘锝剑锝ゃ�併�冦�嬨�屻�嶃�庛�忋�愩�戙�斻�曘�栥�椼�樸�欍�氥�涖�溿�濄�炪�熴�般�俱�库�撯�斺�樷�欌�涒�溾�濃�炩�熲�︹�э箯'
+CHINESE_PUNC_LIST = CHINESE_PUNC_STOP + CHINESE_PUNC_NON_STOP
+
+# ================================================================================ #
+# basic class
+# ================================================================================ #
+class ChineseChar(object):
+ """
+ 涓枃瀛楃
+ 姣忎釜瀛楃瀵瑰簲绠�浣撳拰绻佷綋,
+ e.g. 绠�浣� = '璐�', 绻佷綋 = '璨�'
+ 杞崲鏃跺彲杞崲涓虹畝浣撴垨绻佷綋
+ """
+
+ def __init__(self, simplified, traditional):
+ self.simplified = simplified
+ self.traditional = traditional
+ #self.__repr__ = self.__str__
+
+ def __str__(self):
+ return self.simplified or self.traditional or None
+
+ def __repr__(self):
+ return self.__str__()
+
+
+class ChineseNumberUnit(ChineseChar):
+ """
+ 涓枃鏁板瓧/鏁颁綅瀛楃
+ 姣忎釜瀛楃闄ょ箒绠�浣撳杩樻湁涓�涓澶栫殑澶у啓瀛楃
+ e.g. '闄�' 鍜� '闄�'
+ """
+
+ def __init__(self, power, simplified, traditional, big_s, big_t):
+ super(ChineseNumberUnit, self).__init__(simplified, traditional)
+ self.power = power
+ self.big_s = big_s
+ self.big_t = big_t
+
+ def __str__(self):
+ return '10^{}'.format(self.power)
+
+ @classmethod
+ def create(cls, index, value, numbering_type=NUMBERING_TYPES[1], small_unit=False):
+
+ if small_unit:
+ return ChineseNumberUnit(power=index + 1,
+ simplified=value[0], traditional=value[1], big_s=value[1], big_t=value[1])
+ elif numbering_type == NUMBERING_TYPES[0]:
+ return ChineseNumberUnit(power=index + 8,
+ simplified=value[0], traditional=value[1], big_s=value[0], big_t=value[1])
+ elif numbering_type == NUMBERING_TYPES[1]:
+ return ChineseNumberUnit(power=(index + 2) * 4,
+ simplified=value[0], traditional=value[1], big_s=value[0], big_t=value[1])
+ elif numbering_type == NUMBERING_TYPES[2]:
+ return ChineseNumberUnit(power=pow(2, index + 3),
+ simplified=value[0], traditional=value[1], big_s=value[0], big_t=value[1])
+ else:
+ raise ValueError(
+ 'Counting type should be in {0} ({1} provided).'.format(NUMBERING_TYPES, numbering_type))
+
+
+class ChineseNumberDigit(ChineseChar):
+ """
+ 涓枃鏁板瓧瀛楃
+ """
+
+ def __init__(self, value, simplified, traditional, big_s, big_t, alt_s=None, alt_t=None):
+ super(ChineseNumberDigit, self).__init__(simplified, traditional)
+ self.value = value
+ self.big_s = big_s
+ self.big_t = big_t
+ self.alt_s = alt_s
+ self.alt_t = alt_t
+
+ def __str__(self):
+ return str(self.value)
+
+ @classmethod
+ def create(cls, i, v):
+ return ChineseNumberDigit(i, v[0], v[1], v[2], v[3])
+
+
+class ChineseMath(ChineseChar):
+ """
+ 涓枃鏁颁綅瀛楃
+ """
+
+ def __init__(self, simplified, traditional, symbol, expression=None):
+ super(ChineseMath, self).__init__(simplified, traditional)
+ self.symbol = symbol
+ self.expression = expression
+ self.big_s = simplified
+ self.big_t = traditional
+
+
+CC, CNU, CND, CM = ChineseChar, ChineseNumberUnit, ChineseNumberDigit, ChineseMath
+
+
+class NumberSystem(object):
+ """
+ 涓枃鏁板瓧绯荤粺
+ """
+ pass
+
+
+class MathSymbol(object):
+ """
+ 鐢ㄤ簬涓枃鏁板瓧绯荤粺鐨勬暟瀛︾鍙� (绻�/绠�浣�), e.g.
+ positive = ['姝�', '姝�']
+ negative = ['璐�', '璨�']
+ point = ['鐐�', '榛�']
+ """
+
+ def __init__(self, positive, negative, point):
+ self.positive = positive
+ self.negative = negative
+ self.point = point
+
+ def __iter__(self):
+ for v in self.__dict__.values():
+ yield v
+
+
+# class OtherSymbol(object):
+# """
+# 鍏朵粬绗﹀彿
+# """
+#
+# def __init__(self, sil):
+# self.sil = sil
+#
+# def __iter__(self):
+# for v in self.__dict__.values():
+# yield v
+
+
+# ================================================================================ #
+# basic utils
+# ================================================================================ #
+def create_system(numbering_type=NUMBERING_TYPES[1]):
+ """
+ 鏍规嵁鏁板瓧绯荤粺绫诲瀷杩斿洖鍒涘缓鐩稿簲鐨勬暟瀛楃郴缁燂紝榛樿涓� mid
+ NUMBERING_TYPES = ['low', 'mid', 'high']: 涓枃鏁板瓧绯荤粺绫诲瀷
+ low: '鍏�' = '浜�' * '鍗�' = $10^{9}$, '浜�' = '鍏�' * '鍗�', etc.
+ mid: '鍏�' = '浜�' * '涓�' = $10^{12}$, '浜�' = '鍏�' * '涓�', etc.
+ high: '鍏�' = '浜�' * '浜�' = $10^{16}$, '浜�' = '鍏�' * '鍏�', etc.
+ 杩斿洖瀵瑰簲鐨勬暟瀛楃郴缁�
+ """
+
+ # chinese number units of '浜�' and larger
+ all_larger_units = zip(
+ LARGER_CHINESE_NUMERING_UNITS_SIMPLIFIED, LARGER_CHINESE_NUMERING_UNITS_TRADITIONAL)
+ larger_units = [CNU.create(i, v, numbering_type, False)
+ for i, v in enumerate(all_larger_units)]
+ # chinese number units of '鍗�, 鐧�, 鍗�, 涓�'
+ all_smaller_units = zip(
+ SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED, SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL)
+ smaller_units = [CNU.create(i, v, small_unit=True)
+ for i, v in enumerate(all_smaller_units)]
+ # digis
+ chinese_digis = zip(CHINESE_DIGIS, CHINESE_DIGIS,
+ BIG_CHINESE_DIGIS_SIMPLIFIED, BIG_CHINESE_DIGIS_TRADITIONAL)
+ digits = [CND.create(i, v) for i, v in enumerate(chinese_digis)]
+ digits[0].alt_s, digits[0].alt_t = ZERO_ALT, ZERO_ALT
+ digits[1].alt_s, digits[1].alt_t = ONE_ALT, ONE_ALT
+ digits[2].alt_s, digits[2].alt_t = TWO_ALTS[0], TWO_ALTS[1]
+
+ # symbols
+ positive_cn = CM(POSITIVE[0], POSITIVE[1], '+', lambda x: x)
+ negative_cn = CM(NEGATIVE[0], NEGATIVE[1], '-', lambda x: -x)
+ point_cn = CM(POINT[0], POINT[1], '.', lambda x,
+ y: float(str(x) + '.' + str(y)))
+ # sil_cn = CM(SIL[0], SIL[1], '-', lambda x, y: float(str(x) + '-' + str(y)))
+ system = NumberSystem()
+ system.units = smaller_units + larger_units
+ system.digits = digits
+ system.math = MathSymbol(positive_cn, negative_cn, point_cn)
+ # system.symbols = OtherSymbol(sil_cn)
+ return system
+
+
+def chn2num(chinese_string, numbering_type=NUMBERING_TYPES[1]):
+
+ def get_symbol(char, system):
+ for u in system.units:
+ if char in [u.traditional, u.simplified, u.big_s, u.big_t]:
+ return u
+ for d in system.digits:
+ if char in [d.traditional, d.simplified, d.big_s, d.big_t, d.alt_s, d.alt_t]:
+ return d
+ for m in system.math:
+ if char in [m.traditional, m.simplified]:
+ return m
+
+ def string2symbols(chinese_string, system):
+ int_string, dec_string = chinese_string, ''
+ for p in [system.math.point.simplified, system.math.point.traditional]:
+ if p in chinese_string:
+ int_string, dec_string = chinese_string.split(p)
+ break
+ return [get_symbol(c, system) for c in int_string], \
+ [get_symbol(c, system) for c in dec_string]
+
+ def correct_symbols(integer_symbols, system):
+ """
+ 涓�鐧惧叓 to 涓�鐧惧叓鍗�
+ 涓�浜夸竴鍗冧笁鐧句竾 to 涓�浜� 涓�鍗冧竾 涓夌櫨涓�
+ """
+
+ if integer_symbols and isinstance(integer_symbols[0], CNU):
+ if integer_symbols[0].power == 1:
+ integer_symbols = [system.digits[1]] + integer_symbols
+
+ if len(integer_symbols) > 1:
+ if isinstance(integer_symbols[-1], CND) and isinstance(integer_symbols[-2], CNU):
+ integer_symbols.append(
+ CNU(integer_symbols[-2].power - 1, None, None, None, None))
+
+ result = []
+ unit_count = 0
+ for s in integer_symbols:
+ if isinstance(s, CND):
+ result.append(s)
+ unit_count = 0
+ elif isinstance(s, CNU):
+ current_unit = CNU(s.power, None, None, None, None)
+ unit_count += 1
+
+ if unit_count == 1:
+ result.append(current_unit)
+ elif unit_count > 1:
+ for i in range(len(result)):
+ if isinstance(result[-i - 1], CNU) and result[-i - 1].power < current_unit.power:
+ result[-i - 1] = CNU(result[-i - 1].power +
+ current_unit.power, None, None, None, None)
+ return result
+
+ def compute_value(integer_symbols):
+ """
+ Compute the value.
+ When current unit is larger than previous unit, current unit * all previous units will be used as all previous units.
+ e.g. '涓ゅ崈涓�' = 2000 * 10000 not 2000 + 10000
+ """
+ value = [0]
+ last_power = 0
+ for s in integer_symbols:
+ if isinstance(s, CND):
+ value[-1] = s.value
+ elif isinstance(s, CNU):
+ value[-1] *= pow(10, s.power)
+ if s.power > last_power:
+ value[:-1] = list(map(lambda v: v *
+ pow(10, s.power), value[:-1]))
+ last_power = s.power
+ value.append(0)
+ return sum(value)
+
+ system = create_system(numbering_type)
+ int_part, dec_part = string2symbols(chinese_string, system)
+ int_part = correct_symbols(int_part, system)
+ int_str = str(compute_value(int_part))
+ dec_str = ''.join([str(d.value) for d in dec_part])
+ if dec_part:
+ return '{0}.{1}'.format(int_str, dec_str)
+ else:
+ return int_str
+
+
+def num2chn(number_string, numbering_type=NUMBERING_TYPES[1], big=False,
+ traditional=False, alt_zero=False, alt_one=False, alt_two=True,
+ use_zeros=True, use_units=True):
+
+ def get_value(value_string, use_zeros=True):
+
+ striped_string = value_string.lstrip('0')
+
+ # record nothing if all zeros
+ if not striped_string:
+ return []
+
+ # record one digits
+ elif len(striped_string) == 1:
+ if use_zeros and len(value_string) != len(striped_string):
+ return [system.digits[0], system.digits[int(striped_string)]]
+ else:
+ return [system.digits[int(striped_string)]]
+
+ # recursively record multiple digits
+ else:
+ result_unit = next(u for u in reversed(
+ system.units) if u.power < len(striped_string))
+ result_string = value_string[:-result_unit.power]
+ return get_value(result_string) + [result_unit] + get_value(striped_string[-result_unit.power:])
+
+ system = create_system(numbering_type)
+
+ int_dec = number_string.split('.')
+ if len(int_dec) == 1:
+ int_string = int_dec[0]
+ dec_string = ""
+ elif len(int_dec) == 2:
+ int_string = int_dec[0]
+ dec_string = int_dec[1]
+ else:
+ raise ValueError(
+ "invalid input num string with more than one dot: {}".format(number_string))
+
+ if use_units and len(int_string) > 1:
+ result_symbols = get_value(int_string)
+ else:
+ result_symbols = [system.digits[int(c)] for c in int_string]
+ dec_symbols = [system.digits[int(c)] for c in dec_string]
+ if dec_string:
+ result_symbols += [system.math.point] + dec_symbols
+
+ if alt_two:
+ liang = CND(2, system.digits[2].alt_s, system.digits[2].alt_t,
+ system.digits[2].big_s, system.digits[2].big_t)
+ for i, v in enumerate(result_symbols):
+ if isinstance(v, CND) and v.value == 2:
+ next_symbol = result_symbols[i +
+ 1] if i < len(result_symbols) - 1 else None
+ previous_symbol = result_symbols[i - 1] if i > 0 else None
+ if isinstance(next_symbol, CNU) and isinstance(previous_symbol, (CNU, type(None))):
+ if next_symbol.power != 1 and ((previous_symbol is None) or (previous_symbol.power != 1)):
+ result_symbols[i] = liang
+
+ # if big is True, '涓�' will not be used and `alt_two` has no impact on output
+ if big:
+ attr_name = 'big_'
+ if traditional:
+ attr_name += 't'
+ else:
+ attr_name += 's'
+ else:
+ if traditional:
+ attr_name = 'traditional'
+ else:
+ attr_name = 'simplified'
+
+ result = ''.join([getattr(s, attr_name) for s in result_symbols])
+
+ # if not use_zeros:
+ # result = result.strip(getattr(system.digits[0], attr_name))
+
+ if alt_zero:
+ result = result.replace(
+ getattr(system.digits[0], attr_name), system.digits[0].alt_s)
+
+ if alt_one:
+ result = result.replace(
+ getattr(system.digits[1], attr_name), system.digits[1].alt_s)
+
+ for i, p in enumerate(POINT):
+ if result.startswith(p):
+ return CHINESE_DIGIS[0] + result
+
+ # ^10, 11, .., 19
+ if len(result) >= 2 and result[1] in [SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED[0],
+ SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL[0]] and \
+ result[0] in [CHINESE_DIGIS[1], BIG_CHINESE_DIGIS_SIMPLIFIED[1], BIG_CHINESE_DIGIS_TRADITIONAL[1]]:
+ result = result[1:]
+
+ return result
+
+
+# ================================================================================ #
+# different types of rewriters
+# ================================================================================ #
+class Cardinal:
+ """
+ CARDINAL绫�
+ """
+
+ def __init__(self, cardinal=None, chntext=None):
+ self.cardinal = cardinal
+ self.chntext = chntext
+
+ def chntext2cardinal(self):
+ return chn2num(self.chntext)
+
+ def cardinal2chntext(self):
+ return num2chn(self.cardinal)
+
+class Digit:
+ """
+ DIGIT绫�
+ """
+
+ def __init__(self, digit=None, chntext=None):
+ self.digit = digit
+ self.chntext = chntext
+
+ # def chntext2digit(self):
+ # return chn2num(self.chntext)
+
+ def digit2chntext(self):
+ return num2chn(self.digit, alt_two=False, use_units=False)
+
+
+class TelePhone:
+ """
+ TELEPHONE绫�
+ """
+
+ def __init__(self, telephone=None, raw_chntext=None, chntext=None):
+ self.telephone = telephone
+ self.raw_chntext = raw_chntext
+ self.chntext = chntext
+
+ # def chntext2telephone(self):
+ # sil_parts = self.raw_chntext.split('<SIL>')
+ # self.telephone = '-'.join([
+ # str(chn2num(p)) for p in sil_parts
+ # ])
+ # return self.telephone
+
+ def telephone2chntext(self, fixed=False):
+
+ if fixed:
+ sil_parts = self.telephone.split('-')
+ self.raw_chntext = '<SIL>'.join([
+ num2chn(part, alt_two=False, use_units=False) for part in sil_parts
+ ])
+ self.chntext = self.raw_chntext.replace('<SIL>', '')
+ else:
+ sp_parts = self.telephone.strip('+').split()
+ self.raw_chntext = '<SP>'.join([
+ num2chn(part, alt_two=False, use_units=False) for part in sp_parts
+ ])
+ self.chntext = self.raw_chntext.replace('<SP>', '')
+ return self.chntext
+
+
+class Fraction:
+ """
+ FRACTION绫�
+ """
+
+ def __init__(self, fraction=None, chntext=None):
+ self.fraction = fraction
+ self.chntext = chntext
+
+ def chntext2fraction(self):
+ denominator, numerator = self.chntext.split('鍒嗕箣')
+ return chn2num(numerator) + '/' + chn2num(denominator)
+
+ def fraction2chntext(self):
+ numerator, denominator = self.fraction.split('/')
+ return num2chn(denominator) + '鍒嗕箣' + num2chn(numerator)
+
+
+class Date:
+ """
+ DATE绫�
+ """
+
+ def __init__(self, date=None, chntext=None):
+ self.date = date
+ self.chntext = chntext
+
+ # def chntext2date(self):
+ # chntext = self.chntext
+ # try:
+ # year, other = chntext.strip().split('骞�', maxsplit=1)
+ # year = Digit(chntext=year).digit2chntext() + '骞�'
+ # except ValueError:
+ # other = chntext
+ # year = ''
+ # if other:
+ # try:
+ # month, day = other.strip().split('鏈�', maxsplit=1)
+ # month = Cardinal(chntext=month).chntext2cardinal() + '鏈�'
+ # except ValueError:
+ # day = chntext
+ # month = ''
+ # if day:
+ # day = Cardinal(chntext=day[:-1]).chntext2cardinal() + day[-1]
+ # else:
+ # month = ''
+ # day = ''
+ # date = year + month + day
+ # self.date = date
+ # return self.date
+
+ def date2chntext(self):
+ date = self.date
+ try:
+ year, other = date.strip().split('骞�', 1)
+ year = Digit(digit=year).digit2chntext() + '骞�'
+ except ValueError:
+ other = date
+ year = ''
+ if other:
+ try:
+ month, day = other.strip().split('鏈�', 1)
+ month = Cardinal(cardinal=month).cardinal2chntext() + '鏈�'
+ except ValueError:
+ day = date
+ month = ''
+ if day:
+ day = Cardinal(cardinal=day[:-1]).cardinal2chntext() + day[-1]
+ else:
+ month = ''
+ day = ''
+ chntext = year + month + day
+ self.chntext = chntext
+ return self.chntext
+
+
+class Money:
+ """
+ MONEY绫�
+ """
+
+ def __init__(self, money=None, chntext=None):
+ self.money = money
+ self.chntext = chntext
+
+ # def chntext2money(self):
+ # return self.money
+
+ def money2chntext(self):
+ money = self.money
+ pattern = re.compile(r'(\d+(\.\d+)?)')
+ matchers = pattern.findall(money)
+ if matchers:
+ for matcher in matchers:
+ money = money.replace(matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext())
+ self.chntext = money
+ return self.chntext
+
+
+class Percentage:
+ """
+ PERCENTAGE绫�
+ """
+
+ def __init__(self, percentage=None, chntext=None):
+ self.percentage = percentage
+ self.chntext = chntext
+
+ def chntext2percentage(self):
+ return chn2num(self.chntext.strip().strip('鐧惧垎涔�')) + '%'
+
+ def percentage2chntext(self):
+ return '鐧惧垎涔�' + num2chn(self.percentage.strip().strip('%'))
+
+
+def remove_erhua(text, er_whitelist):
+ """
+ 鍘婚櫎鍎垮寲闊宠瘝涓殑鍎�:
+ 浠栧コ鍎垮湪閭h竟鍎� -> 浠栧コ鍎垮湪閭h竟
+ """
+
+ er_pattern = re.compile(er_whitelist)
+ new_str=''
+ while re.search('鍎�',text):
+ a = re.search('鍎�',text).span()
+ remove_er_flag = 0
+
+ if er_pattern.search(text):
+ b = er_pattern.search(text).span()
+ if b[0] <= a[0]:
+ remove_er_flag = 1
+
+ if remove_er_flag == 0 :
+ new_str = new_str + text[0:a[0]]
+ text = text[a[1]:]
+ else:
+ new_str = new_str + text[0:b[1]]
+ text = text[b[1]:]
+
+ text = new_str + text
+ return text
+
+# ================================================================================ #
+# NSW Normalizer
+# ================================================================================ #
+class NSWNormalizer:
+ def __init__(self, raw_text):
+ self.raw_text = '^' + raw_text + '$'
+ self.norm_text = ''
+
+ def _particular(self):
+ text = self.norm_text
+ pattern = re.compile(r"(([a-zA-Z]+)浜�([a-zA-Z]+))")
+ matchers = pattern.findall(text)
+ if matchers:
+ # print('particular')
+ for matcher in matchers:
+ text = text.replace(matcher[0], matcher[1]+'2'+matcher[2], 1)
+ self.norm_text = text
+ return self.norm_text
+
+ def normalize(self):
+ text = self.raw_text
+
+ # 瑙勮寖鍖栨棩鏈�
+ pattern = re.compile(r"\D+((([089]\d|(19|20)\d{2})骞�)?(\d{1,2}鏈�(\d{1,2}[鏃ュ彿])?)?)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('date')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Date(date=matcher[0]).date2chntext(), 1)
+
+ # 瑙勮寖鍖栭噾閽�
+ pattern = re.compile(r"\D+((\d+(\.\d+)?)[澶氫綑鍑燷?" + CURRENCY_UNITS + r"(\d" + CURRENCY_UNITS + r"?)?)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('money')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Money(money=matcher[0]).money2chntext(), 1)
+
+ # 瑙勮寖鍖栧浐璇�/鎵嬫満鍙风爜
+ # 鎵嬫満
+ # http://www.jihaoba.com/news/show/13680
+ # 绉诲姩锛�139銆�138銆�137銆�136銆�135銆�134銆�159銆�158銆�157銆�150銆�151銆�152銆�188銆�187銆�182銆�183銆�184銆�178銆�198
+ # 鑱旈�氾細130銆�131銆�132銆�156銆�155銆�186銆�185銆�176
+ # 鐢典俊锛�133銆�153銆�189銆�180銆�181銆�177
+ pattern = re.compile(r"\D((\+?86 ?)?1([38]\d|5[0-35-9]|7[678]|9[89])\d{8})\D")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('telephone')
+ for matcher in matchers:
+ text = text.replace(matcher[0], TelePhone(telephone=matcher[0]).telephone2chntext(), 1)
+ # 鍥鸿瘽
+ pattern = re.compile(r"\D((0(10|2[1-3]|[3-9]\d{2})-?)?[1-9]\d{6,7})\D")
+ matchers = pattern.findall(text)
+ if matchers:
+ # print('fixed telephone')
+ for matcher in matchers:
+ text = text.replace(matcher[0], TelePhone(telephone=matcher[0]).telephone2chntext(fixed=True), 1)
+
+ # 瑙勮寖鍖栧垎鏁�
+ pattern = re.compile(r"(\d+/\d+)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('fraction')
+ for matcher in matchers:
+ text = text.replace(matcher, Fraction(fraction=matcher).fraction2chntext(), 1)
+
+ # 瑙勮寖鍖栫櫨鍒嗘暟
+ text = text.replace('锛�', '%')
+ pattern = re.compile(r"(\d+(\.\d+)?%)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('percentage')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Percentage(percentage=matcher[0]).percentage2chntext(), 1)
+
+ # 瑙勮寖鍖栫函鏁�+閲忚瘝
+ pattern = re.compile(r"(\d+(\.\d+)?)[澶氫綑鍑燷?" + COM_QUANTIFIERS)
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('cardinal+quantifier')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext(), 1)
+
+ # 瑙勮寖鍖栨暟瀛楃紪鍙�
+ pattern = re.compile(r"(\d{4,32})")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('digit')
+ for matcher in matchers:
+ text = text.replace(matcher, Digit(digit=matcher).digit2chntext(), 1)
+
+ # 瑙勮寖鍖栫函鏁�
+ pattern = re.compile(r"(\d+(\.\d+)?)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('cardinal')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext(), 1)
+
+ self.norm_text = text
+ self._particular()
+
+ return self.norm_text.lstrip('^').rstrip('$')
+
+
+def nsw_test_case(raw_text):
+ print('I:' + raw_text)
+ print('O:' + NSWNormalizer(raw_text).normalize())
+ print('')
+
+
+def nsw_test():
+ nsw_test_case('鍥鸿瘽锛�0595-23865596鎴�23880880銆�')
+ nsw_test_case('鍥鸿瘽锛�0595-23865596鎴�23880880銆�')
+ nsw_test_case('鎵嬫満锛�+86 19859213959鎴�15659451527銆�')
+ nsw_test_case('鍒嗘暟锛�32477/76391銆�')
+ nsw_test_case('鐧惧垎鏁帮細80.03%銆�')
+ nsw_test_case('缂栧彿锛�31520181154418銆�')
+ nsw_test_case('绾暟锛�2983.07鍏嬫垨12345.60绫炽��')
+ nsw_test_case('鏃ユ湡锛�1999骞�2鏈�20鏃ユ垨09骞�3鏈�15鍙枫��')
+ nsw_test_case('閲戦挶锛�12鍧�5锛�34.5鍏冿紝20.1涓�')
+ nsw_test_case('鐗规畩锛歄2O鎴朆2C銆�')
+ nsw_test_case('3456涓囧惃')
+ nsw_test_case('2938涓�')
+ nsw_test_case('938')
+ nsw_test_case('浠婂ぉ鍚冧簡115涓皬绗煎寘231涓澶�')
+ nsw_test_case('鏈�62锛呯殑姒傜巼')
+
+
+if __name__ == '__main__':
+ #nsw_test()
+
+ p = argparse.ArgumentParser()
+ p.add_argument('ifile', help='input filename, assume utf-8 encoding')
+ p.add_argument('ofile', help='output filename')
+ p.add_argument('--to_upper', action='store_true', help='convert to upper case')
+ p.add_argument('--to_lower', action='store_true', help='convert to lower case')
+ p.add_argument('--has_key', action='store_true', help="input text has Kaldi's key as first field.")
+ p.add_argument('--remove_fillers', type=bool, default=True, help='remove filler chars such as "鍛�, 鍟�"')
+ p.add_argument('--remove_erhua', type=bool, default=True, help='remove erhua chars such as "杩欏効"')
+ p.add_argument('--log_interval', type=int, default=10000, help='log interval in number of processed lines')
+ args = p.parse_args()
+
+ ifile = codecs.open(args.ifile, 'r', 'utf8')
+ ofile = codecs.open(args.ofile, 'w+', 'utf8')
+
+ n = 0
+ for l in ifile:
+ key = ''
+ text = ''
+ if args.has_key:
+ cols = l.split(maxsplit=1)
+ key = cols[0]
+ if len(cols) == 2:
+ text = cols[1].strip()
+ else:
+ text = ''
+ else:
+ text = l.strip()
+
+ # cases
+ if args.to_upper and args.to_lower:
+ sys.stderr.write('text norm: to_upper OR to_lower?')
+ exit(1)
+ if args.to_upper:
+ text = text.upper()
+ if args.to_lower:
+ text = text.lower()
+
+ # Filler chars removal
+ if args.remove_fillers:
+ for ch in FILLER_CHARS:
+ text = text.replace(ch, '')
+
+ if args.remove_erhua:
+ text = remove_erhua(text, ER_WHITELIST)
+
+ # NSW(Non-Standard-Word) normalization
+ text = NSWNormalizer(text).normalize()
+
+ # Punctuations removal
+ old_chars = CHINESE_PUNC_LIST + string.punctuation # includes all CN and EN punctuations
+ new_chars = ' ' * len(old_chars)
+ del_chars = ''
+ text = text.translate(str.maketrans(old_chars, new_chars, del_chars))
+
+ #
+ if args.has_key:
+ ofile.write(key + '\t' + text + '\n')
+ else:
+ ofile.write(text + '\n')
+
+ n += 1
+ if n % args.log_interval == 0:
+ sys.stderr.write("text norm: {} lines done.\n".format(n))
+
+ sys.stderr.write("text norm: {} lines done in total.\n".format(n))
+
+ ifile.close()
+ ofile.close()
diff --git a/egs_modelscope/aishell2/paraformer/conf/train_asr_paraformer_sanm_50e_16d_2048_512_lfr6.yaml b/egs_modelscope/aishell2/paraformer/conf/train_asr_paraformer_sanm_50e_16d_2048_512_lfr6.yaml
index e9210f3..c8990eb 100644
--- a/egs_modelscope/aishell2/paraformer/conf/train_asr_paraformer_sanm_50e_16d_2048_512_lfr6.yaml
+++ b/egs_modelscope/aishell2/paraformer/conf/train_asr_paraformer_sanm_50e_16d_2048_512_lfr6.yaml
@@ -27,6 +27,7 @@
predictor_bias: 1
sampling_ratio: 0.75
+
# minibatch related
# dataset_type: small
batch_type: length
diff --git a/egs_modelscope/aishell2/paraformer/modelscope_utils b/egs_modelscope/aishell2/paraformer/modelscope_utils
deleted file mode 120000
index fc97768..0000000
--- a/egs_modelscope/aishell2/paraformer/modelscope_utils
+++ /dev/null
@@ -1 +0,0 @@
-../../common/modelscope_utils
\ No newline at end of file
diff --git a/egs_modelscope/aishell2/paraformer/modelscope_utils/download_model.py b/egs_modelscope/aishell2/paraformer/modelscope_utils/download_model.py
new file mode 100755
index 0000000..51ba6b8
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/modelscope_utils/download_model.py
@@ -0,0 +1,25 @@
+#!/usr/bin/env python3
+import argparse
+
+from modelscope.pipelines import pipeline
+from modelscope.utils.constant import Tasks
+
+if __name__ == '__main__':
+ parser = argparse.ArgumentParser(
+ description="download model configs",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument("--model_name",
+ type=str,
+ default="speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-pytorch",
+ help="model name in modelscope")
+ parser.add_argument("--model_revision",
+ type=str,
+ default="v1.0.3",
+ help="model revision in modelscope")
+ args = parser.parse_args()
+
+ inference_pipeline = pipeline(
+ task=Tasks.auto_speech_recognition,
+ model='damo/{}'.format(args.model_name),
+ model_revision=args.model_revision)
diff --git a/egs_modelscope/aishell2/paraformer/modelscope_utils/modelscope_infer.sh b/egs_modelscope/aishell2/paraformer/modelscope_utils/modelscope_infer.sh
new file mode 100755
index 0000000..a0c606f
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/modelscope_utils/modelscope_infer.sh
@@ -0,0 +1,88 @@
+#!/usr/bin/env bash
+
+set -e
+set -u
+set -o pipefail
+
+data_dir=
+exp_dir=
+model_name=
+model_revision=
+inference_nj=32
+gpuid_list="0,1,2,3"
+njob=32
+gpu_inference=true
+
+test_sets="dev test"
+decode_cmd=utils/run.pl
+
+# LM configs
+use_lm=false
+beam_size=1
+lm_weight=0.0
+
+. utils/parse_options.sh
+
+if ${gpu_inference}; then
+ _ngpu=1
+else
+ _ngpu=0
+fi
+
+# download model from modelscope
+python modelscope_utils/download_model.py \
+ --model_name ${model_name} --model_revision ${model_revision}
+
+modelscope_dir=${HOME}/.cache/modelscope/hub/damo/${model_name}
+
+
+for dset in ${test_sets}; do
+ _dir=${exp_dir}/${model_name}/decode_asr/${dset}
+ _logdir=${_dir}/logdir
+ _data=${data_dir}/${dset}
+ if [ -d ${_dir} ]; then
+ echo "${_dir} is already exists. if you want to decode again, please delete ${_dir} first."
+ exit 1
+ else
+ mkdir -p "${_dir}"
+ mkdir -p "${_logdir}"
+ fi
+
+ if "${use_lm}"; then
+ cp ${modelscope_dir}/decoding.yaml ${modelscope_dir}/decoding.yaml.back
+ sed -i "s#beam_size: [0-9]*#beam_size: `echo $beam_size`#g" ${modelscope_dir}/decoding.yaml
+ sed -i "s#lm_weight: 0.[0-9]*#lm_weight: `echo $lm_weight`#g" ${modelscope_dir}/decoding.yaml
+ fi
+
+ for n in $(seq "${inference_nj}"); do
+ split_scps+=" ${_logdir}/keys.${n}.scp"
+ done
+ # shellcheck disable=SC2086
+ utils/split_scp.pl "${data_dir}/${dset}/wav.scp" ${split_scps}
+
+ echo "Decoding started... log: '${_logdir}/asr_inference.*.log'"
+ # shellcheck disable=SC2086
+ ${decode_cmd} --max-jobs-run "${inference_nj}" JOB=1:"${inference_nj}" "${_logdir}"/asr_inference.JOB.log \
+ python -m funasr.bin.modelscope_infer \
+ --model_name ${model_name} \
+ --model_revision ${model_revision} \
+ --wav_list ${_logdir}/keys.JOB.scp \
+ --output_file ${_logdir}/text.JOB \
+ --gpuid_list ${gpuid_list} \
+ --njob ${njob} \
+ --ngpu ${_ngpu} \
+
+ for i in $(seq ${inference_nj}); do
+ cat ${_logdir}/text.${i}
+ done | sort -k1 >${_dir}/text
+
+ python utils/proce_text.py ${_dir}/text ${_dir}/text.proc
+ python utils/proce_text.py ${_data}/text ${_data}/text.proc
+ python utils/compute_wer.py ${_data}/text.proc ${_dir}/text.proc ${_dir}/text.cer
+ tail -n 3 ${_dir}/text.cer > ${_dir}/text.cer.txt
+ cat ${_dir}/text.cer.txt
+done
+
+if "${use_lm}"; then
+ mv ${modelscope_dir}/decoding.yaml.back ${modelscope_dir}/decoding.yaml
+fi
diff --git a/egs_modelscope/aishell2/paraformer/modelscope_utils/update_config.py b/egs_modelscope/aishell2/paraformer/modelscope_utils/update_config.py
new file mode 100644
index 0000000..88466ed
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/modelscope_utils/update_config.py
@@ -0,0 +1,41 @@
+import yaml
+import argparse
+
+def update_dct(fin_configs, root):
+ if root == {}:
+ return {}
+ for root_key, root_value in root.items():
+ if not isinstance(root[root_key],dict):
+ fin_configs[root_key] = root[root_key]
+ else:
+ result = update_dct(fin_configs[root_key], root[root_key])
+ fin_configs[root_key] = result
+ return fin_configs
+
+if __name__ == '__main__':
+ parser = argparse.ArgumentParser(
+ description="update configs",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument("--modelscope_config",
+ type=str,
+ help="modelscope config file")
+ parser.add_argument("--finetune_config",
+ type=str,
+ help="finetune config file")
+ parser.add_argument("--output_config",
+ type=str,
+ help="output config file")
+ args = parser.parse_args()
+
+ with open(args.modelscope_config) as f:
+ modelscope_configs = yaml.safe_load(f)
+
+ with open(args.finetune_config) as f:
+ finetune_configs = yaml.safe_load(f)
+
+ # update configs, e.g., lr, batch_size, ...
+ modelscope_configs = update_dct(modelscope_configs, finetune_configs)
+
+ with open(args.output_config, "w") as f:
+ yaml.dump(modelscope_configs, f, indent=4)
diff --git a/egs_modelscope/aishell2/paraformer/paraformer_large_finetune.sh b/egs_modelscope/aishell2/paraformer/paraformer_large_finetune.sh
index b864370..1c02267 100755
--- a/egs_modelscope/aishell2/paraformer/paraformer_large_finetune.sh
+++ b/egs_modelscope/aishell2/paraformer/paraformer_large_finetune.sh
@@ -9,6 +9,7 @@
gpu_inference=true # Whether to perform gpu decoding, set false for cpu decoding
njob=4 # the number of jobs for each gpu
train_cmd=utils/run.pl
+infer_cmd=utils/run.pl
# general configuration
feats_dir="../DATA" #feature output dictionary, for large data
@@ -32,7 +33,7 @@
lfr_n=6
init_model_name=speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-pytorch # pre-trained model, download from modelscope during fine-tuning
-model_revision="v1.0.3" # please do not modify the model revision
+model_revision="v1.0.4" # please do not modify the model revision
cmvn_file=init_model/${init_model_name}/am.mvn
seg_file=init_model/${init_model_name}/seg_dict
vocab=init_model/${init_model_name}/tokens.txt
@@ -82,7 +83,14 @@
# you can set gpu num for decoding here
gpuid_list=$CUDA_VISIBLE_DEVICES # set gpus for decoding, the same as training stage by default
ngpu=$(echo $gpuid_list | awk -F "," '{print NF}')
-inference_nj=$[${ngpu}*${njob}]
+
+if ${gpu_inference}; then
+ inference_nj=$[${ngpu}*${njob}]
+ _ngpu=1
+else
+ inference_nj=$njob
+ _ngpu=0
+fi
if [ ${stage} -le 0 ] && [ ${stop_stage} -ge 0 ]; then
echo "stage 0: Data preparation"
@@ -106,7 +114,7 @@
feat_train_dir=${feats_dir}/${dumpdir}/${train_set}; mkdir -p ${feat_train_dir}
feat_dev_dir=${feats_dir}/${dumpdir}/${valid_set}; mkdir -p ${feat_dev_dir}
if [ ${stage} -le 1 ] && [ ${stop_stage} -ge 1 ]; then
- echo "Feature Generation"
+ echo "stage 1: Feature Generation"
# compute fbank features
fbankdir=${feats_dir}/fbank
utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj --speed_perturb ${speed_perturb} \
@@ -167,6 +175,7 @@
# Training Stage
world_size=$gpu_num # run on one machine
if [ ${stage} -le 3 ] && [ ${stop_stage} -ge 3 ]; then
+ echo "stage 3: Training"
# update asr train config.yaml
python modelscope_utils/update_config.py --modelscope_config init_model/${init_model_name}/finetune.yaml --finetune_config ${asr_config} --output_config init_model/${init_model_name}/asr_finetune_config.yaml
finetune_config=init_model/${init_model_name}/asr_finetune_config.yaml
@@ -216,25 +225,57 @@
# Testing Stage
if [ ${stage} -le 4 ] && [ ${stop_stage} -ge 4 ]; then
- ./utils/easy_asr_infer.sh \
- --lang zh \
- --datadir ${feats_dir} \
- --feats_type ${feats_type} \
- --feats_dim ${feats_dim} \
- --token_type ${token_type} \
- --gpu_inference ${gpu_inference} \
- --inference_config "${inference_config}" \
- --test_sets "${test_sets}" \
- --token_list $token_list \
- --asr_exp ${exp_dir}/exp/${model_dir} \
- --stage 12 \
- --stop_stage 12 \
- --scp $scp \
- --text text \
- --inference_nj $inference_nj \
- --njob $njob \
- --inference_asr_model $inference_asr_model \
- --gpuid_list $gpuid_list \
- --mode paraformer
+ echo "stage 4: Inference"
+ for dset in ${test_sets}; do
+ asr_exp=${exp_dir}/exp/${model_dir}
+ inference_tag="$(basename "${inference_config}" .yaml)"
+ _dir="${asr_exp}/${inference_tag}/${inference_asr_model}/${dset}"
+ _logdir="${_dir}/logdir"
+ if [ -d ${_dir} ]; then
+ echo "${_dir} is already exists. if you want to decode again, please delete this dir first."
+ exit 0
+ fi
+ mkdir -p "${_logdir}"
+ _data="${feats_dir}/${dumpdir}/${dset}"
+ key_file=${_data}/${scp}
+ num_scp_file="$(<${key_file} wc -l)"
+ _nj=$([ $inference_nj -le $num_scp_file ] && echo "$inference_nj" || echo "$num_scp_file")
+ split_scps=
+ for n in $(seq "${_nj}"); do
+ split_scps+=" ${_logdir}/keys.${n}.scp"
+ done
+ # shellcheck disable=SC2086
+ utils/split_scp.pl "${key_file}" ${split_scps}
+ _opts=
+ if [ -n "${inference_config}" ]; then
+ _opts+="--config ${inference_config} "
+ fi
+ ${infer_cmd} --gpu "${_ngpu}" --max-jobs-run "${_nj}" JOB=1:"${_nj}" "${_logdir}"/asr_inference.JOB.log \
+ python -m funasr.bin.asr_inference_launch \
+ --batch_size 1 \
+ --ngpu "${_ngpu}" \
+ --njob ${njob} \
+ --gpuid_list ${gpuid_list} \
+ --data_path_and_name_and_type "${_data}/${scp},speech,${type}" \
+ --key_file "${_logdir}"/keys.JOB.scp \
+ --asr_train_config "${asr_exp}"/config.yaml \
+ --asr_model_file "${asr_exp}"/"${inference_asr_model}" \
+ --output_dir "${_logdir}"/output.JOB \
+ --mode paraformer \
+ ${_opts}
+
+ for f in token token_int score text; do
+ if [ -f "${_logdir}/output.1/1best_recog/${f}" ]; then
+ for i in $(seq "${_nj}"); do
+ cat "${_logdir}/output.${i}/1best_recog/${f}"
+ done | sort -k1 >"${_dir}/${f}"
+ fi
+ done
+ python utils/proce_text.py ${_dir}/text ${_dir}/text.proc
+ python utils/proce_text.py ${_data}/text ${_data}/text.proc
+ python utils/compute_wer.py ${_data}/text.proc ${_dir}/text.proc ${_dir}/text.cer
+ tail -n 3 ${_dir}/text.cer > ${_dir}/text.cer.txt
+ cat ${_dir}/text.cer.txt
+ done
fi
diff --git a/egs_modelscope/aishell2/paraformer/paraformer_large_infer.sh b/egs_modelscope/aishell2/paraformer/paraformer_large_infer.sh
index 2d5afd8..462dd6c 100755
--- a/egs_modelscope/aishell2/paraformer/paraformer_large_infer.sh
+++ b/egs_modelscope/aishell2/paraformer/paraformer_large_infer.sh
@@ -8,7 +8,7 @@
data_dir=
exp_dir=
model_name=speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-pytorch
-model_revision="v1.0.3" # please do not modify the model revision
+model_revision="v1.0.4" # please do not modify the model revision
inference_nj=32
gpuid_list="0,1" # set gpus, e.g., gpuid_list="0,1"
ngpu=$(echo $gpuid_list | awk -F "," '{print NF}')
diff --git a/egs_modelscope/aishell2/paraformer/utils b/egs_modelscope/aishell2/paraformer/utils
deleted file mode 120000
index 37d9761..0000000
--- a/egs_modelscope/aishell2/paraformer/utils
+++ /dev/null
@@ -1 +0,0 @@
-../../../egs/aishell/tranformer/utils/
\ No newline at end of file
diff --git a/egs_modelscope/aishell2/paraformer/utils/__init__.py b/egs_modelscope/aishell2/paraformer/utils/__init__.py
new file mode 100644
index 0000000..e69de29
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/__init__.py
diff --git a/egs_modelscope/aishell2/paraformer/utils/apply_cmvn.py b/egs_modelscope/aishell2/paraformer/utils/apply_cmvn.py
new file mode 100755
index 0000000..b5c5086
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/apply_cmvn.py
@@ -0,0 +1,79 @@
+from kaldiio import ReadHelper
+from kaldiio import WriteHelper
+
+import argparse
+import json
+import math
+import numpy as np
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="apply cmvn",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--ark-file",
+ "-a",
+ default=False,
+ required=True,
+ type=str,
+ help="fbank ark file",
+ )
+ parser.add_argument(
+ "--cmvn-file",
+ "-c",
+ default=False,
+ required=True,
+ type=str,
+ help="cmvn file",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ ark_file = args.output_dir + "/feats." + str(args.ark_index) + ".ark"
+ scp_file = args.output_dir + "/feats." + str(args.ark_index) + ".scp"
+ ark_writer = WriteHelper('ark,scp:{},{}'.format(ark_file, scp_file))
+
+ with open(args.cmvn_file) as f:
+ cmvn_stats = json.load(f)
+
+ means = cmvn_stats['mean_stats']
+ vars = cmvn_stats['var_stats']
+ total_frames = cmvn_stats['total_frames']
+
+ for i in range(len(means)):
+ means[i] /= total_frames
+ vars[i] = vars[i] / total_frames - means[i] * means[i]
+ if vars[i] < 1.0e-20:
+ vars[i] = 1.0e-20
+ vars[i] = 1.0 / math.sqrt(vars[i])
+
+ with ReadHelper('ark:{}'.format(args.ark_file)) as ark_reader:
+ for key, mat in ark_reader:
+ mat = (mat - means) * vars
+ ark_writer(key, mat)
+
+
+if __name__ == '__main__':
+ main()
diff --git a/egs_modelscope/aishell2/paraformer/utils/apply_cmvn.sh b/egs_modelscope/aishell2/paraformer/utils/apply_cmvn.sh
new file mode 100755
index 0000000..f8fd1d1
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/apply_cmvn.sh
@@ -0,0 +1,29 @@
+#!/usr/bin/env bash
+
+. ./path.sh || exit 1;
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+fbankdir=$1
+cmvn_file=$2
+logdir=$3
+output_dir=$4
+
+dump_dir=${output_dir}/ark; mkdir -p ${dump_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/apply_cmvn.JOB.log \
+ python utils/apply_cmvn.py -a $fbankdir/ark/feats.JOB.ark \
+ -c $cmvn_file -i JOB -o ${dump_dir} \
+ || exit 1;
+
+for n in $(seq $nj); do
+ cat ${dump_dir}/feats.$n.scp || exit 1
+done > ${output_dir}/feats.scp || exit 1
+
+echo "$0: Succeeded apply cmvn"
diff --git a/egs_modelscope/aishell2/paraformer/utils/apply_lfr_and_cmvn.py b/egs_modelscope/aishell2/paraformer/utils/apply_lfr_and_cmvn.py
new file mode 100755
index 0000000..50d18d1
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/apply_lfr_and_cmvn.py
@@ -0,0 +1,143 @@
+from kaldiio import ReadHelper, WriteHelper
+
+import argparse
+import numpy as np
+
+
+def build_LFR_features(inputs, m=7, n=6):
+ LFR_inputs = []
+ T = inputs.shape[0]
+ T_lfr = int(np.ceil(T / n))
+ left_padding = np.tile(inputs[0], ((m - 1) // 2, 1))
+ inputs = np.vstack((left_padding, inputs))
+ T = T + (m - 1) // 2
+ for i in range(T_lfr):
+ if m <= T - i * n:
+ LFR_inputs.append(np.hstack(inputs[i * n:i * n + m]))
+ else:
+ num_padding = m - (T - i * n)
+ frame = np.hstack(inputs[i * n:])
+ for _ in range(num_padding):
+ frame = np.hstack((frame, inputs[-1]))
+ LFR_inputs.append(frame)
+ return np.vstack(LFR_inputs)
+
+
+def build_CMVN_features(inputs, mvn_file): # noqa
+ with open(mvn_file, 'r', encoding='utf-8') as f:
+ lines = f.readlines()
+
+ add_shift_list = []
+ rescale_list = []
+ for i in range(len(lines)):
+ line_item = lines[i].split()
+ if line_item[0] == '<AddShift>':
+ line_item = lines[i + 1].split()
+ if line_item[0] == '<LearnRateCoef>':
+ add_shift_line = line_item[3:(len(line_item) - 1)]
+ add_shift_list = list(add_shift_line)
+ continue
+ elif line_item[0] == '<Rescale>':
+ line_item = lines[i + 1].split()
+ if line_item[0] == '<LearnRateCoef>':
+ rescale_line = line_item[3:(len(line_item) - 1)]
+ rescale_list = list(rescale_line)
+ continue
+
+ for j in range(inputs.shape[0]):
+ for k in range(inputs.shape[1]):
+ add_shift_value = add_shift_list[k]
+ rescale_value = rescale_list[k]
+ inputs[j, k] = float(inputs[j, k]) + float(add_shift_value)
+ inputs[j, k] = float(inputs[j, k]) * float(rescale_value)
+
+ return inputs
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="apply low_frame_rate and cmvn",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--ark-file",
+ "-a",
+ default=False,
+ required=True,
+ type=str,
+ help="fbank ark file",
+ )
+ parser.add_argument(
+ "--lfr",
+ "-f",
+ default=True,
+ type=str,
+ help="low frame rate",
+ )
+ parser.add_argument(
+ "--lfr-m",
+ "-m",
+ default=7,
+ type=int,
+ help="number of frames to stack",
+ )
+ parser.add_argument(
+ "--lfr-n",
+ "-n",
+ default=6,
+ type=int,
+ help="number of frames to skip",
+ )
+ parser.add_argument(
+ "--cmvn-file",
+ "-c",
+ default=False,
+ required=True,
+ type=str,
+ help="global cmvn file",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ dump_ark_file = args.output_dir + "/feats." + str(args.ark_index) + ".ark"
+ dump_scp_file = args.output_dir + "/feats." + str(args.ark_index) + ".scp"
+ shape_file = args.output_dir + "/len." + str(args.ark_index)
+ ark_writer = WriteHelper('ark,scp:{},{}'.format(dump_ark_file, dump_scp_file))
+
+ shape_writer = open(shape_file, 'w')
+ with ReadHelper('ark:{}'.format(args.ark_file)) as ark_reader:
+ for key, mat in ark_reader:
+ if args.lfr:
+ lfr = build_LFR_features(mat, args.lfr_m, args.lfr_n)
+ else:
+ lfr = mat
+ cmvn = build_CMVN_features(lfr, args.cmvn_file)
+ dims = cmvn.shape[1]
+ lens = cmvn.shape[0]
+ shape_writer.write(key + " " + str(lens) + "," + str(dims) + '\n')
+ ark_writer(key, cmvn)
+
+
+if __name__ == '__main__':
+ main()
+
diff --git a/egs_modelscope/aishell2/paraformer/utils/apply_lfr_and_cmvn.sh b/egs_modelscope/aishell2/paraformer/utils/apply_lfr_and_cmvn.sh
new file mode 100755
index 0000000..3119fdb
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/apply_lfr_and_cmvn.sh
@@ -0,0 +1,38 @@
+#!/usr/bin/env bash
+
+
+# Begin configuration section.
+nj=32
+cmd=utils/run.pl
+
+# feature configuration
+lfr=True
+lfr_m=7
+lfr_n=6
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+fbankdir=$1
+cmvn_file=$2
+logdir=$3
+output_dir=$4
+
+dump_dir=${output_dir}/ark; mkdir -p ${dump_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/apply_lfr_and_cmvn.JOB.log \
+ python utils/apply_lfr_and_cmvn.py -a $fbankdir/ark/feats.JOB.ark \
+ -f $lfr -m $lfr_m -n $lfr_n -c $cmvn_file -i JOB -o ${dump_dir} \
+ || exit 1;
+
+for n in $(seq $nj); do
+ cat ${dump_dir}/feats.$n.scp || exit 1
+done > ${output_dir}/feats.scp || exit 1
+
+for n in $(seq $nj); do
+ cat ${dump_dir}/len.$n || exit 1
+done > ${output_dir}/speech_shape || exit 1
+
+echo "$0: Succeeded apply low frame rate and cmvn"
diff --git a/egs_modelscope/aishell2/paraformer/utils/combine_cmvn_file.py b/egs_modelscope/aishell2/paraformer/utils/combine_cmvn_file.py
new file mode 100755
index 0000000..b2974a4
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/combine_cmvn_file.py
@@ -0,0 +1,73 @@
+import argparse
+import json
+import numpy as np
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="combine cmvn file",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--dims",
+ "-d",
+ default=80,
+ type=int,
+ help="feature dims",
+ )
+ parser.add_argument(
+ "--cmvn-dir",
+ "-c",
+ default=False,
+ required=True,
+ type=str,
+ help="cmvn dir",
+ )
+
+ parser.add_argument(
+ "--nj",
+ "-n",
+ default=1,
+ required=True,
+ type=int,
+ help="num of cmvn file",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ total_means = np.zeros(args.dims)
+ total_vars = np.zeros(args.dims)
+ total_frames = 0
+
+ cmvn_file = args.output_dir + "/cmvn.json"
+
+ for i in range(1, args.nj+1):
+ with open(args.cmvn_dir + "/cmvn." + str(i) + ".json", "r") as fin:
+ cmvn_stats = json.load(fin)
+
+ total_means += np.array(cmvn_stats["mean_stats"])
+ total_vars += np.array(cmvn_stats["var_stats"])
+ total_frames += cmvn_stats["total_frames"]
+
+ cmvn_info = {
+ 'mean_stats': list(total_means.tolist()),
+ 'var_stats': list(total_vars.tolist()),
+ 'total_frames': total_frames
+ }
+ with open(cmvn_file, 'w') as fout:
+ fout.write(json.dumps(cmvn_info))
+
+
+if __name__ == '__main__':
+ main()
diff --git a/egs_modelscope/aishell2/paraformer/utils/compute_cmvn.py b/egs_modelscope/aishell2/paraformer/utils/compute_cmvn.py
new file mode 100755
index 0000000..2b96e26
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/compute_cmvn.py
@@ -0,0 +1,74 @@
+from kaldiio import ReadHelper
+
+import argparse
+import numpy as np
+import json
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="computer global cmvn",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--dims",
+ "-d",
+ default=80,
+ type=int,
+ help="feature dims",
+ )
+ parser.add_argument(
+ "--ark-file",
+ "-a",
+ default=False,
+ required=True,
+ type=str,
+ help="fbank ark file",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ ark_file = args.ark_file + "/feats." + str(args.ark_index) + ".ark"
+ cmvn_file = args.output_dir + "/cmvn." + str(args.ark_index) + ".json"
+
+ mean_stats = np.zeros(args.dims)
+ var_stats = np.zeros(args.dims)
+ total_frames = 0
+
+ with ReadHelper('ark:{}'.format(ark_file)) as ark_reader:
+ for key, mat in ark_reader:
+ mean_stats += np.sum(mat, axis=0)
+ var_stats += np.sum(np.square(mat), axis=0)
+ total_frames += mat.shape[0]
+
+ cmvn_info = {
+ 'mean_stats': list(mean_stats.tolist()),
+ 'var_stats': list(var_stats.tolist()),
+ 'total_frames': total_frames
+ }
+ with open(cmvn_file, 'w') as fout:
+ fout.write(json.dumps(cmvn_info))
+
+
+if __name__ == '__main__':
+ main()
diff --git a/egs_modelscope/aishell2/paraformer/utils/compute_cmvn.sh b/egs_modelscope/aishell2/paraformer/utils/compute_cmvn.sh
new file mode 100755
index 0000000..12173ee
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/compute_cmvn.sh
@@ -0,0 +1,25 @@
+#!/usr/bin/env bash
+
+. ./path.sh || exit 1;
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+feats_dim=80
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+fbankdir=$1
+logdir=$2
+
+output_dir=${fbankdir}/cmvn; mkdir -p ${output_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/cmvn.JOB.log \
+ python utils/compute_cmvn.py -d ${feats_dim} -a $fbankdir/ark -i JOB -o ${output_dir} \
+ || exit 1;
+
+python utils/combine_cmvn_file.py -d ${feats_dim} -c ${output_dir} -n $nj -o $fbankdir
+
+echo "$0: Succeeded compute global cmvn"
diff --git a/egs_modelscope/aishell2/paraformer/utils/compute_fbank.py b/egs_modelscope/aishell2/paraformer/utils/compute_fbank.py
new file mode 100755
index 0000000..d03b5a8
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/compute_fbank.py
@@ -0,0 +1,153 @@
+from kaldiio import WriteHelper
+
+import argparse
+import numpy as np
+import json
+import torch
+import torchaudio
+import torchaudio.compliance.kaldi as kaldi
+
+
+def compute_fbank(wav_file,
+ num_mel_bins=80,
+ frame_length=25,
+ frame_shift=10,
+ dither=0.0,
+ resample_rate=16000,
+ speed=1.0):
+
+ waveform, sample_rate = torchaudio.load(wav_file)
+ if resample_rate != sample_rate:
+ waveform = torchaudio.transforms.Resample(orig_freq=sample_rate,
+ new_freq=resample_rate)(waveform)
+ if speed != 1.0:
+ waveform, _ = torchaudio.sox_effects.apply_effects_tensor(
+ waveform, resample_rate,
+ [['speed', str(speed)], ['rate', str(resample_rate)]]
+ )
+
+ waveform = waveform * (1 << 15)
+ mat = kaldi.fbank(waveform,
+ num_mel_bins=num_mel_bins,
+ frame_length=frame_length,
+ frame_shift=frame_shift,
+ dither=dither,
+ energy_floor=0.0,
+ window_type='hamming',
+ sample_frequency=resample_rate)
+
+ return mat.numpy()
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="computer features",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--wav-lists",
+ "-w",
+ default=False,
+ required=True,
+ type=str,
+ help="input wav lists",
+ )
+ parser.add_argument(
+ "--text-files",
+ "-t",
+ default=False,
+ required=True,
+ type=str,
+ help="input text files",
+ )
+ parser.add_argument(
+ "--dims",
+ "-d",
+ default=80,
+ type=int,
+ help="feature dims",
+ )
+ parser.add_argument(
+ "--sample-frequency",
+ "-s",
+ default=16000,
+ type=int,
+ help="sample frequency",
+ )
+ parser.add_argument(
+ "--speed-perturb",
+ "-p",
+ default="1.0",
+ type=str,
+ help="speed perturb",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-a",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ ark_file = args.output_dir + "/ark/feats." + str(args.ark_index) + ".ark"
+ scp_file = args.output_dir + "/ark/feats." + str(args.ark_index) + ".scp"
+ text_file = args.output_dir + "/txt/text." + str(args.ark_index) + ".txt"
+ feats_shape_file = args.output_dir + "/ark/len." + str(args.ark_index)
+ text_shape_file = args.output_dir + "/txt/len." + str(args.ark_index)
+
+ ark_writer = WriteHelper('ark,scp:{},{}'.format(ark_file, scp_file))
+ text_writer = open(text_file, 'w')
+ feats_shape_writer = open(feats_shape_file, 'w')
+ text_shape_writer = open(text_shape_file, 'w')
+
+ speed_perturb_list = args.speed_perturb.split(',')
+
+ for speed in speed_perturb_list:
+ with open(args.wav_lists, 'r', encoding='utf-8') as wavfile:
+ with open(args.text_files, 'r', encoding='utf-8') as textfile:
+ for wav, text in zip(wavfile, textfile):
+ s_w = wav.strip().split()
+ wav_id = s_w[0]
+ wav_file = s_w[1]
+
+ s_t = text.strip().split()
+ text_id = s_t[0]
+ txt = s_t[1:]
+ fbank = compute_fbank(wav_file,
+ num_mel_bins=args.dims,
+ resample_rate=args.sample_frequency,
+ speed=float(speed)
+ )
+ feats_dims = fbank.shape[1]
+ feats_lens = fbank.shape[0]
+ txt_lens = len(txt)
+ if speed == "1.0":
+ wav_id_sp = wav_id
+ else:
+ wav_id_sp = wav_id + "_sp" + speed
+
+ feats_shape_writer.write(wav_id_sp + " " + str(feats_lens) + "," + str(feats_dims) + '\n')
+ text_shape_writer.write(wav_id_sp + " " + str(txt_lens) + '\n')
+
+ text_writer.write(wav_id_sp + " " + " ".join(txt) + '\n')
+ ark_writer(wav_id_sp, fbank)
+
+
+if __name__ == '__main__':
+ main()
+
diff --git a/egs_modelscope/aishell2/paraformer/utils/compute_fbank.sh b/egs_modelscope/aishell2/paraformer/utils/compute_fbank.sh
new file mode 100755
index 0000000..92a4fe6
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/compute_fbank.sh
@@ -0,0 +1,51 @@
+#!/usr/bin/env bash
+
+. ./path.sh || exit 1;
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+
+# feature configuration
+feats_dim=80
+sample_frequency=16000
+speed_perturb="1.0"
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+data=$1
+logdir=$2
+fbankdir=$3
+
+[ ! -f $data/wav.scp ] && echo "$0: no such file $data/wav.scp" && exit 1;
+[ ! -f $data/text ] && echo "$0: no such file $data/text" && exit 1;
+
+python utils/split_data.py $data $data $nj
+
+ark_dir=${fbankdir}/ark; mkdir -p ${ark_dir}
+text_dir=${fbankdir}/txt; mkdir -p ${text_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/make_fbank.JOB.log \
+ python utils/compute_fbank.py -w $data/split${nj}/JOB/wav.scp -t $data/split${nj}/JOB/text \
+ -d $feats_dim -s $sample_frequency -p ${speed_perturb} -a JOB -o ${fbankdir} \
+ || exit 1;
+
+for n in $(seq $nj); do
+ cat ${ark_dir}/feats.$n.scp || exit 1
+done > $fbankdir/feats.scp || exit 1
+
+for n in $(seq $nj); do
+ cat ${text_dir}/text.$n.txt || exit 1
+done > $fbankdir/text || exit 1
+
+for n in $(seq $nj); do
+ cat ${ark_dir}/len.$n || exit 1
+done > $fbankdir/speech_shape || exit 1
+
+for n in $(seq $nj); do
+ cat ${text_dir}/len.$n || exit 1
+done > $fbankdir/text_shape || exit 1
+
+echo "$0: Succeeded compute FBANK features"
diff --git a/egs_modelscope/aishell2/paraformer/utils/compute_wer.py b/egs_modelscope/aishell2/paraformer/utils/compute_wer.py
new file mode 100755
index 0000000..349a3f6
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/compute_wer.py
@@ -0,0 +1,157 @@
+import os
+import numpy as np
+import sys
+
+def compute_wer(ref_file,
+ hyp_file,
+ cer_detail_file):
+ rst = {
+ 'Wrd': 0,
+ 'Corr': 0,
+ 'Ins': 0,
+ 'Del': 0,
+ 'Sub': 0,
+ 'Snt': 0,
+ 'Err': 0.0,
+ 'S.Err': 0.0,
+ 'wrong_words': 0,
+ 'wrong_sentences': 0
+ }
+
+ hyp_dict = {}
+ ref_dict = {}
+ with open(hyp_file, 'r') as hyp_reader:
+ for line in hyp_reader:
+ key = line.strip().split()[0]
+ value = line.strip().split()[1:]
+ hyp_dict[key] = value
+ with open(ref_file, 'r') as ref_reader:
+ for line in ref_reader:
+ key = line.strip().split()[0]
+ value = line.strip().split()[1:]
+ ref_dict[key] = value
+
+ cer_detail_writer = open(cer_detail_file, 'w')
+ for hyp_key in hyp_dict:
+ if hyp_key in ref_dict:
+ out_item = compute_wer_by_line(hyp_dict[hyp_key], ref_dict[hyp_key])
+ rst['Wrd'] += out_item['nwords']
+ rst['Corr'] += out_item['cor']
+ rst['wrong_words'] += out_item['wrong']
+ rst['Ins'] += out_item['ins']
+ rst['Del'] += out_item['del']
+ rst['Sub'] += out_item['sub']
+ rst['Snt'] += 1
+ if out_item['wrong'] > 0:
+ rst['wrong_sentences'] += 1
+ cer_detail_writer.write(hyp_key + print_cer_detail(out_item) + '\n')
+ cer_detail_writer.write("ref:" + '\t' + "".join(ref_dict[hyp_key]) + '\n')
+ cer_detail_writer.write("hyp:" + '\t' + "".join(hyp_dict[hyp_key]) + '\n')
+
+ if rst['Wrd'] > 0:
+ rst['Err'] = round(rst['wrong_words'] * 100 / rst['Wrd'], 2)
+ if rst['Snt'] > 0:
+ rst['S.Err'] = round(rst['wrong_sentences'] * 100 / rst['Snt'], 2)
+
+ cer_detail_writer.write('\n')
+ cer_detail_writer.write("%WER " + str(rst['Err']) + " [ " + str(rst['wrong_words'])+ " / " + str(rst['Wrd']) +
+ ", " + str(rst['Ins']) + " ins, " + str(rst['Del']) + " del, " + str(rst['Sub']) + " sub ]" + '\n')
+ cer_detail_writer.write("%SER " + str(rst['S.Err']) + " [ " + str(rst['wrong_sentences']) + " / " + str(rst['Snt']) + " ]" + '\n')
+ cer_detail_writer.write("Scored " + str(len(hyp_dict)) + " sentences, " + str(len(hyp_dict) - rst['Snt']) + " not present in hyp." + '\n')
+
+
+def compute_wer_by_line(hyp,
+ ref):
+ hyp = list(map(lambda x: x.lower(), hyp))
+ ref = list(map(lambda x: x.lower(), ref))
+
+ len_hyp = len(hyp)
+ len_ref = len(ref)
+
+ cost_matrix = np.zeros((len_hyp + 1, len_ref + 1), dtype=np.int16)
+
+ ops_matrix = np.zeros((len_hyp + 1, len_ref + 1), dtype=np.int8)
+
+ for i in range(len_hyp + 1):
+ cost_matrix[i][0] = i
+ for j in range(len_ref + 1):
+ cost_matrix[0][j] = j
+
+ for i in range(1, len_hyp + 1):
+ for j in range(1, len_ref + 1):
+ if hyp[i - 1] == ref[j - 1]:
+ cost_matrix[i][j] = cost_matrix[i - 1][j - 1]
+ else:
+ substitution = cost_matrix[i - 1][j - 1] + 1
+ insertion = cost_matrix[i - 1][j] + 1
+ deletion = cost_matrix[i][j - 1] + 1
+
+ compare_val = [substitution, insertion, deletion]
+
+ min_val = min(compare_val)
+ operation_idx = compare_val.index(min_val) + 1
+ cost_matrix[i][j] = min_val
+ ops_matrix[i][j] = operation_idx
+
+ match_idx = []
+ i = len_hyp
+ j = len_ref
+ rst = {
+ 'nwords': len_ref,
+ 'cor': 0,
+ 'wrong': 0,
+ 'ins': 0,
+ 'del': 0,
+ 'sub': 0
+ }
+ while i >= 0 or j >= 0:
+ i_idx = max(0, i)
+ j_idx = max(0, j)
+
+ if ops_matrix[i_idx][j_idx] == 0: # correct
+ if i - 1 >= 0 and j - 1 >= 0:
+ match_idx.append((j - 1, i - 1))
+ rst['cor'] += 1
+
+ i -= 1
+ j -= 1
+
+ elif ops_matrix[i_idx][j_idx] == 2: # insert
+ i -= 1
+ rst['ins'] += 1
+
+ elif ops_matrix[i_idx][j_idx] == 3: # delete
+ j -= 1
+ rst['del'] += 1
+
+ elif ops_matrix[i_idx][j_idx] == 1: # substitute
+ i -= 1
+ j -= 1
+ rst['sub'] += 1
+
+ if i < 0 and j >= 0:
+ rst['del'] += 1
+ elif j < 0 and i >= 0:
+ rst['ins'] += 1
+
+ match_idx.reverse()
+ wrong_cnt = cost_matrix[len_hyp][len_ref]
+ rst['wrong'] = wrong_cnt
+
+ return rst
+
+def print_cer_detail(rst):
+ return ("(" + "nwords=" + str(rst['nwords']) + ",cor=" + str(rst['cor'])
+ + ",ins=" + str(rst['ins']) + ",del=" + str(rst['del']) + ",sub="
+ + str(rst['sub']) + ") corr:" + '{:.2%}'.format(rst['cor']/rst['nwords'])
+ + ",cer:" + '{:.2%}'.format(rst['wrong']/rst['nwords']))
+
+if __name__ == '__main__':
+ if len(sys.argv) != 4:
+ print("usage : python compute-wer.py test.ref test.hyp test.wer")
+ sys.exit(0)
+
+ ref_file = sys.argv[1]
+ hyp_file = sys.argv[2]
+ cer_detail_file = sys.argv[3]
+ compute_wer(ref_file, hyp_file, cer_detail_file)
diff --git a/egs_modelscope/aishell2/paraformer/utils/error_rate_zh b/egs_modelscope/aishell2/paraformer/utils/error_rate_zh
new file mode 100755
index 0000000..6871a07
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/error_rate_zh
@@ -0,0 +1,370 @@
+#!/usr/bin/env python3
+# coding=utf8
+
+# Copyright 2021 Jiayu DU
+
+import sys
+import argparse
+import json
+import logging
+logging.basicConfig(stream=sys.stderr, level=logging.INFO, format='[%(levelname)s] %(message)s')
+
+DEBUG = None
+
+def GetEditType(ref_token, hyp_token):
+ if ref_token == None and hyp_token != None:
+ return 'I'
+ elif ref_token != None and hyp_token == None:
+ return 'D'
+ elif ref_token == hyp_token:
+ return 'C'
+ elif ref_token != hyp_token:
+ return 'S'
+ else:
+ raise RuntimeError
+
+class AlignmentArc:
+ def __init__(self, src, dst, ref, hyp):
+ self.src = src
+ self.dst = dst
+ self.ref = ref
+ self.hyp = hyp
+ self.edit_type = GetEditType(ref, hyp)
+
+def similarity_score_function(ref_token, hyp_token):
+ return 0 if (ref_token == hyp_token) else -1.0
+
+def insertion_score_function(token):
+ return -1.0
+
+def deletion_score_function(token):
+ return -1.0
+
+def EditDistance(
+ ref,
+ hyp,
+ similarity_score_function = similarity_score_function,
+ insertion_score_function = insertion_score_function,
+ deletion_score_function = deletion_score_function):
+ assert(len(ref) != 0)
+ class DPState:
+ def __init__(self):
+ self.score = -float('inf')
+ # backpointer
+ self.prev_r = None
+ self.prev_h = None
+
+ def print_search_grid(S, R, H, fstream):
+ print(file=fstream)
+ for r in range(R):
+ for h in range(H):
+ print(F'[{r},{h}]:{S[r][h].score:4.3f}:({S[r][h].prev_r},{S[r][h].prev_h}) ', end='', file=fstream)
+ print(file=fstream)
+
+ R = len(ref) + 1
+ H = len(hyp) + 1
+
+ # Construct DP search space, a (R x H) grid
+ S = [ [] for r in range(R) ]
+ for r in range(R):
+ S[r] = [ DPState() for x in range(H) ]
+
+ # initialize DP search grid origin, S(r = 0, h = 0)
+ S[0][0].score = 0.0
+ S[0][0].prev_r = None
+ S[0][0].prev_h = None
+
+ # initialize REF axis
+ for r in range(1, R):
+ S[r][0].score = S[r-1][0].score + deletion_score_function(ref[r-1])
+ S[r][0].prev_r = r-1
+ S[r][0].prev_h = 0
+
+ # initialize HYP axis
+ for h in range(1, H):
+ S[0][h].score = S[0][h-1].score + insertion_score_function(hyp[h-1])
+ S[0][h].prev_r = 0
+ S[0][h].prev_h = h-1
+
+ best_score = S[0][0].score
+ best_state = (0, 0)
+
+ for r in range(1, R):
+ for h in range(1, H):
+ sub_or_cor_score = similarity_score_function(ref[r-1], hyp[h-1])
+ new_score = S[r-1][h-1].score + sub_or_cor_score
+ if new_score >= S[r][h].score:
+ S[r][h].score = new_score
+ S[r][h].prev_r = r-1
+ S[r][h].prev_h = h-1
+
+ del_score = deletion_score_function(ref[r-1])
+ new_score = S[r-1][h].score + del_score
+ if new_score >= S[r][h].score:
+ S[r][h].score = new_score
+ S[r][h].prev_r = r - 1
+ S[r][h].prev_h = h
+
+ ins_score = insertion_score_function(hyp[h-1])
+ new_score = S[r][h-1].score + ins_score
+ if new_score >= S[r][h].score:
+ S[r][h].score = new_score
+ S[r][h].prev_r = r
+ S[r][h].prev_h = h-1
+
+ best_score = S[R-1][H-1].score
+ best_state = (R-1, H-1)
+
+ if DEBUG:
+ print_search_grid(S, R, H, sys.stderr)
+
+ # Backtracing best alignment path, i.e. a list of arcs
+ # arc = (src, dst, ref, hyp, edit_type)
+ # src/dst = (r, h), where r/h refers to search grid state-id along Ref/Hyp axis
+ best_path = []
+ r, h = best_state[0], best_state[1]
+ prev_r, prev_h = S[r][h].prev_r, S[r][h].prev_h
+ score = S[r][h].score
+ # loop invariant:
+ # 1. (prev_r, prev_h) -> (r, h) is a "forward arc" on best alignment path
+ # 2. score is the value of point(r, h) on DP search grid
+ while prev_r != None or prev_h != None:
+ src = (prev_r, prev_h)
+ dst = (r, h)
+ if (r == prev_r + 1 and h == prev_h + 1): # Substitution or correct
+ arc = AlignmentArc(src, dst, ref[prev_r], hyp[prev_h])
+ elif (r == prev_r + 1 and h == prev_h): # Deletion
+ arc = AlignmentArc(src, dst, ref[prev_r], None)
+ elif (r == prev_r and h == prev_h + 1): # Insertion
+ arc = AlignmentArc(src, dst, None, hyp[prev_h])
+ else:
+ raise RuntimeError
+ best_path.append(arc)
+ r, h = prev_r, prev_h
+ prev_r, prev_h = S[r][h].prev_r, S[r][h].prev_h
+ score = S[r][h].score
+
+ best_path.reverse()
+ return (best_path, best_score)
+
+def PrettyPrintAlignment(alignment, stream = sys.stderr):
+ def get_token_str(token):
+ if token == None:
+ return "*"
+ return token
+
+ def is_double_width_char(ch):
+ if (ch >= '\u4e00') and (ch <= '\u9fa5'): # codepoint ranges for Chinese chars
+ return True
+ # TODO: support other double-width-char language such as Japanese, Korean
+ else:
+ return False
+
+ def display_width(token_str):
+ m = 0
+ for c in token_str:
+ if is_double_width_char(c):
+ m += 2
+ else:
+ m += 1
+ return m
+
+ R = ' REF : '
+ H = ' HYP : '
+ E = ' EDIT : '
+ for arc in alignment:
+ r = get_token_str(arc.ref)
+ h = get_token_str(arc.hyp)
+ e = arc.edit_type if arc.edit_type != 'C' else ''
+
+ nr, nh, ne = display_width(r), display_width(h), display_width(e)
+ n = max(nr, nh, ne) + 1
+
+ R += r + ' ' * (n-nr)
+ H += h + ' ' * (n-nh)
+ E += e + ' ' * (n-ne)
+
+ print(R, file=stream)
+ print(H, file=stream)
+ print(E, file=stream)
+
+def CountEdits(alignment):
+ c, s, i, d = 0, 0, 0, 0
+ for arc in alignment:
+ if arc.edit_type == 'C':
+ c += 1
+ elif arc.edit_type == 'S':
+ s += 1
+ elif arc.edit_type == 'I':
+ i += 1
+ elif arc.edit_type == 'D':
+ d += 1
+ else:
+ raise RuntimeError
+ return (c, s, i, d)
+
+def ComputeTokenErrorRate(c, s, i, d):
+ return 100.0 * (s + d + i) / (s + d + c)
+
+def ComputeSentenceErrorRate(num_err_utts, num_utts):
+ assert(num_utts != 0)
+ return 100.0 * num_err_utts / num_utts
+
+
+class EvaluationResult:
+ def __init__(self):
+ self.num_ref_utts = 0
+ self.num_hyp_utts = 0
+ self.num_eval_utts = 0 # seen in both ref & hyp
+ self.num_hyp_without_ref = 0
+
+ self.C = 0
+ self.S = 0
+ self.I = 0
+ self.D = 0
+ self.token_error_rate = 0.0
+
+ self.num_utts_with_error = 0
+ self.sentence_error_rate = 0.0
+
+ def to_json(self):
+ return json.dumps(self.__dict__)
+
+ def to_kaldi(self):
+ info = (
+ F'%WER {self.token_error_rate:.2f} [ {self.S + self.D + self.I} / {self.C + self.S + self.D}, {self.I} ins, {self.D} del, {self.S} sub ]\n'
+ F'%SER {self.sentence_error_rate:.2f} [ {self.num_utts_with_error} / {self.num_eval_utts} ]\n'
+ )
+ return info
+
+ def to_sclite(self):
+ return "TODO"
+
+ def to_espnet(self):
+ return "TODO"
+
+ def to_summary(self):
+ #return json.dumps(self.__dict__, indent=4)
+ summary = (
+ '==================== Overall Statistics ====================\n'
+ F'num_ref_utts: {self.num_ref_utts}\n'
+ F'num_hyp_utts: {self.num_hyp_utts}\n'
+ F'num_hyp_without_ref: {self.num_hyp_without_ref}\n'
+ F'num_eval_utts: {self.num_eval_utts}\n'
+ F'sentence_error_rate: {self.sentence_error_rate:.2f}%\n'
+ F'token_error_rate: {self.token_error_rate:.2f}%\n'
+ F'token_stats:\n'
+ F' - tokens:{self.C + self.S + self.D:>7}\n'
+ F' - edits: {self.S + self.I + self.D:>7}\n'
+ F' - cor: {self.C:>7}\n'
+ F' - sub: {self.S:>7}\n'
+ F' - ins: {self.I:>7}\n'
+ F' - del: {self.D:>7}\n'
+ '============================================================\n'
+ )
+ return summary
+
+
+class Utterance:
+ def __init__(self, uid, text):
+ self.uid = uid
+ self.text = text
+
+
+def LoadUtterances(filepath, format):
+ utts = {}
+ if format == 'text': # utt_id word1 word2 ...
+ with open(filepath, 'r', encoding='utf8') as f:
+ for line in f:
+ line = line.strip()
+ if line:
+ cols = line.split(maxsplit=1)
+ assert(len(cols) == 2 or len(cols) == 1)
+ uid = cols[0]
+ text = cols[1] if len(cols) == 2 else ''
+ if utts.get(uid) != None:
+ raise RuntimeError(F'Found duplicated utterence id {uid}')
+ utts[uid] = Utterance(uid, text)
+ else:
+ raise RuntimeError(F'Unsupported text format {format}')
+ return utts
+
+
+def tokenize_text(text, tokenizer):
+ if tokenizer == 'whitespace':
+ return text.split()
+ elif tokenizer == 'char':
+ return [ ch for ch in ''.join(text.split()) ]
+ else:
+ raise RuntimeError(F'ERROR: Unsupported tokenizer {tokenizer}')
+
+
+if __name__ == '__main__':
+ parser = argparse.ArgumentParser()
+ # optional
+ parser.add_argument('--tokenizer', choices=['whitespace', 'char'], default='whitespace', help='whitespace for WER, char for CER')
+ parser.add_argument('--ref-format', choices=['text'], default='text', help='reference format, first col is utt_id, the rest is text')
+ parser.add_argument('--hyp-format', choices=['text'], default='text', help='hypothesis format, first col is utt_id, the rest is text')
+ # required
+ parser.add_argument('--ref', type=str, required=True, help='input reference file')
+ parser.add_argument('--hyp', type=str, required=True, help='input hypothesis file')
+
+ parser.add_argument('result_file', type=str)
+ args = parser.parse_args()
+ logging.info(args)
+
+ ref_utts = LoadUtterances(args.ref, args.ref_format)
+ hyp_utts = LoadUtterances(args.hyp, args.hyp_format)
+
+ r = EvaluationResult()
+
+ # check valid utterances in hyp that have matched non-empty reference
+ eval_utts = []
+ r.num_hyp_without_ref = 0
+ for uid in sorted(hyp_utts.keys()):
+ if uid in ref_utts.keys(): # TODO: efficiency
+ if ref_utts[uid].text.strip(): # non-empty reference
+ eval_utts.append(uid)
+ else:
+ logging.warn(F'Found {uid} with empty reference, skipping...')
+ else:
+ logging.warn(F'Found {uid} without reference, skipping...')
+ r.num_hyp_without_ref += 1
+
+ r.num_hyp_utts = len(hyp_utts)
+ r.num_ref_utts = len(ref_utts)
+ r.num_eval_utts = len(eval_utts)
+
+ with open(args.result_file, 'w+', encoding='utf8') as fo:
+ for uid in eval_utts:
+ ref = ref_utts[uid]
+ hyp = hyp_utts[uid]
+
+ alignment, score = EditDistance(
+ tokenize_text(ref.text, args.tokenizer),
+ tokenize_text(hyp.text, args.tokenizer)
+ )
+
+ c, s, i, d = CountEdits(alignment)
+ utt_ter = ComputeTokenErrorRate(c, s, i, d)
+
+ # utt-level evaluation result
+ print(F'{{"uid":{uid}, "score":{score}, "ter":{utt_ter:.2f}, "cor":{c}, "sub":{s}, "ins":{i}, "del":{d}}}', file=fo)
+ PrettyPrintAlignment(alignment, fo)
+
+ r.C += c
+ r.S += s
+ r.I += i
+ r.D += d
+
+ if utt_ter > 0:
+ r.num_utts_with_error += 1
+
+ # corpus level evaluation result
+ r.sentence_error_rate = ComputeSentenceErrorRate(r.num_utts_with_error, r.num_eval_utts)
+ r.token_error_rate = ComputeTokenErrorRate(r.C, r.S, r.I, r.D)
+
+ print(r.to_summary(), file=fo)
+
+ print(r.to_json())
+ print(r.to_kaldi())
diff --git a/egs_modelscope/aishell2/paraformer/utils/extract_embeds.py b/egs_modelscope/aishell2/paraformer/utils/extract_embeds.py
new file mode 100755
index 0000000..7b817d8
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/extract_embeds.py
@@ -0,0 +1,47 @@
+from transformers import AutoTokenizer, AutoModel, pipeline
+import numpy as np
+import sys
+import os
+import torch
+from kaldiio import WriteHelper
+import re
+text_file_json = sys.argv[1]
+out_ark = sys.argv[2]
+out_scp = sys.argv[3]
+out_shape = sys.argv[4]
+device = int(sys.argv[5])
+model_path = sys.argv[6]
+
+model = AutoModel.from_pretrained(model_path)
+tokenizer = AutoTokenizer.from_pretrained(model_path)
+extractor = pipeline(task="feature-extraction", model=model, tokenizer=tokenizer, device=device)
+
+with open(text_file_json, 'r') as f:
+ js = f.readlines()
+
+
+f_shape = open(out_shape, "w")
+with WriteHelper('ark,scp:{},{}'.format(out_ark, out_scp)) as writer:
+ with torch.no_grad():
+ for idx, line in enumerate(js):
+ id, tokens = line.strip().split(" ", 1)
+ tokens = re.sub(" ", "", tokens.strip())
+ tokens = ' '.join([j for j in tokens])
+ token_num = len(tokens.split(" "))
+ outputs = extractor(tokens)
+ outputs = np.array(outputs)
+ embeds = outputs[0, 1:-1, :]
+
+ token_num_embeds, dim = embeds.shape
+ if token_num == token_num_embeds:
+ writer(id, embeds)
+ shape_line = "{} {},{}\n".format(id, token_num_embeds, dim)
+ f_shape.write(shape_line)
+ else:
+ print("{}, size has changed, {}, {}, {}".format(id, token_num, token_num_embeds, tokens))
+
+
+
+f_shape.close()
+
+
diff --git a/egs_modelscope/aishell2/paraformer/utils/filter_scp.pl b/egs_modelscope/aishell2/paraformer/utils/filter_scp.pl
new file mode 100755
index 0000000..003530d
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/filter_scp.pl
@@ -0,0 +1,87 @@
+#!/usr/bin/env perl
+# Copyright 2010-2012 Microsoft Corporation
+# Johns Hopkins University (author: Daniel Povey)
+
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+# This script takes a list of utterance-ids or any file whose first field
+# of each line is an utterance-id, and filters an scp
+# file (or any file whose "n-th" field is an utterance id), printing
+# out only those lines whose "n-th" field is in id_list. The index of
+# the "n-th" field is 1, by default, but can be changed by using
+# the -f <n> switch
+
+$exclude = 0;
+$field = 1;
+$shifted = 0;
+
+do {
+ $shifted=0;
+ if ($ARGV[0] eq "--exclude") {
+ $exclude = 1;
+ shift @ARGV;
+ $shifted=1;
+ }
+ if ($ARGV[0] eq "-f") {
+ $field = $ARGV[1];
+ shift @ARGV; shift @ARGV;
+ $shifted=1
+ }
+} while ($shifted);
+
+if(@ARGV < 1 || @ARGV > 2) {
+ die "Usage: filter_scp.pl [--exclude] [-f <field-to-filter-on>] id_list [in.scp] > out.scp \n" .
+ "Prints only the input lines whose f'th field (default: first) is in 'id_list'.\n" .
+ "Note: only the first field of each line in id_list matters. With --exclude, prints\n" .
+ "only the lines that were *not* in id_list.\n" .
+ "Caution: previously, the -f option was interpreted as a zero-based field index.\n" .
+ "If your older scripts (written before Oct 2014) stopped working and you used the\n" .
+ "-f option, add 1 to the argument.\n" .
+ "See also: scripts/filter_scp.pl .\n";
+}
+
+
+$idlist = shift @ARGV;
+open(F, "<$idlist") || die "Could not open id-list file $idlist";
+while(<F>) {
+ @A = split;
+ @A>=1 || die "Invalid id-list file line $_";
+ $seen{$A[0]} = 1;
+}
+
+if ($field == 1) { # Treat this as special case, since it is common.
+ while(<>) {
+ $_ =~ m/\s*(\S+)\s*/ || die "Bad line $_, could not get first field.";
+ # $1 is what we filter on.
+ if ((!$exclude && $seen{$1}) || ($exclude && !defined $seen{$1})) {
+ print $_;
+ }
+ }
+} else {
+ while(<>) {
+ @A = split;
+ @A > 0 || die "Invalid scp file line $_";
+ @A >= $field || die "Invalid scp file line $_";
+ if ((!$exclude && $seen{$A[$field-1]}) || ($exclude && !defined $seen{$A[$field-1]})) {
+ print $_;
+ }
+ }
+}
+
+# tests:
+# the following should print "foo 1"
+# ( echo foo 1; echo bar 2 ) | scripts/filter_scp.pl <(echo foo)
+# the following should print "bar 2".
+# ( echo foo 1; echo bar 2 ) | scripts/filter_scp.pl -f 2 <(echo 2)
diff --git a/egs_modelscope/aishell2/paraformer/utils/fix_data.sh b/egs_modelscope/aishell2/paraformer/utils/fix_data.sh
new file mode 100755
index 0000000..32cdde5
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/fix_data.sh
@@ -0,0 +1,35 @@
+#!/usr/bin/env bash
+
+echo "$0 $@"
+data_dir=$1
+
+if [ ! -f ${data_dir}/wav.scp ]; then
+ echo "$0: wav.scp is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/text ]; then
+ echo "$0: text is not found"
+ exit 1;
+fi
+
+
+
+mkdir -p ${data_dir}/.backup
+
+awk '{print $1}' ${data_dir}/wav.scp > ${data_dir}/.backup/wav_id
+awk '{print $1}' ${data_dir}/text > ${data_dir}/.backup/text_id
+
+sort ${data_dir}/.backup/wav_id ${data_dir}/.backup/text_id | uniq -d > ${data_dir}/.backup/id
+
+cp ${data_dir}/wav.scp ${data_dir}/.backup/wav.scp
+cp ${data_dir}/text ${data_dir}/.backup/text
+
+mv ${data_dir}/wav.scp ${data_dir}/wav.scp.bak
+mv ${data_dir}/text ${data_dir}/text.bak
+
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/wav.scp.bak > ${data_dir}/wav.scp
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/text.bak > ${data_dir}/text
+
+rm ${data_dir}/wav.scp.bak
+rm ${data_dir}/text.bak
diff --git a/egs_modelscope/aishell2/paraformer/utils/fix_data_feat.sh b/egs_modelscope/aishell2/paraformer/utils/fix_data_feat.sh
new file mode 100755
index 0000000..2c92d7f
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/fix_data_feat.sh
@@ -0,0 +1,52 @@
+#!/usr/bin/env bash
+
+echo "$0 $@"
+data_dir=$1
+
+if [ ! -f ${data_dir}/feats.scp ]; then
+ echo "$0: feats.scp is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/text ]; then
+ echo "$0: text is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/speech_shape ]; then
+ echo "$0: feature lengths is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/text_shape ]; then
+ echo "$0: text lengths is not found"
+ exit 1;
+fi
+
+mkdir -p ${data_dir}/.backup
+
+awk '{print $1}' ${data_dir}/feats.scp > ${data_dir}/.backup/wav_id
+awk '{print $1}' ${data_dir}/text > ${data_dir}/.backup/text_id
+
+sort ${data_dir}/.backup/wav_id ${data_dir}/.backup/text_id | uniq -d > ${data_dir}/.backup/id
+
+cp ${data_dir}/feats.scp ${data_dir}/.backup/feats.scp
+cp ${data_dir}/text ${data_dir}/.backup/text
+cp ${data_dir}/speech_shape ${data_dir}/.backup/speech_shape
+cp ${data_dir}/text_shape ${data_dir}/.backup/text_shape
+
+mv ${data_dir}/feats.scp ${data_dir}/feats.scp.bak
+mv ${data_dir}/text ${data_dir}/text.bak
+mv ${data_dir}/speech_shape ${data_dir}/speech_shape.bak
+mv ${data_dir}/text_shape ${data_dir}/text_shape.bak
+
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/feats.scp.bak > ${data_dir}/feats.scp
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/text.bak > ${data_dir}/text
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/speech_shape.bak > ${data_dir}/speech_shape
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/text_shape.bak > ${data_dir}/text_shape
+
+rm ${data_dir}/feats.scp.bak
+rm ${data_dir}/text.bak
+rm ${data_dir}/speech_shape.bak
+rm ${data_dir}/text_shape.bak
+
diff --git a/egs_modelscope/aishell2/paraformer/utils/gen_ark_list.sh b/egs_modelscope/aishell2/paraformer/utils/gen_ark_list.sh
new file mode 100755
index 0000000..aebf356
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/gen_ark_list.sh
@@ -0,0 +1,22 @@
+#!/usr/bin/env bash
+
+
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+ark_dir=$1
+txt_dir=$2
+output_dir=$3
+
+[ ! -d ${ark_dir}/ark ] && echo "$0: ark data is required" && exit 1;
+[ ! -d ${txt_dir}/txt ] && echo "$0: txt data is required" && exit 1;
+
+for n in $(seq $nj); do
+ echo "${ark_dir}/ark/feats.$n.ark ${txt_dir}/txt/text.$n.txt" || exit 1
+done > ${output_dir}/ark_txt.scp || exit 1
+
diff --git a/egs_modelscope/aishell2/paraformer/utils/parse_options.sh b/egs_modelscope/aishell2/paraformer/utils/parse_options.sh
new file mode 100755
index 0000000..71fb9e5
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/parse_options.sh
@@ -0,0 +1,97 @@
+#!/usr/bin/env bash
+
+# Copyright 2012 Johns Hopkins University (Author: Daniel Povey);
+# Arnab Ghoshal, Karel Vesely
+
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+# Parse command-line options.
+# To be sourced by another script (as in ". parse_options.sh").
+# Option format is: --option-name arg
+# and shell variable "option_name" gets set to value "arg."
+# The exception is --help, which takes no arguments, but prints the
+# $help_message variable (if defined).
+
+
+###
+### The --config file options have lower priority to command line
+### options, so we need to import them first...
+###
+
+# Now import all the configs specified by command-line, in left-to-right order
+for ((argpos=1; argpos<$#; argpos++)); do
+ if [ "${!argpos}" == "--config" ]; then
+ argpos_plus1=$((argpos+1))
+ config=${!argpos_plus1}
+ [ ! -r $config ] && echo "$0: missing config '$config'" && exit 1
+ . $config # source the config file.
+ fi
+done
+
+
+###
+### Now we process the command line options
+###
+while true; do
+ [ -z "${1:-}" ] && break; # break if there are no arguments
+ case "$1" in
+ # If the enclosing script is called with --help option, print the help
+ # message and exit. Scripts should put help messages in $help_message
+ --help|-h) if [ -z "$help_message" ]; then echo "No help found." 1>&2;
+ else printf "$help_message\n" 1>&2 ; fi;
+ exit 0 ;;
+ --*=*) echo "$0: options to scripts must be of the form --name value, got '$1'"
+ exit 1 ;;
+ # If the first command-line argument begins with "--" (e.g. --foo-bar),
+ # then work out the variable name as $name, which will equal "foo_bar".
+ --*) name=`echo "$1" | sed s/^--// | sed s/-/_/g`;
+ # Next we test whether the variable in question is undefned-- if so it's
+ # an invalid option and we die. Note: $0 evaluates to the name of the
+ # enclosing script.
+ # The test [ -z ${foo_bar+xxx} ] will return true if the variable foo_bar
+ # is undefined. We then have to wrap this test inside "eval" because
+ # foo_bar is itself inside a variable ($name).
+ eval '[ -z "${'$name'+xxx}" ]' && echo "$0: invalid option $1" 1>&2 && exit 1;
+
+ oldval="`eval echo \\$$name`";
+ # Work out whether we seem to be expecting a Boolean argument.
+ if [ "$oldval" == "true" ] || [ "$oldval" == "false" ]; then
+ was_bool=true;
+ else
+ was_bool=false;
+ fi
+
+ # Set the variable to the right value-- the escaped quotes make it work if
+ # the option had spaces, like --cmd "queue.pl -sync y"
+ eval $name=\"$2\";
+
+ # Check that Boolean-valued arguments are really Boolean.
+ if $was_bool && [[ "$2" != "true" && "$2" != "false" ]]; then
+ echo "$0: expected \"true\" or \"false\": $1 $2" 1>&2
+ exit 1;
+ fi
+ shift 2;
+ ;;
+ *) break;
+ esac
+done
+
+
+# Check for an empty argument to the --cmd option, which can easily occur as a
+# result of scripting errors.
+[ ! -z "${cmd+xxx}" ] && [ -z "$cmd" ] && echo "$0: empty argument to --cmd option" 1>&2 && exit 1;
+
+
+true; # so this script returns exit code 0.
diff --git a/egs_modelscope/aishell2/paraformer/utils/print_args.py b/egs_modelscope/aishell2/paraformer/utils/print_args.py
new file mode 100755
index 0000000..b0c61e5
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/print_args.py
@@ -0,0 +1,45 @@
+#!/usr/bin/env python
+import sys
+
+
+def get_commandline_args(no_executable=True):
+ extra_chars = [
+ " ",
+ ";",
+ "&",
+ "|",
+ "<",
+ ">",
+ "?",
+ "*",
+ "~",
+ "`",
+ '"',
+ "'",
+ "\\",
+ "{",
+ "}",
+ "(",
+ ")",
+ ]
+
+ # Escape the extra characters for shell
+ argv = [
+ arg.replace("'", "'\\''")
+ if all(char not in arg for char in extra_chars)
+ else "'" + arg.replace("'", "'\\''") + "'"
+ for arg in sys.argv
+ ]
+
+ if no_executable:
+ return " ".join(argv[1:])
+ else:
+ return sys.executable + " " + " ".join(argv)
+
+
+def main():
+ print(get_commandline_args())
+
+
+if __name__ == "__main__":
+ main()
diff --git a/egs_modelscope/aishell2/paraformer/utils/proc_conf_oss.py b/egs_modelscope/aishell2/paraformer/utils/proc_conf_oss.py
new file mode 100755
index 0000000..c4a90c5
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/proc_conf_oss.py
@@ -0,0 +1,35 @@
+from pathlib import Path
+
+import torch
+import yaml
+
+
+class NoAliasSafeDumper(yaml.SafeDumper):
+ # Disable anchor/alias in yaml because looks ugly
+ def ignore_aliases(self, data):
+ return True
+
+
+def yaml_no_alias_safe_dump(data, stream=None, **kwargs):
+ """Safe-dump in yaml with no anchor/alias"""
+ return yaml.dump(
+ data, stream, allow_unicode=True, Dumper=NoAliasSafeDumper, **kwargs
+ )
+
+
+def gen_conf(file, out_dir):
+ conf = torch.load(file)["config"]
+ conf["oss_bucket"] = "null"
+ print(conf)
+ output_dir = Path(out_dir)
+ output_dir.mkdir(parents=True, exist_ok=True)
+ with (output_dir / "config.yaml").open("w", encoding="utf-8") as f:
+ yaml_no_alias_safe_dump(conf, f, indent=4, sort_keys=False)
+
+
+if __name__ == "__main__":
+ import sys
+
+ in_f = sys.argv[1]
+ out_f = sys.argv[2]
+ gen_conf(in_f, out_f)
diff --git a/egs_modelscope/aishell2/paraformer/utils/proce_text.py b/egs_modelscope/aishell2/paraformer/utils/proce_text.py
new file mode 100755
index 0000000..9e517a4
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/proce_text.py
@@ -0,0 +1,31 @@
+
+import sys
+import re
+
+in_f = sys.argv[1]
+out_f = sys.argv[2]
+
+
+with open(in_f, "r", encoding="utf-8") as f:
+ lines = f.readlines()
+
+with open(out_f, "w", encoding="utf-8") as f:
+ for line in lines:
+ outs = line.strip().split(" ", 1)
+ if len(outs) == 2:
+ idx, text = outs
+ text = re.sub("</s>", "", text)
+ text = re.sub("<s>", "", text)
+ text = re.sub("@@", "", text)
+ text = re.sub("@", "", text)
+ text = re.sub("<unk>", "", text)
+ text = re.sub(" ", "", text)
+ text = text.lower()
+ else:
+ idx = outs[0]
+ text = " "
+
+ text = [x for x in text]
+ text = " ".join(text)
+ out = "{} {}\n".format(idx, text)
+ f.write(out)
diff --git a/egs_modelscope/aishell2/paraformer/utils/run.pl b/egs_modelscope/aishell2/paraformer/utils/run.pl
new file mode 100755
index 0000000..483f95b
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/run.pl
@@ -0,0 +1,356 @@
+#!/usr/bin/env perl
+use warnings; #sed replacement for -w perl parameter
+# In general, doing
+# run.pl some.log a b c is like running the command a b c in
+# the bash shell, and putting the standard error and output into some.log.
+# To run parallel jobs (backgrounded on the host machine), you can do (e.g.)
+# run.pl JOB=1:4 some.JOB.log a b c JOB is like running the command a b c JOB
+# and putting it in some.JOB.log, for each one. [Note: JOB can be any identifier].
+# If any of the jobs fails, this script will fail.
+
+# A typical example is:
+# run.pl some.log my-prog "--opt=foo bar" foo \| other-prog baz
+# and run.pl will run something like:
+# ( my-prog '--opt=foo bar' foo | other-prog baz ) >& some.log
+#
+# Basically it takes the command-line arguments, quotes them
+# as necessary to preserve spaces, and evaluates them with bash.
+# In addition it puts the command line at the top of the log, and
+# the start and end times of the command at the beginning and end.
+# The reason why this is useful is so that we can create a different
+# version of this program that uses a queueing system instead.
+
+#use Data::Dumper;
+
+@ARGV < 2 && die "usage: run.pl log-file command-line arguments...";
+
+#print STDERR "COMMAND-LINE: " . Dumper(\@ARGV) . "\n";
+$job_pick = 'all';
+$max_jobs_run = -1;
+$jobstart = 1;
+$jobend = 1;
+$ignored_opts = ""; # These will be ignored.
+
+# First parse an option like JOB=1:4, and any
+# options that would normally be given to
+# queue.pl, which we will just discard.
+
+for (my $x = 1; $x <= 2; $x++) { # This for-loop is to
+ # allow the JOB=1:n option to be interleaved with the
+ # options to qsub.
+ while (@ARGV >= 2 && $ARGV[0] =~ m:^-:) {
+ # parse any options that would normally go to qsub, but which will be ignored here.
+ my $switch = shift @ARGV;
+ if ($switch eq "-V") {
+ $ignored_opts .= "-V ";
+ } elsif ($switch eq "--max-jobs-run" || $switch eq "-tc") {
+ # we do support the option --max-jobs-run n, and its GridEngine form -tc n.
+ # if the command appears multiple times uses the smallest option.
+ if ( $max_jobs_run <= 0 ) {
+ $max_jobs_run = shift @ARGV;
+ } else {
+ my $new_constraint = shift @ARGV;
+ if ( ($new_constraint < $max_jobs_run) ) {
+ $max_jobs_run = $new_constraint;
+ }
+ }
+
+ if (! ($max_jobs_run > 0)) {
+ die "run.pl: invalid option --max-jobs-run $max_jobs_run";
+ }
+ } else {
+ my $argument = shift @ARGV;
+ if ($argument =~ m/^--/) {
+ print STDERR "run.pl: WARNING: suspicious argument '$argument' to $switch; starts with '-'\n";
+ }
+ if ($switch eq "-sync" && $argument =~ m/^[yY]/) {
+ $ignored_opts .= "-sync "; # Note: in the
+ # corresponding code in queue.pl it says instead, just "$sync = 1;".
+ } elsif ($switch eq "-pe") { # e.g. -pe smp 5
+ my $argument2 = shift @ARGV;
+ $ignored_opts .= "$switch $argument $argument2 ";
+ } elsif ($switch eq "--gpu") {
+ $using_gpu = $argument;
+ } elsif ($switch eq "--pick") {
+ if($argument =~ m/^(all|failed|incomplete)$/) {
+ $job_pick = $argument;
+ } else {
+ print STDERR "run.pl: ERROR: --pick argument must be one of 'all', 'failed' or 'incomplete'"
+ }
+ } else {
+ # Ignore option.
+ $ignored_opts .= "$switch $argument ";
+ }
+ }
+ }
+ if ($ARGV[0] =~ m/^([\w_][\w\d_]*)+=(\d+):(\d+)$/) { # e.g. JOB=1:20
+ $jobname = $1;
+ $jobstart = $2;
+ $jobend = $3;
+ if ($jobstart > $jobend) {
+ die "run.pl: invalid job range $ARGV[0]";
+ }
+ if ($jobstart <= 0) {
+ die "run.pl: invalid job range $ARGV[0], start must be strictly positive (this is required for GridEngine compatibility).";
+ }
+ shift;
+ } elsif ($ARGV[0] =~ m/^([\w_][\w\d_]*)+=(\d+)$/) { # e.g. JOB=1.
+ $jobname = $1;
+ $jobstart = $2;
+ $jobend = $2;
+ shift;
+ } elsif ($ARGV[0] =~ m/.+\=.*\:.*$/) {
+ print STDERR "run.pl: Warning: suspicious first argument to run.pl: $ARGV[0]\n";
+ }
+}
+
+# Users found this message confusing so we are removing it.
+# if ($ignored_opts ne "") {
+# print STDERR "run.pl: Warning: ignoring options \"$ignored_opts\"\n";
+# }
+
+if ($max_jobs_run == -1) { # If --max-jobs-run option not set,
+ # then work out the number of processors if possible,
+ # and set it based on that.
+ $max_jobs_run = 0;
+ if ($using_gpu) {
+ if (open(P, "nvidia-smi -L |")) {
+ $max_jobs_run++ while (<P>);
+ close(P);
+ }
+ if ($max_jobs_run == 0) {
+ $max_jobs_run = 1;
+ print STDERR "run.pl: Warning: failed to detect number of GPUs from nvidia-smi, using ${max_jobs_run}\n";
+ }
+ } elsif (open(P, "</proc/cpuinfo")) { # Linux
+ while (<P>) { if (m/^processor/) { $max_jobs_run++; } }
+ if ($max_jobs_run == 0) {
+ print STDERR "run.pl: Warning: failed to detect any processors from /proc/cpuinfo\n";
+ $max_jobs_run = 10; # reasonable default.
+ }
+ close(P);
+ } elsif (open(P, "sysctl -a |")) { # BSD/Darwin
+ while (<P>) {
+ if (m/hw\.ncpu\s*[:=]\s*(\d+)/) { # hw.ncpu = 4, or hw.ncpu: 4
+ $max_jobs_run = $1;
+ last;
+ }
+ }
+ close(P);
+ if ($max_jobs_run == 0) {
+ print STDERR "run.pl: Warning: failed to detect any processors from sysctl -a\n";
+ $max_jobs_run = 10; # reasonable default.
+ }
+ } else {
+ # allow at most 32 jobs at once, on non-UNIX systems; change this code
+ # if you need to change this default.
+ $max_jobs_run = 32;
+ }
+ # The just-computed value of $max_jobs_run is just the number of processors
+ # (or our best guess); and if it happens that the number of jobs we need to
+ # run is just slightly above $max_jobs_run, it will make sense to increase
+ # $max_jobs_run to equal the number of jobs, so we don't have a small number
+ # of leftover jobs.
+ $num_jobs = $jobend - $jobstart + 1;
+ if (!$using_gpu &&
+ $num_jobs > $max_jobs_run && $num_jobs < 1.4 * $max_jobs_run) {
+ $max_jobs_run = $num_jobs;
+ }
+}
+
+sub pick_or_exit {
+ # pick_or_exit ( $logfile )
+ # Invoked before each job is started helps to run jobs selectively.
+ #
+ # Given the name of the output logfile decides whether the job must be
+ # executed (by returning from the subroutine) or not (by terminating the
+ # process calling exit)
+ #
+ # PRE: $job_pick is a global variable set by command line switch --pick
+ # and indicates which class of jobs must be executed.
+ #
+ # 1) If a failed job is not executed the process exit code will indicate
+ # failure, just as if the task was just executed and failed.
+ #
+ # 2) If a task is incomplete it will be executed. Incomplete may be either
+ # a job whose log file does not contain the accounting notes in the end,
+ # or a job whose log file does not exist.
+ #
+ # 3) If the $job_pick is set to 'all' (default behavior) a task will be
+ # executed regardless of the result of previous attempts.
+ #
+ # This logic could have been implemented in the main execution loop
+ # but a subroutine to preserve the current level of readability of
+ # that part of the code.
+ #
+ # Alexandre Felipe, (o.alexandre.felipe@gmail.com) 14th of August of 2020
+ #
+ if($job_pick eq 'all'){
+ return; # no need to bother with the previous log
+ }
+ open my $fh, "<", $_[0] or return; # job not executed yet
+ my $log_line;
+ my $cur_line;
+ while ($cur_line = <$fh>) {
+ if( $cur_line =~ m/# Ended \(code .*/ ) {
+ $log_line = $cur_line;
+ }
+ }
+ close $fh;
+ if (! defined($log_line)){
+ return; # incomplete
+ }
+ if ( $log_line =~ m/# Ended \(code 0\).*/ ) {
+ exit(0); # complete
+ } elsif ( $log_line =~ m/# Ended \(code \d+(; signal \d+)?\).*/ ){
+ if ($job_pick !~ m/^(failed|all)$/) {
+ exit(1); # failed but not going to run
+ } else {
+ return; # failed
+ }
+ } elsif ( $log_line =~ m/.*\S.*/ ) {
+ return; # incomplete jobs are always run
+ }
+}
+
+
+$logfile = shift @ARGV;
+
+if (defined $jobname && $logfile !~ m/$jobname/ &&
+ $jobend > $jobstart) {
+ print STDERR "run.pl: you are trying to run a parallel job but "
+ . "you are putting the output into just one log file ($logfile)\n";
+ exit(1);
+}
+
+$cmd = "";
+
+foreach $x (@ARGV) {
+ if ($x =~ m/^\S+$/) { $cmd .= $x . " "; }
+ elsif ($x =~ m:\":) { $cmd .= "'$x' "; }
+ else { $cmd .= "\"$x\" "; }
+}
+
+#$Data::Dumper::Indent=0;
+$ret = 0;
+$numfail = 0;
+%active_pids=();
+
+use POSIX ":sys_wait_h";
+for ($jobid = $jobstart; $jobid <= $jobend; $jobid++) {
+ if (scalar(keys %active_pids) >= $max_jobs_run) {
+
+ # Lets wait for a change in any child's status
+ # Then we have to work out which child finished
+ $r = waitpid(-1, 0);
+ $code = $?;
+ if ($r < 0 ) { die "run.pl: Error waiting for child process"; } # should never happen.
+ if ( defined $active_pids{$r} ) {
+ $jid=$active_pids{$r};
+ $fail[$jid]=$code;
+ if ($code !=0) { $numfail++;}
+ delete $active_pids{$r};
+ # print STDERR "Finished: $r/$jid " . Dumper(\%active_pids) . "\n";
+ } else {
+ die "run.pl: Cannot find the PID of the child process that just finished.";
+ }
+
+ # In theory we could do a non-blocking waitpid over all jobs running just
+ # to find out if only one or more jobs finished during the previous waitpid()
+ # However, we just omit this and will reap the next one in the next pass
+ # through the for(;;) cycle
+ }
+ $childpid = fork();
+ if (!defined $childpid) { die "run.pl: Error forking in run.pl (writing to $logfile)"; }
+ if ($childpid == 0) { # We're in the child... this branch
+ # executes the job and returns (possibly with an error status).
+ if (defined $jobname) {
+ $cmd =~ s/$jobname/$jobid/g;
+ $logfile =~ s/$jobname/$jobid/g;
+ }
+ # exit if the job does not need to be executed
+ pick_or_exit( $logfile );
+
+ system("mkdir -p `dirname $logfile` 2>/dev/null");
+ open(F, ">$logfile") || die "run.pl: Error opening log file $logfile";
+ print F "# " . $cmd . "\n";
+ print F "# Started at " . `date`;
+ $starttime = `date +'%s'`;
+ print F "#\n";
+ close(F);
+
+ # Pipe into bash.. make sure we're not using any other shell.
+ open(B, "|bash") || die "run.pl: Error opening shell command";
+ print B "( " . $cmd . ") 2>>$logfile >> $logfile";
+ close(B); # If there was an error, exit status is in $?
+ $ret = $?;
+
+ $lowbits = $ret & 127;
+ $highbits = $ret >> 8;
+ if ($lowbits != 0) { $return_str = "code $highbits; signal $lowbits" }
+ else { $return_str = "code $highbits"; }
+
+ $endtime = `date +'%s'`;
+ open(F, ">>$logfile") || die "run.pl: Error opening log file $logfile (again)";
+ $enddate = `date`;
+ chop $enddate;
+ print F "# Accounting: time=" . ($endtime - $starttime) . " threads=1\n";
+ print F "# Ended ($return_str) at " . $enddate . ", elapsed time " . ($endtime-$starttime) . " seconds\n";
+ close(F);
+ exit($ret == 0 ? 0 : 1);
+ } else {
+ $pid[$jobid] = $childpid;
+ $active_pids{$childpid} = $jobid;
+ # print STDERR "Queued: " . Dumper(\%active_pids) . "\n";
+ }
+}
+
+# Now we have submitted all the jobs, lets wait until all the jobs finish
+foreach $child (keys %active_pids) {
+ $jobid=$active_pids{$child};
+ $r = waitpid($pid[$jobid], 0);
+ $code = $?;
+ if ($r == -1) { die "run.pl: Error waiting for child process"; } # should never happen.
+ if ($r != 0) { $fail[$jobid]=$code; $numfail++ if $code!=0; } # Completed successfully
+}
+
+# Some sanity checks:
+# The $fail array should not contain undefined codes
+# The number of non-zeros in that array should be equal to $numfail
+# We cannot do foreach() here, as the JOB ids do not start at zero
+$failed_jids=0;
+for ($jobid = $jobstart; $jobid <= $jobend; $jobid++) {
+ $job_return = $fail[$jobid];
+ if (not defined $job_return ) {
+ # print Dumper(\@fail);
+
+ die "run.pl: Sanity check failed: we have indication that some jobs are running " .
+ "even after we waited for all jobs to finish" ;
+ }
+ if ($job_return != 0 ){ $failed_jids++;}
+}
+if ($failed_jids != $numfail) {
+ die "run.pl: Sanity check failed: cannot find out how many jobs failed ($failed_jids x $numfail)."
+}
+if ($numfail > 0) { $ret = 1; }
+
+if ($ret != 0) {
+ $njobs = $jobend - $jobstart + 1;
+ if ($njobs == 1) {
+ if (defined $jobname) {
+ $logfile =~ s/$jobname/$jobstart/; # only one numbered job, so replace name with
+ # that job.
+ }
+ print STDERR "run.pl: job failed, log is in $logfile\n";
+ if ($logfile =~ m/JOB/) {
+ print STDERR "run.pl: probably you forgot to put JOB=1:\$nj in your script.";
+ }
+ }
+ else {
+ $logfile =~ s/$jobname/*/g;
+ print STDERR "run.pl: $numfail / $njobs failed, log is in $logfile\n";
+ }
+}
+
+
+exit ($ret);
diff --git a/egs_modelscope/aishell2/paraformer/utils/shuffle_list.pl b/egs_modelscope/aishell2/paraformer/utils/shuffle_list.pl
new file mode 100755
index 0000000..a116200
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/shuffle_list.pl
@@ -0,0 +1,44 @@
+#!/usr/bin/env perl
+
+# Copyright 2013 Johns Hopkins University (author: Daniel Povey)
+
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+if ($ARGV[0] eq "--srand") {
+ $n = $ARGV[1];
+ $n =~ m/\d+/ || die "Bad argument to --srand option: \"$n\"";
+ srand($ARGV[1]);
+ shift;
+ shift;
+} else {
+ srand(0); # Gives inconsistent behavior if we don't seed.
+}
+
+if (@ARGV > 1 || $ARGV[0] =~ m/^-.+/) { # >1 args, or an option we
+ # don't understand.
+ print "Usage: shuffle_list.pl [--srand N] [input file] > output\n";
+ print "randomizes the order of lines of input.\n";
+ exit(1);
+}
+
+@lines;
+while (<>) {
+ push @lines, [ (rand(), $_)] ;
+}
+
+@lines = sort { $a->[0] cmp $b->[0] } @lines;
+foreach $l (@lines) {
+ print $l->[1];
+}
\ No newline at end of file
diff --git a/egs_modelscope/aishell2/paraformer/utils/split_data.py b/egs_modelscope/aishell2/paraformer/utils/split_data.py
new file mode 100755
index 0000000..060eae6
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/split_data.py
@@ -0,0 +1,60 @@
+import os
+import sys
+import random
+
+
+in_dir = sys.argv[1]
+out_dir = sys.argv[2]
+num_split = sys.argv[3]
+
+
+def split_scp(scp, num):
+ assert len(scp) >= num
+ avg = len(scp) // num
+ out = []
+ begin = 0
+
+ for i in range(num):
+ if i == num - 1:
+ out.append(scp[begin:])
+ else:
+ out.append(scp[begin:begin+avg])
+ begin += avg
+
+ return out
+
+
+os.path.exists("{}/wav.scp".format(in_dir))
+os.path.exists("{}/text".format(in_dir))
+
+with open("{}/wav.scp".format(in_dir), 'r') as infile:
+ wav_list = infile.readlines()
+
+with open("{}/text".format(in_dir), 'r') as infile:
+ text_list = infile.readlines()
+
+assert len(wav_list) == len(text_list)
+
+x = list(zip(wav_list, text_list))
+random.shuffle(x)
+wav_shuffle_list, text_shuffle_list = zip(*x)
+
+num_split = int(num_split)
+wav_split_list = split_scp(wav_shuffle_list, num_split)
+text_split_list = split_scp(text_shuffle_list, num_split)
+
+for idx, wav_list in enumerate(wav_split_list, 1):
+ path = out_dir + "/split" + str(num_split) + "/" + str(idx)
+ if not os.path.exists(path):
+ os.makedirs(path)
+ with open("{}/wav.scp".format(path), 'w') as wav_writer:
+ for line in wav_list:
+ wav_writer.write(line)
+
+for idx, text_list in enumerate(text_split_list, 1):
+ path = out_dir + "/split" + str(num_split) + "/" + str(idx)
+ if not os.path.exists(path):
+ os.makedirs(path)
+ with open("{}/text".format(path), 'w') as text_writer:
+ for line in text_list:
+ text_writer.write(line)
diff --git a/egs_modelscope/aishell2/paraformer/utils/split_scp.pl b/egs_modelscope/aishell2/paraformer/utils/split_scp.pl
new file mode 100755
index 0000000..0876dcb
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/split_scp.pl
@@ -0,0 +1,246 @@
+#!/usr/bin/env perl
+
+# Copyright 2010-2011 Microsoft Corporation
+
+# See ../../COPYING for clarification regarding multiple authors
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+# This program splits up any kind of .scp or archive-type file.
+# If there is no utt2spk option it will work on any text file and
+# will split it up with an approximately equal number of lines in
+# each but.
+# With the --utt2spk option it will work on anything that has the
+# utterance-id as the first entry on each line; the utt2spk file is
+# of the form "utterance speaker" (on each line).
+# It splits it into equal size chunks as far as it can. If you use the utt2spk
+# option it will make sure these chunks coincide with speaker boundaries. In
+# this case, if there are more chunks than speakers (and in some other
+# circumstances), some of the resulting chunks will be empty and it will print
+# an error message and exit with nonzero status.
+# You will normally call this like:
+# split_scp.pl scp scp.1 scp.2 scp.3 ...
+# or
+# split_scp.pl --utt2spk=utt2spk scp scp.1 scp.2 scp.3 ...
+# Note that you can use this script to split the utt2spk file itself,
+# e.g. split_scp.pl --utt2spk=utt2spk utt2spk utt2spk.1 utt2spk.2 ...
+
+# You can also call the scripts like:
+# split_scp.pl -j 3 0 scp scp.0
+# [note: with this option, it assumes zero-based indexing of the split parts,
+# i.e. the second number must be 0 <= n < num-jobs.]
+
+use warnings;
+
+$num_jobs = 0;
+$job_id = 0;
+$utt2spk_file = "";
+$one_based = 0;
+
+for ($x = 1; $x <= 3 && @ARGV > 0; $x++) {
+ if ($ARGV[0] eq "-j") {
+ shift @ARGV;
+ $num_jobs = shift @ARGV;
+ $job_id = shift @ARGV;
+ }
+ if ($ARGV[0] =~ /--utt2spk=(.+)/) {
+ $utt2spk_file=$1;
+ shift;
+ }
+ if ($ARGV[0] eq '--one-based') {
+ $one_based = 1;
+ shift @ARGV;
+ }
+}
+
+if ($num_jobs != 0 && ($num_jobs < 0 || $job_id - $one_based < 0 ||
+ $job_id - $one_based >= $num_jobs)) {
+ die "$0: Invalid job number/index values for '-j $num_jobs $job_id" .
+ ($one_based ? " --one-based" : "") . "'\n"
+}
+
+$one_based
+ and $job_id--;
+
+if(($num_jobs == 0 && @ARGV < 2) || ($num_jobs > 0 && (@ARGV < 1 || @ARGV > 2))) {
+ die
+"Usage: split_scp.pl [--utt2spk=<utt2spk_file>] in.scp out1.scp out2.scp ...
+ or: split_scp.pl -j num-jobs job-id [--one-based] [--utt2spk=<utt2spk_file>] in.scp [out.scp]
+ ... where 0 <= job-id < num-jobs, or 1 <= job-id <- num-jobs if --one-based.\n";
+}
+
+$error = 0;
+$inscp = shift @ARGV;
+if ($num_jobs == 0) { # without -j option
+ @OUTPUTS = @ARGV;
+} else {
+ for ($j = 0; $j < $num_jobs; $j++) {
+ if ($j == $job_id) {
+ if (@ARGV > 0) { push @OUTPUTS, $ARGV[0]; }
+ else { push @OUTPUTS, "-"; }
+ } else {
+ push @OUTPUTS, "/dev/null";
+ }
+ }
+}
+
+if ($utt2spk_file ne "") { # We have the --utt2spk option...
+ open($u_fh, '<', $utt2spk_file) || die "$0: Error opening utt2spk file $utt2spk_file: $!\n";
+ while(<$u_fh>) {
+ @A = split;
+ @A == 2 || die "$0: Bad line $_ in utt2spk file $utt2spk_file\n";
+ ($u,$s) = @A;
+ $utt2spk{$u} = $s;
+ }
+ close $u_fh;
+ open($i_fh, '<', $inscp) || die "$0: Error opening input scp file $inscp: $!\n";
+ @spkrs = ();
+ while(<$i_fh>) {
+ @A = split;
+ if(@A == 0) { die "$0: Empty or space-only line in scp file $inscp\n"; }
+ $u = $A[0];
+ $s = $utt2spk{$u};
+ defined $s || die "$0: No utterance $u in utt2spk file $utt2spk_file\n";
+ if(!defined $spk_count{$s}) {
+ push @spkrs, $s;
+ $spk_count{$s} = 0;
+ $spk_data{$s} = []; # ref to new empty array.
+ }
+ $spk_count{$s}++;
+ push @{$spk_data{$s}}, $_;
+ }
+ # Now split as equally as possible ..
+ # First allocate spks to files by allocating an approximately
+ # equal number of speakers.
+ $numspks = @spkrs; # number of speakers.
+ $numscps = @OUTPUTS; # number of output files.
+ if ($numspks < $numscps) {
+ die "$0: Refusing to split data because number of speakers $numspks " .
+ "is less than the number of output .scp files $numscps\n";
+ }
+ for($scpidx = 0; $scpidx < $numscps; $scpidx++) {
+ $scparray[$scpidx] = []; # [] is array reference.
+ }
+ for ($spkidx = 0; $spkidx < $numspks; $spkidx++) {
+ $scpidx = int(($spkidx*$numscps) / $numspks);
+ $spk = $spkrs[$spkidx];
+ push @{$scparray[$scpidx]}, $spk;
+ $scpcount[$scpidx] += $spk_count{$spk};
+ }
+
+ # Now will try to reassign beginning + ending speakers
+ # to different scp's and see if it gets more balanced.
+ # Suppose objf we're minimizing is sum_i (num utts in scp[i] - average)^2.
+ # We can show that if considering changing just 2 scp's, we minimize
+ # this by minimizing the squared difference in sizes. This is
+ # equivalent to minimizing the absolute difference in sizes. This
+ # shows this method is bound to converge.
+
+ $changed = 1;
+ while($changed) {
+ $changed = 0;
+ for($scpidx = 0; $scpidx < $numscps; $scpidx++) {
+ # First try to reassign ending spk of this scp.
+ if($scpidx < $numscps-1) {
+ $sz = @{$scparray[$scpidx]};
+ if($sz > 0) {
+ $spk = $scparray[$scpidx]->[$sz-1];
+ $count = $spk_count{$spk};
+ $nutt1 = $scpcount[$scpidx];
+ $nutt2 = $scpcount[$scpidx+1];
+ if( abs( ($nutt2+$count) - ($nutt1-$count))
+ < abs($nutt2 - $nutt1)) { # Would decrease
+ # size-diff by reassigning spk...
+ $scpcount[$scpidx+1] += $count;
+ $scpcount[$scpidx] -= $count;
+ pop @{$scparray[$scpidx]};
+ unshift @{$scparray[$scpidx+1]}, $spk;
+ $changed = 1;
+ }
+ }
+ }
+ if($scpidx > 0 && @{$scparray[$scpidx]} > 0) {
+ $spk = $scparray[$scpidx]->[0];
+ $count = $spk_count{$spk};
+ $nutt1 = $scpcount[$scpidx-1];
+ $nutt2 = $scpcount[$scpidx];
+ if( abs( ($nutt2-$count) - ($nutt1+$count))
+ < abs($nutt2 - $nutt1)) { # Would decrease
+ # size-diff by reassigning spk...
+ $scpcount[$scpidx-1] += $count;
+ $scpcount[$scpidx] -= $count;
+ shift @{$scparray[$scpidx]};
+ push @{$scparray[$scpidx-1]}, $spk;
+ $changed = 1;
+ }
+ }
+ }
+ }
+ # Now print out the files...
+ for($scpidx = 0; $scpidx < $numscps; $scpidx++) {
+ $scpfile = $OUTPUTS[$scpidx];
+ ($scpfile ne '-' ? open($f_fh, '>', $scpfile)
+ : open($f_fh, '>&', \*STDOUT)) ||
+ die "$0: Could not open scp file $scpfile for writing: $!\n";
+ $count = 0;
+ if(@{$scparray[$scpidx]} == 0) {
+ print STDERR "$0: eError: split_scp.pl producing empty .scp file " .
+ "$scpfile (too many splits and too few speakers?)\n";
+ $error = 1;
+ } else {
+ foreach $spk ( @{$scparray[$scpidx]} ) {
+ print $f_fh @{$spk_data{$spk}};
+ $count += $spk_count{$spk};
+ }
+ $count == $scpcount[$scpidx] || die "Count mismatch [code error]";
+ }
+ close($f_fh);
+ }
+} else {
+ # This block is the "normal" case where there is no --utt2spk
+ # option and we just break into equal size chunks.
+
+ open($i_fh, '<', $inscp) || die "$0: Error opening input scp file $inscp: $!\n";
+
+ $numscps = @OUTPUTS; # size of array.
+ @F = ();
+ while(<$i_fh>) {
+ push @F, $_;
+ }
+ $numlines = @F;
+ if($numlines == 0) {
+ print STDERR "$0: error: empty input scp file $inscp\n";
+ $error = 1;
+ }
+ $linesperscp = int( $numlines / $numscps); # the "whole part"..
+ $linesperscp >= 1 || die "$0: You are splitting into too many pieces! [reduce \$nj ($numscps) to be smaller than the number of lines ($numlines) in $inscp]\n";
+ $remainder = $numlines - ($linesperscp * $numscps);
+ ($remainder >= 0 && $remainder < $numlines) || die "bad remainder $remainder";
+ # [just doing int() rounds down].
+ $n = 0;
+ for($scpidx = 0; $scpidx < @OUTPUTS; $scpidx++) {
+ $scpfile = $OUTPUTS[$scpidx];
+ ($scpfile ne '-' ? open($o_fh, '>', $scpfile)
+ : open($o_fh, '>&', \*STDOUT)) ||
+ die "$0: Could not open scp file $scpfile for writing: $!\n";
+ for($k = 0; $k < $linesperscp + ($scpidx < $remainder ? 1 : 0); $k++) {
+ print $o_fh $F[$n++];
+ }
+ close($o_fh) || die "$0: Eror closing scp file $scpfile: $!\n";
+ }
+ $n == $numlines || die "$n != $numlines [code error]";
+}
+
+exit ($error);
diff --git a/egs_modelscope/aishell2/paraformer/utils/subset_data_dir_tr_cv.sh b/egs_modelscope/aishell2/paraformer/utils/subset_data_dir_tr_cv.sh
new file mode 100755
index 0000000..e16cebd
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/subset_data_dir_tr_cv.sh
@@ -0,0 +1,30 @@
+#!/usr/bin/env bash
+
+dev_num_utt=1000
+
+echo "$0 $@"
+. utils/parse_options.sh || exit 1;
+
+train_data=$1
+out_dir=$2
+
+[ ! -f ${train_data}/wav.scp ] && echo "$0: no such file ${train_data}/wav.scp" && exit 1;
+[ ! -f ${train_data}/text ] && echo "$0: no such file ${train_data}/text" && exit 1;
+
+mkdir -p ${out_dir}/train && mkdir -p ${out_dir}/dev
+
+cp ${train_data}/wav.scp ${out_dir}/train/wav.scp.bak
+cp ${train_data}/text ${out_dir}/train/text.bak
+
+num_utt=$(wc -l <${out_dir}/train/wav.scp.bak)
+
+utils/shuffle_list.pl --srand 1 ${out_dir}/train/wav.scp.bak > ${out_dir}/train/wav.scp.shuf
+head -n ${dev_num_utt} ${out_dir}/train/wav.scp.shuf > ${out_dir}/dev/wav.scp
+tail -n $((${num_utt}-${dev_num_utt})) ${out_dir}/train/wav.scp.shuf > ${out_dir}/train/wav.scp
+
+utils/shuffle_list.pl --srand 1 ${out_dir}/train/text.bak > ${out_dir}/train/text.shuf
+head -n ${dev_num_utt} ${out_dir}/train/text.shuf > ${out_dir}/dev/text
+tail -n $((${num_utt}-${dev_num_utt})) ${out_dir}/train/text.shuf > ${out_dir}/train/text
+
+rm ${out_dir}/train/wav.scp.bak ${out_dir}/train/text.bak
+rm ${out_dir}/train/wav.scp.shuf ${out_dir}/train/text.shuf
diff --git a/egs_modelscope/aishell2/paraformer/utils/text2token.py b/egs_modelscope/aishell2/paraformer/utils/text2token.py
new file mode 100755
index 0000000..56c3913
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/text2token.py
@@ -0,0 +1,135 @@
+#!/usr/bin/env python3
+
+# Copyright 2017 Johns Hopkins University (Shinji Watanabe)
+# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
+
+
+import argparse
+import codecs
+import re
+import sys
+
+is_python2 = sys.version_info[0] == 2
+
+
+def exist_or_not(i, match_pos):
+ start_pos = None
+ end_pos = None
+ for pos in match_pos:
+ if pos[0] <= i < pos[1]:
+ start_pos = pos[0]
+ end_pos = pos[1]
+ break
+
+ return start_pos, end_pos
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="convert raw text to tokenized text",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--nchar",
+ "-n",
+ default=1,
+ type=int,
+ help="number of characters to split, i.e., \
+ aabb -> a a b b with -n 1 and aa bb with -n 2",
+ )
+ parser.add_argument(
+ "--skip-ncols", "-s", default=0, type=int, help="skip first n columns"
+ )
+ parser.add_argument("--space", default="<space>", type=str, help="space symbol")
+ parser.add_argument(
+ "--non-lang-syms",
+ "-l",
+ default=None,
+ type=str,
+ help="list of non-linguistic symobles, e.g., <NOISE> etc.",
+ )
+ parser.add_argument("text", type=str, default=False, nargs="?", help="input text")
+ parser.add_argument(
+ "--trans_type",
+ "-t",
+ type=str,
+ default="char",
+ choices=["char", "phn"],
+ help="""Transcript type. char/phn. e.g., for TIMIT FADG0_SI1279 -
+ If trans_type is char,
+ read from SI1279.WRD file -> "bricks are an alternative"
+ Else if trans_type is phn,
+ read from SI1279.PHN file -> "sil b r ih sil k s aa r er n aa l
+ sil t er n ih sil t ih v sil" """,
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ rs = []
+ if args.non_lang_syms is not None:
+ with codecs.open(args.non_lang_syms, "r", encoding="utf-8") as f:
+ nls = [x.rstrip() for x in f.readlines()]
+ rs = [re.compile(re.escape(x)) for x in nls]
+
+ if args.text:
+ f = codecs.open(args.text, encoding="utf-8")
+ else:
+ f = codecs.getreader("utf-8")(sys.stdin if is_python2 else sys.stdin.buffer)
+
+ sys.stdout = codecs.getwriter("utf-8")(
+ sys.stdout if is_python2 else sys.stdout.buffer
+ )
+ line = f.readline()
+ n = args.nchar
+ while line:
+ x = line.split()
+ print(" ".join(x[: args.skip_ncols]), end=" ")
+ a = " ".join(x[args.skip_ncols :])
+
+ # get all matched positions
+ match_pos = []
+ for r in rs:
+ i = 0
+ while i >= 0:
+ m = r.search(a, i)
+ if m:
+ match_pos.append([m.start(), m.end()])
+ i = m.end()
+ else:
+ break
+
+ if args.trans_type == "phn":
+ a = a.split(" ")
+ else:
+ if len(match_pos) > 0:
+ chars = []
+ i = 0
+ while i < len(a):
+ start_pos, end_pos = exist_or_not(i, match_pos)
+ if start_pos is not None:
+ chars.append(a[start_pos:end_pos])
+ i = end_pos
+ else:
+ chars.append(a[i])
+ i += 1
+ a = chars
+
+ a = [a[j : j + n] for j in range(0, len(a), n)]
+
+ a_flat = []
+ for z in a:
+ a_flat.append("".join(z))
+
+ a_chars = [z.replace(" ", args.space) for z in a_flat]
+ if args.trans_type == "phn":
+ a_chars = [z.replace("sil", args.space) for z in a_chars]
+ print(" ".join(a_chars))
+ line = f.readline()
+
+
+if __name__ == "__main__":
+ main()
diff --git a/egs_modelscope/aishell2/paraformer/utils/text_tokenize.py b/egs_modelscope/aishell2/paraformer/utils/text_tokenize.py
new file mode 100755
index 0000000..962ea11
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/text_tokenize.py
@@ -0,0 +1,106 @@
+import re
+import argparse
+
+
+def load_dict(seg_file):
+ seg_dict = {}
+ with open(seg_file, 'r') as infile:
+ for line in infile:
+ s = line.strip().split()
+ key = s[0]
+ value = s[1:]
+ seg_dict[key] = " ".join(value)
+ return seg_dict
+
+
+def forward_segment(text, dic):
+ word_list = []
+ i = 0
+ while i < len(text):
+ longest_word = text[i]
+ for j in range(i + 1, len(text) + 1):
+ word = text[i:j]
+ if word in dic:
+ if len(word) > len(longest_word):
+ longest_word = word
+ word_list.append(longest_word)
+ i += len(longest_word)
+ return word_list
+
+
+def tokenize(txt,
+ seg_dict):
+ out_txt = ""
+ pattern = re.compile(r"([\u4E00-\u9FA5A-Za-z0-9])")
+ for word in txt:
+ if pattern.match(word):
+ if word in seg_dict:
+ out_txt += seg_dict[word] + " "
+ else:
+ out_txt += "<unk>" + " "
+ else:
+ continue
+ return out_txt.strip()
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="text tokenize",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--text-file",
+ "-t",
+ default=False,
+ required=True,
+ type=str,
+ help="input text",
+ )
+ parser.add_argument(
+ "--seg-file",
+ "-s",
+ default=False,
+ required=True,
+ type=str,
+ help="seg file",
+ )
+ parser.add_argument(
+ "--txt-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="txt index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ txt_writer = open("{}/text.{}.txt".format(args.output_dir, args.txt_index), 'w')
+ shape_writer = open("{}/len.{}".format(args.output_dir, args.txt_index), 'w')
+ seg_dict = load_dict(args.seg_file)
+ with open(args.text_file, 'r') as infile:
+ for line in infile:
+ s = line.strip().split()
+ text_id = s[0]
+ text_list = forward_segment("".join(s[1:]).lower(), seg_dict)
+ text = tokenize(text_list, seg_dict)
+ lens = len(text.strip().split())
+ txt_writer.write(text_id + " " + text + '\n')
+ shape_writer.write(text_id + " " + str(lens) + '\n')
+
+
+if __name__ == '__main__':
+ main()
+
diff --git a/egs_modelscope/aishell2/paraformer/utils/text_tokenize.sh b/egs_modelscope/aishell2/paraformer/utils/text_tokenize.sh
new file mode 100755
index 0000000..6b74fef
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/text_tokenize.sh
@@ -0,0 +1,35 @@
+#!/usr/bin/env bash
+
+
+# Begin configuration section.
+nj=32
+cmd=utils/run.pl
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+# tokenize configuration
+text_dir=$1
+seg_file=$2
+logdir=$3
+output_dir=$4
+
+txt_dir=${output_dir}/txt; mkdir -p ${output_dir}/txt
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/text_tokenize.JOB.log \
+ python utils/text_tokenize.py -t ${text_dir}/txt/text.JOB.txt \
+ -s ${seg_file} -i JOB -o ${txt_dir} \
+ || exit 1;
+
+# concatenate the text files together.
+for n in $(seq $nj); do
+ cat ${txt_dir}/text.$n.txt || exit 1
+done > ${output_dir}/text || exit 1
+
+for n in $(seq $nj); do
+ cat ${txt_dir}/len.$n || exit 1
+done > ${output_dir}/text_shape || exit 1
+
+echo "$0: Succeeded text tokenize"
diff --git a/egs_modelscope/aishell2/paraformer/utils/textnorm_zh.py b/egs_modelscope/aishell2/paraformer/utils/textnorm_zh.py
new file mode 100755
index 0000000..79feb83
--- /dev/null
+++ b/egs_modelscope/aishell2/paraformer/utils/textnorm_zh.py
@@ -0,0 +1,834 @@
+#!/usr/bin/env python3
+# coding=utf-8
+
+# Authors:
+# 2019.5 Zhiyang Zhou (https://github.com/Joee1995/chn_text_norm.git)
+# 2019.9 Jiayu DU
+#
+# requirements:
+# - python 3.X
+# notes: python 2.X WILL fail or produce misleading results
+
+import sys, os, argparse, codecs, string, re
+
+# ================================================================================ #
+# basic constant
+# ================================================================================ #
+CHINESE_DIGIS = u'闆朵竴浜屼笁鍥涗簲鍏竷鍏節'
+BIG_CHINESE_DIGIS_SIMPLIFIED = u'闆跺9璐板弫鑲嗕紞闄嗘煉鎹岀帠'
+BIG_CHINESE_DIGIS_TRADITIONAL = u'闆跺9璨冲弮鑲嗕紞闄告煉鎹岀帠'
+SMALLER_BIG_CHINESE_UNITS_SIMPLIFIED = u'鍗佺櫨鍗冧竾'
+SMALLER_BIG_CHINESE_UNITS_TRADITIONAL = u'鎷句桨浠熻惉'
+LARGER_CHINESE_NUMERING_UNITS_SIMPLIFIED = u'浜垮厗浜灀绉┌娌熸锭姝h浇'
+LARGER_CHINESE_NUMERING_UNITS_TRADITIONAL = u'鍎勫厗浜灀绉┌婧濇緱姝h級'
+SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED = u'鍗佺櫨鍗冧竾'
+SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL = u'鎷句桨浠熻惉'
+
+ZERO_ALT = u'銆�'
+ONE_ALT = u'骞�'
+TWO_ALTS = [u'涓�', u'鍏�']
+
+POSITIVE = [u'姝�', u'姝�']
+NEGATIVE = [u'璐�', u'璨�']
+POINT = [u'鐐�', u'榛�']
+# PLUS = [u'鍔�', u'鍔�']
+# SIL = [u'鏉�', u'妲�']
+
+FILLER_CHARS = ['鍛�', '鍟�']
+ER_WHITELIST = '(鍎垮コ|鍎垮瓙|鍎垮瓩|濂冲効|鍎垮|濡诲効|' \
+ '鑳庡効|濠村効|鏂扮敓鍎縷濠村辜鍎縷骞煎効|灏戝効|灏忓効|鍎挎瓕|鍎跨|鍎跨|鎵樺効鎵�|瀛ゅ効|' \
+ '鍎挎垙|鍎垮寲|鍙板効搴剕楣垮効宀泑姝e効鍏粡|鍚婂効閮庡綋|鐢熷効鑲插コ|鎵樺効甯﹀コ|鍏诲効闃茶�亅鐥村効鍛嗗コ|' \
+ '浣冲効浣冲|鍎挎�滃吔鎵皘鍎挎棤甯哥埗|鍎夸笉瀚屾瘝涓憒鍎胯鍗冮噷姣嶆媴蹇鍎垮ぇ涓嶇敱鐖穦鑻忎篂鍎�)'
+
+# 涓枃鏁板瓧绯荤粺绫诲瀷
+NUMBERING_TYPES = ['low', 'mid', 'high']
+
+CURRENCY_NAMES = '(浜烘皯甯亅缇庡厓|鏃ュ厓|鑻遍晳|娆у厓|椹厠|娉曢儙|鍔犳嬁澶у厓|婢冲厓|娓竵|鍏堜护|鑺叞椹厠|鐖卞皵鍏伴晳|' \
+ '閲屾媺|鑽峰叞鐩緗鍩冩柉搴撳|姣斿濉攟鍗板凹鐩緗鏋楀悏鐗箌鏂拌タ鍏板厓|姣旂储|鍗㈠竷|鏂板姞鍧″厓|闊╁厓|娉伴摙)'
+CURRENCY_UNITS = '((浜縷鍗冧竾|鐧句竾|涓噟鍗億鐧�)|(浜縷鍗冧竾|鐧句竾|涓噟鍗億鐧緗)鍏億(浜縷鍗冧竾|鐧句竾|涓噟鍗億鐧緗)鍧梶瑙抾姣泑鍒�)'
+COM_QUANTIFIERS = '(鍖箌寮爘搴鍥瀨鍦簗灏緗鏉涓獆棣東闃檤闃祙缃憒鐐畖椤秥涓榺妫祙鍙獆鏀瘄琚瓅杈唡鎸憒鎷厊棰梶澹硘绐爘鏇瞸澧檤缇鑵攟' \
+ '鐮搴瀹璐瘄鎵巪鎹唡鍒�|浠鎵搢鎵媩缃梶鍧灞眧宀瓅姹焲婧獆閽焲闃焲鍗晐鍙寍瀵箌鍑簗鍙澶磡鑴殀鏉縷璺硘鏋潀浠秥璐磡' \
+ '閽坾绾縷绠鍚峾浣峾韬珅鍫倈璇緗鏈瑋椤祙瀹秥鎴穦灞倈涓潀姣珅鍘榺鍒唡閽眧涓鏂鎷厊閾鐭硘閽閿眧蹇絴(鍗億姣珅寰�)鍏媩' \
+ '姣珅鍘榺鍒唡瀵竱灏簗涓坾閲寍瀵粅甯竱閾簗绋媩(鍗億鍒唡鍘榺姣珅寰�)绫硘鎾畖鍕簗鍚坾鍗噟鏂梶鐭硘鐩榺纰梶纰焲鍙爘妗秥绗紎鐩唡' \
+ '鐩抾鏉瘄閽焲鏂泑閿厊绨媩绡畖鐩榺妗秥缃恷鐡秥澹秥鍗畖鐩弢绠﹟绠眧鐓瞸鍟東琚媩閽祙骞磡鏈坾鏃瀛鍒粅鏃秥鍛▅澶﹟绉抾鍒唡鏃瑋' \
+ '绾獆宀亅涓東鏇磡澶渱鏄澶弢绉媩鍐瑋浠浼弢杈坾涓竱娉绮抾棰梶骞鍫唡鏉鏍箌鏀瘄閬搢闈鐗噟寮爘棰梶鍧�)'
+
+# punctuation information are based on Zhon project (https://github.com/tsroten/zhon.git)
+CHINESE_PUNC_STOP = '锛侊紵锝°��'
+CHINESE_PUNC_NON_STOP = '锛傦純锛勶紖锛嗭紘锛堬級锛婏紜锛岋紞锛忥細锛涳紲锛濓紴锛狅蓟锛硷冀锛撅伎锝�锝涳綔锝濓綖锝燂綘锝剑锝ゃ�併�冦�嬨�屻�嶃�庛�忋�愩�戙�斻�曘�栥�椼�樸�欍�氥�涖�溿�濄�炪�熴�般�俱�库�撯�斺�樷�欌�涒�溾�濃�炩�熲�︹�э箯'
+CHINESE_PUNC_LIST = CHINESE_PUNC_STOP + CHINESE_PUNC_NON_STOP
+
+# ================================================================================ #
+# basic class
+# ================================================================================ #
+class ChineseChar(object):
+ """
+ 涓枃瀛楃
+ 姣忎釜瀛楃瀵瑰簲绠�浣撳拰绻佷綋,
+ e.g. 绠�浣� = '璐�', 绻佷綋 = '璨�'
+ 杞崲鏃跺彲杞崲涓虹畝浣撴垨绻佷綋
+ """
+
+ def __init__(self, simplified, traditional):
+ self.simplified = simplified
+ self.traditional = traditional
+ #self.__repr__ = self.__str__
+
+ def __str__(self):
+ return self.simplified or self.traditional or None
+
+ def __repr__(self):
+ return self.__str__()
+
+
+class ChineseNumberUnit(ChineseChar):
+ """
+ 涓枃鏁板瓧/鏁颁綅瀛楃
+ 姣忎釜瀛楃闄ょ箒绠�浣撳杩樻湁涓�涓澶栫殑澶у啓瀛楃
+ e.g. '闄�' 鍜� '闄�'
+ """
+
+ def __init__(self, power, simplified, traditional, big_s, big_t):
+ super(ChineseNumberUnit, self).__init__(simplified, traditional)
+ self.power = power
+ self.big_s = big_s
+ self.big_t = big_t
+
+ def __str__(self):
+ return '10^{}'.format(self.power)
+
+ @classmethod
+ def create(cls, index, value, numbering_type=NUMBERING_TYPES[1], small_unit=False):
+
+ if small_unit:
+ return ChineseNumberUnit(power=index + 1,
+ simplified=value[0], traditional=value[1], big_s=value[1], big_t=value[1])
+ elif numbering_type == NUMBERING_TYPES[0]:
+ return ChineseNumberUnit(power=index + 8,
+ simplified=value[0], traditional=value[1], big_s=value[0], big_t=value[1])
+ elif numbering_type == NUMBERING_TYPES[1]:
+ return ChineseNumberUnit(power=(index + 2) * 4,
+ simplified=value[0], traditional=value[1], big_s=value[0], big_t=value[1])
+ elif numbering_type == NUMBERING_TYPES[2]:
+ return ChineseNumberUnit(power=pow(2, index + 3),
+ simplified=value[0], traditional=value[1], big_s=value[0], big_t=value[1])
+ else:
+ raise ValueError(
+ 'Counting type should be in {0} ({1} provided).'.format(NUMBERING_TYPES, numbering_type))
+
+
+class ChineseNumberDigit(ChineseChar):
+ """
+ 涓枃鏁板瓧瀛楃
+ """
+
+ def __init__(self, value, simplified, traditional, big_s, big_t, alt_s=None, alt_t=None):
+ super(ChineseNumberDigit, self).__init__(simplified, traditional)
+ self.value = value
+ self.big_s = big_s
+ self.big_t = big_t
+ self.alt_s = alt_s
+ self.alt_t = alt_t
+
+ def __str__(self):
+ return str(self.value)
+
+ @classmethod
+ def create(cls, i, v):
+ return ChineseNumberDigit(i, v[0], v[1], v[2], v[3])
+
+
+class ChineseMath(ChineseChar):
+ """
+ 涓枃鏁颁綅瀛楃
+ """
+
+ def __init__(self, simplified, traditional, symbol, expression=None):
+ super(ChineseMath, self).__init__(simplified, traditional)
+ self.symbol = symbol
+ self.expression = expression
+ self.big_s = simplified
+ self.big_t = traditional
+
+
+CC, CNU, CND, CM = ChineseChar, ChineseNumberUnit, ChineseNumberDigit, ChineseMath
+
+
+class NumberSystem(object):
+ """
+ 涓枃鏁板瓧绯荤粺
+ """
+ pass
+
+
+class MathSymbol(object):
+ """
+ 鐢ㄤ簬涓枃鏁板瓧绯荤粺鐨勬暟瀛︾鍙� (绻�/绠�浣�), e.g.
+ positive = ['姝�', '姝�']
+ negative = ['璐�', '璨�']
+ point = ['鐐�', '榛�']
+ """
+
+ def __init__(self, positive, negative, point):
+ self.positive = positive
+ self.negative = negative
+ self.point = point
+
+ def __iter__(self):
+ for v in self.__dict__.values():
+ yield v
+
+
+# class OtherSymbol(object):
+# """
+# 鍏朵粬绗﹀彿
+# """
+#
+# def __init__(self, sil):
+# self.sil = sil
+#
+# def __iter__(self):
+# for v in self.__dict__.values():
+# yield v
+
+
+# ================================================================================ #
+# basic utils
+# ================================================================================ #
+def create_system(numbering_type=NUMBERING_TYPES[1]):
+ """
+ 鏍规嵁鏁板瓧绯荤粺绫诲瀷杩斿洖鍒涘缓鐩稿簲鐨勬暟瀛楃郴缁燂紝榛樿涓� mid
+ NUMBERING_TYPES = ['low', 'mid', 'high']: 涓枃鏁板瓧绯荤粺绫诲瀷
+ low: '鍏�' = '浜�' * '鍗�' = $10^{9}$, '浜�' = '鍏�' * '鍗�', etc.
+ mid: '鍏�' = '浜�' * '涓�' = $10^{12}$, '浜�' = '鍏�' * '涓�', etc.
+ high: '鍏�' = '浜�' * '浜�' = $10^{16}$, '浜�' = '鍏�' * '鍏�', etc.
+ 杩斿洖瀵瑰簲鐨勬暟瀛楃郴缁�
+ """
+
+ # chinese number units of '浜�' and larger
+ all_larger_units = zip(
+ LARGER_CHINESE_NUMERING_UNITS_SIMPLIFIED, LARGER_CHINESE_NUMERING_UNITS_TRADITIONAL)
+ larger_units = [CNU.create(i, v, numbering_type, False)
+ for i, v in enumerate(all_larger_units)]
+ # chinese number units of '鍗�, 鐧�, 鍗�, 涓�'
+ all_smaller_units = zip(
+ SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED, SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL)
+ smaller_units = [CNU.create(i, v, small_unit=True)
+ for i, v in enumerate(all_smaller_units)]
+ # digis
+ chinese_digis = zip(CHINESE_DIGIS, CHINESE_DIGIS,
+ BIG_CHINESE_DIGIS_SIMPLIFIED, BIG_CHINESE_DIGIS_TRADITIONAL)
+ digits = [CND.create(i, v) for i, v in enumerate(chinese_digis)]
+ digits[0].alt_s, digits[0].alt_t = ZERO_ALT, ZERO_ALT
+ digits[1].alt_s, digits[1].alt_t = ONE_ALT, ONE_ALT
+ digits[2].alt_s, digits[2].alt_t = TWO_ALTS[0], TWO_ALTS[1]
+
+ # symbols
+ positive_cn = CM(POSITIVE[0], POSITIVE[1], '+', lambda x: x)
+ negative_cn = CM(NEGATIVE[0], NEGATIVE[1], '-', lambda x: -x)
+ point_cn = CM(POINT[0], POINT[1], '.', lambda x,
+ y: float(str(x) + '.' + str(y)))
+ # sil_cn = CM(SIL[0], SIL[1], '-', lambda x, y: float(str(x) + '-' + str(y)))
+ system = NumberSystem()
+ system.units = smaller_units + larger_units
+ system.digits = digits
+ system.math = MathSymbol(positive_cn, negative_cn, point_cn)
+ # system.symbols = OtherSymbol(sil_cn)
+ return system
+
+
+def chn2num(chinese_string, numbering_type=NUMBERING_TYPES[1]):
+
+ def get_symbol(char, system):
+ for u in system.units:
+ if char in [u.traditional, u.simplified, u.big_s, u.big_t]:
+ return u
+ for d in system.digits:
+ if char in [d.traditional, d.simplified, d.big_s, d.big_t, d.alt_s, d.alt_t]:
+ return d
+ for m in system.math:
+ if char in [m.traditional, m.simplified]:
+ return m
+
+ def string2symbols(chinese_string, system):
+ int_string, dec_string = chinese_string, ''
+ for p in [system.math.point.simplified, system.math.point.traditional]:
+ if p in chinese_string:
+ int_string, dec_string = chinese_string.split(p)
+ break
+ return [get_symbol(c, system) for c in int_string], \
+ [get_symbol(c, system) for c in dec_string]
+
+ def correct_symbols(integer_symbols, system):
+ """
+ 涓�鐧惧叓 to 涓�鐧惧叓鍗�
+ 涓�浜夸竴鍗冧笁鐧句竾 to 涓�浜� 涓�鍗冧竾 涓夌櫨涓�
+ """
+
+ if integer_symbols and isinstance(integer_symbols[0], CNU):
+ if integer_symbols[0].power == 1:
+ integer_symbols = [system.digits[1]] + integer_symbols
+
+ if len(integer_symbols) > 1:
+ if isinstance(integer_symbols[-1], CND) and isinstance(integer_symbols[-2], CNU):
+ integer_symbols.append(
+ CNU(integer_symbols[-2].power - 1, None, None, None, None))
+
+ result = []
+ unit_count = 0
+ for s in integer_symbols:
+ if isinstance(s, CND):
+ result.append(s)
+ unit_count = 0
+ elif isinstance(s, CNU):
+ current_unit = CNU(s.power, None, None, None, None)
+ unit_count += 1
+
+ if unit_count == 1:
+ result.append(current_unit)
+ elif unit_count > 1:
+ for i in range(len(result)):
+ if isinstance(result[-i - 1], CNU) and result[-i - 1].power < current_unit.power:
+ result[-i - 1] = CNU(result[-i - 1].power +
+ current_unit.power, None, None, None, None)
+ return result
+
+ def compute_value(integer_symbols):
+ """
+ Compute the value.
+ When current unit is larger than previous unit, current unit * all previous units will be used as all previous units.
+ e.g. '涓ゅ崈涓�' = 2000 * 10000 not 2000 + 10000
+ """
+ value = [0]
+ last_power = 0
+ for s in integer_symbols:
+ if isinstance(s, CND):
+ value[-1] = s.value
+ elif isinstance(s, CNU):
+ value[-1] *= pow(10, s.power)
+ if s.power > last_power:
+ value[:-1] = list(map(lambda v: v *
+ pow(10, s.power), value[:-1]))
+ last_power = s.power
+ value.append(0)
+ return sum(value)
+
+ system = create_system(numbering_type)
+ int_part, dec_part = string2symbols(chinese_string, system)
+ int_part = correct_symbols(int_part, system)
+ int_str = str(compute_value(int_part))
+ dec_str = ''.join([str(d.value) for d in dec_part])
+ if dec_part:
+ return '{0}.{1}'.format(int_str, dec_str)
+ else:
+ return int_str
+
+
+def num2chn(number_string, numbering_type=NUMBERING_TYPES[1], big=False,
+ traditional=False, alt_zero=False, alt_one=False, alt_two=True,
+ use_zeros=True, use_units=True):
+
+ def get_value(value_string, use_zeros=True):
+
+ striped_string = value_string.lstrip('0')
+
+ # record nothing if all zeros
+ if not striped_string:
+ return []
+
+ # record one digits
+ elif len(striped_string) == 1:
+ if use_zeros and len(value_string) != len(striped_string):
+ return [system.digits[0], system.digits[int(striped_string)]]
+ else:
+ return [system.digits[int(striped_string)]]
+
+ # recursively record multiple digits
+ else:
+ result_unit = next(u for u in reversed(
+ system.units) if u.power < len(striped_string))
+ result_string = value_string[:-result_unit.power]
+ return get_value(result_string) + [result_unit] + get_value(striped_string[-result_unit.power:])
+
+ system = create_system(numbering_type)
+
+ int_dec = number_string.split('.')
+ if len(int_dec) == 1:
+ int_string = int_dec[0]
+ dec_string = ""
+ elif len(int_dec) == 2:
+ int_string = int_dec[0]
+ dec_string = int_dec[1]
+ else:
+ raise ValueError(
+ "invalid input num string with more than one dot: {}".format(number_string))
+
+ if use_units and len(int_string) > 1:
+ result_symbols = get_value(int_string)
+ else:
+ result_symbols = [system.digits[int(c)] for c in int_string]
+ dec_symbols = [system.digits[int(c)] for c in dec_string]
+ if dec_string:
+ result_symbols += [system.math.point] + dec_symbols
+
+ if alt_two:
+ liang = CND(2, system.digits[2].alt_s, system.digits[2].alt_t,
+ system.digits[2].big_s, system.digits[2].big_t)
+ for i, v in enumerate(result_symbols):
+ if isinstance(v, CND) and v.value == 2:
+ next_symbol = result_symbols[i +
+ 1] if i < len(result_symbols) - 1 else None
+ previous_symbol = result_symbols[i - 1] if i > 0 else None
+ if isinstance(next_symbol, CNU) and isinstance(previous_symbol, (CNU, type(None))):
+ if next_symbol.power != 1 and ((previous_symbol is None) or (previous_symbol.power != 1)):
+ result_symbols[i] = liang
+
+ # if big is True, '涓�' will not be used and `alt_two` has no impact on output
+ if big:
+ attr_name = 'big_'
+ if traditional:
+ attr_name += 't'
+ else:
+ attr_name += 's'
+ else:
+ if traditional:
+ attr_name = 'traditional'
+ else:
+ attr_name = 'simplified'
+
+ result = ''.join([getattr(s, attr_name) for s in result_symbols])
+
+ # if not use_zeros:
+ # result = result.strip(getattr(system.digits[0], attr_name))
+
+ if alt_zero:
+ result = result.replace(
+ getattr(system.digits[0], attr_name), system.digits[0].alt_s)
+
+ if alt_one:
+ result = result.replace(
+ getattr(system.digits[1], attr_name), system.digits[1].alt_s)
+
+ for i, p in enumerate(POINT):
+ if result.startswith(p):
+ return CHINESE_DIGIS[0] + result
+
+ # ^10, 11, .., 19
+ if len(result) >= 2 and result[1] in [SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED[0],
+ SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL[0]] and \
+ result[0] in [CHINESE_DIGIS[1], BIG_CHINESE_DIGIS_SIMPLIFIED[1], BIG_CHINESE_DIGIS_TRADITIONAL[1]]:
+ result = result[1:]
+
+ return result
+
+
+# ================================================================================ #
+# different types of rewriters
+# ================================================================================ #
+class Cardinal:
+ """
+ CARDINAL绫�
+ """
+
+ def __init__(self, cardinal=None, chntext=None):
+ self.cardinal = cardinal
+ self.chntext = chntext
+
+ def chntext2cardinal(self):
+ return chn2num(self.chntext)
+
+ def cardinal2chntext(self):
+ return num2chn(self.cardinal)
+
+class Digit:
+ """
+ DIGIT绫�
+ """
+
+ def __init__(self, digit=None, chntext=None):
+ self.digit = digit
+ self.chntext = chntext
+
+ # def chntext2digit(self):
+ # return chn2num(self.chntext)
+
+ def digit2chntext(self):
+ return num2chn(self.digit, alt_two=False, use_units=False)
+
+
+class TelePhone:
+ """
+ TELEPHONE绫�
+ """
+
+ def __init__(self, telephone=None, raw_chntext=None, chntext=None):
+ self.telephone = telephone
+ self.raw_chntext = raw_chntext
+ self.chntext = chntext
+
+ # def chntext2telephone(self):
+ # sil_parts = self.raw_chntext.split('<SIL>')
+ # self.telephone = '-'.join([
+ # str(chn2num(p)) for p in sil_parts
+ # ])
+ # return self.telephone
+
+ def telephone2chntext(self, fixed=False):
+
+ if fixed:
+ sil_parts = self.telephone.split('-')
+ self.raw_chntext = '<SIL>'.join([
+ num2chn(part, alt_two=False, use_units=False) for part in sil_parts
+ ])
+ self.chntext = self.raw_chntext.replace('<SIL>', '')
+ else:
+ sp_parts = self.telephone.strip('+').split()
+ self.raw_chntext = '<SP>'.join([
+ num2chn(part, alt_two=False, use_units=False) for part in sp_parts
+ ])
+ self.chntext = self.raw_chntext.replace('<SP>', '')
+ return self.chntext
+
+
+class Fraction:
+ """
+ FRACTION绫�
+ """
+
+ def __init__(self, fraction=None, chntext=None):
+ self.fraction = fraction
+ self.chntext = chntext
+
+ def chntext2fraction(self):
+ denominator, numerator = self.chntext.split('鍒嗕箣')
+ return chn2num(numerator) + '/' + chn2num(denominator)
+
+ def fraction2chntext(self):
+ numerator, denominator = self.fraction.split('/')
+ return num2chn(denominator) + '鍒嗕箣' + num2chn(numerator)
+
+
+class Date:
+ """
+ DATE绫�
+ """
+
+ def __init__(self, date=None, chntext=None):
+ self.date = date
+ self.chntext = chntext
+
+ # def chntext2date(self):
+ # chntext = self.chntext
+ # try:
+ # year, other = chntext.strip().split('骞�', maxsplit=1)
+ # year = Digit(chntext=year).digit2chntext() + '骞�'
+ # except ValueError:
+ # other = chntext
+ # year = ''
+ # if other:
+ # try:
+ # month, day = other.strip().split('鏈�', maxsplit=1)
+ # month = Cardinal(chntext=month).chntext2cardinal() + '鏈�'
+ # except ValueError:
+ # day = chntext
+ # month = ''
+ # if day:
+ # day = Cardinal(chntext=day[:-1]).chntext2cardinal() + day[-1]
+ # else:
+ # month = ''
+ # day = ''
+ # date = year + month + day
+ # self.date = date
+ # return self.date
+
+ def date2chntext(self):
+ date = self.date
+ try:
+ year, other = date.strip().split('骞�', 1)
+ year = Digit(digit=year).digit2chntext() + '骞�'
+ except ValueError:
+ other = date
+ year = ''
+ if other:
+ try:
+ month, day = other.strip().split('鏈�', 1)
+ month = Cardinal(cardinal=month).cardinal2chntext() + '鏈�'
+ except ValueError:
+ day = date
+ month = ''
+ if day:
+ day = Cardinal(cardinal=day[:-1]).cardinal2chntext() + day[-1]
+ else:
+ month = ''
+ day = ''
+ chntext = year + month + day
+ self.chntext = chntext
+ return self.chntext
+
+
+class Money:
+ """
+ MONEY绫�
+ """
+
+ def __init__(self, money=None, chntext=None):
+ self.money = money
+ self.chntext = chntext
+
+ # def chntext2money(self):
+ # return self.money
+
+ def money2chntext(self):
+ money = self.money
+ pattern = re.compile(r'(\d+(\.\d+)?)')
+ matchers = pattern.findall(money)
+ if matchers:
+ for matcher in matchers:
+ money = money.replace(matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext())
+ self.chntext = money
+ return self.chntext
+
+
+class Percentage:
+ """
+ PERCENTAGE绫�
+ """
+
+ def __init__(self, percentage=None, chntext=None):
+ self.percentage = percentage
+ self.chntext = chntext
+
+ def chntext2percentage(self):
+ return chn2num(self.chntext.strip().strip('鐧惧垎涔�')) + '%'
+
+ def percentage2chntext(self):
+ return '鐧惧垎涔�' + num2chn(self.percentage.strip().strip('%'))
+
+
+def remove_erhua(text, er_whitelist):
+ """
+ 鍘婚櫎鍎垮寲闊宠瘝涓殑鍎�:
+ 浠栧コ鍎垮湪閭h竟鍎� -> 浠栧コ鍎垮湪閭h竟
+ """
+
+ er_pattern = re.compile(er_whitelist)
+ new_str=''
+ while re.search('鍎�',text):
+ a = re.search('鍎�',text).span()
+ remove_er_flag = 0
+
+ if er_pattern.search(text):
+ b = er_pattern.search(text).span()
+ if b[0] <= a[0]:
+ remove_er_flag = 1
+
+ if remove_er_flag == 0 :
+ new_str = new_str + text[0:a[0]]
+ text = text[a[1]:]
+ else:
+ new_str = new_str + text[0:b[1]]
+ text = text[b[1]:]
+
+ text = new_str + text
+ return text
+
+# ================================================================================ #
+# NSW Normalizer
+# ================================================================================ #
+class NSWNormalizer:
+ def __init__(self, raw_text):
+ self.raw_text = '^' + raw_text + '$'
+ self.norm_text = ''
+
+ def _particular(self):
+ text = self.norm_text
+ pattern = re.compile(r"(([a-zA-Z]+)浜�([a-zA-Z]+))")
+ matchers = pattern.findall(text)
+ if matchers:
+ # print('particular')
+ for matcher in matchers:
+ text = text.replace(matcher[0], matcher[1]+'2'+matcher[2], 1)
+ self.norm_text = text
+ return self.norm_text
+
+ def normalize(self):
+ text = self.raw_text
+
+ # 瑙勮寖鍖栨棩鏈�
+ pattern = re.compile(r"\D+((([089]\d|(19|20)\d{2})骞�)?(\d{1,2}鏈�(\d{1,2}[鏃ュ彿])?)?)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('date')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Date(date=matcher[0]).date2chntext(), 1)
+
+ # 瑙勮寖鍖栭噾閽�
+ pattern = re.compile(r"\D+((\d+(\.\d+)?)[澶氫綑鍑燷?" + CURRENCY_UNITS + r"(\d" + CURRENCY_UNITS + r"?)?)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('money')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Money(money=matcher[0]).money2chntext(), 1)
+
+ # 瑙勮寖鍖栧浐璇�/鎵嬫満鍙风爜
+ # 鎵嬫満
+ # http://www.jihaoba.com/news/show/13680
+ # 绉诲姩锛�139銆�138銆�137銆�136銆�135銆�134銆�159銆�158銆�157銆�150銆�151銆�152銆�188銆�187銆�182銆�183銆�184銆�178銆�198
+ # 鑱旈�氾細130銆�131銆�132銆�156銆�155銆�186銆�185銆�176
+ # 鐢典俊锛�133銆�153銆�189銆�180銆�181銆�177
+ pattern = re.compile(r"\D((\+?86 ?)?1([38]\d|5[0-35-9]|7[678]|9[89])\d{8})\D")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('telephone')
+ for matcher in matchers:
+ text = text.replace(matcher[0], TelePhone(telephone=matcher[0]).telephone2chntext(), 1)
+ # 鍥鸿瘽
+ pattern = re.compile(r"\D((0(10|2[1-3]|[3-9]\d{2})-?)?[1-9]\d{6,7})\D")
+ matchers = pattern.findall(text)
+ if matchers:
+ # print('fixed telephone')
+ for matcher in matchers:
+ text = text.replace(matcher[0], TelePhone(telephone=matcher[0]).telephone2chntext(fixed=True), 1)
+
+ # 瑙勮寖鍖栧垎鏁�
+ pattern = re.compile(r"(\d+/\d+)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('fraction')
+ for matcher in matchers:
+ text = text.replace(matcher, Fraction(fraction=matcher).fraction2chntext(), 1)
+
+ # 瑙勮寖鍖栫櫨鍒嗘暟
+ text = text.replace('锛�', '%')
+ pattern = re.compile(r"(\d+(\.\d+)?%)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('percentage')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Percentage(percentage=matcher[0]).percentage2chntext(), 1)
+
+ # 瑙勮寖鍖栫函鏁�+閲忚瘝
+ pattern = re.compile(r"(\d+(\.\d+)?)[澶氫綑鍑燷?" + COM_QUANTIFIERS)
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('cardinal+quantifier')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext(), 1)
+
+ # 瑙勮寖鍖栨暟瀛楃紪鍙�
+ pattern = re.compile(r"(\d{4,32})")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('digit')
+ for matcher in matchers:
+ text = text.replace(matcher, Digit(digit=matcher).digit2chntext(), 1)
+
+ # 瑙勮寖鍖栫函鏁�
+ pattern = re.compile(r"(\d+(\.\d+)?)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('cardinal')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext(), 1)
+
+ self.norm_text = text
+ self._particular()
+
+ return self.norm_text.lstrip('^').rstrip('$')
+
+
+def nsw_test_case(raw_text):
+ print('I:' + raw_text)
+ print('O:' + NSWNormalizer(raw_text).normalize())
+ print('')
+
+
+def nsw_test():
+ nsw_test_case('鍥鸿瘽锛�0595-23865596鎴�23880880銆�')
+ nsw_test_case('鍥鸿瘽锛�0595-23865596鎴�23880880銆�')
+ nsw_test_case('鎵嬫満锛�+86 19859213959鎴�15659451527銆�')
+ nsw_test_case('鍒嗘暟锛�32477/76391銆�')
+ nsw_test_case('鐧惧垎鏁帮細80.03%銆�')
+ nsw_test_case('缂栧彿锛�31520181154418銆�')
+ nsw_test_case('绾暟锛�2983.07鍏嬫垨12345.60绫炽��')
+ nsw_test_case('鏃ユ湡锛�1999骞�2鏈�20鏃ユ垨09骞�3鏈�15鍙枫��')
+ nsw_test_case('閲戦挶锛�12鍧�5锛�34.5鍏冿紝20.1涓�')
+ nsw_test_case('鐗规畩锛歄2O鎴朆2C銆�')
+ nsw_test_case('3456涓囧惃')
+ nsw_test_case('2938涓�')
+ nsw_test_case('938')
+ nsw_test_case('浠婂ぉ鍚冧簡115涓皬绗煎寘231涓澶�')
+ nsw_test_case('鏈�62锛呯殑姒傜巼')
+
+
+if __name__ == '__main__':
+ #nsw_test()
+
+ p = argparse.ArgumentParser()
+ p.add_argument('ifile', help='input filename, assume utf-8 encoding')
+ p.add_argument('ofile', help='output filename')
+ p.add_argument('--to_upper', action='store_true', help='convert to upper case')
+ p.add_argument('--to_lower', action='store_true', help='convert to lower case')
+ p.add_argument('--has_key', action='store_true', help="input text has Kaldi's key as first field.")
+ p.add_argument('--remove_fillers', type=bool, default=True, help='remove filler chars such as "鍛�, 鍟�"')
+ p.add_argument('--remove_erhua', type=bool, default=True, help='remove erhua chars such as "杩欏効"')
+ p.add_argument('--log_interval', type=int, default=10000, help='log interval in number of processed lines')
+ args = p.parse_args()
+
+ ifile = codecs.open(args.ifile, 'r', 'utf8')
+ ofile = codecs.open(args.ofile, 'w+', 'utf8')
+
+ n = 0
+ for l in ifile:
+ key = ''
+ text = ''
+ if args.has_key:
+ cols = l.split(maxsplit=1)
+ key = cols[0]
+ if len(cols) == 2:
+ text = cols[1].strip()
+ else:
+ text = ''
+ else:
+ text = l.strip()
+
+ # cases
+ if args.to_upper and args.to_lower:
+ sys.stderr.write('text norm: to_upper OR to_lower?')
+ exit(1)
+ if args.to_upper:
+ text = text.upper()
+ if args.to_lower:
+ text = text.lower()
+
+ # Filler chars removal
+ if args.remove_fillers:
+ for ch in FILLER_CHARS:
+ text = text.replace(ch, '')
+
+ if args.remove_erhua:
+ text = remove_erhua(text, ER_WHITELIST)
+
+ # NSW(Non-Standard-Word) normalization
+ text = NSWNormalizer(text).normalize()
+
+ # Punctuations removal
+ old_chars = CHINESE_PUNC_LIST + string.punctuation # includes all CN and EN punctuations
+ new_chars = ' ' * len(old_chars)
+ del_chars = ''
+ text = text.translate(str.maketrans(old_chars, new_chars, del_chars))
+
+ #
+ if args.has_key:
+ ofile.write(key + '\t' + text + '\n')
+ else:
+ ofile.write(text + '\n')
+
+ n += 1
+ if n % args.log_interval == 0:
+ sys.stderr.write("text norm: {} lines done.\n".format(n))
+
+ sys.stderr.write("text norm: {} lines done in total.\n".format(n))
+
+ ifile.close()
+ ofile.close()
diff --git a/egs_modelscope/common/modelscope_common_finetune.sh b/egs_modelscope/common/modelscope_common_finetune.sh
index 39cd03e..63143c8 100755
--- a/egs_modelscope/common/modelscope_common_finetune.sh
+++ b/egs_modelscope/common/modelscope_common_finetune.sh
@@ -9,6 +9,7 @@
gpu_inference=true # Whether to perform gpu decoding, set false for cpu decoding
njob=4 # the number of jobs for each gpu
train_cmd=utils/run.pl
+infer_cmd=utils/run.pl
# general configuration
feats_dir="../DATA" #feature output dictionary, for large data
@@ -32,7 +33,7 @@
lfr_n=6
init_model_name=speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-pytorch # pre-trained model, download from modelscope during fine-tuning
-model_revision="v1.0.3" # please do not modify the model revision
+model_revision="v1.0.4" # please do not modify the model revision
cmvn_file=init_model/${init_model_name}/am.mvn
seg_file=init_model/${init_model_name}/seg_dict
vocab=init_model/${init_model_name}/tokens.txt
@@ -81,7 +82,14 @@
# you can set gpu num for decoding here
gpuid_list=$CUDA_VISIBLE_DEVICES # set gpus for decoding, the same as training stage by default
ngpu=$(echo $gpuid_list | awk -F "," '{print NF}')
-inference_nj=$[${ngpu}*${njob}]
+
+if ${gpu_inference}; then
+ inference_nj=$[${ngpu}*${njob}]
+ _ngpu=1
+else
+ inference_nj=$njob
+ _ngpu=0
+fi
[ ! -d ${dataset} ] && echo "$0: Training data is required" && exit 1;
[ ! -f ${dataset}/train/wav.scp ] && [ ! -f ${dataset}/train/text ] && echo "$0: Training data wav.scp or text is not found" && exit 1;
@@ -99,17 +107,17 @@
feat_test_dir=${feats_dir}/${dumpdir}/test; mkdir -p ${feat_test_dir}
if [ ${stage} -le 1 ] && [ ${stop_stage} -ge 1 ]; then
- echo "Feature Generation"
+ echo "stage 1: Feature Generation"
# compute fbank features
fbankdir=${feats_dir}/fbank
- utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj --speed_perturb ${speed_perturb} \
+ utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj --sample_frequency ${sample_frequency} --speed_perturb ${speed_perturb} \
${dataset}/train ${exp_dir}/exp/make_fbank/train ${fbankdir}/train
utils/fix_data_feat.sh ${fbankdir}/train
- utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj \
+ utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj --sample_frequency ${sample_frequency} \
${dataset}/dev ${exp_dir}/exp/make_fbank/dev ${fbankdir}/dev
utils/fix_data_feat.sh ${fbankdir}/dev
if [ -d "${dataset}/test" ]; then
- utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj \
+ utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj --sample_frequency ${sample_frequency} \
${dataset}/test ${exp_dir}/exp/make_fbank/test ${fbankdir}/test
utils/fix_data_feat.sh ${fbankdir}/test
fi
@@ -158,13 +166,14 @@
# Training Stage
world_size=$gpu_num # run on one machine
if [ ${stage} -le 3 ] && [ ${stop_stage} -ge 3 ]; then
+ echo "stage 3: Training"
# update asr train config.yaml
python modelscope_utils/update_config.py --modelscope_config init_model/${init_model_name}/finetune.yaml --finetune_config ${asr_config} --output_config init_model/${init_model_name}/asr_finetune_config.yaml
finetune_config=init_model/${init_model_name}/asr_finetune_config.yaml
mkdir -p ${exp_dir}/exp/${model_dir}
mkdir -p ${exp_dir}/exp/${model_dir}/log
- INIT_FILE=$exp_dir/ddp_init
+ INIT_FILE=${exp_dir}/exp/${model_dir}/ddp_init
if [ -f $INIT_FILE ];then
rm -f $INIT_FILE
fi
@@ -206,26 +215,58 @@
fi
# Testing Stage
+# Testing Stage
if [ ${stage} -le 4 ] && [ ${stop_stage} -ge 4 ]; then
- ./utils/easy_asr_infer.sh \
- --lang zh \
- --datadir ${feats_dir} \
- --feats_type ${feats_type} \
- --feats_dim ${feats_dim} \
- --token_type ${token_type} \
- --gpu_inference ${gpu_inference} \
- --inference_config "${inference_config}" \
- --test_sets "${test_sets}" \
- --token_list $token_list \
- --asr_exp ${exp_dir}/exp/${model_dir} \
- --stage 12 \
- --stop_stage 12 \
- --scp $scp \
- --text text \
- --inference_nj $inference_nj \
- --njob $njob \
- --inference_asr_model $inference_asr_model \
- --gpuid_list $gpuid_list \
- --mode paraformer
-fi
+ echo "stage 4: Inference"
+ for dset in ${test_sets}; do
+ asr_exp=${exp_dir}/exp/${model_dir}
+ inference_tag="$(basename "${inference_config}" .yaml)"
+ _dir="${asr_exp}/${inference_tag}/${inference_asr_model}/${dset}"
+ _logdir="${_dir}/logdir"
+ if [ -d ${_dir} ]; then
+ echo "${_dir} is already exists. if you want to decode again, please delete this dir first."
+ exit 0
+ fi
+ mkdir -p "${_logdir}"
+ _data="${feats_dir}/${dumpdir}/${dset}"
+ key_file=${_data}/${scp}
+ num_scp_file="$(<${key_file} wc -l)"
+ _nj=$([ $inference_nj -le $num_scp_file ] && echo "$inference_nj" || echo "$num_scp_file")
+ split_scps=
+ for n in $(seq "${_nj}"); do
+ split_scps+=" ${_logdir}/keys.${n}.scp"
+ done
+ # shellcheck disable=SC2086
+ utils/split_scp.pl "${key_file}" ${split_scps}
+ _opts=
+ if [ -n "${inference_config}" ]; then
+ _opts+="--config ${inference_config} "
+ fi
+ ${infer_cmd} --gpu "${_ngpu}" --max-jobs-run "${_nj}" JOB=1:"${_nj}" "${_logdir}"/asr_inference.JOB.log \
+ python -m funasr.bin.asr_inference_launch \
+ --batch_size 1 \
+ --ngpu "${_ngpu}" \
+ --njob ${njob} \
+ --gpuid_list ${gpuid_list} \
+ --data_path_and_name_and_type "${_data}/${scp},speech,${type}" \
+ --key_file "${_logdir}"/keys.JOB.scp \
+ --asr_train_config "${asr_exp}"/config.yaml \
+ --asr_model_file "${asr_exp}"/"${inference_asr_model}" \
+ --output_dir "${_logdir}"/output.JOB \
+ --mode paraformer \
+ ${_opts}
+ for f in token token_int score text; do
+ if [ -f "${_logdir}/output.1/1best_recog/${f}" ]; then
+ for i in $(seq "${_nj}"); do
+ cat "${_logdir}/output.${i}/1best_recog/${f}"
+ done | sort -k1 >"${_dir}/${f}"
+ fi
+ done
+ python utils/proce_text.py ${_dir}/text ${_dir}/text.proc
+ python utils/proce_text.py ${_data}/text ${_data}/text.proc
+ python utils/compute_wer.py ${_data}/text.proc ${_dir}/text.proc ${_dir}/text.cer
+ tail -n 3 ${_dir}/text.cer > ${_dir}/text.cer.txt
+ cat ${_dir}/text.cer.txt
+ done
+fi
diff --git a/egs_modelscope/common/modelscope_common_infer.sh b/egs_modelscope/common/modelscope_common_infer.sh
index 9c607cb..f0dc459 100755
--- a/egs_modelscope/common/modelscope_common_infer.sh
+++ b/egs_modelscope/common/modelscope_common_infer.sh
@@ -5,7 +5,7 @@
set -o pipefail
model_name=speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-pytorch # pre-trained model, download from modelscope
-model_revision="v1.0.3" # please do not modify the model revision
+model_revision="v1.0.4" # please do not modify the model revision
data_dir= # wav list, ${data_dir}/wav.scp
exp_dir="exp"
gpuid_list="0,1"
@@ -71,5 +71,6 @@
cat ${_logdir}/text.${i}
done | sort -k1 >${_dir}/text
-mv ${exp_dir}/${model_name}/decoding.yaml.back ${exp_dir}/${model_name}/decoding.yaml
-
+if "${use_lm}"; then
+ mv ${exp_dir}/${model_name}/decoding.yaml.back ${exp_dir}/${model_name}/decoding.yaml
+fi
diff --git a/egs_modelscope/common/utils b/egs_modelscope/common/utils
deleted file mode 120000
index cbef564..0000000
--- a/egs_modelscope/common/utils
+++ /dev/null
@@ -1 +0,0 @@
-../../egs/aishell/tranformer/utils/
\ No newline at end of file
diff --git a/egs_modelscope/common/utils/__init__.py b/egs_modelscope/common/utils/__init__.py
new file mode 100644
index 0000000..e69de29
--- /dev/null
+++ b/egs_modelscope/common/utils/__init__.py
diff --git a/egs_modelscope/common/utils/apply_cmvn.py b/egs_modelscope/common/utils/apply_cmvn.py
new file mode 100755
index 0000000..b5c5086
--- /dev/null
+++ b/egs_modelscope/common/utils/apply_cmvn.py
@@ -0,0 +1,79 @@
+from kaldiio import ReadHelper
+from kaldiio import WriteHelper
+
+import argparse
+import json
+import math
+import numpy as np
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="apply cmvn",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--ark-file",
+ "-a",
+ default=False,
+ required=True,
+ type=str,
+ help="fbank ark file",
+ )
+ parser.add_argument(
+ "--cmvn-file",
+ "-c",
+ default=False,
+ required=True,
+ type=str,
+ help="cmvn file",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ ark_file = args.output_dir + "/feats." + str(args.ark_index) + ".ark"
+ scp_file = args.output_dir + "/feats." + str(args.ark_index) + ".scp"
+ ark_writer = WriteHelper('ark,scp:{},{}'.format(ark_file, scp_file))
+
+ with open(args.cmvn_file) as f:
+ cmvn_stats = json.load(f)
+
+ means = cmvn_stats['mean_stats']
+ vars = cmvn_stats['var_stats']
+ total_frames = cmvn_stats['total_frames']
+
+ for i in range(len(means)):
+ means[i] /= total_frames
+ vars[i] = vars[i] / total_frames - means[i] * means[i]
+ if vars[i] < 1.0e-20:
+ vars[i] = 1.0e-20
+ vars[i] = 1.0 / math.sqrt(vars[i])
+
+ with ReadHelper('ark:{}'.format(args.ark_file)) as ark_reader:
+ for key, mat in ark_reader:
+ mat = (mat - means) * vars
+ ark_writer(key, mat)
+
+
+if __name__ == '__main__':
+ main()
diff --git a/egs_modelscope/common/utils/apply_cmvn.sh b/egs_modelscope/common/utils/apply_cmvn.sh
new file mode 100755
index 0000000..f8fd1d1
--- /dev/null
+++ b/egs_modelscope/common/utils/apply_cmvn.sh
@@ -0,0 +1,29 @@
+#!/usr/bin/env bash
+
+. ./path.sh || exit 1;
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+fbankdir=$1
+cmvn_file=$2
+logdir=$3
+output_dir=$4
+
+dump_dir=${output_dir}/ark; mkdir -p ${dump_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/apply_cmvn.JOB.log \
+ python utils/apply_cmvn.py -a $fbankdir/ark/feats.JOB.ark \
+ -c $cmvn_file -i JOB -o ${dump_dir} \
+ || exit 1;
+
+for n in $(seq $nj); do
+ cat ${dump_dir}/feats.$n.scp || exit 1
+done > ${output_dir}/feats.scp || exit 1
+
+echo "$0: Succeeded apply cmvn"
diff --git a/egs_modelscope/common/utils/apply_lfr_and_cmvn.py b/egs_modelscope/common/utils/apply_lfr_and_cmvn.py
new file mode 100755
index 0000000..50d18d1
--- /dev/null
+++ b/egs_modelscope/common/utils/apply_lfr_and_cmvn.py
@@ -0,0 +1,143 @@
+from kaldiio import ReadHelper, WriteHelper
+
+import argparse
+import numpy as np
+
+
+def build_LFR_features(inputs, m=7, n=6):
+ LFR_inputs = []
+ T = inputs.shape[0]
+ T_lfr = int(np.ceil(T / n))
+ left_padding = np.tile(inputs[0], ((m - 1) // 2, 1))
+ inputs = np.vstack((left_padding, inputs))
+ T = T + (m - 1) // 2
+ for i in range(T_lfr):
+ if m <= T - i * n:
+ LFR_inputs.append(np.hstack(inputs[i * n:i * n + m]))
+ else:
+ num_padding = m - (T - i * n)
+ frame = np.hstack(inputs[i * n:])
+ for _ in range(num_padding):
+ frame = np.hstack((frame, inputs[-1]))
+ LFR_inputs.append(frame)
+ return np.vstack(LFR_inputs)
+
+
+def build_CMVN_features(inputs, mvn_file): # noqa
+ with open(mvn_file, 'r', encoding='utf-8') as f:
+ lines = f.readlines()
+
+ add_shift_list = []
+ rescale_list = []
+ for i in range(len(lines)):
+ line_item = lines[i].split()
+ if line_item[0] == '<AddShift>':
+ line_item = lines[i + 1].split()
+ if line_item[0] == '<LearnRateCoef>':
+ add_shift_line = line_item[3:(len(line_item) - 1)]
+ add_shift_list = list(add_shift_line)
+ continue
+ elif line_item[0] == '<Rescale>':
+ line_item = lines[i + 1].split()
+ if line_item[0] == '<LearnRateCoef>':
+ rescale_line = line_item[3:(len(line_item) - 1)]
+ rescale_list = list(rescale_line)
+ continue
+
+ for j in range(inputs.shape[0]):
+ for k in range(inputs.shape[1]):
+ add_shift_value = add_shift_list[k]
+ rescale_value = rescale_list[k]
+ inputs[j, k] = float(inputs[j, k]) + float(add_shift_value)
+ inputs[j, k] = float(inputs[j, k]) * float(rescale_value)
+
+ return inputs
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="apply low_frame_rate and cmvn",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--ark-file",
+ "-a",
+ default=False,
+ required=True,
+ type=str,
+ help="fbank ark file",
+ )
+ parser.add_argument(
+ "--lfr",
+ "-f",
+ default=True,
+ type=str,
+ help="low frame rate",
+ )
+ parser.add_argument(
+ "--lfr-m",
+ "-m",
+ default=7,
+ type=int,
+ help="number of frames to stack",
+ )
+ parser.add_argument(
+ "--lfr-n",
+ "-n",
+ default=6,
+ type=int,
+ help="number of frames to skip",
+ )
+ parser.add_argument(
+ "--cmvn-file",
+ "-c",
+ default=False,
+ required=True,
+ type=str,
+ help="global cmvn file",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ dump_ark_file = args.output_dir + "/feats." + str(args.ark_index) + ".ark"
+ dump_scp_file = args.output_dir + "/feats." + str(args.ark_index) + ".scp"
+ shape_file = args.output_dir + "/len." + str(args.ark_index)
+ ark_writer = WriteHelper('ark,scp:{},{}'.format(dump_ark_file, dump_scp_file))
+
+ shape_writer = open(shape_file, 'w')
+ with ReadHelper('ark:{}'.format(args.ark_file)) as ark_reader:
+ for key, mat in ark_reader:
+ if args.lfr:
+ lfr = build_LFR_features(mat, args.lfr_m, args.lfr_n)
+ else:
+ lfr = mat
+ cmvn = build_CMVN_features(lfr, args.cmvn_file)
+ dims = cmvn.shape[1]
+ lens = cmvn.shape[0]
+ shape_writer.write(key + " " + str(lens) + "," + str(dims) + '\n')
+ ark_writer(key, cmvn)
+
+
+if __name__ == '__main__':
+ main()
+
diff --git a/egs_modelscope/common/utils/apply_lfr_and_cmvn.sh b/egs_modelscope/common/utils/apply_lfr_and_cmvn.sh
new file mode 100755
index 0000000..3119fdb
--- /dev/null
+++ b/egs_modelscope/common/utils/apply_lfr_and_cmvn.sh
@@ -0,0 +1,38 @@
+#!/usr/bin/env bash
+
+
+# Begin configuration section.
+nj=32
+cmd=utils/run.pl
+
+# feature configuration
+lfr=True
+lfr_m=7
+lfr_n=6
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+fbankdir=$1
+cmvn_file=$2
+logdir=$3
+output_dir=$4
+
+dump_dir=${output_dir}/ark; mkdir -p ${dump_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/apply_lfr_and_cmvn.JOB.log \
+ python utils/apply_lfr_and_cmvn.py -a $fbankdir/ark/feats.JOB.ark \
+ -f $lfr -m $lfr_m -n $lfr_n -c $cmvn_file -i JOB -o ${dump_dir} \
+ || exit 1;
+
+for n in $(seq $nj); do
+ cat ${dump_dir}/feats.$n.scp || exit 1
+done > ${output_dir}/feats.scp || exit 1
+
+for n in $(seq $nj); do
+ cat ${dump_dir}/len.$n || exit 1
+done > ${output_dir}/speech_shape || exit 1
+
+echo "$0: Succeeded apply low frame rate and cmvn"
diff --git a/egs_modelscope/common/utils/combine_cmvn_file.py b/egs_modelscope/common/utils/combine_cmvn_file.py
new file mode 100755
index 0000000..b2974a4
--- /dev/null
+++ b/egs_modelscope/common/utils/combine_cmvn_file.py
@@ -0,0 +1,73 @@
+import argparse
+import json
+import numpy as np
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="combine cmvn file",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--dims",
+ "-d",
+ default=80,
+ type=int,
+ help="feature dims",
+ )
+ parser.add_argument(
+ "--cmvn-dir",
+ "-c",
+ default=False,
+ required=True,
+ type=str,
+ help="cmvn dir",
+ )
+
+ parser.add_argument(
+ "--nj",
+ "-n",
+ default=1,
+ required=True,
+ type=int,
+ help="num of cmvn file",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ total_means = np.zeros(args.dims)
+ total_vars = np.zeros(args.dims)
+ total_frames = 0
+
+ cmvn_file = args.output_dir + "/cmvn.json"
+
+ for i in range(1, args.nj+1):
+ with open(args.cmvn_dir + "/cmvn." + str(i) + ".json", "r") as fin:
+ cmvn_stats = json.load(fin)
+
+ total_means += np.array(cmvn_stats["mean_stats"])
+ total_vars += np.array(cmvn_stats["var_stats"])
+ total_frames += cmvn_stats["total_frames"]
+
+ cmvn_info = {
+ 'mean_stats': list(total_means.tolist()),
+ 'var_stats': list(total_vars.tolist()),
+ 'total_frames': total_frames
+ }
+ with open(cmvn_file, 'w') as fout:
+ fout.write(json.dumps(cmvn_info))
+
+
+if __name__ == '__main__':
+ main()
diff --git a/egs_modelscope/common/utils/compute_cmvn.py b/egs_modelscope/common/utils/compute_cmvn.py
new file mode 100755
index 0000000..2b96e26
--- /dev/null
+++ b/egs_modelscope/common/utils/compute_cmvn.py
@@ -0,0 +1,74 @@
+from kaldiio import ReadHelper
+
+import argparse
+import numpy as np
+import json
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="computer global cmvn",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--dims",
+ "-d",
+ default=80,
+ type=int,
+ help="feature dims",
+ )
+ parser.add_argument(
+ "--ark-file",
+ "-a",
+ default=False,
+ required=True,
+ type=str,
+ help="fbank ark file",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ ark_file = args.ark_file + "/feats." + str(args.ark_index) + ".ark"
+ cmvn_file = args.output_dir + "/cmvn." + str(args.ark_index) + ".json"
+
+ mean_stats = np.zeros(args.dims)
+ var_stats = np.zeros(args.dims)
+ total_frames = 0
+
+ with ReadHelper('ark:{}'.format(ark_file)) as ark_reader:
+ for key, mat in ark_reader:
+ mean_stats += np.sum(mat, axis=0)
+ var_stats += np.sum(np.square(mat), axis=0)
+ total_frames += mat.shape[0]
+
+ cmvn_info = {
+ 'mean_stats': list(mean_stats.tolist()),
+ 'var_stats': list(var_stats.tolist()),
+ 'total_frames': total_frames
+ }
+ with open(cmvn_file, 'w') as fout:
+ fout.write(json.dumps(cmvn_info))
+
+
+if __name__ == '__main__':
+ main()
diff --git a/egs_modelscope/common/utils/compute_cmvn.sh b/egs_modelscope/common/utils/compute_cmvn.sh
new file mode 100755
index 0000000..12173ee
--- /dev/null
+++ b/egs_modelscope/common/utils/compute_cmvn.sh
@@ -0,0 +1,25 @@
+#!/usr/bin/env bash
+
+. ./path.sh || exit 1;
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+feats_dim=80
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+fbankdir=$1
+logdir=$2
+
+output_dir=${fbankdir}/cmvn; mkdir -p ${output_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/cmvn.JOB.log \
+ python utils/compute_cmvn.py -d ${feats_dim} -a $fbankdir/ark -i JOB -o ${output_dir} \
+ || exit 1;
+
+python utils/combine_cmvn_file.py -d ${feats_dim} -c ${output_dir} -n $nj -o $fbankdir
+
+echo "$0: Succeeded compute global cmvn"
diff --git a/egs_modelscope/common/utils/compute_fbank.py b/egs_modelscope/common/utils/compute_fbank.py
new file mode 100755
index 0000000..d03b5a8
--- /dev/null
+++ b/egs_modelscope/common/utils/compute_fbank.py
@@ -0,0 +1,153 @@
+from kaldiio import WriteHelper
+
+import argparse
+import numpy as np
+import json
+import torch
+import torchaudio
+import torchaudio.compliance.kaldi as kaldi
+
+
+def compute_fbank(wav_file,
+ num_mel_bins=80,
+ frame_length=25,
+ frame_shift=10,
+ dither=0.0,
+ resample_rate=16000,
+ speed=1.0):
+
+ waveform, sample_rate = torchaudio.load(wav_file)
+ if resample_rate != sample_rate:
+ waveform = torchaudio.transforms.Resample(orig_freq=sample_rate,
+ new_freq=resample_rate)(waveform)
+ if speed != 1.0:
+ waveform, _ = torchaudio.sox_effects.apply_effects_tensor(
+ waveform, resample_rate,
+ [['speed', str(speed)], ['rate', str(resample_rate)]]
+ )
+
+ waveform = waveform * (1 << 15)
+ mat = kaldi.fbank(waveform,
+ num_mel_bins=num_mel_bins,
+ frame_length=frame_length,
+ frame_shift=frame_shift,
+ dither=dither,
+ energy_floor=0.0,
+ window_type='hamming',
+ sample_frequency=resample_rate)
+
+ return mat.numpy()
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="computer features",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--wav-lists",
+ "-w",
+ default=False,
+ required=True,
+ type=str,
+ help="input wav lists",
+ )
+ parser.add_argument(
+ "--text-files",
+ "-t",
+ default=False,
+ required=True,
+ type=str,
+ help="input text files",
+ )
+ parser.add_argument(
+ "--dims",
+ "-d",
+ default=80,
+ type=int,
+ help="feature dims",
+ )
+ parser.add_argument(
+ "--sample-frequency",
+ "-s",
+ default=16000,
+ type=int,
+ help="sample frequency",
+ )
+ parser.add_argument(
+ "--speed-perturb",
+ "-p",
+ default="1.0",
+ type=str,
+ help="speed perturb",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-a",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ ark_file = args.output_dir + "/ark/feats." + str(args.ark_index) + ".ark"
+ scp_file = args.output_dir + "/ark/feats." + str(args.ark_index) + ".scp"
+ text_file = args.output_dir + "/txt/text." + str(args.ark_index) + ".txt"
+ feats_shape_file = args.output_dir + "/ark/len." + str(args.ark_index)
+ text_shape_file = args.output_dir + "/txt/len." + str(args.ark_index)
+
+ ark_writer = WriteHelper('ark,scp:{},{}'.format(ark_file, scp_file))
+ text_writer = open(text_file, 'w')
+ feats_shape_writer = open(feats_shape_file, 'w')
+ text_shape_writer = open(text_shape_file, 'w')
+
+ speed_perturb_list = args.speed_perturb.split(',')
+
+ for speed in speed_perturb_list:
+ with open(args.wav_lists, 'r', encoding='utf-8') as wavfile:
+ with open(args.text_files, 'r', encoding='utf-8') as textfile:
+ for wav, text in zip(wavfile, textfile):
+ s_w = wav.strip().split()
+ wav_id = s_w[0]
+ wav_file = s_w[1]
+
+ s_t = text.strip().split()
+ text_id = s_t[0]
+ txt = s_t[1:]
+ fbank = compute_fbank(wav_file,
+ num_mel_bins=args.dims,
+ resample_rate=args.sample_frequency,
+ speed=float(speed)
+ )
+ feats_dims = fbank.shape[1]
+ feats_lens = fbank.shape[0]
+ txt_lens = len(txt)
+ if speed == "1.0":
+ wav_id_sp = wav_id
+ else:
+ wav_id_sp = wav_id + "_sp" + speed
+
+ feats_shape_writer.write(wav_id_sp + " " + str(feats_lens) + "," + str(feats_dims) + '\n')
+ text_shape_writer.write(wav_id_sp + " " + str(txt_lens) + '\n')
+
+ text_writer.write(wav_id_sp + " " + " ".join(txt) + '\n')
+ ark_writer(wav_id_sp, fbank)
+
+
+if __name__ == '__main__':
+ main()
+
diff --git a/egs_modelscope/common/utils/compute_fbank.sh b/egs_modelscope/common/utils/compute_fbank.sh
new file mode 100755
index 0000000..92a4fe6
--- /dev/null
+++ b/egs_modelscope/common/utils/compute_fbank.sh
@@ -0,0 +1,51 @@
+#!/usr/bin/env bash
+
+. ./path.sh || exit 1;
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+
+# feature configuration
+feats_dim=80
+sample_frequency=16000
+speed_perturb="1.0"
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+data=$1
+logdir=$2
+fbankdir=$3
+
+[ ! -f $data/wav.scp ] && echo "$0: no such file $data/wav.scp" && exit 1;
+[ ! -f $data/text ] && echo "$0: no such file $data/text" && exit 1;
+
+python utils/split_data.py $data $data $nj
+
+ark_dir=${fbankdir}/ark; mkdir -p ${ark_dir}
+text_dir=${fbankdir}/txt; mkdir -p ${text_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/make_fbank.JOB.log \
+ python utils/compute_fbank.py -w $data/split${nj}/JOB/wav.scp -t $data/split${nj}/JOB/text \
+ -d $feats_dim -s $sample_frequency -p ${speed_perturb} -a JOB -o ${fbankdir} \
+ || exit 1;
+
+for n in $(seq $nj); do
+ cat ${ark_dir}/feats.$n.scp || exit 1
+done > $fbankdir/feats.scp || exit 1
+
+for n in $(seq $nj); do
+ cat ${text_dir}/text.$n.txt || exit 1
+done > $fbankdir/text || exit 1
+
+for n in $(seq $nj); do
+ cat ${ark_dir}/len.$n || exit 1
+done > $fbankdir/speech_shape || exit 1
+
+for n in $(seq $nj); do
+ cat ${text_dir}/len.$n || exit 1
+done > $fbankdir/text_shape || exit 1
+
+echo "$0: Succeeded compute FBANK features"
diff --git a/egs_modelscope/common/utils/compute_wer.py b/egs_modelscope/common/utils/compute_wer.py
new file mode 100755
index 0000000..349a3f6
--- /dev/null
+++ b/egs_modelscope/common/utils/compute_wer.py
@@ -0,0 +1,157 @@
+import os
+import numpy as np
+import sys
+
+def compute_wer(ref_file,
+ hyp_file,
+ cer_detail_file):
+ rst = {
+ 'Wrd': 0,
+ 'Corr': 0,
+ 'Ins': 0,
+ 'Del': 0,
+ 'Sub': 0,
+ 'Snt': 0,
+ 'Err': 0.0,
+ 'S.Err': 0.0,
+ 'wrong_words': 0,
+ 'wrong_sentences': 0
+ }
+
+ hyp_dict = {}
+ ref_dict = {}
+ with open(hyp_file, 'r') as hyp_reader:
+ for line in hyp_reader:
+ key = line.strip().split()[0]
+ value = line.strip().split()[1:]
+ hyp_dict[key] = value
+ with open(ref_file, 'r') as ref_reader:
+ for line in ref_reader:
+ key = line.strip().split()[0]
+ value = line.strip().split()[1:]
+ ref_dict[key] = value
+
+ cer_detail_writer = open(cer_detail_file, 'w')
+ for hyp_key in hyp_dict:
+ if hyp_key in ref_dict:
+ out_item = compute_wer_by_line(hyp_dict[hyp_key], ref_dict[hyp_key])
+ rst['Wrd'] += out_item['nwords']
+ rst['Corr'] += out_item['cor']
+ rst['wrong_words'] += out_item['wrong']
+ rst['Ins'] += out_item['ins']
+ rst['Del'] += out_item['del']
+ rst['Sub'] += out_item['sub']
+ rst['Snt'] += 1
+ if out_item['wrong'] > 0:
+ rst['wrong_sentences'] += 1
+ cer_detail_writer.write(hyp_key + print_cer_detail(out_item) + '\n')
+ cer_detail_writer.write("ref:" + '\t' + "".join(ref_dict[hyp_key]) + '\n')
+ cer_detail_writer.write("hyp:" + '\t' + "".join(hyp_dict[hyp_key]) + '\n')
+
+ if rst['Wrd'] > 0:
+ rst['Err'] = round(rst['wrong_words'] * 100 / rst['Wrd'], 2)
+ if rst['Snt'] > 0:
+ rst['S.Err'] = round(rst['wrong_sentences'] * 100 / rst['Snt'], 2)
+
+ cer_detail_writer.write('\n')
+ cer_detail_writer.write("%WER " + str(rst['Err']) + " [ " + str(rst['wrong_words'])+ " / " + str(rst['Wrd']) +
+ ", " + str(rst['Ins']) + " ins, " + str(rst['Del']) + " del, " + str(rst['Sub']) + " sub ]" + '\n')
+ cer_detail_writer.write("%SER " + str(rst['S.Err']) + " [ " + str(rst['wrong_sentences']) + " / " + str(rst['Snt']) + " ]" + '\n')
+ cer_detail_writer.write("Scored " + str(len(hyp_dict)) + " sentences, " + str(len(hyp_dict) - rst['Snt']) + " not present in hyp." + '\n')
+
+
+def compute_wer_by_line(hyp,
+ ref):
+ hyp = list(map(lambda x: x.lower(), hyp))
+ ref = list(map(lambda x: x.lower(), ref))
+
+ len_hyp = len(hyp)
+ len_ref = len(ref)
+
+ cost_matrix = np.zeros((len_hyp + 1, len_ref + 1), dtype=np.int16)
+
+ ops_matrix = np.zeros((len_hyp + 1, len_ref + 1), dtype=np.int8)
+
+ for i in range(len_hyp + 1):
+ cost_matrix[i][0] = i
+ for j in range(len_ref + 1):
+ cost_matrix[0][j] = j
+
+ for i in range(1, len_hyp + 1):
+ for j in range(1, len_ref + 1):
+ if hyp[i - 1] == ref[j - 1]:
+ cost_matrix[i][j] = cost_matrix[i - 1][j - 1]
+ else:
+ substitution = cost_matrix[i - 1][j - 1] + 1
+ insertion = cost_matrix[i - 1][j] + 1
+ deletion = cost_matrix[i][j - 1] + 1
+
+ compare_val = [substitution, insertion, deletion]
+
+ min_val = min(compare_val)
+ operation_idx = compare_val.index(min_val) + 1
+ cost_matrix[i][j] = min_val
+ ops_matrix[i][j] = operation_idx
+
+ match_idx = []
+ i = len_hyp
+ j = len_ref
+ rst = {
+ 'nwords': len_ref,
+ 'cor': 0,
+ 'wrong': 0,
+ 'ins': 0,
+ 'del': 0,
+ 'sub': 0
+ }
+ while i >= 0 or j >= 0:
+ i_idx = max(0, i)
+ j_idx = max(0, j)
+
+ if ops_matrix[i_idx][j_idx] == 0: # correct
+ if i - 1 >= 0 and j - 1 >= 0:
+ match_idx.append((j - 1, i - 1))
+ rst['cor'] += 1
+
+ i -= 1
+ j -= 1
+
+ elif ops_matrix[i_idx][j_idx] == 2: # insert
+ i -= 1
+ rst['ins'] += 1
+
+ elif ops_matrix[i_idx][j_idx] == 3: # delete
+ j -= 1
+ rst['del'] += 1
+
+ elif ops_matrix[i_idx][j_idx] == 1: # substitute
+ i -= 1
+ j -= 1
+ rst['sub'] += 1
+
+ if i < 0 and j >= 0:
+ rst['del'] += 1
+ elif j < 0 and i >= 0:
+ rst['ins'] += 1
+
+ match_idx.reverse()
+ wrong_cnt = cost_matrix[len_hyp][len_ref]
+ rst['wrong'] = wrong_cnt
+
+ return rst
+
+def print_cer_detail(rst):
+ return ("(" + "nwords=" + str(rst['nwords']) + ",cor=" + str(rst['cor'])
+ + ",ins=" + str(rst['ins']) + ",del=" + str(rst['del']) + ",sub="
+ + str(rst['sub']) + ") corr:" + '{:.2%}'.format(rst['cor']/rst['nwords'])
+ + ",cer:" + '{:.2%}'.format(rst['wrong']/rst['nwords']))
+
+if __name__ == '__main__':
+ if len(sys.argv) != 4:
+ print("usage : python compute-wer.py test.ref test.hyp test.wer")
+ sys.exit(0)
+
+ ref_file = sys.argv[1]
+ hyp_file = sys.argv[2]
+ cer_detail_file = sys.argv[3]
+ compute_wer(ref_file, hyp_file, cer_detail_file)
diff --git a/egs_modelscope/common/utils/error_rate_zh b/egs_modelscope/common/utils/error_rate_zh
new file mode 100755
index 0000000..6871a07
--- /dev/null
+++ b/egs_modelscope/common/utils/error_rate_zh
@@ -0,0 +1,370 @@
+#!/usr/bin/env python3
+# coding=utf8
+
+# Copyright 2021 Jiayu DU
+
+import sys
+import argparse
+import json
+import logging
+logging.basicConfig(stream=sys.stderr, level=logging.INFO, format='[%(levelname)s] %(message)s')
+
+DEBUG = None
+
+def GetEditType(ref_token, hyp_token):
+ if ref_token == None and hyp_token != None:
+ return 'I'
+ elif ref_token != None and hyp_token == None:
+ return 'D'
+ elif ref_token == hyp_token:
+ return 'C'
+ elif ref_token != hyp_token:
+ return 'S'
+ else:
+ raise RuntimeError
+
+class AlignmentArc:
+ def __init__(self, src, dst, ref, hyp):
+ self.src = src
+ self.dst = dst
+ self.ref = ref
+ self.hyp = hyp
+ self.edit_type = GetEditType(ref, hyp)
+
+def similarity_score_function(ref_token, hyp_token):
+ return 0 if (ref_token == hyp_token) else -1.0
+
+def insertion_score_function(token):
+ return -1.0
+
+def deletion_score_function(token):
+ return -1.0
+
+def EditDistance(
+ ref,
+ hyp,
+ similarity_score_function = similarity_score_function,
+ insertion_score_function = insertion_score_function,
+ deletion_score_function = deletion_score_function):
+ assert(len(ref) != 0)
+ class DPState:
+ def __init__(self):
+ self.score = -float('inf')
+ # backpointer
+ self.prev_r = None
+ self.prev_h = None
+
+ def print_search_grid(S, R, H, fstream):
+ print(file=fstream)
+ for r in range(R):
+ for h in range(H):
+ print(F'[{r},{h}]:{S[r][h].score:4.3f}:({S[r][h].prev_r},{S[r][h].prev_h}) ', end='', file=fstream)
+ print(file=fstream)
+
+ R = len(ref) + 1
+ H = len(hyp) + 1
+
+ # Construct DP search space, a (R x H) grid
+ S = [ [] for r in range(R) ]
+ for r in range(R):
+ S[r] = [ DPState() for x in range(H) ]
+
+ # initialize DP search grid origin, S(r = 0, h = 0)
+ S[0][0].score = 0.0
+ S[0][0].prev_r = None
+ S[0][0].prev_h = None
+
+ # initialize REF axis
+ for r in range(1, R):
+ S[r][0].score = S[r-1][0].score + deletion_score_function(ref[r-1])
+ S[r][0].prev_r = r-1
+ S[r][0].prev_h = 0
+
+ # initialize HYP axis
+ for h in range(1, H):
+ S[0][h].score = S[0][h-1].score + insertion_score_function(hyp[h-1])
+ S[0][h].prev_r = 0
+ S[0][h].prev_h = h-1
+
+ best_score = S[0][0].score
+ best_state = (0, 0)
+
+ for r in range(1, R):
+ for h in range(1, H):
+ sub_or_cor_score = similarity_score_function(ref[r-1], hyp[h-1])
+ new_score = S[r-1][h-1].score + sub_or_cor_score
+ if new_score >= S[r][h].score:
+ S[r][h].score = new_score
+ S[r][h].prev_r = r-1
+ S[r][h].prev_h = h-1
+
+ del_score = deletion_score_function(ref[r-1])
+ new_score = S[r-1][h].score + del_score
+ if new_score >= S[r][h].score:
+ S[r][h].score = new_score
+ S[r][h].prev_r = r - 1
+ S[r][h].prev_h = h
+
+ ins_score = insertion_score_function(hyp[h-1])
+ new_score = S[r][h-1].score + ins_score
+ if new_score >= S[r][h].score:
+ S[r][h].score = new_score
+ S[r][h].prev_r = r
+ S[r][h].prev_h = h-1
+
+ best_score = S[R-1][H-1].score
+ best_state = (R-1, H-1)
+
+ if DEBUG:
+ print_search_grid(S, R, H, sys.stderr)
+
+ # Backtracing best alignment path, i.e. a list of arcs
+ # arc = (src, dst, ref, hyp, edit_type)
+ # src/dst = (r, h), where r/h refers to search grid state-id along Ref/Hyp axis
+ best_path = []
+ r, h = best_state[0], best_state[1]
+ prev_r, prev_h = S[r][h].prev_r, S[r][h].prev_h
+ score = S[r][h].score
+ # loop invariant:
+ # 1. (prev_r, prev_h) -> (r, h) is a "forward arc" on best alignment path
+ # 2. score is the value of point(r, h) on DP search grid
+ while prev_r != None or prev_h != None:
+ src = (prev_r, prev_h)
+ dst = (r, h)
+ if (r == prev_r + 1 and h == prev_h + 1): # Substitution or correct
+ arc = AlignmentArc(src, dst, ref[prev_r], hyp[prev_h])
+ elif (r == prev_r + 1 and h == prev_h): # Deletion
+ arc = AlignmentArc(src, dst, ref[prev_r], None)
+ elif (r == prev_r and h == prev_h + 1): # Insertion
+ arc = AlignmentArc(src, dst, None, hyp[prev_h])
+ else:
+ raise RuntimeError
+ best_path.append(arc)
+ r, h = prev_r, prev_h
+ prev_r, prev_h = S[r][h].prev_r, S[r][h].prev_h
+ score = S[r][h].score
+
+ best_path.reverse()
+ return (best_path, best_score)
+
+def PrettyPrintAlignment(alignment, stream = sys.stderr):
+ def get_token_str(token):
+ if token == None:
+ return "*"
+ return token
+
+ def is_double_width_char(ch):
+ if (ch >= '\u4e00') and (ch <= '\u9fa5'): # codepoint ranges for Chinese chars
+ return True
+ # TODO: support other double-width-char language such as Japanese, Korean
+ else:
+ return False
+
+ def display_width(token_str):
+ m = 0
+ for c in token_str:
+ if is_double_width_char(c):
+ m += 2
+ else:
+ m += 1
+ return m
+
+ R = ' REF : '
+ H = ' HYP : '
+ E = ' EDIT : '
+ for arc in alignment:
+ r = get_token_str(arc.ref)
+ h = get_token_str(arc.hyp)
+ e = arc.edit_type if arc.edit_type != 'C' else ''
+
+ nr, nh, ne = display_width(r), display_width(h), display_width(e)
+ n = max(nr, nh, ne) + 1
+
+ R += r + ' ' * (n-nr)
+ H += h + ' ' * (n-nh)
+ E += e + ' ' * (n-ne)
+
+ print(R, file=stream)
+ print(H, file=stream)
+ print(E, file=stream)
+
+def CountEdits(alignment):
+ c, s, i, d = 0, 0, 0, 0
+ for arc in alignment:
+ if arc.edit_type == 'C':
+ c += 1
+ elif arc.edit_type == 'S':
+ s += 1
+ elif arc.edit_type == 'I':
+ i += 1
+ elif arc.edit_type == 'D':
+ d += 1
+ else:
+ raise RuntimeError
+ return (c, s, i, d)
+
+def ComputeTokenErrorRate(c, s, i, d):
+ return 100.0 * (s + d + i) / (s + d + c)
+
+def ComputeSentenceErrorRate(num_err_utts, num_utts):
+ assert(num_utts != 0)
+ return 100.0 * num_err_utts / num_utts
+
+
+class EvaluationResult:
+ def __init__(self):
+ self.num_ref_utts = 0
+ self.num_hyp_utts = 0
+ self.num_eval_utts = 0 # seen in both ref & hyp
+ self.num_hyp_without_ref = 0
+
+ self.C = 0
+ self.S = 0
+ self.I = 0
+ self.D = 0
+ self.token_error_rate = 0.0
+
+ self.num_utts_with_error = 0
+ self.sentence_error_rate = 0.0
+
+ def to_json(self):
+ return json.dumps(self.__dict__)
+
+ def to_kaldi(self):
+ info = (
+ F'%WER {self.token_error_rate:.2f} [ {self.S + self.D + self.I} / {self.C + self.S + self.D}, {self.I} ins, {self.D} del, {self.S} sub ]\n'
+ F'%SER {self.sentence_error_rate:.2f} [ {self.num_utts_with_error} / {self.num_eval_utts} ]\n'
+ )
+ return info
+
+ def to_sclite(self):
+ return "TODO"
+
+ def to_espnet(self):
+ return "TODO"
+
+ def to_summary(self):
+ #return json.dumps(self.__dict__, indent=4)
+ summary = (
+ '==================== Overall Statistics ====================\n'
+ F'num_ref_utts: {self.num_ref_utts}\n'
+ F'num_hyp_utts: {self.num_hyp_utts}\n'
+ F'num_hyp_without_ref: {self.num_hyp_without_ref}\n'
+ F'num_eval_utts: {self.num_eval_utts}\n'
+ F'sentence_error_rate: {self.sentence_error_rate:.2f}%\n'
+ F'token_error_rate: {self.token_error_rate:.2f}%\n'
+ F'token_stats:\n'
+ F' - tokens:{self.C + self.S + self.D:>7}\n'
+ F' - edits: {self.S + self.I + self.D:>7}\n'
+ F' - cor: {self.C:>7}\n'
+ F' - sub: {self.S:>7}\n'
+ F' - ins: {self.I:>7}\n'
+ F' - del: {self.D:>7}\n'
+ '============================================================\n'
+ )
+ return summary
+
+
+class Utterance:
+ def __init__(self, uid, text):
+ self.uid = uid
+ self.text = text
+
+
+def LoadUtterances(filepath, format):
+ utts = {}
+ if format == 'text': # utt_id word1 word2 ...
+ with open(filepath, 'r', encoding='utf8') as f:
+ for line in f:
+ line = line.strip()
+ if line:
+ cols = line.split(maxsplit=1)
+ assert(len(cols) == 2 or len(cols) == 1)
+ uid = cols[0]
+ text = cols[1] if len(cols) == 2 else ''
+ if utts.get(uid) != None:
+ raise RuntimeError(F'Found duplicated utterence id {uid}')
+ utts[uid] = Utterance(uid, text)
+ else:
+ raise RuntimeError(F'Unsupported text format {format}')
+ return utts
+
+
+def tokenize_text(text, tokenizer):
+ if tokenizer == 'whitespace':
+ return text.split()
+ elif tokenizer == 'char':
+ return [ ch for ch in ''.join(text.split()) ]
+ else:
+ raise RuntimeError(F'ERROR: Unsupported tokenizer {tokenizer}')
+
+
+if __name__ == '__main__':
+ parser = argparse.ArgumentParser()
+ # optional
+ parser.add_argument('--tokenizer', choices=['whitespace', 'char'], default='whitespace', help='whitespace for WER, char for CER')
+ parser.add_argument('--ref-format', choices=['text'], default='text', help='reference format, first col is utt_id, the rest is text')
+ parser.add_argument('--hyp-format', choices=['text'], default='text', help='hypothesis format, first col is utt_id, the rest is text')
+ # required
+ parser.add_argument('--ref', type=str, required=True, help='input reference file')
+ parser.add_argument('--hyp', type=str, required=True, help='input hypothesis file')
+
+ parser.add_argument('result_file', type=str)
+ args = parser.parse_args()
+ logging.info(args)
+
+ ref_utts = LoadUtterances(args.ref, args.ref_format)
+ hyp_utts = LoadUtterances(args.hyp, args.hyp_format)
+
+ r = EvaluationResult()
+
+ # check valid utterances in hyp that have matched non-empty reference
+ eval_utts = []
+ r.num_hyp_without_ref = 0
+ for uid in sorted(hyp_utts.keys()):
+ if uid in ref_utts.keys(): # TODO: efficiency
+ if ref_utts[uid].text.strip(): # non-empty reference
+ eval_utts.append(uid)
+ else:
+ logging.warn(F'Found {uid} with empty reference, skipping...')
+ else:
+ logging.warn(F'Found {uid} without reference, skipping...')
+ r.num_hyp_without_ref += 1
+
+ r.num_hyp_utts = len(hyp_utts)
+ r.num_ref_utts = len(ref_utts)
+ r.num_eval_utts = len(eval_utts)
+
+ with open(args.result_file, 'w+', encoding='utf8') as fo:
+ for uid in eval_utts:
+ ref = ref_utts[uid]
+ hyp = hyp_utts[uid]
+
+ alignment, score = EditDistance(
+ tokenize_text(ref.text, args.tokenizer),
+ tokenize_text(hyp.text, args.tokenizer)
+ )
+
+ c, s, i, d = CountEdits(alignment)
+ utt_ter = ComputeTokenErrorRate(c, s, i, d)
+
+ # utt-level evaluation result
+ print(F'{{"uid":{uid}, "score":{score}, "ter":{utt_ter:.2f}, "cor":{c}, "sub":{s}, "ins":{i}, "del":{d}}}', file=fo)
+ PrettyPrintAlignment(alignment, fo)
+
+ r.C += c
+ r.S += s
+ r.I += i
+ r.D += d
+
+ if utt_ter > 0:
+ r.num_utts_with_error += 1
+
+ # corpus level evaluation result
+ r.sentence_error_rate = ComputeSentenceErrorRate(r.num_utts_with_error, r.num_eval_utts)
+ r.token_error_rate = ComputeTokenErrorRate(r.C, r.S, r.I, r.D)
+
+ print(r.to_summary(), file=fo)
+
+ print(r.to_json())
+ print(r.to_kaldi())
diff --git a/egs_modelscope/common/utils/extract_embeds.py b/egs_modelscope/common/utils/extract_embeds.py
new file mode 100755
index 0000000..7b817d8
--- /dev/null
+++ b/egs_modelscope/common/utils/extract_embeds.py
@@ -0,0 +1,47 @@
+from transformers import AutoTokenizer, AutoModel, pipeline
+import numpy as np
+import sys
+import os
+import torch
+from kaldiio import WriteHelper
+import re
+text_file_json = sys.argv[1]
+out_ark = sys.argv[2]
+out_scp = sys.argv[3]
+out_shape = sys.argv[4]
+device = int(sys.argv[5])
+model_path = sys.argv[6]
+
+model = AutoModel.from_pretrained(model_path)
+tokenizer = AutoTokenizer.from_pretrained(model_path)
+extractor = pipeline(task="feature-extraction", model=model, tokenizer=tokenizer, device=device)
+
+with open(text_file_json, 'r') as f:
+ js = f.readlines()
+
+
+f_shape = open(out_shape, "w")
+with WriteHelper('ark,scp:{},{}'.format(out_ark, out_scp)) as writer:
+ with torch.no_grad():
+ for idx, line in enumerate(js):
+ id, tokens = line.strip().split(" ", 1)
+ tokens = re.sub(" ", "", tokens.strip())
+ tokens = ' '.join([j for j in tokens])
+ token_num = len(tokens.split(" "))
+ outputs = extractor(tokens)
+ outputs = np.array(outputs)
+ embeds = outputs[0, 1:-1, :]
+
+ token_num_embeds, dim = embeds.shape
+ if token_num == token_num_embeds:
+ writer(id, embeds)
+ shape_line = "{} {},{}\n".format(id, token_num_embeds, dim)
+ f_shape.write(shape_line)
+ else:
+ print("{}, size has changed, {}, {}, {}".format(id, token_num, token_num_embeds, tokens))
+
+
+
+f_shape.close()
+
+
diff --git a/egs_modelscope/common/utils/filter_scp.pl b/egs_modelscope/common/utils/filter_scp.pl
new file mode 100755
index 0000000..003530d
--- /dev/null
+++ b/egs_modelscope/common/utils/filter_scp.pl
@@ -0,0 +1,87 @@
+#!/usr/bin/env perl
+# Copyright 2010-2012 Microsoft Corporation
+# Johns Hopkins University (author: Daniel Povey)
+
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+# This script takes a list of utterance-ids or any file whose first field
+# of each line is an utterance-id, and filters an scp
+# file (or any file whose "n-th" field is an utterance id), printing
+# out only those lines whose "n-th" field is in id_list. The index of
+# the "n-th" field is 1, by default, but can be changed by using
+# the -f <n> switch
+
+$exclude = 0;
+$field = 1;
+$shifted = 0;
+
+do {
+ $shifted=0;
+ if ($ARGV[0] eq "--exclude") {
+ $exclude = 1;
+ shift @ARGV;
+ $shifted=1;
+ }
+ if ($ARGV[0] eq "-f") {
+ $field = $ARGV[1];
+ shift @ARGV; shift @ARGV;
+ $shifted=1
+ }
+} while ($shifted);
+
+if(@ARGV < 1 || @ARGV > 2) {
+ die "Usage: filter_scp.pl [--exclude] [-f <field-to-filter-on>] id_list [in.scp] > out.scp \n" .
+ "Prints only the input lines whose f'th field (default: first) is in 'id_list'.\n" .
+ "Note: only the first field of each line in id_list matters. With --exclude, prints\n" .
+ "only the lines that were *not* in id_list.\n" .
+ "Caution: previously, the -f option was interpreted as a zero-based field index.\n" .
+ "If your older scripts (written before Oct 2014) stopped working and you used the\n" .
+ "-f option, add 1 to the argument.\n" .
+ "See also: scripts/filter_scp.pl .\n";
+}
+
+
+$idlist = shift @ARGV;
+open(F, "<$idlist") || die "Could not open id-list file $idlist";
+while(<F>) {
+ @A = split;
+ @A>=1 || die "Invalid id-list file line $_";
+ $seen{$A[0]} = 1;
+}
+
+if ($field == 1) { # Treat this as special case, since it is common.
+ while(<>) {
+ $_ =~ m/\s*(\S+)\s*/ || die "Bad line $_, could not get first field.";
+ # $1 is what we filter on.
+ if ((!$exclude && $seen{$1}) || ($exclude && !defined $seen{$1})) {
+ print $_;
+ }
+ }
+} else {
+ while(<>) {
+ @A = split;
+ @A > 0 || die "Invalid scp file line $_";
+ @A >= $field || die "Invalid scp file line $_";
+ if ((!$exclude && $seen{$A[$field-1]}) || ($exclude && !defined $seen{$A[$field-1]})) {
+ print $_;
+ }
+ }
+}
+
+# tests:
+# the following should print "foo 1"
+# ( echo foo 1; echo bar 2 ) | scripts/filter_scp.pl <(echo foo)
+# the following should print "bar 2".
+# ( echo foo 1; echo bar 2 ) | scripts/filter_scp.pl -f 2 <(echo 2)
diff --git a/egs_modelscope/common/utils/fix_data.sh b/egs_modelscope/common/utils/fix_data.sh
new file mode 100755
index 0000000..32cdde5
--- /dev/null
+++ b/egs_modelscope/common/utils/fix_data.sh
@@ -0,0 +1,35 @@
+#!/usr/bin/env bash
+
+echo "$0 $@"
+data_dir=$1
+
+if [ ! -f ${data_dir}/wav.scp ]; then
+ echo "$0: wav.scp is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/text ]; then
+ echo "$0: text is not found"
+ exit 1;
+fi
+
+
+
+mkdir -p ${data_dir}/.backup
+
+awk '{print $1}' ${data_dir}/wav.scp > ${data_dir}/.backup/wav_id
+awk '{print $1}' ${data_dir}/text > ${data_dir}/.backup/text_id
+
+sort ${data_dir}/.backup/wav_id ${data_dir}/.backup/text_id | uniq -d > ${data_dir}/.backup/id
+
+cp ${data_dir}/wav.scp ${data_dir}/.backup/wav.scp
+cp ${data_dir}/text ${data_dir}/.backup/text
+
+mv ${data_dir}/wav.scp ${data_dir}/wav.scp.bak
+mv ${data_dir}/text ${data_dir}/text.bak
+
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/wav.scp.bak > ${data_dir}/wav.scp
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/text.bak > ${data_dir}/text
+
+rm ${data_dir}/wav.scp.bak
+rm ${data_dir}/text.bak
diff --git a/egs_modelscope/common/utils/fix_data_feat.sh b/egs_modelscope/common/utils/fix_data_feat.sh
new file mode 100755
index 0000000..2c92d7f
--- /dev/null
+++ b/egs_modelscope/common/utils/fix_data_feat.sh
@@ -0,0 +1,52 @@
+#!/usr/bin/env bash
+
+echo "$0 $@"
+data_dir=$1
+
+if [ ! -f ${data_dir}/feats.scp ]; then
+ echo "$0: feats.scp is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/text ]; then
+ echo "$0: text is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/speech_shape ]; then
+ echo "$0: feature lengths is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/text_shape ]; then
+ echo "$0: text lengths is not found"
+ exit 1;
+fi
+
+mkdir -p ${data_dir}/.backup
+
+awk '{print $1}' ${data_dir}/feats.scp > ${data_dir}/.backup/wav_id
+awk '{print $1}' ${data_dir}/text > ${data_dir}/.backup/text_id
+
+sort ${data_dir}/.backup/wav_id ${data_dir}/.backup/text_id | uniq -d > ${data_dir}/.backup/id
+
+cp ${data_dir}/feats.scp ${data_dir}/.backup/feats.scp
+cp ${data_dir}/text ${data_dir}/.backup/text
+cp ${data_dir}/speech_shape ${data_dir}/.backup/speech_shape
+cp ${data_dir}/text_shape ${data_dir}/.backup/text_shape
+
+mv ${data_dir}/feats.scp ${data_dir}/feats.scp.bak
+mv ${data_dir}/text ${data_dir}/text.bak
+mv ${data_dir}/speech_shape ${data_dir}/speech_shape.bak
+mv ${data_dir}/text_shape ${data_dir}/text_shape.bak
+
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/feats.scp.bak > ${data_dir}/feats.scp
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/text.bak > ${data_dir}/text
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/speech_shape.bak > ${data_dir}/speech_shape
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/text_shape.bak > ${data_dir}/text_shape
+
+rm ${data_dir}/feats.scp.bak
+rm ${data_dir}/text.bak
+rm ${data_dir}/speech_shape.bak
+rm ${data_dir}/text_shape.bak
+
diff --git a/egs_modelscope/common/utils/gen_ark_list.sh b/egs_modelscope/common/utils/gen_ark_list.sh
new file mode 100755
index 0000000..aebf356
--- /dev/null
+++ b/egs_modelscope/common/utils/gen_ark_list.sh
@@ -0,0 +1,22 @@
+#!/usr/bin/env bash
+
+
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+ark_dir=$1
+txt_dir=$2
+output_dir=$3
+
+[ ! -d ${ark_dir}/ark ] && echo "$0: ark data is required" && exit 1;
+[ ! -d ${txt_dir}/txt ] && echo "$0: txt data is required" && exit 1;
+
+for n in $(seq $nj); do
+ echo "${ark_dir}/ark/feats.$n.ark ${txt_dir}/txt/text.$n.txt" || exit 1
+done > ${output_dir}/ark_txt.scp || exit 1
+
diff --git a/egs_modelscope/common/utils/parse_options.sh b/egs_modelscope/common/utils/parse_options.sh
new file mode 100755
index 0000000..71fb9e5
--- /dev/null
+++ b/egs_modelscope/common/utils/parse_options.sh
@@ -0,0 +1,97 @@
+#!/usr/bin/env bash
+
+# Copyright 2012 Johns Hopkins University (Author: Daniel Povey);
+# Arnab Ghoshal, Karel Vesely
+
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+# Parse command-line options.
+# To be sourced by another script (as in ". parse_options.sh").
+# Option format is: --option-name arg
+# and shell variable "option_name" gets set to value "arg."
+# The exception is --help, which takes no arguments, but prints the
+# $help_message variable (if defined).
+
+
+###
+### The --config file options have lower priority to command line
+### options, so we need to import them first...
+###
+
+# Now import all the configs specified by command-line, in left-to-right order
+for ((argpos=1; argpos<$#; argpos++)); do
+ if [ "${!argpos}" == "--config" ]; then
+ argpos_plus1=$((argpos+1))
+ config=${!argpos_plus1}
+ [ ! -r $config ] && echo "$0: missing config '$config'" && exit 1
+ . $config # source the config file.
+ fi
+done
+
+
+###
+### Now we process the command line options
+###
+while true; do
+ [ -z "${1:-}" ] && break; # break if there are no arguments
+ case "$1" in
+ # If the enclosing script is called with --help option, print the help
+ # message and exit. Scripts should put help messages in $help_message
+ --help|-h) if [ -z "$help_message" ]; then echo "No help found." 1>&2;
+ else printf "$help_message\n" 1>&2 ; fi;
+ exit 0 ;;
+ --*=*) echo "$0: options to scripts must be of the form --name value, got '$1'"
+ exit 1 ;;
+ # If the first command-line argument begins with "--" (e.g. --foo-bar),
+ # then work out the variable name as $name, which will equal "foo_bar".
+ --*) name=`echo "$1" | sed s/^--// | sed s/-/_/g`;
+ # Next we test whether the variable in question is undefned-- if so it's
+ # an invalid option and we die. Note: $0 evaluates to the name of the
+ # enclosing script.
+ # The test [ -z ${foo_bar+xxx} ] will return true if the variable foo_bar
+ # is undefined. We then have to wrap this test inside "eval" because
+ # foo_bar is itself inside a variable ($name).
+ eval '[ -z "${'$name'+xxx}" ]' && echo "$0: invalid option $1" 1>&2 && exit 1;
+
+ oldval="`eval echo \\$$name`";
+ # Work out whether we seem to be expecting a Boolean argument.
+ if [ "$oldval" == "true" ] || [ "$oldval" == "false" ]; then
+ was_bool=true;
+ else
+ was_bool=false;
+ fi
+
+ # Set the variable to the right value-- the escaped quotes make it work if
+ # the option had spaces, like --cmd "queue.pl -sync y"
+ eval $name=\"$2\";
+
+ # Check that Boolean-valued arguments are really Boolean.
+ if $was_bool && [[ "$2" != "true" && "$2" != "false" ]]; then
+ echo "$0: expected \"true\" or \"false\": $1 $2" 1>&2
+ exit 1;
+ fi
+ shift 2;
+ ;;
+ *) break;
+ esac
+done
+
+
+# Check for an empty argument to the --cmd option, which can easily occur as a
+# result of scripting errors.
+[ ! -z "${cmd+xxx}" ] && [ -z "$cmd" ] && echo "$0: empty argument to --cmd option" 1>&2 && exit 1;
+
+
+true; # so this script returns exit code 0.
diff --git a/egs_modelscope/common/utils/print_args.py b/egs_modelscope/common/utils/print_args.py
new file mode 100755
index 0000000..b0c61e5
--- /dev/null
+++ b/egs_modelscope/common/utils/print_args.py
@@ -0,0 +1,45 @@
+#!/usr/bin/env python
+import sys
+
+
+def get_commandline_args(no_executable=True):
+ extra_chars = [
+ " ",
+ ";",
+ "&",
+ "|",
+ "<",
+ ">",
+ "?",
+ "*",
+ "~",
+ "`",
+ '"',
+ "'",
+ "\\",
+ "{",
+ "}",
+ "(",
+ ")",
+ ]
+
+ # Escape the extra characters for shell
+ argv = [
+ arg.replace("'", "'\\''")
+ if all(char not in arg for char in extra_chars)
+ else "'" + arg.replace("'", "'\\''") + "'"
+ for arg in sys.argv
+ ]
+
+ if no_executable:
+ return " ".join(argv[1:])
+ else:
+ return sys.executable + " " + " ".join(argv)
+
+
+def main():
+ print(get_commandline_args())
+
+
+if __name__ == "__main__":
+ main()
diff --git a/egs_modelscope/common/utils/proc_conf_oss.py b/egs_modelscope/common/utils/proc_conf_oss.py
new file mode 100755
index 0000000..c4a90c5
--- /dev/null
+++ b/egs_modelscope/common/utils/proc_conf_oss.py
@@ -0,0 +1,35 @@
+from pathlib import Path
+
+import torch
+import yaml
+
+
+class NoAliasSafeDumper(yaml.SafeDumper):
+ # Disable anchor/alias in yaml because looks ugly
+ def ignore_aliases(self, data):
+ return True
+
+
+def yaml_no_alias_safe_dump(data, stream=None, **kwargs):
+ """Safe-dump in yaml with no anchor/alias"""
+ return yaml.dump(
+ data, stream, allow_unicode=True, Dumper=NoAliasSafeDumper, **kwargs
+ )
+
+
+def gen_conf(file, out_dir):
+ conf = torch.load(file)["config"]
+ conf["oss_bucket"] = "null"
+ print(conf)
+ output_dir = Path(out_dir)
+ output_dir.mkdir(parents=True, exist_ok=True)
+ with (output_dir / "config.yaml").open("w", encoding="utf-8") as f:
+ yaml_no_alias_safe_dump(conf, f, indent=4, sort_keys=False)
+
+
+if __name__ == "__main__":
+ import sys
+
+ in_f = sys.argv[1]
+ out_f = sys.argv[2]
+ gen_conf(in_f, out_f)
diff --git a/egs_modelscope/common/utils/proce_text.py b/egs_modelscope/common/utils/proce_text.py
new file mode 100755
index 0000000..9e517a4
--- /dev/null
+++ b/egs_modelscope/common/utils/proce_text.py
@@ -0,0 +1,31 @@
+
+import sys
+import re
+
+in_f = sys.argv[1]
+out_f = sys.argv[2]
+
+
+with open(in_f, "r", encoding="utf-8") as f:
+ lines = f.readlines()
+
+with open(out_f, "w", encoding="utf-8") as f:
+ for line in lines:
+ outs = line.strip().split(" ", 1)
+ if len(outs) == 2:
+ idx, text = outs
+ text = re.sub("</s>", "", text)
+ text = re.sub("<s>", "", text)
+ text = re.sub("@@", "", text)
+ text = re.sub("@", "", text)
+ text = re.sub("<unk>", "", text)
+ text = re.sub(" ", "", text)
+ text = text.lower()
+ else:
+ idx = outs[0]
+ text = " "
+
+ text = [x for x in text]
+ text = " ".join(text)
+ out = "{} {}\n".format(idx, text)
+ f.write(out)
diff --git a/egs_modelscope/common/utils/run.pl b/egs_modelscope/common/utils/run.pl
new file mode 100755
index 0000000..483f95b
--- /dev/null
+++ b/egs_modelscope/common/utils/run.pl
@@ -0,0 +1,356 @@
+#!/usr/bin/env perl
+use warnings; #sed replacement for -w perl parameter
+# In general, doing
+# run.pl some.log a b c is like running the command a b c in
+# the bash shell, and putting the standard error and output into some.log.
+# To run parallel jobs (backgrounded on the host machine), you can do (e.g.)
+# run.pl JOB=1:4 some.JOB.log a b c JOB is like running the command a b c JOB
+# and putting it in some.JOB.log, for each one. [Note: JOB can be any identifier].
+# If any of the jobs fails, this script will fail.
+
+# A typical example is:
+# run.pl some.log my-prog "--opt=foo bar" foo \| other-prog baz
+# and run.pl will run something like:
+# ( my-prog '--opt=foo bar' foo | other-prog baz ) >& some.log
+#
+# Basically it takes the command-line arguments, quotes them
+# as necessary to preserve spaces, and evaluates them with bash.
+# In addition it puts the command line at the top of the log, and
+# the start and end times of the command at the beginning and end.
+# The reason why this is useful is so that we can create a different
+# version of this program that uses a queueing system instead.
+
+#use Data::Dumper;
+
+@ARGV < 2 && die "usage: run.pl log-file command-line arguments...";
+
+#print STDERR "COMMAND-LINE: " . Dumper(\@ARGV) . "\n";
+$job_pick = 'all';
+$max_jobs_run = -1;
+$jobstart = 1;
+$jobend = 1;
+$ignored_opts = ""; # These will be ignored.
+
+# First parse an option like JOB=1:4, and any
+# options that would normally be given to
+# queue.pl, which we will just discard.
+
+for (my $x = 1; $x <= 2; $x++) { # This for-loop is to
+ # allow the JOB=1:n option to be interleaved with the
+ # options to qsub.
+ while (@ARGV >= 2 && $ARGV[0] =~ m:^-:) {
+ # parse any options that would normally go to qsub, but which will be ignored here.
+ my $switch = shift @ARGV;
+ if ($switch eq "-V") {
+ $ignored_opts .= "-V ";
+ } elsif ($switch eq "--max-jobs-run" || $switch eq "-tc") {
+ # we do support the option --max-jobs-run n, and its GridEngine form -tc n.
+ # if the command appears multiple times uses the smallest option.
+ if ( $max_jobs_run <= 0 ) {
+ $max_jobs_run = shift @ARGV;
+ } else {
+ my $new_constraint = shift @ARGV;
+ if ( ($new_constraint < $max_jobs_run) ) {
+ $max_jobs_run = $new_constraint;
+ }
+ }
+
+ if (! ($max_jobs_run > 0)) {
+ die "run.pl: invalid option --max-jobs-run $max_jobs_run";
+ }
+ } else {
+ my $argument = shift @ARGV;
+ if ($argument =~ m/^--/) {
+ print STDERR "run.pl: WARNING: suspicious argument '$argument' to $switch; starts with '-'\n";
+ }
+ if ($switch eq "-sync" && $argument =~ m/^[yY]/) {
+ $ignored_opts .= "-sync "; # Note: in the
+ # corresponding code in queue.pl it says instead, just "$sync = 1;".
+ } elsif ($switch eq "-pe") { # e.g. -pe smp 5
+ my $argument2 = shift @ARGV;
+ $ignored_opts .= "$switch $argument $argument2 ";
+ } elsif ($switch eq "--gpu") {
+ $using_gpu = $argument;
+ } elsif ($switch eq "--pick") {
+ if($argument =~ m/^(all|failed|incomplete)$/) {
+ $job_pick = $argument;
+ } else {
+ print STDERR "run.pl: ERROR: --pick argument must be one of 'all', 'failed' or 'incomplete'"
+ }
+ } else {
+ # Ignore option.
+ $ignored_opts .= "$switch $argument ";
+ }
+ }
+ }
+ if ($ARGV[0] =~ m/^([\w_][\w\d_]*)+=(\d+):(\d+)$/) { # e.g. JOB=1:20
+ $jobname = $1;
+ $jobstart = $2;
+ $jobend = $3;
+ if ($jobstart > $jobend) {
+ die "run.pl: invalid job range $ARGV[0]";
+ }
+ if ($jobstart <= 0) {
+ die "run.pl: invalid job range $ARGV[0], start must be strictly positive (this is required for GridEngine compatibility).";
+ }
+ shift;
+ } elsif ($ARGV[0] =~ m/^([\w_][\w\d_]*)+=(\d+)$/) { # e.g. JOB=1.
+ $jobname = $1;
+ $jobstart = $2;
+ $jobend = $2;
+ shift;
+ } elsif ($ARGV[0] =~ m/.+\=.*\:.*$/) {
+ print STDERR "run.pl: Warning: suspicious first argument to run.pl: $ARGV[0]\n";
+ }
+}
+
+# Users found this message confusing so we are removing it.
+# if ($ignored_opts ne "") {
+# print STDERR "run.pl: Warning: ignoring options \"$ignored_opts\"\n";
+# }
+
+if ($max_jobs_run == -1) { # If --max-jobs-run option not set,
+ # then work out the number of processors if possible,
+ # and set it based on that.
+ $max_jobs_run = 0;
+ if ($using_gpu) {
+ if (open(P, "nvidia-smi -L |")) {
+ $max_jobs_run++ while (<P>);
+ close(P);
+ }
+ if ($max_jobs_run == 0) {
+ $max_jobs_run = 1;
+ print STDERR "run.pl: Warning: failed to detect number of GPUs from nvidia-smi, using ${max_jobs_run}\n";
+ }
+ } elsif (open(P, "</proc/cpuinfo")) { # Linux
+ while (<P>) { if (m/^processor/) { $max_jobs_run++; } }
+ if ($max_jobs_run == 0) {
+ print STDERR "run.pl: Warning: failed to detect any processors from /proc/cpuinfo\n";
+ $max_jobs_run = 10; # reasonable default.
+ }
+ close(P);
+ } elsif (open(P, "sysctl -a |")) { # BSD/Darwin
+ while (<P>) {
+ if (m/hw\.ncpu\s*[:=]\s*(\d+)/) { # hw.ncpu = 4, or hw.ncpu: 4
+ $max_jobs_run = $1;
+ last;
+ }
+ }
+ close(P);
+ if ($max_jobs_run == 0) {
+ print STDERR "run.pl: Warning: failed to detect any processors from sysctl -a\n";
+ $max_jobs_run = 10; # reasonable default.
+ }
+ } else {
+ # allow at most 32 jobs at once, on non-UNIX systems; change this code
+ # if you need to change this default.
+ $max_jobs_run = 32;
+ }
+ # The just-computed value of $max_jobs_run is just the number of processors
+ # (or our best guess); and if it happens that the number of jobs we need to
+ # run is just slightly above $max_jobs_run, it will make sense to increase
+ # $max_jobs_run to equal the number of jobs, so we don't have a small number
+ # of leftover jobs.
+ $num_jobs = $jobend - $jobstart + 1;
+ if (!$using_gpu &&
+ $num_jobs > $max_jobs_run && $num_jobs < 1.4 * $max_jobs_run) {
+ $max_jobs_run = $num_jobs;
+ }
+}
+
+sub pick_or_exit {
+ # pick_or_exit ( $logfile )
+ # Invoked before each job is started helps to run jobs selectively.
+ #
+ # Given the name of the output logfile decides whether the job must be
+ # executed (by returning from the subroutine) or not (by terminating the
+ # process calling exit)
+ #
+ # PRE: $job_pick is a global variable set by command line switch --pick
+ # and indicates which class of jobs must be executed.
+ #
+ # 1) If a failed job is not executed the process exit code will indicate
+ # failure, just as if the task was just executed and failed.
+ #
+ # 2) If a task is incomplete it will be executed. Incomplete may be either
+ # a job whose log file does not contain the accounting notes in the end,
+ # or a job whose log file does not exist.
+ #
+ # 3) If the $job_pick is set to 'all' (default behavior) a task will be
+ # executed regardless of the result of previous attempts.
+ #
+ # This logic could have been implemented in the main execution loop
+ # but a subroutine to preserve the current level of readability of
+ # that part of the code.
+ #
+ # Alexandre Felipe, (o.alexandre.felipe@gmail.com) 14th of August of 2020
+ #
+ if($job_pick eq 'all'){
+ return; # no need to bother with the previous log
+ }
+ open my $fh, "<", $_[0] or return; # job not executed yet
+ my $log_line;
+ my $cur_line;
+ while ($cur_line = <$fh>) {
+ if( $cur_line =~ m/# Ended \(code .*/ ) {
+ $log_line = $cur_line;
+ }
+ }
+ close $fh;
+ if (! defined($log_line)){
+ return; # incomplete
+ }
+ if ( $log_line =~ m/# Ended \(code 0\).*/ ) {
+ exit(0); # complete
+ } elsif ( $log_line =~ m/# Ended \(code \d+(; signal \d+)?\).*/ ){
+ if ($job_pick !~ m/^(failed|all)$/) {
+ exit(1); # failed but not going to run
+ } else {
+ return; # failed
+ }
+ } elsif ( $log_line =~ m/.*\S.*/ ) {
+ return; # incomplete jobs are always run
+ }
+}
+
+
+$logfile = shift @ARGV;
+
+if (defined $jobname && $logfile !~ m/$jobname/ &&
+ $jobend > $jobstart) {
+ print STDERR "run.pl: you are trying to run a parallel job but "
+ . "you are putting the output into just one log file ($logfile)\n";
+ exit(1);
+}
+
+$cmd = "";
+
+foreach $x (@ARGV) {
+ if ($x =~ m/^\S+$/) { $cmd .= $x . " "; }
+ elsif ($x =~ m:\":) { $cmd .= "'$x' "; }
+ else { $cmd .= "\"$x\" "; }
+}
+
+#$Data::Dumper::Indent=0;
+$ret = 0;
+$numfail = 0;
+%active_pids=();
+
+use POSIX ":sys_wait_h";
+for ($jobid = $jobstart; $jobid <= $jobend; $jobid++) {
+ if (scalar(keys %active_pids) >= $max_jobs_run) {
+
+ # Lets wait for a change in any child's status
+ # Then we have to work out which child finished
+ $r = waitpid(-1, 0);
+ $code = $?;
+ if ($r < 0 ) { die "run.pl: Error waiting for child process"; } # should never happen.
+ if ( defined $active_pids{$r} ) {
+ $jid=$active_pids{$r};
+ $fail[$jid]=$code;
+ if ($code !=0) { $numfail++;}
+ delete $active_pids{$r};
+ # print STDERR "Finished: $r/$jid " . Dumper(\%active_pids) . "\n";
+ } else {
+ die "run.pl: Cannot find the PID of the child process that just finished.";
+ }
+
+ # In theory we could do a non-blocking waitpid over all jobs running just
+ # to find out if only one or more jobs finished during the previous waitpid()
+ # However, we just omit this and will reap the next one in the next pass
+ # through the for(;;) cycle
+ }
+ $childpid = fork();
+ if (!defined $childpid) { die "run.pl: Error forking in run.pl (writing to $logfile)"; }
+ if ($childpid == 0) { # We're in the child... this branch
+ # executes the job and returns (possibly with an error status).
+ if (defined $jobname) {
+ $cmd =~ s/$jobname/$jobid/g;
+ $logfile =~ s/$jobname/$jobid/g;
+ }
+ # exit if the job does not need to be executed
+ pick_or_exit( $logfile );
+
+ system("mkdir -p `dirname $logfile` 2>/dev/null");
+ open(F, ">$logfile") || die "run.pl: Error opening log file $logfile";
+ print F "# " . $cmd . "\n";
+ print F "# Started at " . `date`;
+ $starttime = `date +'%s'`;
+ print F "#\n";
+ close(F);
+
+ # Pipe into bash.. make sure we're not using any other shell.
+ open(B, "|bash") || die "run.pl: Error opening shell command";
+ print B "( " . $cmd . ") 2>>$logfile >> $logfile";
+ close(B); # If there was an error, exit status is in $?
+ $ret = $?;
+
+ $lowbits = $ret & 127;
+ $highbits = $ret >> 8;
+ if ($lowbits != 0) { $return_str = "code $highbits; signal $lowbits" }
+ else { $return_str = "code $highbits"; }
+
+ $endtime = `date +'%s'`;
+ open(F, ">>$logfile") || die "run.pl: Error opening log file $logfile (again)";
+ $enddate = `date`;
+ chop $enddate;
+ print F "# Accounting: time=" . ($endtime - $starttime) . " threads=1\n";
+ print F "# Ended ($return_str) at " . $enddate . ", elapsed time " . ($endtime-$starttime) . " seconds\n";
+ close(F);
+ exit($ret == 0 ? 0 : 1);
+ } else {
+ $pid[$jobid] = $childpid;
+ $active_pids{$childpid} = $jobid;
+ # print STDERR "Queued: " . Dumper(\%active_pids) . "\n";
+ }
+}
+
+# Now we have submitted all the jobs, lets wait until all the jobs finish
+foreach $child (keys %active_pids) {
+ $jobid=$active_pids{$child};
+ $r = waitpid($pid[$jobid], 0);
+ $code = $?;
+ if ($r == -1) { die "run.pl: Error waiting for child process"; } # should never happen.
+ if ($r != 0) { $fail[$jobid]=$code; $numfail++ if $code!=0; } # Completed successfully
+}
+
+# Some sanity checks:
+# The $fail array should not contain undefined codes
+# The number of non-zeros in that array should be equal to $numfail
+# We cannot do foreach() here, as the JOB ids do not start at zero
+$failed_jids=0;
+for ($jobid = $jobstart; $jobid <= $jobend; $jobid++) {
+ $job_return = $fail[$jobid];
+ if (not defined $job_return ) {
+ # print Dumper(\@fail);
+
+ die "run.pl: Sanity check failed: we have indication that some jobs are running " .
+ "even after we waited for all jobs to finish" ;
+ }
+ if ($job_return != 0 ){ $failed_jids++;}
+}
+if ($failed_jids != $numfail) {
+ die "run.pl: Sanity check failed: cannot find out how many jobs failed ($failed_jids x $numfail)."
+}
+if ($numfail > 0) { $ret = 1; }
+
+if ($ret != 0) {
+ $njobs = $jobend - $jobstart + 1;
+ if ($njobs == 1) {
+ if (defined $jobname) {
+ $logfile =~ s/$jobname/$jobstart/; # only one numbered job, so replace name with
+ # that job.
+ }
+ print STDERR "run.pl: job failed, log is in $logfile\n";
+ if ($logfile =~ m/JOB/) {
+ print STDERR "run.pl: probably you forgot to put JOB=1:\$nj in your script.";
+ }
+ }
+ else {
+ $logfile =~ s/$jobname/*/g;
+ print STDERR "run.pl: $numfail / $njobs failed, log is in $logfile\n";
+ }
+}
+
+
+exit ($ret);
diff --git a/egs_modelscope/common/utils/shuffle_list.pl b/egs_modelscope/common/utils/shuffle_list.pl
new file mode 100755
index 0000000..a116200
--- /dev/null
+++ b/egs_modelscope/common/utils/shuffle_list.pl
@@ -0,0 +1,44 @@
+#!/usr/bin/env perl
+
+# Copyright 2013 Johns Hopkins University (author: Daniel Povey)
+
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+if ($ARGV[0] eq "--srand") {
+ $n = $ARGV[1];
+ $n =~ m/\d+/ || die "Bad argument to --srand option: \"$n\"";
+ srand($ARGV[1]);
+ shift;
+ shift;
+} else {
+ srand(0); # Gives inconsistent behavior if we don't seed.
+}
+
+if (@ARGV > 1 || $ARGV[0] =~ m/^-.+/) { # >1 args, or an option we
+ # don't understand.
+ print "Usage: shuffle_list.pl [--srand N] [input file] > output\n";
+ print "randomizes the order of lines of input.\n";
+ exit(1);
+}
+
+@lines;
+while (<>) {
+ push @lines, [ (rand(), $_)] ;
+}
+
+@lines = sort { $a->[0] cmp $b->[0] } @lines;
+foreach $l (@lines) {
+ print $l->[1];
+}
\ No newline at end of file
diff --git a/egs_modelscope/common/utils/split_data.py b/egs_modelscope/common/utils/split_data.py
new file mode 100755
index 0000000..060eae6
--- /dev/null
+++ b/egs_modelscope/common/utils/split_data.py
@@ -0,0 +1,60 @@
+import os
+import sys
+import random
+
+
+in_dir = sys.argv[1]
+out_dir = sys.argv[2]
+num_split = sys.argv[3]
+
+
+def split_scp(scp, num):
+ assert len(scp) >= num
+ avg = len(scp) // num
+ out = []
+ begin = 0
+
+ for i in range(num):
+ if i == num - 1:
+ out.append(scp[begin:])
+ else:
+ out.append(scp[begin:begin+avg])
+ begin += avg
+
+ return out
+
+
+os.path.exists("{}/wav.scp".format(in_dir))
+os.path.exists("{}/text".format(in_dir))
+
+with open("{}/wav.scp".format(in_dir), 'r') as infile:
+ wav_list = infile.readlines()
+
+with open("{}/text".format(in_dir), 'r') as infile:
+ text_list = infile.readlines()
+
+assert len(wav_list) == len(text_list)
+
+x = list(zip(wav_list, text_list))
+random.shuffle(x)
+wav_shuffle_list, text_shuffle_list = zip(*x)
+
+num_split = int(num_split)
+wav_split_list = split_scp(wav_shuffle_list, num_split)
+text_split_list = split_scp(text_shuffle_list, num_split)
+
+for idx, wav_list in enumerate(wav_split_list, 1):
+ path = out_dir + "/split" + str(num_split) + "/" + str(idx)
+ if not os.path.exists(path):
+ os.makedirs(path)
+ with open("{}/wav.scp".format(path), 'w') as wav_writer:
+ for line in wav_list:
+ wav_writer.write(line)
+
+for idx, text_list in enumerate(text_split_list, 1):
+ path = out_dir + "/split" + str(num_split) + "/" + str(idx)
+ if not os.path.exists(path):
+ os.makedirs(path)
+ with open("{}/text".format(path), 'w') as text_writer:
+ for line in text_list:
+ text_writer.write(line)
diff --git a/egs_modelscope/common/utils/split_scp.pl b/egs_modelscope/common/utils/split_scp.pl
new file mode 100755
index 0000000..0876dcb
--- /dev/null
+++ b/egs_modelscope/common/utils/split_scp.pl
@@ -0,0 +1,246 @@
+#!/usr/bin/env perl
+
+# Copyright 2010-2011 Microsoft Corporation
+
+# See ../../COPYING for clarification regarding multiple authors
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+# This program splits up any kind of .scp or archive-type file.
+# If there is no utt2spk option it will work on any text file and
+# will split it up with an approximately equal number of lines in
+# each but.
+# With the --utt2spk option it will work on anything that has the
+# utterance-id as the first entry on each line; the utt2spk file is
+# of the form "utterance speaker" (on each line).
+# It splits it into equal size chunks as far as it can. If you use the utt2spk
+# option it will make sure these chunks coincide with speaker boundaries. In
+# this case, if there are more chunks than speakers (and in some other
+# circumstances), some of the resulting chunks will be empty and it will print
+# an error message and exit with nonzero status.
+# You will normally call this like:
+# split_scp.pl scp scp.1 scp.2 scp.3 ...
+# or
+# split_scp.pl --utt2spk=utt2spk scp scp.1 scp.2 scp.3 ...
+# Note that you can use this script to split the utt2spk file itself,
+# e.g. split_scp.pl --utt2spk=utt2spk utt2spk utt2spk.1 utt2spk.2 ...
+
+# You can also call the scripts like:
+# split_scp.pl -j 3 0 scp scp.0
+# [note: with this option, it assumes zero-based indexing of the split parts,
+# i.e. the second number must be 0 <= n < num-jobs.]
+
+use warnings;
+
+$num_jobs = 0;
+$job_id = 0;
+$utt2spk_file = "";
+$one_based = 0;
+
+for ($x = 1; $x <= 3 && @ARGV > 0; $x++) {
+ if ($ARGV[0] eq "-j") {
+ shift @ARGV;
+ $num_jobs = shift @ARGV;
+ $job_id = shift @ARGV;
+ }
+ if ($ARGV[0] =~ /--utt2spk=(.+)/) {
+ $utt2spk_file=$1;
+ shift;
+ }
+ if ($ARGV[0] eq '--one-based') {
+ $one_based = 1;
+ shift @ARGV;
+ }
+}
+
+if ($num_jobs != 0 && ($num_jobs < 0 || $job_id - $one_based < 0 ||
+ $job_id - $one_based >= $num_jobs)) {
+ die "$0: Invalid job number/index values for '-j $num_jobs $job_id" .
+ ($one_based ? " --one-based" : "") . "'\n"
+}
+
+$one_based
+ and $job_id--;
+
+if(($num_jobs == 0 && @ARGV < 2) || ($num_jobs > 0 && (@ARGV < 1 || @ARGV > 2))) {
+ die
+"Usage: split_scp.pl [--utt2spk=<utt2spk_file>] in.scp out1.scp out2.scp ...
+ or: split_scp.pl -j num-jobs job-id [--one-based] [--utt2spk=<utt2spk_file>] in.scp [out.scp]
+ ... where 0 <= job-id < num-jobs, or 1 <= job-id <- num-jobs if --one-based.\n";
+}
+
+$error = 0;
+$inscp = shift @ARGV;
+if ($num_jobs == 0) { # without -j option
+ @OUTPUTS = @ARGV;
+} else {
+ for ($j = 0; $j < $num_jobs; $j++) {
+ if ($j == $job_id) {
+ if (@ARGV > 0) { push @OUTPUTS, $ARGV[0]; }
+ else { push @OUTPUTS, "-"; }
+ } else {
+ push @OUTPUTS, "/dev/null";
+ }
+ }
+}
+
+if ($utt2spk_file ne "") { # We have the --utt2spk option...
+ open($u_fh, '<', $utt2spk_file) || die "$0: Error opening utt2spk file $utt2spk_file: $!\n";
+ while(<$u_fh>) {
+ @A = split;
+ @A == 2 || die "$0: Bad line $_ in utt2spk file $utt2spk_file\n";
+ ($u,$s) = @A;
+ $utt2spk{$u} = $s;
+ }
+ close $u_fh;
+ open($i_fh, '<', $inscp) || die "$0: Error opening input scp file $inscp: $!\n";
+ @spkrs = ();
+ while(<$i_fh>) {
+ @A = split;
+ if(@A == 0) { die "$0: Empty or space-only line in scp file $inscp\n"; }
+ $u = $A[0];
+ $s = $utt2spk{$u};
+ defined $s || die "$0: No utterance $u in utt2spk file $utt2spk_file\n";
+ if(!defined $spk_count{$s}) {
+ push @spkrs, $s;
+ $spk_count{$s} = 0;
+ $spk_data{$s} = []; # ref to new empty array.
+ }
+ $spk_count{$s}++;
+ push @{$spk_data{$s}}, $_;
+ }
+ # Now split as equally as possible ..
+ # First allocate spks to files by allocating an approximately
+ # equal number of speakers.
+ $numspks = @spkrs; # number of speakers.
+ $numscps = @OUTPUTS; # number of output files.
+ if ($numspks < $numscps) {
+ die "$0: Refusing to split data because number of speakers $numspks " .
+ "is less than the number of output .scp files $numscps\n";
+ }
+ for($scpidx = 0; $scpidx < $numscps; $scpidx++) {
+ $scparray[$scpidx] = []; # [] is array reference.
+ }
+ for ($spkidx = 0; $spkidx < $numspks; $spkidx++) {
+ $scpidx = int(($spkidx*$numscps) / $numspks);
+ $spk = $spkrs[$spkidx];
+ push @{$scparray[$scpidx]}, $spk;
+ $scpcount[$scpidx] += $spk_count{$spk};
+ }
+
+ # Now will try to reassign beginning + ending speakers
+ # to different scp's and see if it gets more balanced.
+ # Suppose objf we're minimizing is sum_i (num utts in scp[i] - average)^2.
+ # We can show that if considering changing just 2 scp's, we minimize
+ # this by minimizing the squared difference in sizes. This is
+ # equivalent to minimizing the absolute difference in sizes. This
+ # shows this method is bound to converge.
+
+ $changed = 1;
+ while($changed) {
+ $changed = 0;
+ for($scpidx = 0; $scpidx < $numscps; $scpidx++) {
+ # First try to reassign ending spk of this scp.
+ if($scpidx < $numscps-1) {
+ $sz = @{$scparray[$scpidx]};
+ if($sz > 0) {
+ $spk = $scparray[$scpidx]->[$sz-1];
+ $count = $spk_count{$spk};
+ $nutt1 = $scpcount[$scpidx];
+ $nutt2 = $scpcount[$scpidx+1];
+ if( abs( ($nutt2+$count) - ($nutt1-$count))
+ < abs($nutt2 - $nutt1)) { # Would decrease
+ # size-diff by reassigning spk...
+ $scpcount[$scpidx+1] += $count;
+ $scpcount[$scpidx] -= $count;
+ pop @{$scparray[$scpidx]};
+ unshift @{$scparray[$scpidx+1]}, $spk;
+ $changed = 1;
+ }
+ }
+ }
+ if($scpidx > 0 && @{$scparray[$scpidx]} > 0) {
+ $spk = $scparray[$scpidx]->[0];
+ $count = $spk_count{$spk};
+ $nutt1 = $scpcount[$scpidx-1];
+ $nutt2 = $scpcount[$scpidx];
+ if( abs( ($nutt2-$count) - ($nutt1+$count))
+ < abs($nutt2 - $nutt1)) { # Would decrease
+ # size-diff by reassigning spk...
+ $scpcount[$scpidx-1] += $count;
+ $scpcount[$scpidx] -= $count;
+ shift @{$scparray[$scpidx]};
+ push @{$scparray[$scpidx-1]}, $spk;
+ $changed = 1;
+ }
+ }
+ }
+ }
+ # Now print out the files...
+ for($scpidx = 0; $scpidx < $numscps; $scpidx++) {
+ $scpfile = $OUTPUTS[$scpidx];
+ ($scpfile ne '-' ? open($f_fh, '>', $scpfile)
+ : open($f_fh, '>&', \*STDOUT)) ||
+ die "$0: Could not open scp file $scpfile for writing: $!\n";
+ $count = 0;
+ if(@{$scparray[$scpidx]} == 0) {
+ print STDERR "$0: eError: split_scp.pl producing empty .scp file " .
+ "$scpfile (too many splits and too few speakers?)\n";
+ $error = 1;
+ } else {
+ foreach $spk ( @{$scparray[$scpidx]} ) {
+ print $f_fh @{$spk_data{$spk}};
+ $count += $spk_count{$spk};
+ }
+ $count == $scpcount[$scpidx] || die "Count mismatch [code error]";
+ }
+ close($f_fh);
+ }
+} else {
+ # This block is the "normal" case where there is no --utt2spk
+ # option and we just break into equal size chunks.
+
+ open($i_fh, '<', $inscp) || die "$0: Error opening input scp file $inscp: $!\n";
+
+ $numscps = @OUTPUTS; # size of array.
+ @F = ();
+ while(<$i_fh>) {
+ push @F, $_;
+ }
+ $numlines = @F;
+ if($numlines == 0) {
+ print STDERR "$0: error: empty input scp file $inscp\n";
+ $error = 1;
+ }
+ $linesperscp = int( $numlines / $numscps); # the "whole part"..
+ $linesperscp >= 1 || die "$0: You are splitting into too many pieces! [reduce \$nj ($numscps) to be smaller than the number of lines ($numlines) in $inscp]\n";
+ $remainder = $numlines - ($linesperscp * $numscps);
+ ($remainder >= 0 && $remainder < $numlines) || die "bad remainder $remainder";
+ # [just doing int() rounds down].
+ $n = 0;
+ for($scpidx = 0; $scpidx < @OUTPUTS; $scpidx++) {
+ $scpfile = $OUTPUTS[$scpidx];
+ ($scpfile ne '-' ? open($o_fh, '>', $scpfile)
+ : open($o_fh, '>&', \*STDOUT)) ||
+ die "$0: Could not open scp file $scpfile for writing: $!\n";
+ for($k = 0; $k < $linesperscp + ($scpidx < $remainder ? 1 : 0); $k++) {
+ print $o_fh $F[$n++];
+ }
+ close($o_fh) || die "$0: Eror closing scp file $scpfile: $!\n";
+ }
+ $n == $numlines || die "$n != $numlines [code error]";
+}
+
+exit ($error);
diff --git a/egs_modelscope/common/utils/subset_data_dir_tr_cv.sh b/egs_modelscope/common/utils/subset_data_dir_tr_cv.sh
new file mode 100755
index 0000000..e16cebd
--- /dev/null
+++ b/egs_modelscope/common/utils/subset_data_dir_tr_cv.sh
@@ -0,0 +1,30 @@
+#!/usr/bin/env bash
+
+dev_num_utt=1000
+
+echo "$0 $@"
+. utils/parse_options.sh || exit 1;
+
+train_data=$1
+out_dir=$2
+
+[ ! -f ${train_data}/wav.scp ] && echo "$0: no such file ${train_data}/wav.scp" && exit 1;
+[ ! -f ${train_data}/text ] && echo "$0: no such file ${train_data}/text" && exit 1;
+
+mkdir -p ${out_dir}/train && mkdir -p ${out_dir}/dev
+
+cp ${train_data}/wav.scp ${out_dir}/train/wav.scp.bak
+cp ${train_data}/text ${out_dir}/train/text.bak
+
+num_utt=$(wc -l <${out_dir}/train/wav.scp.bak)
+
+utils/shuffle_list.pl --srand 1 ${out_dir}/train/wav.scp.bak > ${out_dir}/train/wav.scp.shuf
+head -n ${dev_num_utt} ${out_dir}/train/wav.scp.shuf > ${out_dir}/dev/wav.scp
+tail -n $((${num_utt}-${dev_num_utt})) ${out_dir}/train/wav.scp.shuf > ${out_dir}/train/wav.scp
+
+utils/shuffle_list.pl --srand 1 ${out_dir}/train/text.bak > ${out_dir}/train/text.shuf
+head -n ${dev_num_utt} ${out_dir}/train/text.shuf > ${out_dir}/dev/text
+tail -n $((${num_utt}-${dev_num_utt})) ${out_dir}/train/text.shuf > ${out_dir}/train/text
+
+rm ${out_dir}/train/wav.scp.bak ${out_dir}/train/text.bak
+rm ${out_dir}/train/wav.scp.shuf ${out_dir}/train/text.shuf
diff --git a/egs_modelscope/common/utils/text2token.py b/egs_modelscope/common/utils/text2token.py
new file mode 100755
index 0000000..56c3913
--- /dev/null
+++ b/egs_modelscope/common/utils/text2token.py
@@ -0,0 +1,135 @@
+#!/usr/bin/env python3
+
+# Copyright 2017 Johns Hopkins University (Shinji Watanabe)
+# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
+
+
+import argparse
+import codecs
+import re
+import sys
+
+is_python2 = sys.version_info[0] == 2
+
+
+def exist_or_not(i, match_pos):
+ start_pos = None
+ end_pos = None
+ for pos in match_pos:
+ if pos[0] <= i < pos[1]:
+ start_pos = pos[0]
+ end_pos = pos[1]
+ break
+
+ return start_pos, end_pos
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="convert raw text to tokenized text",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--nchar",
+ "-n",
+ default=1,
+ type=int,
+ help="number of characters to split, i.e., \
+ aabb -> a a b b with -n 1 and aa bb with -n 2",
+ )
+ parser.add_argument(
+ "--skip-ncols", "-s", default=0, type=int, help="skip first n columns"
+ )
+ parser.add_argument("--space", default="<space>", type=str, help="space symbol")
+ parser.add_argument(
+ "--non-lang-syms",
+ "-l",
+ default=None,
+ type=str,
+ help="list of non-linguistic symobles, e.g., <NOISE> etc.",
+ )
+ parser.add_argument("text", type=str, default=False, nargs="?", help="input text")
+ parser.add_argument(
+ "--trans_type",
+ "-t",
+ type=str,
+ default="char",
+ choices=["char", "phn"],
+ help="""Transcript type. char/phn. e.g., for TIMIT FADG0_SI1279 -
+ If trans_type is char,
+ read from SI1279.WRD file -> "bricks are an alternative"
+ Else if trans_type is phn,
+ read from SI1279.PHN file -> "sil b r ih sil k s aa r er n aa l
+ sil t er n ih sil t ih v sil" """,
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ rs = []
+ if args.non_lang_syms is not None:
+ with codecs.open(args.non_lang_syms, "r", encoding="utf-8") as f:
+ nls = [x.rstrip() for x in f.readlines()]
+ rs = [re.compile(re.escape(x)) for x in nls]
+
+ if args.text:
+ f = codecs.open(args.text, encoding="utf-8")
+ else:
+ f = codecs.getreader("utf-8")(sys.stdin if is_python2 else sys.stdin.buffer)
+
+ sys.stdout = codecs.getwriter("utf-8")(
+ sys.stdout if is_python2 else sys.stdout.buffer
+ )
+ line = f.readline()
+ n = args.nchar
+ while line:
+ x = line.split()
+ print(" ".join(x[: args.skip_ncols]), end=" ")
+ a = " ".join(x[args.skip_ncols :])
+
+ # get all matched positions
+ match_pos = []
+ for r in rs:
+ i = 0
+ while i >= 0:
+ m = r.search(a, i)
+ if m:
+ match_pos.append([m.start(), m.end()])
+ i = m.end()
+ else:
+ break
+
+ if args.trans_type == "phn":
+ a = a.split(" ")
+ else:
+ if len(match_pos) > 0:
+ chars = []
+ i = 0
+ while i < len(a):
+ start_pos, end_pos = exist_or_not(i, match_pos)
+ if start_pos is not None:
+ chars.append(a[start_pos:end_pos])
+ i = end_pos
+ else:
+ chars.append(a[i])
+ i += 1
+ a = chars
+
+ a = [a[j : j + n] for j in range(0, len(a), n)]
+
+ a_flat = []
+ for z in a:
+ a_flat.append("".join(z))
+
+ a_chars = [z.replace(" ", args.space) for z in a_flat]
+ if args.trans_type == "phn":
+ a_chars = [z.replace("sil", args.space) for z in a_chars]
+ print(" ".join(a_chars))
+ line = f.readline()
+
+
+if __name__ == "__main__":
+ main()
diff --git a/egs_modelscope/common/utils/text_tokenize.py b/egs_modelscope/common/utils/text_tokenize.py
new file mode 100755
index 0000000..962ea11
--- /dev/null
+++ b/egs_modelscope/common/utils/text_tokenize.py
@@ -0,0 +1,106 @@
+import re
+import argparse
+
+
+def load_dict(seg_file):
+ seg_dict = {}
+ with open(seg_file, 'r') as infile:
+ for line in infile:
+ s = line.strip().split()
+ key = s[0]
+ value = s[1:]
+ seg_dict[key] = " ".join(value)
+ return seg_dict
+
+
+def forward_segment(text, dic):
+ word_list = []
+ i = 0
+ while i < len(text):
+ longest_word = text[i]
+ for j in range(i + 1, len(text) + 1):
+ word = text[i:j]
+ if word in dic:
+ if len(word) > len(longest_word):
+ longest_word = word
+ word_list.append(longest_word)
+ i += len(longest_word)
+ return word_list
+
+
+def tokenize(txt,
+ seg_dict):
+ out_txt = ""
+ pattern = re.compile(r"([\u4E00-\u9FA5A-Za-z0-9])")
+ for word in txt:
+ if pattern.match(word):
+ if word in seg_dict:
+ out_txt += seg_dict[word] + " "
+ else:
+ out_txt += "<unk>" + " "
+ else:
+ continue
+ return out_txt.strip()
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="text tokenize",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--text-file",
+ "-t",
+ default=False,
+ required=True,
+ type=str,
+ help="input text",
+ )
+ parser.add_argument(
+ "--seg-file",
+ "-s",
+ default=False,
+ required=True,
+ type=str,
+ help="seg file",
+ )
+ parser.add_argument(
+ "--txt-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="txt index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ txt_writer = open("{}/text.{}.txt".format(args.output_dir, args.txt_index), 'w')
+ shape_writer = open("{}/len.{}".format(args.output_dir, args.txt_index), 'w')
+ seg_dict = load_dict(args.seg_file)
+ with open(args.text_file, 'r') as infile:
+ for line in infile:
+ s = line.strip().split()
+ text_id = s[0]
+ text_list = forward_segment("".join(s[1:]).lower(), seg_dict)
+ text = tokenize(text_list, seg_dict)
+ lens = len(text.strip().split())
+ txt_writer.write(text_id + " " + text + '\n')
+ shape_writer.write(text_id + " " + str(lens) + '\n')
+
+
+if __name__ == '__main__':
+ main()
+
diff --git a/egs_modelscope/common/utils/text_tokenize.sh b/egs_modelscope/common/utils/text_tokenize.sh
new file mode 100755
index 0000000..6b74fef
--- /dev/null
+++ b/egs_modelscope/common/utils/text_tokenize.sh
@@ -0,0 +1,35 @@
+#!/usr/bin/env bash
+
+
+# Begin configuration section.
+nj=32
+cmd=utils/run.pl
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+# tokenize configuration
+text_dir=$1
+seg_file=$2
+logdir=$3
+output_dir=$4
+
+txt_dir=${output_dir}/txt; mkdir -p ${output_dir}/txt
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/text_tokenize.JOB.log \
+ python utils/text_tokenize.py -t ${text_dir}/txt/text.JOB.txt \
+ -s ${seg_file} -i JOB -o ${txt_dir} \
+ || exit 1;
+
+# concatenate the text files together.
+for n in $(seq $nj); do
+ cat ${txt_dir}/text.$n.txt || exit 1
+done > ${output_dir}/text || exit 1
+
+for n in $(seq $nj); do
+ cat ${txt_dir}/len.$n || exit 1
+done > ${output_dir}/text_shape || exit 1
+
+echo "$0: Succeeded text tokenize"
diff --git a/egs_modelscope/common/utils/textnorm_zh.py b/egs_modelscope/common/utils/textnorm_zh.py
new file mode 100755
index 0000000..79feb83
--- /dev/null
+++ b/egs_modelscope/common/utils/textnorm_zh.py
@@ -0,0 +1,834 @@
+#!/usr/bin/env python3
+# coding=utf-8
+
+# Authors:
+# 2019.5 Zhiyang Zhou (https://github.com/Joee1995/chn_text_norm.git)
+# 2019.9 Jiayu DU
+#
+# requirements:
+# - python 3.X
+# notes: python 2.X WILL fail or produce misleading results
+
+import sys, os, argparse, codecs, string, re
+
+# ================================================================================ #
+# basic constant
+# ================================================================================ #
+CHINESE_DIGIS = u'闆朵竴浜屼笁鍥涗簲鍏竷鍏節'
+BIG_CHINESE_DIGIS_SIMPLIFIED = u'闆跺9璐板弫鑲嗕紞闄嗘煉鎹岀帠'
+BIG_CHINESE_DIGIS_TRADITIONAL = u'闆跺9璨冲弮鑲嗕紞闄告煉鎹岀帠'
+SMALLER_BIG_CHINESE_UNITS_SIMPLIFIED = u'鍗佺櫨鍗冧竾'
+SMALLER_BIG_CHINESE_UNITS_TRADITIONAL = u'鎷句桨浠熻惉'
+LARGER_CHINESE_NUMERING_UNITS_SIMPLIFIED = u'浜垮厗浜灀绉┌娌熸锭姝h浇'
+LARGER_CHINESE_NUMERING_UNITS_TRADITIONAL = u'鍎勫厗浜灀绉┌婧濇緱姝h級'
+SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED = u'鍗佺櫨鍗冧竾'
+SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL = u'鎷句桨浠熻惉'
+
+ZERO_ALT = u'銆�'
+ONE_ALT = u'骞�'
+TWO_ALTS = [u'涓�', u'鍏�']
+
+POSITIVE = [u'姝�', u'姝�']
+NEGATIVE = [u'璐�', u'璨�']
+POINT = [u'鐐�', u'榛�']
+# PLUS = [u'鍔�', u'鍔�']
+# SIL = [u'鏉�', u'妲�']
+
+FILLER_CHARS = ['鍛�', '鍟�']
+ER_WHITELIST = '(鍎垮コ|鍎垮瓙|鍎垮瓩|濂冲効|鍎垮|濡诲効|' \
+ '鑳庡効|濠村効|鏂扮敓鍎縷濠村辜鍎縷骞煎効|灏戝効|灏忓効|鍎挎瓕|鍎跨|鍎跨|鎵樺効鎵�|瀛ゅ効|' \
+ '鍎挎垙|鍎垮寲|鍙板効搴剕楣垮効宀泑姝e効鍏粡|鍚婂効閮庡綋|鐢熷効鑲插コ|鎵樺効甯﹀コ|鍏诲効闃茶�亅鐥村効鍛嗗コ|' \
+ '浣冲効浣冲|鍎挎�滃吔鎵皘鍎挎棤甯哥埗|鍎夸笉瀚屾瘝涓憒鍎胯鍗冮噷姣嶆媴蹇鍎垮ぇ涓嶇敱鐖穦鑻忎篂鍎�)'
+
+# 涓枃鏁板瓧绯荤粺绫诲瀷
+NUMBERING_TYPES = ['low', 'mid', 'high']
+
+CURRENCY_NAMES = '(浜烘皯甯亅缇庡厓|鏃ュ厓|鑻遍晳|娆у厓|椹厠|娉曢儙|鍔犳嬁澶у厓|婢冲厓|娓竵|鍏堜护|鑺叞椹厠|鐖卞皵鍏伴晳|' \
+ '閲屾媺|鑽峰叞鐩緗鍩冩柉搴撳|姣斿濉攟鍗板凹鐩緗鏋楀悏鐗箌鏂拌タ鍏板厓|姣旂储|鍗㈠竷|鏂板姞鍧″厓|闊╁厓|娉伴摙)'
+CURRENCY_UNITS = '((浜縷鍗冧竾|鐧句竾|涓噟鍗億鐧�)|(浜縷鍗冧竾|鐧句竾|涓噟鍗億鐧緗)鍏億(浜縷鍗冧竾|鐧句竾|涓噟鍗億鐧緗)鍧梶瑙抾姣泑鍒�)'
+COM_QUANTIFIERS = '(鍖箌寮爘搴鍥瀨鍦簗灏緗鏉涓獆棣東闃檤闃祙缃憒鐐畖椤秥涓榺妫祙鍙獆鏀瘄琚瓅杈唡鎸憒鎷厊棰梶澹硘绐爘鏇瞸澧檤缇鑵攟' \
+ '鐮搴瀹璐瘄鎵巪鎹唡鍒�|浠鎵搢鎵媩缃梶鍧灞眧宀瓅姹焲婧獆閽焲闃焲鍗晐鍙寍瀵箌鍑簗鍙澶磡鑴殀鏉縷璺硘鏋潀浠秥璐磡' \
+ '閽坾绾縷绠鍚峾浣峾韬珅鍫倈璇緗鏈瑋椤祙瀹秥鎴穦灞倈涓潀姣珅鍘榺鍒唡閽眧涓鏂鎷厊閾鐭硘閽閿眧蹇絴(鍗億姣珅寰�)鍏媩' \
+ '姣珅鍘榺鍒唡瀵竱灏簗涓坾閲寍瀵粅甯竱閾簗绋媩(鍗億鍒唡鍘榺姣珅寰�)绫硘鎾畖鍕簗鍚坾鍗噟鏂梶鐭硘鐩榺纰梶纰焲鍙爘妗秥绗紎鐩唡' \
+ '鐩抾鏉瘄閽焲鏂泑閿厊绨媩绡畖鐩榺妗秥缃恷鐡秥澹秥鍗畖鐩弢绠﹟绠眧鐓瞸鍟東琚媩閽祙骞磡鏈坾鏃瀛鍒粅鏃秥鍛▅澶﹟绉抾鍒唡鏃瑋' \
+ '绾獆宀亅涓東鏇磡澶渱鏄澶弢绉媩鍐瑋浠浼弢杈坾涓竱娉绮抾棰梶骞鍫唡鏉鏍箌鏀瘄閬搢闈鐗噟寮爘棰梶鍧�)'
+
+# punctuation information are based on Zhon project (https://github.com/tsroten/zhon.git)
+CHINESE_PUNC_STOP = '锛侊紵锝°��'
+CHINESE_PUNC_NON_STOP = '锛傦純锛勶紖锛嗭紘锛堬級锛婏紜锛岋紞锛忥細锛涳紲锛濓紴锛狅蓟锛硷冀锛撅伎锝�锝涳綔锝濓綖锝燂綘锝剑锝ゃ�併�冦�嬨�屻�嶃�庛�忋�愩�戙�斻�曘�栥�椼�樸�欍�氥�涖�溿�濄�炪�熴�般�俱�库�撯�斺�樷�欌�涒�溾�濃�炩�熲�︹�э箯'
+CHINESE_PUNC_LIST = CHINESE_PUNC_STOP + CHINESE_PUNC_NON_STOP
+
+# ================================================================================ #
+# basic class
+# ================================================================================ #
+class ChineseChar(object):
+ """
+ 涓枃瀛楃
+ 姣忎釜瀛楃瀵瑰簲绠�浣撳拰绻佷綋,
+ e.g. 绠�浣� = '璐�', 绻佷綋 = '璨�'
+ 杞崲鏃跺彲杞崲涓虹畝浣撴垨绻佷綋
+ """
+
+ def __init__(self, simplified, traditional):
+ self.simplified = simplified
+ self.traditional = traditional
+ #self.__repr__ = self.__str__
+
+ def __str__(self):
+ return self.simplified or self.traditional or None
+
+ def __repr__(self):
+ return self.__str__()
+
+
+class ChineseNumberUnit(ChineseChar):
+ """
+ 涓枃鏁板瓧/鏁颁綅瀛楃
+ 姣忎釜瀛楃闄ょ箒绠�浣撳杩樻湁涓�涓澶栫殑澶у啓瀛楃
+ e.g. '闄�' 鍜� '闄�'
+ """
+
+ def __init__(self, power, simplified, traditional, big_s, big_t):
+ super(ChineseNumberUnit, self).__init__(simplified, traditional)
+ self.power = power
+ self.big_s = big_s
+ self.big_t = big_t
+
+ def __str__(self):
+ return '10^{}'.format(self.power)
+
+ @classmethod
+ def create(cls, index, value, numbering_type=NUMBERING_TYPES[1], small_unit=False):
+
+ if small_unit:
+ return ChineseNumberUnit(power=index + 1,
+ simplified=value[0], traditional=value[1], big_s=value[1], big_t=value[1])
+ elif numbering_type == NUMBERING_TYPES[0]:
+ return ChineseNumberUnit(power=index + 8,
+ simplified=value[0], traditional=value[1], big_s=value[0], big_t=value[1])
+ elif numbering_type == NUMBERING_TYPES[1]:
+ return ChineseNumberUnit(power=(index + 2) * 4,
+ simplified=value[0], traditional=value[1], big_s=value[0], big_t=value[1])
+ elif numbering_type == NUMBERING_TYPES[2]:
+ return ChineseNumberUnit(power=pow(2, index + 3),
+ simplified=value[0], traditional=value[1], big_s=value[0], big_t=value[1])
+ else:
+ raise ValueError(
+ 'Counting type should be in {0} ({1} provided).'.format(NUMBERING_TYPES, numbering_type))
+
+
+class ChineseNumberDigit(ChineseChar):
+ """
+ 涓枃鏁板瓧瀛楃
+ """
+
+ def __init__(self, value, simplified, traditional, big_s, big_t, alt_s=None, alt_t=None):
+ super(ChineseNumberDigit, self).__init__(simplified, traditional)
+ self.value = value
+ self.big_s = big_s
+ self.big_t = big_t
+ self.alt_s = alt_s
+ self.alt_t = alt_t
+
+ def __str__(self):
+ return str(self.value)
+
+ @classmethod
+ def create(cls, i, v):
+ return ChineseNumberDigit(i, v[0], v[1], v[2], v[3])
+
+
+class ChineseMath(ChineseChar):
+ """
+ 涓枃鏁颁綅瀛楃
+ """
+
+ def __init__(self, simplified, traditional, symbol, expression=None):
+ super(ChineseMath, self).__init__(simplified, traditional)
+ self.symbol = symbol
+ self.expression = expression
+ self.big_s = simplified
+ self.big_t = traditional
+
+
+CC, CNU, CND, CM = ChineseChar, ChineseNumberUnit, ChineseNumberDigit, ChineseMath
+
+
+class NumberSystem(object):
+ """
+ 涓枃鏁板瓧绯荤粺
+ """
+ pass
+
+
+class MathSymbol(object):
+ """
+ 鐢ㄤ簬涓枃鏁板瓧绯荤粺鐨勬暟瀛︾鍙� (绻�/绠�浣�), e.g.
+ positive = ['姝�', '姝�']
+ negative = ['璐�', '璨�']
+ point = ['鐐�', '榛�']
+ """
+
+ def __init__(self, positive, negative, point):
+ self.positive = positive
+ self.negative = negative
+ self.point = point
+
+ def __iter__(self):
+ for v in self.__dict__.values():
+ yield v
+
+
+# class OtherSymbol(object):
+# """
+# 鍏朵粬绗﹀彿
+# """
+#
+# def __init__(self, sil):
+# self.sil = sil
+#
+# def __iter__(self):
+# for v in self.__dict__.values():
+# yield v
+
+
+# ================================================================================ #
+# basic utils
+# ================================================================================ #
+def create_system(numbering_type=NUMBERING_TYPES[1]):
+ """
+ 鏍规嵁鏁板瓧绯荤粺绫诲瀷杩斿洖鍒涘缓鐩稿簲鐨勬暟瀛楃郴缁燂紝榛樿涓� mid
+ NUMBERING_TYPES = ['low', 'mid', 'high']: 涓枃鏁板瓧绯荤粺绫诲瀷
+ low: '鍏�' = '浜�' * '鍗�' = $10^{9}$, '浜�' = '鍏�' * '鍗�', etc.
+ mid: '鍏�' = '浜�' * '涓�' = $10^{12}$, '浜�' = '鍏�' * '涓�', etc.
+ high: '鍏�' = '浜�' * '浜�' = $10^{16}$, '浜�' = '鍏�' * '鍏�', etc.
+ 杩斿洖瀵瑰簲鐨勬暟瀛楃郴缁�
+ """
+
+ # chinese number units of '浜�' and larger
+ all_larger_units = zip(
+ LARGER_CHINESE_NUMERING_UNITS_SIMPLIFIED, LARGER_CHINESE_NUMERING_UNITS_TRADITIONAL)
+ larger_units = [CNU.create(i, v, numbering_type, False)
+ for i, v in enumerate(all_larger_units)]
+ # chinese number units of '鍗�, 鐧�, 鍗�, 涓�'
+ all_smaller_units = zip(
+ SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED, SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL)
+ smaller_units = [CNU.create(i, v, small_unit=True)
+ for i, v in enumerate(all_smaller_units)]
+ # digis
+ chinese_digis = zip(CHINESE_DIGIS, CHINESE_DIGIS,
+ BIG_CHINESE_DIGIS_SIMPLIFIED, BIG_CHINESE_DIGIS_TRADITIONAL)
+ digits = [CND.create(i, v) for i, v in enumerate(chinese_digis)]
+ digits[0].alt_s, digits[0].alt_t = ZERO_ALT, ZERO_ALT
+ digits[1].alt_s, digits[1].alt_t = ONE_ALT, ONE_ALT
+ digits[2].alt_s, digits[2].alt_t = TWO_ALTS[0], TWO_ALTS[1]
+
+ # symbols
+ positive_cn = CM(POSITIVE[0], POSITIVE[1], '+', lambda x: x)
+ negative_cn = CM(NEGATIVE[0], NEGATIVE[1], '-', lambda x: -x)
+ point_cn = CM(POINT[0], POINT[1], '.', lambda x,
+ y: float(str(x) + '.' + str(y)))
+ # sil_cn = CM(SIL[0], SIL[1], '-', lambda x, y: float(str(x) + '-' + str(y)))
+ system = NumberSystem()
+ system.units = smaller_units + larger_units
+ system.digits = digits
+ system.math = MathSymbol(positive_cn, negative_cn, point_cn)
+ # system.symbols = OtherSymbol(sil_cn)
+ return system
+
+
+def chn2num(chinese_string, numbering_type=NUMBERING_TYPES[1]):
+
+ def get_symbol(char, system):
+ for u in system.units:
+ if char in [u.traditional, u.simplified, u.big_s, u.big_t]:
+ return u
+ for d in system.digits:
+ if char in [d.traditional, d.simplified, d.big_s, d.big_t, d.alt_s, d.alt_t]:
+ return d
+ for m in system.math:
+ if char in [m.traditional, m.simplified]:
+ return m
+
+ def string2symbols(chinese_string, system):
+ int_string, dec_string = chinese_string, ''
+ for p in [system.math.point.simplified, system.math.point.traditional]:
+ if p in chinese_string:
+ int_string, dec_string = chinese_string.split(p)
+ break
+ return [get_symbol(c, system) for c in int_string], \
+ [get_symbol(c, system) for c in dec_string]
+
+ def correct_symbols(integer_symbols, system):
+ """
+ 涓�鐧惧叓 to 涓�鐧惧叓鍗�
+ 涓�浜夸竴鍗冧笁鐧句竾 to 涓�浜� 涓�鍗冧竾 涓夌櫨涓�
+ """
+
+ if integer_symbols and isinstance(integer_symbols[0], CNU):
+ if integer_symbols[0].power == 1:
+ integer_symbols = [system.digits[1]] + integer_symbols
+
+ if len(integer_symbols) > 1:
+ if isinstance(integer_symbols[-1], CND) and isinstance(integer_symbols[-2], CNU):
+ integer_symbols.append(
+ CNU(integer_symbols[-2].power - 1, None, None, None, None))
+
+ result = []
+ unit_count = 0
+ for s in integer_symbols:
+ if isinstance(s, CND):
+ result.append(s)
+ unit_count = 0
+ elif isinstance(s, CNU):
+ current_unit = CNU(s.power, None, None, None, None)
+ unit_count += 1
+
+ if unit_count == 1:
+ result.append(current_unit)
+ elif unit_count > 1:
+ for i in range(len(result)):
+ if isinstance(result[-i - 1], CNU) and result[-i - 1].power < current_unit.power:
+ result[-i - 1] = CNU(result[-i - 1].power +
+ current_unit.power, None, None, None, None)
+ return result
+
+ def compute_value(integer_symbols):
+ """
+ Compute the value.
+ When current unit is larger than previous unit, current unit * all previous units will be used as all previous units.
+ e.g. '涓ゅ崈涓�' = 2000 * 10000 not 2000 + 10000
+ """
+ value = [0]
+ last_power = 0
+ for s in integer_symbols:
+ if isinstance(s, CND):
+ value[-1] = s.value
+ elif isinstance(s, CNU):
+ value[-1] *= pow(10, s.power)
+ if s.power > last_power:
+ value[:-1] = list(map(lambda v: v *
+ pow(10, s.power), value[:-1]))
+ last_power = s.power
+ value.append(0)
+ return sum(value)
+
+ system = create_system(numbering_type)
+ int_part, dec_part = string2symbols(chinese_string, system)
+ int_part = correct_symbols(int_part, system)
+ int_str = str(compute_value(int_part))
+ dec_str = ''.join([str(d.value) for d in dec_part])
+ if dec_part:
+ return '{0}.{1}'.format(int_str, dec_str)
+ else:
+ return int_str
+
+
+def num2chn(number_string, numbering_type=NUMBERING_TYPES[1], big=False,
+ traditional=False, alt_zero=False, alt_one=False, alt_two=True,
+ use_zeros=True, use_units=True):
+
+ def get_value(value_string, use_zeros=True):
+
+ striped_string = value_string.lstrip('0')
+
+ # record nothing if all zeros
+ if not striped_string:
+ return []
+
+ # record one digits
+ elif len(striped_string) == 1:
+ if use_zeros and len(value_string) != len(striped_string):
+ return [system.digits[0], system.digits[int(striped_string)]]
+ else:
+ return [system.digits[int(striped_string)]]
+
+ # recursively record multiple digits
+ else:
+ result_unit = next(u for u in reversed(
+ system.units) if u.power < len(striped_string))
+ result_string = value_string[:-result_unit.power]
+ return get_value(result_string) + [result_unit] + get_value(striped_string[-result_unit.power:])
+
+ system = create_system(numbering_type)
+
+ int_dec = number_string.split('.')
+ if len(int_dec) == 1:
+ int_string = int_dec[0]
+ dec_string = ""
+ elif len(int_dec) == 2:
+ int_string = int_dec[0]
+ dec_string = int_dec[1]
+ else:
+ raise ValueError(
+ "invalid input num string with more than one dot: {}".format(number_string))
+
+ if use_units and len(int_string) > 1:
+ result_symbols = get_value(int_string)
+ else:
+ result_symbols = [system.digits[int(c)] for c in int_string]
+ dec_symbols = [system.digits[int(c)] for c in dec_string]
+ if dec_string:
+ result_symbols += [system.math.point] + dec_symbols
+
+ if alt_two:
+ liang = CND(2, system.digits[2].alt_s, system.digits[2].alt_t,
+ system.digits[2].big_s, system.digits[2].big_t)
+ for i, v in enumerate(result_symbols):
+ if isinstance(v, CND) and v.value == 2:
+ next_symbol = result_symbols[i +
+ 1] if i < len(result_symbols) - 1 else None
+ previous_symbol = result_symbols[i - 1] if i > 0 else None
+ if isinstance(next_symbol, CNU) and isinstance(previous_symbol, (CNU, type(None))):
+ if next_symbol.power != 1 and ((previous_symbol is None) or (previous_symbol.power != 1)):
+ result_symbols[i] = liang
+
+ # if big is True, '涓�' will not be used and `alt_two` has no impact on output
+ if big:
+ attr_name = 'big_'
+ if traditional:
+ attr_name += 't'
+ else:
+ attr_name += 's'
+ else:
+ if traditional:
+ attr_name = 'traditional'
+ else:
+ attr_name = 'simplified'
+
+ result = ''.join([getattr(s, attr_name) for s in result_symbols])
+
+ # if not use_zeros:
+ # result = result.strip(getattr(system.digits[0], attr_name))
+
+ if alt_zero:
+ result = result.replace(
+ getattr(system.digits[0], attr_name), system.digits[0].alt_s)
+
+ if alt_one:
+ result = result.replace(
+ getattr(system.digits[1], attr_name), system.digits[1].alt_s)
+
+ for i, p in enumerate(POINT):
+ if result.startswith(p):
+ return CHINESE_DIGIS[0] + result
+
+ # ^10, 11, .., 19
+ if len(result) >= 2 and result[1] in [SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED[0],
+ SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL[0]] and \
+ result[0] in [CHINESE_DIGIS[1], BIG_CHINESE_DIGIS_SIMPLIFIED[1], BIG_CHINESE_DIGIS_TRADITIONAL[1]]:
+ result = result[1:]
+
+ return result
+
+
+# ================================================================================ #
+# different types of rewriters
+# ================================================================================ #
+class Cardinal:
+ """
+ CARDINAL绫�
+ """
+
+ def __init__(self, cardinal=None, chntext=None):
+ self.cardinal = cardinal
+ self.chntext = chntext
+
+ def chntext2cardinal(self):
+ return chn2num(self.chntext)
+
+ def cardinal2chntext(self):
+ return num2chn(self.cardinal)
+
+class Digit:
+ """
+ DIGIT绫�
+ """
+
+ def __init__(self, digit=None, chntext=None):
+ self.digit = digit
+ self.chntext = chntext
+
+ # def chntext2digit(self):
+ # return chn2num(self.chntext)
+
+ def digit2chntext(self):
+ return num2chn(self.digit, alt_two=False, use_units=False)
+
+
+class TelePhone:
+ """
+ TELEPHONE绫�
+ """
+
+ def __init__(self, telephone=None, raw_chntext=None, chntext=None):
+ self.telephone = telephone
+ self.raw_chntext = raw_chntext
+ self.chntext = chntext
+
+ # def chntext2telephone(self):
+ # sil_parts = self.raw_chntext.split('<SIL>')
+ # self.telephone = '-'.join([
+ # str(chn2num(p)) for p in sil_parts
+ # ])
+ # return self.telephone
+
+ def telephone2chntext(self, fixed=False):
+
+ if fixed:
+ sil_parts = self.telephone.split('-')
+ self.raw_chntext = '<SIL>'.join([
+ num2chn(part, alt_two=False, use_units=False) for part in sil_parts
+ ])
+ self.chntext = self.raw_chntext.replace('<SIL>', '')
+ else:
+ sp_parts = self.telephone.strip('+').split()
+ self.raw_chntext = '<SP>'.join([
+ num2chn(part, alt_two=False, use_units=False) for part in sp_parts
+ ])
+ self.chntext = self.raw_chntext.replace('<SP>', '')
+ return self.chntext
+
+
+class Fraction:
+ """
+ FRACTION绫�
+ """
+
+ def __init__(self, fraction=None, chntext=None):
+ self.fraction = fraction
+ self.chntext = chntext
+
+ def chntext2fraction(self):
+ denominator, numerator = self.chntext.split('鍒嗕箣')
+ return chn2num(numerator) + '/' + chn2num(denominator)
+
+ def fraction2chntext(self):
+ numerator, denominator = self.fraction.split('/')
+ return num2chn(denominator) + '鍒嗕箣' + num2chn(numerator)
+
+
+class Date:
+ """
+ DATE绫�
+ """
+
+ def __init__(self, date=None, chntext=None):
+ self.date = date
+ self.chntext = chntext
+
+ # def chntext2date(self):
+ # chntext = self.chntext
+ # try:
+ # year, other = chntext.strip().split('骞�', maxsplit=1)
+ # year = Digit(chntext=year).digit2chntext() + '骞�'
+ # except ValueError:
+ # other = chntext
+ # year = ''
+ # if other:
+ # try:
+ # month, day = other.strip().split('鏈�', maxsplit=1)
+ # month = Cardinal(chntext=month).chntext2cardinal() + '鏈�'
+ # except ValueError:
+ # day = chntext
+ # month = ''
+ # if day:
+ # day = Cardinal(chntext=day[:-1]).chntext2cardinal() + day[-1]
+ # else:
+ # month = ''
+ # day = ''
+ # date = year + month + day
+ # self.date = date
+ # return self.date
+
+ def date2chntext(self):
+ date = self.date
+ try:
+ year, other = date.strip().split('骞�', 1)
+ year = Digit(digit=year).digit2chntext() + '骞�'
+ except ValueError:
+ other = date
+ year = ''
+ if other:
+ try:
+ month, day = other.strip().split('鏈�', 1)
+ month = Cardinal(cardinal=month).cardinal2chntext() + '鏈�'
+ except ValueError:
+ day = date
+ month = ''
+ if day:
+ day = Cardinal(cardinal=day[:-1]).cardinal2chntext() + day[-1]
+ else:
+ month = ''
+ day = ''
+ chntext = year + month + day
+ self.chntext = chntext
+ return self.chntext
+
+
+class Money:
+ """
+ MONEY绫�
+ """
+
+ def __init__(self, money=None, chntext=None):
+ self.money = money
+ self.chntext = chntext
+
+ # def chntext2money(self):
+ # return self.money
+
+ def money2chntext(self):
+ money = self.money
+ pattern = re.compile(r'(\d+(\.\d+)?)')
+ matchers = pattern.findall(money)
+ if matchers:
+ for matcher in matchers:
+ money = money.replace(matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext())
+ self.chntext = money
+ return self.chntext
+
+
+class Percentage:
+ """
+ PERCENTAGE绫�
+ """
+
+ def __init__(self, percentage=None, chntext=None):
+ self.percentage = percentage
+ self.chntext = chntext
+
+ def chntext2percentage(self):
+ return chn2num(self.chntext.strip().strip('鐧惧垎涔�')) + '%'
+
+ def percentage2chntext(self):
+ return '鐧惧垎涔�' + num2chn(self.percentage.strip().strip('%'))
+
+
+def remove_erhua(text, er_whitelist):
+ """
+ 鍘婚櫎鍎垮寲闊宠瘝涓殑鍎�:
+ 浠栧コ鍎垮湪閭h竟鍎� -> 浠栧コ鍎垮湪閭h竟
+ """
+
+ er_pattern = re.compile(er_whitelist)
+ new_str=''
+ while re.search('鍎�',text):
+ a = re.search('鍎�',text).span()
+ remove_er_flag = 0
+
+ if er_pattern.search(text):
+ b = er_pattern.search(text).span()
+ if b[0] <= a[0]:
+ remove_er_flag = 1
+
+ if remove_er_flag == 0 :
+ new_str = new_str + text[0:a[0]]
+ text = text[a[1]:]
+ else:
+ new_str = new_str + text[0:b[1]]
+ text = text[b[1]:]
+
+ text = new_str + text
+ return text
+
+# ================================================================================ #
+# NSW Normalizer
+# ================================================================================ #
+class NSWNormalizer:
+ def __init__(self, raw_text):
+ self.raw_text = '^' + raw_text + '$'
+ self.norm_text = ''
+
+ def _particular(self):
+ text = self.norm_text
+ pattern = re.compile(r"(([a-zA-Z]+)浜�([a-zA-Z]+))")
+ matchers = pattern.findall(text)
+ if matchers:
+ # print('particular')
+ for matcher in matchers:
+ text = text.replace(matcher[0], matcher[1]+'2'+matcher[2], 1)
+ self.norm_text = text
+ return self.norm_text
+
+ def normalize(self):
+ text = self.raw_text
+
+ # 瑙勮寖鍖栨棩鏈�
+ pattern = re.compile(r"\D+((([089]\d|(19|20)\d{2})骞�)?(\d{1,2}鏈�(\d{1,2}[鏃ュ彿])?)?)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('date')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Date(date=matcher[0]).date2chntext(), 1)
+
+ # 瑙勮寖鍖栭噾閽�
+ pattern = re.compile(r"\D+((\d+(\.\d+)?)[澶氫綑鍑燷?" + CURRENCY_UNITS + r"(\d" + CURRENCY_UNITS + r"?)?)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('money')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Money(money=matcher[0]).money2chntext(), 1)
+
+ # 瑙勮寖鍖栧浐璇�/鎵嬫満鍙风爜
+ # 鎵嬫満
+ # http://www.jihaoba.com/news/show/13680
+ # 绉诲姩锛�139銆�138銆�137銆�136銆�135銆�134銆�159銆�158銆�157銆�150銆�151銆�152銆�188銆�187銆�182銆�183銆�184銆�178銆�198
+ # 鑱旈�氾細130銆�131銆�132銆�156銆�155銆�186銆�185銆�176
+ # 鐢典俊锛�133銆�153銆�189銆�180銆�181銆�177
+ pattern = re.compile(r"\D((\+?86 ?)?1([38]\d|5[0-35-9]|7[678]|9[89])\d{8})\D")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('telephone')
+ for matcher in matchers:
+ text = text.replace(matcher[0], TelePhone(telephone=matcher[0]).telephone2chntext(), 1)
+ # 鍥鸿瘽
+ pattern = re.compile(r"\D((0(10|2[1-3]|[3-9]\d{2})-?)?[1-9]\d{6,7})\D")
+ matchers = pattern.findall(text)
+ if matchers:
+ # print('fixed telephone')
+ for matcher in matchers:
+ text = text.replace(matcher[0], TelePhone(telephone=matcher[0]).telephone2chntext(fixed=True), 1)
+
+ # 瑙勮寖鍖栧垎鏁�
+ pattern = re.compile(r"(\d+/\d+)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('fraction')
+ for matcher in matchers:
+ text = text.replace(matcher, Fraction(fraction=matcher).fraction2chntext(), 1)
+
+ # 瑙勮寖鍖栫櫨鍒嗘暟
+ text = text.replace('锛�', '%')
+ pattern = re.compile(r"(\d+(\.\d+)?%)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('percentage')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Percentage(percentage=matcher[0]).percentage2chntext(), 1)
+
+ # 瑙勮寖鍖栫函鏁�+閲忚瘝
+ pattern = re.compile(r"(\d+(\.\d+)?)[澶氫綑鍑燷?" + COM_QUANTIFIERS)
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('cardinal+quantifier')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext(), 1)
+
+ # 瑙勮寖鍖栨暟瀛楃紪鍙�
+ pattern = re.compile(r"(\d{4,32})")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('digit')
+ for matcher in matchers:
+ text = text.replace(matcher, Digit(digit=matcher).digit2chntext(), 1)
+
+ # 瑙勮寖鍖栫函鏁�
+ pattern = re.compile(r"(\d+(\.\d+)?)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('cardinal')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext(), 1)
+
+ self.norm_text = text
+ self._particular()
+
+ return self.norm_text.lstrip('^').rstrip('$')
+
+
+def nsw_test_case(raw_text):
+ print('I:' + raw_text)
+ print('O:' + NSWNormalizer(raw_text).normalize())
+ print('')
+
+
+def nsw_test():
+ nsw_test_case('鍥鸿瘽锛�0595-23865596鎴�23880880銆�')
+ nsw_test_case('鍥鸿瘽锛�0595-23865596鎴�23880880銆�')
+ nsw_test_case('鎵嬫満锛�+86 19859213959鎴�15659451527銆�')
+ nsw_test_case('鍒嗘暟锛�32477/76391銆�')
+ nsw_test_case('鐧惧垎鏁帮細80.03%銆�')
+ nsw_test_case('缂栧彿锛�31520181154418銆�')
+ nsw_test_case('绾暟锛�2983.07鍏嬫垨12345.60绫炽��')
+ nsw_test_case('鏃ユ湡锛�1999骞�2鏈�20鏃ユ垨09骞�3鏈�15鍙枫��')
+ nsw_test_case('閲戦挶锛�12鍧�5锛�34.5鍏冿紝20.1涓�')
+ nsw_test_case('鐗规畩锛歄2O鎴朆2C銆�')
+ nsw_test_case('3456涓囧惃')
+ nsw_test_case('2938涓�')
+ nsw_test_case('938')
+ nsw_test_case('浠婂ぉ鍚冧簡115涓皬绗煎寘231涓澶�')
+ nsw_test_case('鏈�62锛呯殑姒傜巼')
+
+
+if __name__ == '__main__':
+ #nsw_test()
+
+ p = argparse.ArgumentParser()
+ p.add_argument('ifile', help='input filename, assume utf-8 encoding')
+ p.add_argument('ofile', help='output filename')
+ p.add_argument('--to_upper', action='store_true', help='convert to upper case')
+ p.add_argument('--to_lower', action='store_true', help='convert to lower case')
+ p.add_argument('--has_key', action='store_true', help="input text has Kaldi's key as first field.")
+ p.add_argument('--remove_fillers', type=bool, default=True, help='remove filler chars such as "鍛�, 鍟�"')
+ p.add_argument('--remove_erhua', type=bool, default=True, help='remove erhua chars such as "杩欏効"')
+ p.add_argument('--log_interval', type=int, default=10000, help='log interval in number of processed lines')
+ args = p.parse_args()
+
+ ifile = codecs.open(args.ifile, 'r', 'utf8')
+ ofile = codecs.open(args.ofile, 'w+', 'utf8')
+
+ n = 0
+ for l in ifile:
+ key = ''
+ text = ''
+ if args.has_key:
+ cols = l.split(maxsplit=1)
+ key = cols[0]
+ if len(cols) == 2:
+ text = cols[1].strip()
+ else:
+ text = ''
+ else:
+ text = l.strip()
+
+ # cases
+ if args.to_upper and args.to_lower:
+ sys.stderr.write('text norm: to_upper OR to_lower?')
+ exit(1)
+ if args.to_upper:
+ text = text.upper()
+ if args.to_lower:
+ text = text.lower()
+
+ # Filler chars removal
+ if args.remove_fillers:
+ for ch in FILLER_CHARS:
+ text = text.replace(ch, '')
+
+ if args.remove_erhua:
+ text = remove_erhua(text, ER_WHITELIST)
+
+ # NSW(Non-Standard-Word) normalization
+ text = NSWNormalizer(text).normalize()
+
+ # Punctuations removal
+ old_chars = CHINESE_PUNC_LIST + string.punctuation # includes all CN and EN punctuations
+ new_chars = ' ' * len(old_chars)
+ del_chars = ''
+ text = text.translate(str.maketrans(old_chars, new_chars, del_chars))
+
+ #
+ if args.has_key:
+ ofile.write(key + '\t' + text + '\n')
+ else:
+ ofile.write(text + '\n')
+
+ n += 1
+ if n % args.log_interval == 0:
+ sys.stderr.write("text norm: {} lines done.\n".format(n))
+
+ sys.stderr.write("text norm: {} lines done in total.\n".format(n))
+
+ ifile.close()
+ ofile.close()
diff --git a/egs_modelscope/common_uniasr/README.md b/egs_modelscope/common_uniasr/README.md
new file mode 100644
index 0000000..bd14334
--- /dev/null
+++ b/egs_modelscope/common_uniasr/README.md
@@ -0,0 +1,27 @@
+# ModelScope Model
+
+## How to finetune and infer using a pretrained ModelScope Model
+
+### Finetune
+- Modify finetune training related parameters in `conf/train_asr_uniasr_40e1_12d1_20e2_12d2_1280_320_lfr6.yaml`
+- Setting parameters in `modelscope_common_finetune.sh`
+ - <strong>dataset:</strong> the dataset dir needs to include files: train/wav.scp, train/text; optional dev/wav.scp, dev/text, test/wav.scp test/text
+ - <strong>tag:</strong> exp tag
+ - <strong>init_model_name:</strong> speech_UniASR_asr_2pass-zh-cn-8k-common-vocab3445-pytorch-online # pre-trained model, download from modelscope during fine-tuning
+- Then you can run the pipeline to finetune with our model download from modelscope:
+```sh
+ sh ./modelscope_common_finetune.sh
+```
+
+### Inference
+
+Or you can use the finetuned model for inference directly.
+
+- Setting parameters in `modelscope_common_infer.sh`
+ - <strong>data_dir:</strong> # wav list, ${data_dir}/wav.scp
+ - <strong>exp_dir:</strong> the result path
+ - <strong>model_name:</strong> speech_UniASR_asr_2pass-zh-cn-8k-common-vocab3445-pytorch-online # pre-trained model, download from modelscope
+- Then you can run the pipeline to infer with:
+```sh
+ sh ./modelscope_common_infer.sh
+```
diff --git a/egs_modelscope/common_uniasr/conf/decode_asr_uniasr.yaml b/egs_modelscope/common_uniasr/conf/decode_asr_uniasr.yaml
new file mode 100644
index 0000000..f723dd6
--- /dev/null
+++ b/egs_modelscope/common_uniasr/conf/decode_asr_uniasr.yaml
@@ -0,0 +1,9 @@
+beam_size: 5
+penalty: 0.0
+maxlenratio: 0.0
+minlenratio: 0.0
+ctc_weight: 0.0
+lm_weight: 0.0
+token_num_relax: 5
+decoding_ind: 0
+decoding_mode: model2
\ No newline at end of file
diff --git a/egs_modelscope/common_uniasr/conf/train_asr_uniasr_40e1_12d1_20e2_12d2_1280_320_lfr6.yaml b/egs_modelscope/common_uniasr/conf/train_asr_uniasr_40e1_12d1_20e2_12d2_1280_320_lfr6.yaml
new file mode 100644
index 0000000..9885f47
--- /dev/null
+++ b/egs_modelscope/common_uniasr/conf/train_asr_uniasr_40e1_12d1_20e2_12d2_1280_320_lfr6.yaml
@@ -0,0 +1,192 @@
+# encoder related
+encoder: sanm_chunk_opt
+encoder_conf:
+ output_size: 320 # dimension of attention
+ attention_heads: 4
+ linear_units: 1280 # the number of units of position-wise feed forward
+ num_blocks: 40 # the number of encoder blocks
+ dropout_rate: 0.1
+ positional_dropout_rate: 0.1
+ attention_dropout_rate: 0.1
+ input_layer: pe # encoder architecture type
+ pos_enc_class: SinusoidalPositionEncoder
+ normalize_before: true
+ kernel_size: 11
+ sanm_shfit: 0
+ selfattention_layer_type: sanm
+ chunk_size:
+ - 20
+ - 60
+ stride:
+ - 10
+ - 40
+ pad_left:
+ - 5
+ - 10
+ encoder_att_look_back_factor:
+ - 0
+ - 0
+ decoder_att_look_back_factor:
+ - 0
+ - 0
+
+# encoder related
+encoder2: sanm_chunk_opt
+encoder2_conf:
+ output_size: 320 # dimension of attention
+ attention_heads: 4
+ linear_units: 1280 # the number of units of position-wise feed forward
+ num_blocks: 20 # the number of encoder blocks
+ dropout_rate: 0.1
+ positional_dropout_rate: 0.1
+ attention_dropout_rate: 0.1
+ input_layer: pe # encoder architecture type
+ pos_enc_class: SinusoidalPositionEncoder
+ normalize_before: true
+ kernel_size: 21
+ sanm_shfit: 0
+ selfattention_layer_type: sanm
+ chunk_size:
+ - 45
+ - 70
+ stride:
+ - 35
+ - 50
+ pad_left:
+ - 5
+ - 10
+ encoder_att_look_back_factor:
+ - 0
+ - 0
+ decoder_att_look_back_factor:
+ - 0
+ - 0
+
+# decoder related
+decoder: fsmn_scama_opt
+decoder_conf:
+ attention_dim: 256
+ attention_heads: 4
+ linear_units: 1024
+ num_blocks: 12
+ dropout_rate: 0.1
+ positional_dropout_rate: 0.1
+ self_attention_dropout_rate: 0.1
+ src_attention_dropout_rate: 0.1
+ att_layer_num: 6
+ kernel_size: 11
+ concat_embeds: true
+
+# decoder related
+decoder2: fsmn_scama_opt
+decoder2_conf:
+ attention_dim: 320
+ attention_heads: 4
+ linear_units: 1280
+ num_blocks: 12
+ dropout_rate: 0.1
+ positional_dropout_rate: 0.1
+ self_attention_dropout_rate: 0.1
+ src_attention_dropout_rate: 0.1
+ att_layer_num: 6
+ kernel_size: 11
+ concat_embeds: true
+
+stride_conv: stride_conv1d
+stride_conv_conf:
+ kernel_size: 2
+ stride: 2
+ pad:
+ - 0
+ - 1
+
+predictor: cif_predictor_v2
+predictor_conf:
+ idim: 320
+ threshold: 1.0
+ l_order: 1
+ r_order: 1
+
+predictor2: cif_predictor_v2
+predictor2_conf:
+ idim: 320
+ threshold: 1.0
+ l_order: 1
+ r_order: 1
+
+# hybrid CTC/attention
+model: uniasr
+model_conf:
+ ctc_weight: 0.0
+ lsm_weight: 0.1 # label smoothing option
+ length_normalized_loss: true
+ predictor_weight: 1.0
+ decoder_attention_chunk_type: chunk
+ ctc_weight2: 0.0
+ predictor_weight2: 1.0
+ decoder_attention_chunk_type2: chunk
+ loss_weight_model1: 0.5
+ enable_maas_finetune: true
+
+# minibatch related
+batch_type: length
+batch_bins: 2000
+num_workers: 16
+
+dataset_conf:
+ filter_conf:
+ min_length: 10
+ max_length: 250
+ min_token_length: 1
+ max_token_length: 200
+ shuffle: True
+ shuffle_conf:
+ shuffle_size: 10240
+ sort_size: 500
+ batch_conf:
+ batch_type: token
+ batch_size: 2000
+ num_workers: 16
+
+# optimization related
+accum_grad: 1
+grad_clip: 5
+max_epoch: 20
+val_scheduler_criterion:
+ - valid
+ - acc
+best_model_criterion:
+- - valid
+ - acc
+ - max
+keep_nbest_models: 20
+
+optim: adam
+optim_conf:
+ lr: 0.0001
+scheduler: warmuplr
+scheduler_conf:
+ warmup_steps: 30000
+
+specaug: specaug_lfr
+specaug_conf:
+ apply_time_warp: false
+ time_warp_window: 5
+ time_warp_mode: bicubic
+ apply_freq_mask: true
+ freq_mask_width_range:
+ - 0
+ - 30
+ lfr_rate: 6
+ num_freq_mask: 1
+ apply_time_mask: true
+ time_mask_width_range:
+ - 0
+ - 12
+ num_time_mask: 1
+
+
+log_interval: 50
+normalize: None
+split_with_space: true
+unused_parameters: true
diff --git a/egs_modelscope/common_uniasr/modelscope_common_finetune.sh b/egs_modelscope/common_uniasr/modelscope_common_finetune.sh
new file mode 100755
index 0000000..dfc1fdb
--- /dev/null
+++ b/egs_modelscope/common_uniasr/modelscope_common_finetune.sh
@@ -0,0 +1,268 @@
+#!/usr/bin/env bash
+
+. ./path.sh || exit 1;
+
+# machines configuration
+CUDA_VISIBLE_DEVICES="0,1" # set gpus, e.g., CUDA_VISIBLE_DEVICES="0,1"
+gpu_num=2
+count=1
+gpu_inference=true # Whether to perform gpu decoding, set false for cpu decoding
+njob=4 # the number of jobs for each gpu
+train_cmd=utils/run.pl
+
+# general configuration
+feats_dir="../DATA" #feature output dictionary, for large data
+exp_dir="."
+lang=zh
+dumpdir=dump/fbank
+feats_type=fbank
+token_type=char
+scp=feats.scp
+type=kaldi_ark
+stage=1
+stop_stage=4
+
+# feature configuration
+feats_dim=560
+sample_frequency=8000
+nj=32
+speed_perturb="1.0"
+lfr=True
+lfr_m=7
+lfr_n=6
+
+init_model_name=speech_UniASR_asr_2pass-zh-cn-8k-common-vocab3445-pytorch-online # pre-trained model, download from modelscope during fine-tuning
+model_revision="v1.0.0" # please do not modify the model revision
+cmvn_file=init_model/${init_model_name}/am.mvn
+seg_file=init_model/${init_model_name}/seg_dict
+vocab=init_model/${init_model_name}/tokens.txt
+
+# data
+dataset= # dataset (include train/wav.scp, train/text, dev/wav.scp, dev/text, optional test/wav.scp test/text)
+
+# exp tag
+tag=""
+
+# Set bash to 'debug' mode, it will exit on :
+# -e 'error', -u 'undefined variable', -o ... 'error in pipeline', -x 'print commands',
+set -e
+set -u
+set -o pipefail
+
+train_set=train
+valid_set=dev
+test_sets="dev test"
+
+asr_config=conf/train_asr_uniasr_40e1_12d1_20e2_12d2_1280_320_lfr6.yaml
+init_param="init_model/${init_model_name}/model.pb"
+
+inference_config=conf/decode_asr_uniasr.yaml
+inference_asr_model=20epoch.pth
+
+. utils/parse_options.sh || exit 1;
+
+# download model from modelscope
+python modelscope_utils/download_model.py --model_name ${init_model_name} --model_revision ${model_revision}
+
+if [ ! -d ${HOME}/.cache/modelscope/hub/damo/${init_model_name} ]; then
+ echo "${HOME}/.cache/modelscope/hub/damo/${init_model_name} must exist"
+ exit 1
+else
+ if [ -d init_model/${init_model_name} ]; then
+ echo "init_model/${init_model_name} is already exists. if you want to decode again, please delete init_model/${init_model_name} first."
+ else
+ mkdir -p init_model/${init_model_name}
+ cp -r ${HOME}/.cache/modelscope/hub/damo/${init_model_name}/* init_model/${init_model_name}
+ fi
+fi
+
+model_dir="baseline_$(basename "${asr_config}" .yaml)_${feats_type}_${lang}_${token_type}_${tag}"
+
+# you can set gpu num for decoding here
+gpuid_list=$CUDA_VISIBLE_DEVICES # set gpus for decoding, the same as training stage by default
+ngpu=$(echo $gpuid_list | awk -F "," '{print NF}')
+
+if ${gpu_inference}; then
+ inference_nj=$[${ngpu}*${njob}]
+ _ngpu=1
+else
+ inference_nj=$njob
+ _ngpu=0
+fi
+
+[ ! -d ${dataset} ] && echo "$0: Training data is required" && exit 1;
+[ ! -f ${dataset}/train/wav.scp ] && [ ! -f ${dataset}/train/text ] && echo "$0: Training data wav.scp or text is not found" && exit 1;
+
+if [ ! -d "${dataset}/dev" ]; then
+ utils/fix_data.sh ${dataset}/train
+ utils/subset_data_dir_tr_cv.sh --dev-num-utt 1000 ${dataset}/train ${dataset}
+fi
+if [ ! -d "${dataset}/test" ]; then
+ test_sets="dev"
+fi
+
+feat_train_dir=${feats_dir}/${dumpdir}/train; mkdir -p ${feat_train_dir}
+feat_dev_dir=${feats_dir}/${dumpdir}/dev; mkdir -p ${feat_dev_dir}
+feat_test_dir=${feats_dir}/${dumpdir}/test; mkdir -p ${feat_test_dir}
+
+if [ ${stage} -le 1 ] && [ ${stop_stage} -ge 1 ]; then
+ echo "stage 1: Feature Generation"
+ # compute fbank features
+ fbankdir=${feats_dir}/fbank
+ utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj --sample_frequency ${sample_frequency} --speed_perturb ${speed_perturb} \
+ ${dataset}/train ${exp_dir}/exp/make_fbank/train ${fbankdir}/train
+ utils/fix_data_feat.sh ${fbankdir}/train
+ utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj --sample_frequency ${sample_frequency} \
+ ${dataset}/dev ${exp_dir}/exp/make_fbank/dev ${fbankdir}/dev
+ utils/fix_data_feat.sh ${fbankdir}/dev
+ if [ -d "${dataset}/test" ]; then
+ utils/compute_fbank.sh --cmd "$train_cmd" --nj $nj --sample_frequency ${sample_frequency} \
+ ${dataset}/test ${exp_dir}/exp/make_fbank/test ${fbankdir}/test
+ utils/fix_data_feat.sh ${fbankdir}/test
+ fi
+
+ echo "apply low_frame_rate and cmvn"
+ [ ! -f ${cmvn_file} ] && echo "$0: cmvn file is required" && exit 1;
+ utils/apply_lfr_and_cmvn.sh --cmd "$train_cmd" --nj $nj \
+ --lfr $lfr --lfr-m $lfr_m --lfr-n $lfr_n \
+ ${fbankdir}/train ${cmvn_file} ${exp_dir}/exp/make_fbank/train ${feat_train_dir}
+ utils/apply_lfr_and_cmvn.sh --cmd "$train_cmd" --nj $nj \
+ --lfr $lfr --lfr-m $lfr_m --lfr-n $lfr_n \
+ ${fbankdir}/dev ${cmvn_file} ${exp_dir}/exp/make_fbank/dev ${feat_dev_dir}
+ if [ -d "${dataset}/test" ]; then
+ utils/apply_lfr_and_cmvn.sh --cmd "$train_cmd" --nj $nj \
+ --lfr $lfr --lfr-m $lfr_m --lfr-n $lfr_n \
+ ${fbankdir}/test ${cmvn_file} ${exp_dir}/exp/make_fbank/test ${feat_test_dir}
+ fi
+
+ echo "Text Tokenize"
+ # 鎴戠埍reading->鎴� 鐖� read@@ ing
+ utils/text_tokenize.sh --cmd "$train_cmd" --nj $nj ${fbankdir}/train ${seg_file} ${feat_train_dir}/log ${feat_train_dir}
+ utils/fix_data_feat.sh ${feat_train_dir}
+ utils/text_tokenize.sh --cmd "$train_cmd" --nj $nj ${fbankdir}/dev ${seg_file} ${feat_dev_dir}/log ${feat_dev_dir}
+ utils/fix_data_feat.sh ${feat_dev_dir}
+ if [ -d "${dataset}/test" ]; then
+ cp ${fbankdir}/test/text ${feat_test_dir}
+ fi
+fi
+
+token_list=${feats_dir}/data/${lang}_token_list/char/tokens.txt
+echo "dictionary: ${token_list}"
+if [ ${stage} -le 2 ] && [ ${stop_stage} -ge 2 ]; then
+ echo "stage 2: Dictionary Preparation"
+ mkdir -p ${feats_dir}/data/${lang}_token_list/char/
+ cp $vocab ${token_list}
+
+ vocab_size=$(wc -l <${token_list})
+ awk -v v=,${vocab_size} '{print $0v}' ${feat_train_dir}/text_shape > ${feat_train_dir}/text_shape.char
+ awk -v v=,${vocab_size} '{print $0v}' ${feat_dev_dir}/text_shape > ${feat_dev_dir}/text_shape.char
+ mkdir -p ${feats_dir}/asr_stats_fbank_zh_char/train
+ mkdir -p ${feats_dir}/asr_stats_fbank_zh_char/dev
+ cp ${feat_train_dir}/speech_shape ${feat_train_dir}/text_shape ${feat_train_dir}/text_shape.char ${feats_dir}/asr_stats_fbank_zh_char/train
+ cp ${feat_dev_dir}/speech_shape ${feat_dev_dir}/text_shape ${feat_dev_dir}/text_shape.char ${feats_dir}/asr_stats_fbank_zh_char/dev
+fi
+
+# Training Stage
+world_size=$gpu_num # run on one machine
+if [ ${stage} -le 3 ] && [ ${stop_stage} -ge 3 ]; then
+ echo "stage 3: Training"
+ # update asr train config.yaml
+ python modelscope_utils/update_config.py --modelscope_config init_model/${init_model_name}/finetune.yaml --finetune_config ${asr_config} --output_config init_model/${init_model_name}/asr_finetune_config.yaml
+ finetune_config=init_model/${init_model_name}/asr_finetune_config.yaml
+
+ mkdir -p ${exp_dir}/exp/${model_dir}
+ mkdir -p ${exp_dir}/exp/${model_dir}/log
+ INIT_FILE=${exp_dir}/exp/${model_dir}/ddp_init
+ if [ -f $INIT_FILE ];then
+ rm -f $INIT_FILE
+ fi
+ init_method=file://$(readlink -f $INIT_FILE)
+ echo "$0: init method is $init_method"
+ for ((i = 0; i < $gpu_num; ++i)); do
+ {
+ rank=$i
+ local_rank=$i
+ gpu_id=$(echo $CUDA_VISIBLE_DEVICES | cut -d',' -f$[$i+1])
+ asr_train_uniasr.py \
+ --gpu_id $gpu_id \
+ --use_preprocessor true \
+ --token_type $token_type \
+ --token_list $token_list \
+ --train_data_path_and_name_and_type ${feats_dir}/${dumpdir}/${train_set}/${scp},speech,${type} \
+ --train_data_path_and_name_and_type ${feats_dir}/${dumpdir}/${train_set}/text,text,text \
+ --train_shape_file ${feats_dir}/asr_stats_fbank_zh_char/${train_set}/speech_shape \
+ --train_shape_file ${feats_dir}/asr_stats_fbank_zh_char/${train_set}/text_shape.char \
+ --valid_data_path_and_name_and_type ${feats_dir}/${dumpdir}/${valid_set}/${scp},speech,${type} \
+ --valid_data_path_and_name_and_type ${feats_dir}/${dumpdir}/${valid_set}/text,text,text \
+ --valid_shape_file ${feats_dir}/asr_stats_fbank_zh_char/${valid_set}/speech_shape \
+ --valid_shape_file ${feats_dir}/asr_stats_fbank_zh_char/${valid_set}/text_shape.char \
+ --resume true \
+ --output_dir ${exp_dir}/exp/${model_dir} \
+ --init_param $init_param \
+ --config $finetune_config \
+ --input_size $feats_dim \
+ --ngpu $gpu_num \
+ --num_worker_count $count \
+ --multiprocessing_distributed true \
+ --dist_init_method $init_method \
+ --dist_world_size $world_size \
+ --dist_rank $rank \
+ --local_rank $local_rank 1> ${exp_dir}/exp/${model_dir}/log/train.log.$i 2>&1
+ } &
+ done
+ wait
+fi
+
+# Testing Stage
+if [ ${stage} -le 4 ] && [ ${stop_stage} -ge 4 ]; then
+ echo "stage 4: Inference"
+ for dset in ${test_sets}; do
+ asr_exp=${exp_dir}/exp/${model_dir}
+ inference_tag="$(basename "${inference_config}" .yaml)"
+ _dir="${asr_exp}/${inference_tag}/${inference_asr_model}/${dset}"
+ _logdir="${_dir}/logdir"
+
+ mkdir -p "${_logdir}"
+ _data="${feats_dir}/${dumpdir}/${dset}"
+ key_file=${_data}/${scp}
+ num_scp_file="$(<${key_file} wc -l)"
+ _nj=$([ $inference_nj -le $num_scp_file ] && echo "$inference_nj" || echo "$num_scp_file")
+ split_scps=
+ for n in $(seq "${_nj}"); do
+ split_scps+=" ${_logdir}/keys.${n}.scp"
+ done
+ # shellcheck disable=SC2086
+ utils/split_scp.pl "${key_file}" ${split_scps}
+ _opts=
+ if [ -n "${inference_config}" ]; then
+ _opts+="--config ${inference_config} "
+ fi
+ ${train_cmd} --gpu "${_ngpu}" --max-jobs-run "${_nj}" JOB=1:"${_nj}" "${_logdir}"/asr_inference.JOB.log \
+ python -m funasr.bin.asr_inference_launch \
+ --batch_size 1 \
+ --ngpu "${_ngpu}" \
+ --njob ${njob} \
+ --gpuid_list ${gpuid_list} \
+ --data_path_and_name_and_type "${_data}/${scp},speech,${type}" \
+ --key_file "${_logdir}"/keys.JOB.scp \
+ --asr_train_config "${asr_exp}"/config.yaml \
+ --asr_model_file "${asr_exp}"/"${inference_asr_model}" \
+ --output_dir "${_logdir}"/output.JOB \
+ --mode uniasr \
+ ${_opts}
+
+ for f in token token_int score text; do
+ if [ -f "${_logdir}/output.1/1best_recog/${f}" ]; then
+ for i in $(seq "${_nj}"); do
+ cat "${_logdir}/output.${i}/1best_recog/${f}"
+ done | sort -k1 >"${_dir}/${f}"
+ fi
+ done
+ python utils/proce_text.py ${_dir}/text ${_dir}/text.proc
+ python utils/proce_text.py ${_data}/text ${_data}/text.proc
+ python utils/compute_wer.py ${_data}/text.proc ${_dir}/text.proc ${_dir}/text.cer
+ tail -n 3 ${_dir}/text.cer > ${_dir}/text.cer.txt
+ cat ${_dir}/text.cer.txt
+ done
+fi
+
diff --git a/egs_modelscope/common_uniasr/modelscope_common_infer.sh b/egs_modelscope/common_uniasr/modelscope_common_infer.sh
new file mode 100755
index 0000000..4e6e124
--- /dev/null
+++ b/egs_modelscope/common_uniasr/modelscope_common_infer.sh
@@ -0,0 +1,76 @@
+#!/usr/bin/env bash
+
+set -e
+set -u
+set -o pipefail
+
+model_name=speech_UniASR_asr_2pass-zh-cn-8k-common-vocab3445-pytorch-online # pre-trained model, download from modelscope
+model_revision="v1.0.0" # please do not modify the model revision
+data_dir= # wav list, ${data_dir}/wav.scp
+exp_dir="exp"
+gpuid_list="0,1"
+ngpu=$(echo $gpuid_list | awk -F "," '{print NF}')
+njob=4
+gpu_inference=true
+decode_cmd=utils/run.pl
+
+. utils/parse_options.sh
+
+if ${gpu_inference}; then
+ inference_nj=$[${ngpu}*${njob}]
+ _ngpu=1
+else
+ inference_nj=${njob}
+ _ngpu=0
+fi
+
+# LM configs
+use_lm=false
+beam_size=1
+lm_weight=0.0
+
+python modelscope_utils/download_model.py \
+ --model_name ${model_name} --model_revision ${model_revision}
+
+if [ -d ${exp_dir} ]; then
+ echo "${exp_dir} is already exists. if you want to decode again, please delete ${exp_dir} first."
+ exit 1
+else
+ mkdir -p ${exp_dir}/${model_name}
+ cp ${HOME}/.cache/modelscope/hub/damo/${model_name}/* ${exp_dir}/${model_name}/. -r
+ _dir=${exp_dir}/decode_asr
+ _logdir=${_dir}/logdir
+ mkdir -p "${_dir}"
+ mkdir -p "${_logdir}"
+fi
+
+for n in $(seq "${inference_nj}"); do
+ split_scps+=" ${_logdir}/keys.${n}.scp"
+done
+# shellcheck disable=SC2086
+utils/split_scp.pl "${data_dir}/wav.scp" ${split_scps}
+
+if "${use_lm}"; then
+ cp ${exp_dir}/${model_name}/decoding.yaml ${exp_dir}/${model_name}/decoding.yaml.back
+ sed -i "s#beam_size: [0-9]*#beam_size: `echo $beam_size`#g" ${exp_dir}/${model_name}/decoding.yaml
+ sed -i "s#lm_weight: 0.[0-9]*#lm_weight: `echo $lm_weight`#g" ${exp_dir}/${model_name}/decoding.yaml
+fi
+
+echo "Decoding started... log: '${_logdir}/asr_inference.*.log'"
+# shellcheck disable=SC2086
+${decode_cmd} --max-jobs-run "${inference_nj}" JOB=1:"${inference_nj}" "${_logdir}"/asr_inference.JOB.log \
+ python -m funasr.bin.modelscope_infer \
+ --local_model_path ${exp_dir}/${model_name} \
+ --wav_list ${_logdir}/keys.JOB.scp \
+ --output_file ${_logdir}/text.JOB \
+ --gpuid_list ${gpuid_list} \
+ --njob ${njob} \
+ --ngpu ${_ngpu} \
+
+ for i in $(seq ${inference_nj}); do
+ cat ${_logdir}/text.${i}
+ done | sort -k1 >${_dir}/text
+
+if "${use_lm}"; then
+ mv ${exp_dir}/${model_name}/decoding.yaml.back ${exp_dir}/${model_name}/decoding.yaml
+fi
diff --git a/egs_modelscope/common_uniasr/modelscope_common_infer_after_finetune.sh b/egs_modelscope/common_uniasr/modelscope_common_infer_after_finetune.sh
new file mode 100755
index 0000000..e92e0ed
--- /dev/null
+++ b/egs_modelscope/common_uniasr/modelscope_common_infer_after_finetune.sh
@@ -0,0 +1,66 @@
+#!/usr/bin/env bash
+
+set -e
+set -u
+set -o pipefail
+
+pretrained_model_name=speech_UniASR_asr_2pass-zh-cn-8k-common-vocab3445-pytorch-online # pre-trained model, download from modelscope
+data_dir= # wav list, ${data_dir}/wav.scp
+finetune_model_name= # fine-tuning model name
+finetune_exp_dir= # fine-tuning model experiment result path
+gpuid_list="0,1"
+ngpu=$(echo $gpuid_list | awk -F "," '{print NF}')
+njob=4
+gpu_inference=true
+decode_cmd=utils/run.pl
+
+. utils/parse_options.sh
+
+if ${gpu_inference}; then
+ inference_nj=$[${ngpu}*${njob}]
+ _ngpu=1
+else
+ inference_nj=${njob}
+ inference_nj=${njob}
+ _ngpu=0
+fi
+
+if [ ! -d ${HOME}/.cache/modelscope/hub/damo/${pretrained_model_name} ]; then
+ echo "${HOME}/.cache/modelscope/hub/damo/${pretrained_model_name} must exist."
+ exit 1
+else
+ exp_dir=${finetune_exp_dir}/${finetune_model_name}.modelscope
+ mkdir -p $exp_dir
+ cp ${finetune_exp_dir}/${finetune_model_name} ${exp_dir}/${finetune_model_name}.modelscope
+ cp ${HOME}/.cache/modelscope/hub/damo/${pretrained_model_name}/* ${exp_dir}/. -r
+fi
+
+_dir=${exp_dir}/decode_asr
+_logdir=${_dir}/logdir
+if [ -d ${_dir} ]; then
+ echo "${_dir} is already exists. if you want to decode again, please delete ${_dir} first."
+else
+ mkdir -p "${_dir}"
+ mkdir -p "${_logdir}"
+fi
+
+for n in $(seq "${inference_nj}"); do
+ split_scps+=" ${_logdir}/keys.${n}.scp"
+done
+# shellcheck disable=SC2086
+utils/split_scp.pl "${data_dir}/wav.scp" ${split_scps}
+
+echo "Decoding started... log: '${_logdir}/asr_inference.*.log'"
+# shellcheck disable=SC2086
+${decode_cmd} --max-jobs-run "${inference_nj}" JOB=1:"${inference_nj}" "${_logdir}"/asr_inference.JOB.log \
+ python -m funasr.bin.modelscope_infer \
+ --local_model_path ${exp_dir} \
+ --wav_list ${_logdir}/keys.JOB.scp \
+ --output_file ${_logdir}/text.JOB \
+ --gpuid_list ${gpuid_list} \
+ --njob ${njob} \
+ --ngpu ${_ngpu} \
+
+ for i in $(seq ${inference_nj}); do
+ cat ${_logdir}/text.${i}
+ done | sort -k1 >${_dir}/text
\ No newline at end of file
diff --git a/egs_modelscope/common_uniasr/modelscope_utils/download_model.py b/egs_modelscope/common_uniasr/modelscope_utils/download_model.py
new file mode 100755
index 0000000..51ba6b8
--- /dev/null
+++ b/egs_modelscope/common_uniasr/modelscope_utils/download_model.py
@@ -0,0 +1,25 @@
+#!/usr/bin/env python3
+import argparse
+
+from modelscope.pipelines import pipeline
+from modelscope.utils.constant import Tasks
+
+if __name__ == '__main__':
+ parser = argparse.ArgumentParser(
+ description="download model configs",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument("--model_name",
+ type=str,
+ default="speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-pytorch",
+ help="model name in modelscope")
+ parser.add_argument("--model_revision",
+ type=str,
+ default="v1.0.3",
+ help="model revision in modelscope")
+ args = parser.parse_args()
+
+ inference_pipeline = pipeline(
+ task=Tasks.auto_speech_recognition,
+ model='damo/{}'.format(args.model_name),
+ model_revision=args.model_revision)
diff --git a/egs_modelscope/common_uniasr/modelscope_utils/modelscope_infer.sh b/egs_modelscope/common_uniasr/modelscope_utils/modelscope_infer.sh
new file mode 100755
index 0000000..a0c606f
--- /dev/null
+++ b/egs_modelscope/common_uniasr/modelscope_utils/modelscope_infer.sh
@@ -0,0 +1,88 @@
+#!/usr/bin/env bash
+
+set -e
+set -u
+set -o pipefail
+
+data_dir=
+exp_dir=
+model_name=
+model_revision=
+inference_nj=32
+gpuid_list="0,1,2,3"
+njob=32
+gpu_inference=true
+
+test_sets="dev test"
+decode_cmd=utils/run.pl
+
+# LM configs
+use_lm=false
+beam_size=1
+lm_weight=0.0
+
+. utils/parse_options.sh
+
+if ${gpu_inference}; then
+ _ngpu=1
+else
+ _ngpu=0
+fi
+
+# download model from modelscope
+python modelscope_utils/download_model.py \
+ --model_name ${model_name} --model_revision ${model_revision}
+
+modelscope_dir=${HOME}/.cache/modelscope/hub/damo/${model_name}
+
+
+for dset in ${test_sets}; do
+ _dir=${exp_dir}/${model_name}/decode_asr/${dset}
+ _logdir=${_dir}/logdir
+ _data=${data_dir}/${dset}
+ if [ -d ${_dir} ]; then
+ echo "${_dir} is already exists. if you want to decode again, please delete ${_dir} first."
+ exit 1
+ else
+ mkdir -p "${_dir}"
+ mkdir -p "${_logdir}"
+ fi
+
+ if "${use_lm}"; then
+ cp ${modelscope_dir}/decoding.yaml ${modelscope_dir}/decoding.yaml.back
+ sed -i "s#beam_size: [0-9]*#beam_size: `echo $beam_size`#g" ${modelscope_dir}/decoding.yaml
+ sed -i "s#lm_weight: 0.[0-9]*#lm_weight: `echo $lm_weight`#g" ${modelscope_dir}/decoding.yaml
+ fi
+
+ for n in $(seq "${inference_nj}"); do
+ split_scps+=" ${_logdir}/keys.${n}.scp"
+ done
+ # shellcheck disable=SC2086
+ utils/split_scp.pl "${data_dir}/${dset}/wav.scp" ${split_scps}
+
+ echo "Decoding started... log: '${_logdir}/asr_inference.*.log'"
+ # shellcheck disable=SC2086
+ ${decode_cmd} --max-jobs-run "${inference_nj}" JOB=1:"${inference_nj}" "${_logdir}"/asr_inference.JOB.log \
+ python -m funasr.bin.modelscope_infer \
+ --model_name ${model_name} \
+ --model_revision ${model_revision} \
+ --wav_list ${_logdir}/keys.JOB.scp \
+ --output_file ${_logdir}/text.JOB \
+ --gpuid_list ${gpuid_list} \
+ --njob ${njob} \
+ --ngpu ${_ngpu} \
+
+ for i in $(seq ${inference_nj}); do
+ cat ${_logdir}/text.${i}
+ done | sort -k1 >${_dir}/text
+
+ python utils/proce_text.py ${_dir}/text ${_dir}/text.proc
+ python utils/proce_text.py ${_data}/text ${_data}/text.proc
+ python utils/compute_wer.py ${_data}/text.proc ${_dir}/text.proc ${_dir}/text.cer
+ tail -n 3 ${_dir}/text.cer > ${_dir}/text.cer.txt
+ cat ${_dir}/text.cer.txt
+done
+
+if "${use_lm}"; then
+ mv ${modelscope_dir}/decoding.yaml.back ${modelscope_dir}/decoding.yaml
+fi
diff --git a/egs_modelscope/common_uniasr/modelscope_utils/update_config.py b/egs_modelscope/common_uniasr/modelscope_utils/update_config.py
new file mode 100644
index 0000000..88466ed
--- /dev/null
+++ b/egs_modelscope/common_uniasr/modelscope_utils/update_config.py
@@ -0,0 +1,41 @@
+import yaml
+import argparse
+
+def update_dct(fin_configs, root):
+ if root == {}:
+ return {}
+ for root_key, root_value in root.items():
+ if not isinstance(root[root_key],dict):
+ fin_configs[root_key] = root[root_key]
+ else:
+ result = update_dct(fin_configs[root_key], root[root_key])
+ fin_configs[root_key] = result
+ return fin_configs
+
+if __name__ == '__main__':
+ parser = argparse.ArgumentParser(
+ description="update configs",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument("--modelscope_config",
+ type=str,
+ help="modelscope config file")
+ parser.add_argument("--finetune_config",
+ type=str,
+ help="finetune config file")
+ parser.add_argument("--output_config",
+ type=str,
+ help="output config file")
+ args = parser.parse_args()
+
+ with open(args.modelscope_config) as f:
+ modelscope_configs = yaml.safe_load(f)
+
+ with open(args.finetune_config) as f:
+ finetune_configs = yaml.safe_load(f)
+
+ # update configs, e.g., lr, batch_size, ...
+ modelscope_configs = update_dct(modelscope_configs, finetune_configs)
+
+ with open(args.output_config, "w") as f:
+ yaml.dump(modelscope_configs, f, indent=4)
diff --git a/egs_modelscope/common_uniasr/path.sh b/egs_modelscope/common_uniasr/path.sh
new file mode 100755
index 0000000..c340218
--- /dev/null
+++ b/egs_modelscope/common_uniasr/path.sh
@@ -0,0 +1,5 @@
+export FUNASR_DIR=$PWD/../..
+
+# NOTE(kan-bayashi): Use UTF-8 in Python to avoid UnicodeDecodeError when LC_ALL=C
+export PYTHONIOENCODING=UTF-8
+export PATH=$FUNASR_DIR/funasr/bin:$PATH
diff --git a/egs_modelscope/common_uniasr/utils/__init__.py b/egs_modelscope/common_uniasr/utils/__init__.py
new file mode 100644
index 0000000..e69de29
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/__init__.py
diff --git a/egs_modelscope/common_uniasr/utils/apply_cmvn.py b/egs_modelscope/common_uniasr/utils/apply_cmvn.py
new file mode 100755
index 0000000..b5c5086
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/apply_cmvn.py
@@ -0,0 +1,79 @@
+from kaldiio import ReadHelper
+from kaldiio import WriteHelper
+
+import argparse
+import json
+import math
+import numpy as np
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="apply cmvn",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--ark-file",
+ "-a",
+ default=False,
+ required=True,
+ type=str,
+ help="fbank ark file",
+ )
+ parser.add_argument(
+ "--cmvn-file",
+ "-c",
+ default=False,
+ required=True,
+ type=str,
+ help="cmvn file",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ ark_file = args.output_dir + "/feats." + str(args.ark_index) + ".ark"
+ scp_file = args.output_dir + "/feats." + str(args.ark_index) + ".scp"
+ ark_writer = WriteHelper('ark,scp:{},{}'.format(ark_file, scp_file))
+
+ with open(args.cmvn_file) as f:
+ cmvn_stats = json.load(f)
+
+ means = cmvn_stats['mean_stats']
+ vars = cmvn_stats['var_stats']
+ total_frames = cmvn_stats['total_frames']
+
+ for i in range(len(means)):
+ means[i] /= total_frames
+ vars[i] = vars[i] / total_frames - means[i] * means[i]
+ if vars[i] < 1.0e-20:
+ vars[i] = 1.0e-20
+ vars[i] = 1.0 / math.sqrt(vars[i])
+
+ with ReadHelper('ark:{}'.format(args.ark_file)) as ark_reader:
+ for key, mat in ark_reader:
+ mat = (mat - means) * vars
+ ark_writer(key, mat)
+
+
+if __name__ == '__main__':
+ main()
diff --git a/egs_modelscope/common_uniasr/utils/apply_cmvn.sh b/egs_modelscope/common_uniasr/utils/apply_cmvn.sh
new file mode 100755
index 0000000..f8fd1d1
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/apply_cmvn.sh
@@ -0,0 +1,29 @@
+#!/usr/bin/env bash
+
+. ./path.sh || exit 1;
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+fbankdir=$1
+cmvn_file=$2
+logdir=$3
+output_dir=$4
+
+dump_dir=${output_dir}/ark; mkdir -p ${dump_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/apply_cmvn.JOB.log \
+ python utils/apply_cmvn.py -a $fbankdir/ark/feats.JOB.ark \
+ -c $cmvn_file -i JOB -o ${dump_dir} \
+ || exit 1;
+
+for n in $(seq $nj); do
+ cat ${dump_dir}/feats.$n.scp || exit 1
+done > ${output_dir}/feats.scp || exit 1
+
+echo "$0: Succeeded apply cmvn"
diff --git a/egs_modelscope/common_uniasr/utils/apply_lfr_and_cmvn.py b/egs_modelscope/common_uniasr/utils/apply_lfr_and_cmvn.py
new file mode 100755
index 0000000..50d18d1
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/apply_lfr_and_cmvn.py
@@ -0,0 +1,143 @@
+from kaldiio import ReadHelper, WriteHelper
+
+import argparse
+import numpy as np
+
+
+def build_LFR_features(inputs, m=7, n=6):
+ LFR_inputs = []
+ T = inputs.shape[0]
+ T_lfr = int(np.ceil(T / n))
+ left_padding = np.tile(inputs[0], ((m - 1) // 2, 1))
+ inputs = np.vstack((left_padding, inputs))
+ T = T + (m - 1) // 2
+ for i in range(T_lfr):
+ if m <= T - i * n:
+ LFR_inputs.append(np.hstack(inputs[i * n:i * n + m]))
+ else:
+ num_padding = m - (T - i * n)
+ frame = np.hstack(inputs[i * n:])
+ for _ in range(num_padding):
+ frame = np.hstack((frame, inputs[-1]))
+ LFR_inputs.append(frame)
+ return np.vstack(LFR_inputs)
+
+
+def build_CMVN_features(inputs, mvn_file): # noqa
+ with open(mvn_file, 'r', encoding='utf-8') as f:
+ lines = f.readlines()
+
+ add_shift_list = []
+ rescale_list = []
+ for i in range(len(lines)):
+ line_item = lines[i].split()
+ if line_item[0] == '<AddShift>':
+ line_item = lines[i + 1].split()
+ if line_item[0] == '<LearnRateCoef>':
+ add_shift_line = line_item[3:(len(line_item) - 1)]
+ add_shift_list = list(add_shift_line)
+ continue
+ elif line_item[0] == '<Rescale>':
+ line_item = lines[i + 1].split()
+ if line_item[0] == '<LearnRateCoef>':
+ rescale_line = line_item[3:(len(line_item) - 1)]
+ rescale_list = list(rescale_line)
+ continue
+
+ for j in range(inputs.shape[0]):
+ for k in range(inputs.shape[1]):
+ add_shift_value = add_shift_list[k]
+ rescale_value = rescale_list[k]
+ inputs[j, k] = float(inputs[j, k]) + float(add_shift_value)
+ inputs[j, k] = float(inputs[j, k]) * float(rescale_value)
+
+ return inputs
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="apply low_frame_rate and cmvn",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--ark-file",
+ "-a",
+ default=False,
+ required=True,
+ type=str,
+ help="fbank ark file",
+ )
+ parser.add_argument(
+ "--lfr",
+ "-f",
+ default=True,
+ type=str,
+ help="low frame rate",
+ )
+ parser.add_argument(
+ "--lfr-m",
+ "-m",
+ default=7,
+ type=int,
+ help="number of frames to stack",
+ )
+ parser.add_argument(
+ "--lfr-n",
+ "-n",
+ default=6,
+ type=int,
+ help="number of frames to skip",
+ )
+ parser.add_argument(
+ "--cmvn-file",
+ "-c",
+ default=False,
+ required=True,
+ type=str,
+ help="global cmvn file",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ dump_ark_file = args.output_dir + "/feats." + str(args.ark_index) + ".ark"
+ dump_scp_file = args.output_dir + "/feats." + str(args.ark_index) + ".scp"
+ shape_file = args.output_dir + "/len." + str(args.ark_index)
+ ark_writer = WriteHelper('ark,scp:{},{}'.format(dump_ark_file, dump_scp_file))
+
+ shape_writer = open(shape_file, 'w')
+ with ReadHelper('ark:{}'.format(args.ark_file)) as ark_reader:
+ for key, mat in ark_reader:
+ if args.lfr:
+ lfr = build_LFR_features(mat, args.lfr_m, args.lfr_n)
+ else:
+ lfr = mat
+ cmvn = build_CMVN_features(lfr, args.cmvn_file)
+ dims = cmvn.shape[1]
+ lens = cmvn.shape[0]
+ shape_writer.write(key + " " + str(lens) + "," + str(dims) + '\n')
+ ark_writer(key, cmvn)
+
+
+if __name__ == '__main__':
+ main()
+
diff --git a/egs_modelscope/common_uniasr/utils/apply_lfr_and_cmvn.sh b/egs_modelscope/common_uniasr/utils/apply_lfr_and_cmvn.sh
new file mode 100755
index 0000000..3119fdb
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/apply_lfr_and_cmvn.sh
@@ -0,0 +1,38 @@
+#!/usr/bin/env bash
+
+
+# Begin configuration section.
+nj=32
+cmd=utils/run.pl
+
+# feature configuration
+lfr=True
+lfr_m=7
+lfr_n=6
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+fbankdir=$1
+cmvn_file=$2
+logdir=$3
+output_dir=$4
+
+dump_dir=${output_dir}/ark; mkdir -p ${dump_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/apply_lfr_and_cmvn.JOB.log \
+ python utils/apply_lfr_and_cmvn.py -a $fbankdir/ark/feats.JOB.ark \
+ -f $lfr -m $lfr_m -n $lfr_n -c $cmvn_file -i JOB -o ${dump_dir} \
+ || exit 1;
+
+for n in $(seq $nj); do
+ cat ${dump_dir}/feats.$n.scp || exit 1
+done > ${output_dir}/feats.scp || exit 1
+
+for n in $(seq $nj); do
+ cat ${dump_dir}/len.$n || exit 1
+done > ${output_dir}/speech_shape || exit 1
+
+echo "$0: Succeeded apply low frame rate and cmvn"
diff --git a/egs_modelscope/common_uniasr/utils/combine_cmvn_file.py b/egs_modelscope/common_uniasr/utils/combine_cmvn_file.py
new file mode 100755
index 0000000..b2974a4
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/combine_cmvn_file.py
@@ -0,0 +1,73 @@
+import argparse
+import json
+import numpy as np
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="combine cmvn file",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--dims",
+ "-d",
+ default=80,
+ type=int,
+ help="feature dims",
+ )
+ parser.add_argument(
+ "--cmvn-dir",
+ "-c",
+ default=False,
+ required=True,
+ type=str,
+ help="cmvn dir",
+ )
+
+ parser.add_argument(
+ "--nj",
+ "-n",
+ default=1,
+ required=True,
+ type=int,
+ help="num of cmvn file",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ total_means = np.zeros(args.dims)
+ total_vars = np.zeros(args.dims)
+ total_frames = 0
+
+ cmvn_file = args.output_dir + "/cmvn.json"
+
+ for i in range(1, args.nj+1):
+ with open(args.cmvn_dir + "/cmvn." + str(i) + ".json", "r") as fin:
+ cmvn_stats = json.load(fin)
+
+ total_means += np.array(cmvn_stats["mean_stats"])
+ total_vars += np.array(cmvn_stats["var_stats"])
+ total_frames += cmvn_stats["total_frames"]
+
+ cmvn_info = {
+ 'mean_stats': list(total_means.tolist()),
+ 'var_stats': list(total_vars.tolist()),
+ 'total_frames': total_frames
+ }
+ with open(cmvn_file, 'w') as fout:
+ fout.write(json.dumps(cmvn_info))
+
+
+if __name__ == '__main__':
+ main()
diff --git a/egs_modelscope/common_uniasr/utils/compute_cmvn.py b/egs_modelscope/common_uniasr/utils/compute_cmvn.py
new file mode 100755
index 0000000..2b96e26
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/compute_cmvn.py
@@ -0,0 +1,74 @@
+from kaldiio import ReadHelper
+
+import argparse
+import numpy as np
+import json
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="computer global cmvn",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--dims",
+ "-d",
+ default=80,
+ type=int,
+ help="feature dims",
+ )
+ parser.add_argument(
+ "--ark-file",
+ "-a",
+ default=False,
+ required=True,
+ type=str,
+ help="fbank ark file",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ ark_file = args.ark_file + "/feats." + str(args.ark_index) + ".ark"
+ cmvn_file = args.output_dir + "/cmvn." + str(args.ark_index) + ".json"
+
+ mean_stats = np.zeros(args.dims)
+ var_stats = np.zeros(args.dims)
+ total_frames = 0
+
+ with ReadHelper('ark:{}'.format(ark_file)) as ark_reader:
+ for key, mat in ark_reader:
+ mean_stats += np.sum(mat, axis=0)
+ var_stats += np.sum(np.square(mat), axis=0)
+ total_frames += mat.shape[0]
+
+ cmvn_info = {
+ 'mean_stats': list(mean_stats.tolist()),
+ 'var_stats': list(var_stats.tolist()),
+ 'total_frames': total_frames
+ }
+ with open(cmvn_file, 'w') as fout:
+ fout.write(json.dumps(cmvn_info))
+
+
+if __name__ == '__main__':
+ main()
diff --git a/egs_modelscope/common_uniasr/utils/compute_cmvn.sh b/egs_modelscope/common_uniasr/utils/compute_cmvn.sh
new file mode 100755
index 0000000..12173ee
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/compute_cmvn.sh
@@ -0,0 +1,25 @@
+#!/usr/bin/env bash
+
+. ./path.sh || exit 1;
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+feats_dim=80
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+fbankdir=$1
+logdir=$2
+
+output_dir=${fbankdir}/cmvn; mkdir -p ${output_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/cmvn.JOB.log \
+ python utils/compute_cmvn.py -d ${feats_dim} -a $fbankdir/ark -i JOB -o ${output_dir} \
+ || exit 1;
+
+python utils/combine_cmvn_file.py -d ${feats_dim} -c ${output_dir} -n $nj -o $fbankdir
+
+echo "$0: Succeeded compute global cmvn"
diff --git a/egs_modelscope/common_uniasr/utils/compute_fbank.py b/egs_modelscope/common_uniasr/utils/compute_fbank.py
new file mode 100755
index 0000000..d03b5a8
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/compute_fbank.py
@@ -0,0 +1,153 @@
+from kaldiio import WriteHelper
+
+import argparse
+import numpy as np
+import json
+import torch
+import torchaudio
+import torchaudio.compliance.kaldi as kaldi
+
+
+def compute_fbank(wav_file,
+ num_mel_bins=80,
+ frame_length=25,
+ frame_shift=10,
+ dither=0.0,
+ resample_rate=16000,
+ speed=1.0):
+
+ waveform, sample_rate = torchaudio.load(wav_file)
+ if resample_rate != sample_rate:
+ waveform = torchaudio.transforms.Resample(orig_freq=sample_rate,
+ new_freq=resample_rate)(waveform)
+ if speed != 1.0:
+ waveform, _ = torchaudio.sox_effects.apply_effects_tensor(
+ waveform, resample_rate,
+ [['speed', str(speed)], ['rate', str(resample_rate)]]
+ )
+
+ waveform = waveform * (1 << 15)
+ mat = kaldi.fbank(waveform,
+ num_mel_bins=num_mel_bins,
+ frame_length=frame_length,
+ frame_shift=frame_shift,
+ dither=dither,
+ energy_floor=0.0,
+ window_type='hamming',
+ sample_frequency=resample_rate)
+
+ return mat.numpy()
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="computer features",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--wav-lists",
+ "-w",
+ default=False,
+ required=True,
+ type=str,
+ help="input wav lists",
+ )
+ parser.add_argument(
+ "--text-files",
+ "-t",
+ default=False,
+ required=True,
+ type=str,
+ help="input text files",
+ )
+ parser.add_argument(
+ "--dims",
+ "-d",
+ default=80,
+ type=int,
+ help="feature dims",
+ )
+ parser.add_argument(
+ "--sample-frequency",
+ "-s",
+ default=16000,
+ type=int,
+ help="sample frequency",
+ )
+ parser.add_argument(
+ "--speed-perturb",
+ "-p",
+ default="1.0",
+ type=str,
+ help="speed perturb",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-a",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ ark_file = args.output_dir + "/ark/feats." + str(args.ark_index) + ".ark"
+ scp_file = args.output_dir + "/ark/feats." + str(args.ark_index) + ".scp"
+ text_file = args.output_dir + "/txt/text." + str(args.ark_index) + ".txt"
+ feats_shape_file = args.output_dir + "/ark/len." + str(args.ark_index)
+ text_shape_file = args.output_dir + "/txt/len." + str(args.ark_index)
+
+ ark_writer = WriteHelper('ark,scp:{},{}'.format(ark_file, scp_file))
+ text_writer = open(text_file, 'w')
+ feats_shape_writer = open(feats_shape_file, 'w')
+ text_shape_writer = open(text_shape_file, 'w')
+
+ speed_perturb_list = args.speed_perturb.split(',')
+
+ for speed in speed_perturb_list:
+ with open(args.wav_lists, 'r', encoding='utf-8') as wavfile:
+ with open(args.text_files, 'r', encoding='utf-8') as textfile:
+ for wav, text in zip(wavfile, textfile):
+ s_w = wav.strip().split()
+ wav_id = s_w[0]
+ wav_file = s_w[1]
+
+ s_t = text.strip().split()
+ text_id = s_t[0]
+ txt = s_t[1:]
+ fbank = compute_fbank(wav_file,
+ num_mel_bins=args.dims,
+ resample_rate=args.sample_frequency,
+ speed=float(speed)
+ )
+ feats_dims = fbank.shape[1]
+ feats_lens = fbank.shape[0]
+ txt_lens = len(txt)
+ if speed == "1.0":
+ wav_id_sp = wav_id
+ else:
+ wav_id_sp = wav_id + "_sp" + speed
+
+ feats_shape_writer.write(wav_id_sp + " " + str(feats_lens) + "," + str(feats_dims) + '\n')
+ text_shape_writer.write(wav_id_sp + " " + str(txt_lens) + '\n')
+
+ text_writer.write(wav_id_sp + " " + " ".join(txt) + '\n')
+ ark_writer(wav_id_sp, fbank)
+
+
+if __name__ == '__main__':
+ main()
+
diff --git a/egs_modelscope/common_uniasr/utils/compute_fbank.sh b/egs_modelscope/common_uniasr/utils/compute_fbank.sh
new file mode 100755
index 0000000..92a4fe6
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/compute_fbank.sh
@@ -0,0 +1,51 @@
+#!/usr/bin/env bash
+
+. ./path.sh || exit 1;
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+
+# feature configuration
+feats_dim=80
+sample_frequency=16000
+speed_perturb="1.0"
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+data=$1
+logdir=$2
+fbankdir=$3
+
+[ ! -f $data/wav.scp ] && echo "$0: no such file $data/wav.scp" && exit 1;
+[ ! -f $data/text ] && echo "$0: no such file $data/text" && exit 1;
+
+python utils/split_data.py $data $data $nj
+
+ark_dir=${fbankdir}/ark; mkdir -p ${ark_dir}
+text_dir=${fbankdir}/txt; mkdir -p ${text_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/make_fbank.JOB.log \
+ python utils/compute_fbank.py -w $data/split${nj}/JOB/wav.scp -t $data/split${nj}/JOB/text \
+ -d $feats_dim -s $sample_frequency -p ${speed_perturb} -a JOB -o ${fbankdir} \
+ || exit 1;
+
+for n in $(seq $nj); do
+ cat ${ark_dir}/feats.$n.scp || exit 1
+done > $fbankdir/feats.scp || exit 1
+
+for n in $(seq $nj); do
+ cat ${text_dir}/text.$n.txt || exit 1
+done > $fbankdir/text || exit 1
+
+for n in $(seq $nj); do
+ cat ${ark_dir}/len.$n || exit 1
+done > $fbankdir/speech_shape || exit 1
+
+for n in $(seq $nj); do
+ cat ${text_dir}/len.$n || exit 1
+done > $fbankdir/text_shape || exit 1
+
+echo "$0: Succeeded compute FBANK features"
diff --git a/egs_modelscope/common_uniasr/utils/compute_wer.py b/egs_modelscope/common_uniasr/utils/compute_wer.py
new file mode 100755
index 0000000..349a3f6
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/compute_wer.py
@@ -0,0 +1,157 @@
+import os
+import numpy as np
+import sys
+
+def compute_wer(ref_file,
+ hyp_file,
+ cer_detail_file):
+ rst = {
+ 'Wrd': 0,
+ 'Corr': 0,
+ 'Ins': 0,
+ 'Del': 0,
+ 'Sub': 0,
+ 'Snt': 0,
+ 'Err': 0.0,
+ 'S.Err': 0.0,
+ 'wrong_words': 0,
+ 'wrong_sentences': 0
+ }
+
+ hyp_dict = {}
+ ref_dict = {}
+ with open(hyp_file, 'r') as hyp_reader:
+ for line in hyp_reader:
+ key = line.strip().split()[0]
+ value = line.strip().split()[1:]
+ hyp_dict[key] = value
+ with open(ref_file, 'r') as ref_reader:
+ for line in ref_reader:
+ key = line.strip().split()[0]
+ value = line.strip().split()[1:]
+ ref_dict[key] = value
+
+ cer_detail_writer = open(cer_detail_file, 'w')
+ for hyp_key in hyp_dict:
+ if hyp_key in ref_dict:
+ out_item = compute_wer_by_line(hyp_dict[hyp_key], ref_dict[hyp_key])
+ rst['Wrd'] += out_item['nwords']
+ rst['Corr'] += out_item['cor']
+ rst['wrong_words'] += out_item['wrong']
+ rst['Ins'] += out_item['ins']
+ rst['Del'] += out_item['del']
+ rst['Sub'] += out_item['sub']
+ rst['Snt'] += 1
+ if out_item['wrong'] > 0:
+ rst['wrong_sentences'] += 1
+ cer_detail_writer.write(hyp_key + print_cer_detail(out_item) + '\n')
+ cer_detail_writer.write("ref:" + '\t' + "".join(ref_dict[hyp_key]) + '\n')
+ cer_detail_writer.write("hyp:" + '\t' + "".join(hyp_dict[hyp_key]) + '\n')
+
+ if rst['Wrd'] > 0:
+ rst['Err'] = round(rst['wrong_words'] * 100 / rst['Wrd'], 2)
+ if rst['Snt'] > 0:
+ rst['S.Err'] = round(rst['wrong_sentences'] * 100 / rst['Snt'], 2)
+
+ cer_detail_writer.write('\n')
+ cer_detail_writer.write("%WER " + str(rst['Err']) + " [ " + str(rst['wrong_words'])+ " / " + str(rst['Wrd']) +
+ ", " + str(rst['Ins']) + " ins, " + str(rst['Del']) + " del, " + str(rst['Sub']) + " sub ]" + '\n')
+ cer_detail_writer.write("%SER " + str(rst['S.Err']) + " [ " + str(rst['wrong_sentences']) + " / " + str(rst['Snt']) + " ]" + '\n')
+ cer_detail_writer.write("Scored " + str(len(hyp_dict)) + " sentences, " + str(len(hyp_dict) - rst['Snt']) + " not present in hyp." + '\n')
+
+
+def compute_wer_by_line(hyp,
+ ref):
+ hyp = list(map(lambda x: x.lower(), hyp))
+ ref = list(map(lambda x: x.lower(), ref))
+
+ len_hyp = len(hyp)
+ len_ref = len(ref)
+
+ cost_matrix = np.zeros((len_hyp + 1, len_ref + 1), dtype=np.int16)
+
+ ops_matrix = np.zeros((len_hyp + 1, len_ref + 1), dtype=np.int8)
+
+ for i in range(len_hyp + 1):
+ cost_matrix[i][0] = i
+ for j in range(len_ref + 1):
+ cost_matrix[0][j] = j
+
+ for i in range(1, len_hyp + 1):
+ for j in range(1, len_ref + 1):
+ if hyp[i - 1] == ref[j - 1]:
+ cost_matrix[i][j] = cost_matrix[i - 1][j - 1]
+ else:
+ substitution = cost_matrix[i - 1][j - 1] + 1
+ insertion = cost_matrix[i - 1][j] + 1
+ deletion = cost_matrix[i][j - 1] + 1
+
+ compare_val = [substitution, insertion, deletion]
+
+ min_val = min(compare_val)
+ operation_idx = compare_val.index(min_val) + 1
+ cost_matrix[i][j] = min_val
+ ops_matrix[i][j] = operation_idx
+
+ match_idx = []
+ i = len_hyp
+ j = len_ref
+ rst = {
+ 'nwords': len_ref,
+ 'cor': 0,
+ 'wrong': 0,
+ 'ins': 0,
+ 'del': 0,
+ 'sub': 0
+ }
+ while i >= 0 or j >= 0:
+ i_idx = max(0, i)
+ j_idx = max(0, j)
+
+ if ops_matrix[i_idx][j_idx] == 0: # correct
+ if i - 1 >= 0 and j - 1 >= 0:
+ match_idx.append((j - 1, i - 1))
+ rst['cor'] += 1
+
+ i -= 1
+ j -= 1
+
+ elif ops_matrix[i_idx][j_idx] == 2: # insert
+ i -= 1
+ rst['ins'] += 1
+
+ elif ops_matrix[i_idx][j_idx] == 3: # delete
+ j -= 1
+ rst['del'] += 1
+
+ elif ops_matrix[i_idx][j_idx] == 1: # substitute
+ i -= 1
+ j -= 1
+ rst['sub'] += 1
+
+ if i < 0 and j >= 0:
+ rst['del'] += 1
+ elif j < 0 and i >= 0:
+ rst['ins'] += 1
+
+ match_idx.reverse()
+ wrong_cnt = cost_matrix[len_hyp][len_ref]
+ rst['wrong'] = wrong_cnt
+
+ return rst
+
+def print_cer_detail(rst):
+ return ("(" + "nwords=" + str(rst['nwords']) + ",cor=" + str(rst['cor'])
+ + ",ins=" + str(rst['ins']) + ",del=" + str(rst['del']) + ",sub="
+ + str(rst['sub']) + ") corr:" + '{:.2%}'.format(rst['cor']/rst['nwords'])
+ + ",cer:" + '{:.2%}'.format(rst['wrong']/rst['nwords']))
+
+if __name__ == '__main__':
+ if len(sys.argv) != 4:
+ print("usage : python compute-wer.py test.ref test.hyp test.wer")
+ sys.exit(0)
+
+ ref_file = sys.argv[1]
+ hyp_file = sys.argv[2]
+ cer_detail_file = sys.argv[3]
+ compute_wer(ref_file, hyp_file, cer_detail_file)
diff --git a/egs_modelscope/common_uniasr/utils/error_rate_zh b/egs_modelscope/common_uniasr/utils/error_rate_zh
new file mode 100755
index 0000000..6871a07
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/error_rate_zh
@@ -0,0 +1,370 @@
+#!/usr/bin/env python3
+# coding=utf8
+
+# Copyright 2021 Jiayu DU
+
+import sys
+import argparse
+import json
+import logging
+logging.basicConfig(stream=sys.stderr, level=logging.INFO, format='[%(levelname)s] %(message)s')
+
+DEBUG = None
+
+def GetEditType(ref_token, hyp_token):
+ if ref_token == None and hyp_token != None:
+ return 'I'
+ elif ref_token != None and hyp_token == None:
+ return 'D'
+ elif ref_token == hyp_token:
+ return 'C'
+ elif ref_token != hyp_token:
+ return 'S'
+ else:
+ raise RuntimeError
+
+class AlignmentArc:
+ def __init__(self, src, dst, ref, hyp):
+ self.src = src
+ self.dst = dst
+ self.ref = ref
+ self.hyp = hyp
+ self.edit_type = GetEditType(ref, hyp)
+
+def similarity_score_function(ref_token, hyp_token):
+ return 0 if (ref_token == hyp_token) else -1.0
+
+def insertion_score_function(token):
+ return -1.0
+
+def deletion_score_function(token):
+ return -1.0
+
+def EditDistance(
+ ref,
+ hyp,
+ similarity_score_function = similarity_score_function,
+ insertion_score_function = insertion_score_function,
+ deletion_score_function = deletion_score_function):
+ assert(len(ref) != 0)
+ class DPState:
+ def __init__(self):
+ self.score = -float('inf')
+ # backpointer
+ self.prev_r = None
+ self.prev_h = None
+
+ def print_search_grid(S, R, H, fstream):
+ print(file=fstream)
+ for r in range(R):
+ for h in range(H):
+ print(F'[{r},{h}]:{S[r][h].score:4.3f}:({S[r][h].prev_r},{S[r][h].prev_h}) ', end='', file=fstream)
+ print(file=fstream)
+
+ R = len(ref) + 1
+ H = len(hyp) + 1
+
+ # Construct DP search space, a (R x H) grid
+ S = [ [] for r in range(R) ]
+ for r in range(R):
+ S[r] = [ DPState() for x in range(H) ]
+
+ # initialize DP search grid origin, S(r = 0, h = 0)
+ S[0][0].score = 0.0
+ S[0][0].prev_r = None
+ S[0][0].prev_h = None
+
+ # initialize REF axis
+ for r in range(1, R):
+ S[r][0].score = S[r-1][0].score + deletion_score_function(ref[r-1])
+ S[r][0].prev_r = r-1
+ S[r][0].prev_h = 0
+
+ # initialize HYP axis
+ for h in range(1, H):
+ S[0][h].score = S[0][h-1].score + insertion_score_function(hyp[h-1])
+ S[0][h].prev_r = 0
+ S[0][h].prev_h = h-1
+
+ best_score = S[0][0].score
+ best_state = (0, 0)
+
+ for r in range(1, R):
+ for h in range(1, H):
+ sub_or_cor_score = similarity_score_function(ref[r-1], hyp[h-1])
+ new_score = S[r-1][h-1].score + sub_or_cor_score
+ if new_score >= S[r][h].score:
+ S[r][h].score = new_score
+ S[r][h].prev_r = r-1
+ S[r][h].prev_h = h-1
+
+ del_score = deletion_score_function(ref[r-1])
+ new_score = S[r-1][h].score + del_score
+ if new_score >= S[r][h].score:
+ S[r][h].score = new_score
+ S[r][h].prev_r = r - 1
+ S[r][h].prev_h = h
+
+ ins_score = insertion_score_function(hyp[h-1])
+ new_score = S[r][h-1].score + ins_score
+ if new_score >= S[r][h].score:
+ S[r][h].score = new_score
+ S[r][h].prev_r = r
+ S[r][h].prev_h = h-1
+
+ best_score = S[R-1][H-1].score
+ best_state = (R-1, H-1)
+
+ if DEBUG:
+ print_search_grid(S, R, H, sys.stderr)
+
+ # Backtracing best alignment path, i.e. a list of arcs
+ # arc = (src, dst, ref, hyp, edit_type)
+ # src/dst = (r, h), where r/h refers to search grid state-id along Ref/Hyp axis
+ best_path = []
+ r, h = best_state[0], best_state[1]
+ prev_r, prev_h = S[r][h].prev_r, S[r][h].prev_h
+ score = S[r][h].score
+ # loop invariant:
+ # 1. (prev_r, prev_h) -> (r, h) is a "forward arc" on best alignment path
+ # 2. score is the value of point(r, h) on DP search grid
+ while prev_r != None or prev_h != None:
+ src = (prev_r, prev_h)
+ dst = (r, h)
+ if (r == prev_r + 1 and h == prev_h + 1): # Substitution or correct
+ arc = AlignmentArc(src, dst, ref[prev_r], hyp[prev_h])
+ elif (r == prev_r + 1 and h == prev_h): # Deletion
+ arc = AlignmentArc(src, dst, ref[prev_r], None)
+ elif (r == prev_r and h == prev_h + 1): # Insertion
+ arc = AlignmentArc(src, dst, None, hyp[prev_h])
+ else:
+ raise RuntimeError
+ best_path.append(arc)
+ r, h = prev_r, prev_h
+ prev_r, prev_h = S[r][h].prev_r, S[r][h].prev_h
+ score = S[r][h].score
+
+ best_path.reverse()
+ return (best_path, best_score)
+
+def PrettyPrintAlignment(alignment, stream = sys.stderr):
+ def get_token_str(token):
+ if token == None:
+ return "*"
+ return token
+
+ def is_double_width_char(ch):
+ if (ch >= '\u4e00') and (ch <= '\u9fa5'): # codepoint ranges for Chinese chars
+ return True
+ # TODO: support other double-width-char language such as Japanese, Korean
+ else:
+ return False
+
+ def display_width(token_str):
+ m = 0
+ for c in token_str:
+ if is_double_width_char(c):
+ m += 2
+ else:
+ m += 1
+ return m
+
+ R = ' REF : '
+ H = ' HYP : '
+ E = ' EDIT : '
+ for arc in alignment:
+ r = get_token_str(arc.ref)
+ h = get_token_str(arc.hyp)
+ e = arc.edit_type if arc.edit_type != 'C' else ''
+
+ nr, nh, ne = display_width(r), display_width(h), display_width(e)
+ n = max(nr, nh, ne) + 1
+
+ R += r + ' ' * (n-nr)
+ H += h + ' ' * (n-nh)
+ E += e + ' ' * (n-ne)
+
+ print(R, file=stream)
+ print(H, file=stream)
+ print(E, file=stream)
+
+def CountEdits(alignment):
+ c, s, i, d = 0, 0, 0, 0
+ for arc in alignment:
+ if arc.edit_type == 'C':
+ c += 1
+ elif arc.edit_type == 'S':
+ s += 1
+ elif arc.edit_type == 'I':
+ i += 1
+ elif arc.edit_type == 'D':
+ d += 1
+ else:
+ raise RuntimeError
+ return (c, s, i, d)
+
+def ComputeTokenErrorRate(c, s, i, d):
+ return 100.0 * (s + d + i) / (s + d + c)
+
+def ComputeSentenceErrorRate(num_err_utts, num_utts):
+ assert(num_utts != 0)
+ return 100.0 * num_err_utts / num_utts
+
+
+class EvaluationResult:
+ def __init__(self):
+ self.num_ref_utts = 0
+ self.num_hyp_utts = 0
+ self.num_eval_utts = 0 # seen in both ref & hyp
+ self.num_hyp_without_ref = 0
+
+ self.C = 0
+ self.S = 0
+ self.I = 0
+ self.D = 0
+ self.token_error_rate = 0.0
+
+ self.num_utts_with_error = 0
+ self.sentence_error_rate = 0.0
+
+ def to_json(self):
+ return json.dumps(self.__dict__)
+
+ def to_kaldi(self):
+ info = (
+ F'%WER {self.token_error_rate:.2f} [ {self.S + self.D + self.I} / {self.C + self.S + self.D}, {self.I} ins, {self.D} del, {self.S} sub ]\n'
+ F'%SER {self.sentence_error_rate:.2f} [ {self.num_utts_with_error} / {self.num_eval_utts} ]\n'
+ )
+ return info
+
+ def to_sclite(self):
+ return "TODO"
+
+ def to_espnet(self):
+ return "TODO"
+
+ def to_summary(self):
+ #return json.dumps(self.__dict__, indent=4)
+ summary = (
+ '==================== Overall Statistics ====================\n'
+ F'num_ref_utts: {self.num_ref_utts}\n'
+ F'num_hyp_utts: {self.num_hyp_utts}\n'
+ F'num_hyp_without_ref: {self.num_hyp_without_ref}\n'
+ F'num_eval_utts: {self.num_eval_utts}\n'
+ F'sentence_error_rate: {self.sentence_error_rate:.2f}%\n'
+ F'token_error_rate: {self.token_error_rate:.2f}%\n'
+ F'token_stats:\n'
+ F' - tokens:{self.C + self.S + self.D:>7}\n'
+ F' - edits: {self.S + self.I + self.D:>7}\n'
+ F' - cor: {self.C:>7}\n'
+ F' - sub: {self.S:>7}\n'
+ F' - ins: {self.I:>7}\n'
+ F' - del: {self.D:>7}\n'
+ '============================================================\n'
+ )
+ return summary
+
+
+class Utterance:
+ def __init__(self, uid, text):
+ self.uid = uid
+ self.text = text
+
+
+def LoadUtterances(filepath, format):
+ utts = {}
+ if format == 'text': # utt_id word1 word2 ...
+ with open(filepath, 'r', encoding='utf8') as f:
+ for line in f:
+ line = line.strip()
+ if line:
+ cols = line.split(maxsplit=1)
+ assert(len(cols) == 2 or len(cols) == 1)
+ uid = cols[0]
+ text = cols[1] if len(cols) == 2 else ''
+ if utts.get(uid) != None:
+ raise RuntimeError(F'Found duplicated utterence id {uid}')
+ utts[uid] = Utterance(uid, text)
+ else:
+ raise RuntimeError(F'Unsupported text format {format}')
+ return utts
+
+
+def tokenize_text(text, tokenizer):
+ if tokenizer == 'whitespace':
+ return text.split()
+ elif tokenizer == 'char':
+ return [ ch for ch in ''.join(text.split()) ]
+ else:
+ raise RuntimeError(F'ERROR: Unsupported tokenizer {tokenizer}')
+
+
+if __name__ == '__main__':
+ parser = argparse.ArgumentParser()
+ # optional
+ parser.add_argument('--tokenizer', choices=['whitespace', 'char'], default='whitespace', help='whitespace for WER, char for CER')
+ parser.add_argument('--ref-format', choices=['text'], default='text', help='reference format, first col is utt_id, the rest is text')
+ parser.add_argument('--hyp-format', choices=['text'], default='text', help='hypothesis format, first col is utt_id, the rest is text')
+ # required
+ parser.add_argument('--ref', type=str, required=True, help='input reference file')
+ parser.add_argument('--hyp', type=str, required=True, help='input hypothesis file')
+
+ parser.add_argument('result_file', type=str)
+ args = parser.parse_args()
+ logging.info(args)
+
+ ref_utts = LoadUtterances(args.ref, args.ref_format)
+ hyp_utts = LoadUtterances(args.hyp, args.hyp_format)
+
+ r = EvaluationResult()
+
+ # check valid utterances in hyp that have matched non-empty reference
+ eval_utts = []
+ r.num_hyp_without_ref = 0
+ for uid in sorted(hyp_utts.keys()):
+ if uid in ref_utts.keys(): # TODO: efficiency
+ if ref_utts[uid].text.strip(): # non-empty reference
+ eval_utts.append(uid)
+ else:
+ logging.warn(F'Found {uid} with empty reference, skipping...')
+ else:
+ logging.warn(F'Found {uid} without reference, skipping...')
+ r.num_hyp_without_ref += 1
+
+ r.num_hyp_utts = len(hyp_utts)
+ r.num_ref_utts = len(ref_utts)
+ r.num_eval_utts = len(eval_utts)
+
+ with open(args.result_file, 'w+', encoding='utf8') as fo:
+ for uid in eval_utts:
+ ref = ref_utts[uid]
+ hyp = hyp_utts[uid]
+
+ alignment, score = EditDistance(
+ tokenize_text(ref.text, args.tokenizer),
+ tokenize_text(hyp.text, args.tokenizer)
+ )
+
+ c, s, i, d = CountEdits(alignment)
+ utt_ter = ComputeTokenErrorRate(c, s, i, d)
+
+ # utt-level evaluation result
+ print(F'{{"uid":{uid}, "score":{score}, "ter":{utt_ter:.2f}, "cor":{c}, "sub":{s}, "ins":{i}, "del":{d}}}', file=fo)
+ PrettyPrintAlignment(alignment, fo)
+
+ r.C += c
+ r.S += s
+ r.I += i
+ r.D += d
+
+ if utt_ter > 0:
+ r.num_utts_with_error += 1
+
+ # corpus level evaluation result
+ r.sentence_error_rate = ComputeSentenceErrorRate(r.num_utts_with_error, r.num_eval_utts)
+ r.token_error_rate = ComputeTokenErrorRate(r.C, r.S, r.I, r.D)
+
+ print(r.to_summary(), file=fo)
+
+ print(r.to_json())
+ print(r.to_kaldi())
diff --git a/egs_modelscope/common_uniasr/utils/extract_embeds.py b/egs_modelscope/common_uniasr/utils/extract_embeds.py
new file mode 100755
index 0000000..7b817d8
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/extract_embeds.py
@@ -0,0 +1,47 @@
+from transformers import AutoTokenizer, AutoModel, pipeline
+import numpy as np
+import sys
+import os
+import torch
+from kaldiio import WriteHelper
+import re
+text_file_json = sys.argv[1]
+out_ark = sys.argv[2]
+out_scp = sys.argv[3]
+out_shape = sys.argv[4]
+device = int(sys.argv[5])
+model_path = sys.argv[6]
+
+model = AutoModel.from_pretrained(model_path)
+tokenizer = AutoTokenizer.from_pretrained(model_path)
+extractor = pipeline(task="feature-extraction", model=model, tokenizer=tokenizer, device=device)
+
+with open(text_file_json, 'r') as f:
+ js = f.readlines()
+
+
+f_shape = open(out_shape, "w")
+with WriteHelper('ark,scp:{},{}'.format(out_ark, out_scp)) as writer:
+ with torch.no_grad():
+ for idx, line in enumerate(js):
+ id, tokens = line.strip().split(" ", 1)
+ tokens = re.sub(" ", "", tokens.strip())
+ tokens = ' '.join([j for j in tokens])
+ token_num = len(tokens.split(" "))
+ outputs = extractor(tokens)
+ outputs = np.array(outputs)
+ embeds = outputs[0, 1:-1, :]
+
+ token_num_embeds, dim = embeds.shape
+ if token_num == token_num_embeds:
+ writer(id, embeds)
+ shape_line = "{} {},{}\n".format(id, token_num_embeds, dim)
+ f_shape.write(shape_line)
+ else:
+ print("{}, size has changed, {}, {}, {}".format(id, token_num, token_num_embeds, tokens))
+
+
+
+f_shape.close()
+
+
diff --git a/egs_modelscope/common_uniasr/utils/filter_scp.pl b/egs_modelscope/common_uniasr/utils/filter_scp.pl
new file mode 100755
index 0000000..003530d
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/filter_scp.pl
@@ -0,0 +1,87 @@
+#!/usr/bin/env perl
+# Copyright 2010-2012 Microsoft Corporation
+# Johns Hopkins University (author: Daniel Povey)
+
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+# This script takes a list of utterance-ids or any file whose first field
+# of each line is an utterance-id, and filters an scp
+# file (or any file whose "n-th" field is an utterance id), printing
+# out only those lines whose "n-th" field is in id_list. The index of
+# the "n-th" field is 1, by default, but can be changed by using
+# the -f <n> switch
+
+$exclude = 0;
+$field = 1;
+$shifted = 0;
+
+do {
+ $shifted=0;
+ if ($ARGV[0] eq "--exclude") {
+ $exclude = 1;
+ shift @ARGV;
+ $shifted=1;
+ }
+ if ($ARGV[0] eq "-f") {
+ $field = $ARGV[1];
+ shift @ARGV; shift @ARGV;
+ $shifted=1
+ }
+} while ($shifted);
+
+if(@ARGV < 1 || @ARGV > 2) {
+ die "Usage: filter_scp.pl [--exclude] [-f <field-to-filter-on>] id_list [in.scp] > out.scp \n" .
+ "Prints only the input lines whose f'th field (default: first) is in 'id_list'.\n" .
+ "Note: only the first field of each line in id_list matters. With --exclude, prints\n" .
+ "only the lines that were *not* in id_list.\n" .
+ "Caution: previously, the -f option was interpreted as a zero-based field index.\n" .
+ "If your older scripts (written before Oct 2014) stopped working and you used the\n" .
+ "-f option, add 1 to the argument.\n" .
+ "See also: scripts/filter_scp.pl .\n";
+}
+
+
+$idlist = shift @ARGV;
+open(F, "<$idlist") || die "Could not open id-list file $idlist";
+while(<F>) {
+ @A = split;
+ @A>=1 || die "Invalid id-list file line $_";
+ $seen{$A[0]} = 1;
+}
+
+if ($field == 1) { # Treat this as special case, since it is common.
+ while(<>) {
+ $_ =~ m/\s*(\S+)\s*/ || die "Bad line $_, could not get first field.";
+ # $1 is what we filter on.
+ if ((!$exclude && $seen{$1}) || ($exclude && !defined $seen{$1})) {
+ print $_;
+ }
+ }
+} else {
+ while(<>) {
+ @A = split;
+ @A > 0 || die "Invalid scp file line $_";
+ @A >= $field || die "Invalid scp file line $_";
+ if ((!$exclude && $seen{$A[$field-1]}) || ($exclude && !defined $seen{$A[$field-1]})) {
+ print $_;
+ }
+ }
+}
+
+# tests:
+# the following should print "foo 1"
+# ( echo foo 1; echo bar 2 ) | scripts/filter_scp.pl <(echo foo)
+# the following should print "bar 2".
+# ( echo foo 1; echo bar 2 ) | scripts/filter_scp.pl -f 2 <(echo 2)
diff --git a/egs_modelscope/common_uniasr/utils/fix_data.sh b/egs_modelscope/common_uniasr/utils/fix_data.sh
new file mode 100755
index 0000000..32cdde5
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/fix_data.sh
@@ -0,0 +1,35 @@
+#!/usr/bin/env bash
+
+echo "$0 $@"
+data_dir=$1
+
+if [ ! -f ${data_dir}/wav.scp ]; then
+ echo "$0: wav.scp is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/text ]; then
+ echo "$0: text is not found"
+ exit 1;
+fi
+
+
+
+mkdir -p ${data_dir}/.backup
+
+awk '{print $1}' ${data_dir}/wav.scp > ${data_dir}/.backup/wav_id
+awk '{print $1}' ${data_dir}/text > ${data_dir}/.backup/text_id
+
+sort ${data_dir}/.backup/wav_id ${data_dir}/.backup/text_id | uniq -d > ${data_dir}/.backup/id
+
+cp ${data_dir}/wav.scp ${data_dir}/.backup/wav.scp
+cp ${data_dir}/text ${data_dir}/.backup/text
+
+mv ${data_dir}/wav.scp ${data_dir}/wav.scp.bak
+mv ${data_dir}/text ${data_dir}/text.bak
+
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/wav.scp.bak > ${data_dir}/wav.scp
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/text.bak > ${data_dir}/text
+
+rm ${data_dir}/wav.scp.bak
+rm ${data_dir}/text.bak
diff --git a/egs_modelscope/common_uniasr/utils/fix_data_feat.sh b/egs_modelscope/common_uniasr/utils/fix_data_feat.sh
new file mode 100755
index 0000000..2c92d7f
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/fix_data_feat.sh
@@ -0,0 +1,52 @@
+#!/usr/bin/env bash
+
+echo "$0 $@"
+data_dir=$1
+
+if [ ! -f ${data_dir}/feats.scp ]; then
+ echo "$0: feats.scp is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/text ]; then
+ echo "$0: text is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/speech_shape ]; then
+ echo "$0: feature lengths is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/text_shape ]; then
+ echo "$0: text lengths is not found"
+ exit 1;
+fi
+
+mkdir -p ${data_dir}/.backup
+
+awk '{print $1}' ${data_dir}/feats.scp > ${data_dir}/.backup/wav_id
+awk '{print $1}' ${data_dir}/text > ${data_dir}/.backup/text_id
+
+sort ${data_dir}/.backup/wav_id ${data_dir}/.backup/text_id | uniq -d > ${data_dir}/.backup/id
+
+cp ${data_dir}/feats.scp ${data_dir}/.backup/feats.scp
+cp ${data_dir}/text ${data_dir}/.backup/text
+cp ${data_dir}/speech_shape ${data_dir}/.backup/speech_shape
+cp ${data_dir}/text_shape ${data_dir}/.backup/text_shape
+
+mv ${data_dir}/feats.scp ${data_dir}/feats.scp.bak
+mv ${data_dir}/text ${data_dir}/text.bak
+mv ${data_dir}/speech_shape ${data_dir}/speech_shape.bak
+mv ${data_dir}/text_shape ${data_dir}/text_shape.bak
+
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/feats.scp.bak > ${data_dir}/feats.scp
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/text.bak > ${data_dir}/text
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/speech_shape.bak > ${data_dir}/speech_shape
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/text_shape.bak > ${data_dir}/text_shape
+
+rm ${data_dir}/feats.scp.bak
+rm ${data_dir}/text.bak
+rm ${data_dir}/speech_shape.bak
+rm ${data_dir}/text_shape.bak
+
diff --git a/egs_modelscope/common_uniasr/utils/gen_ark_list.sh b/egs_modelscope/common_uniasr/utils/gen_ark_list.sh
new file mode 100755
index 0000000..aebf356
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/gen_ark_list.sh
@@ -0,0 +1,22 @@
+#!/usr/bin/env bash
+
+
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+ark_dir=$1
+txt_dir=$2
+output_dir=$3
+
+[ ! -d ${ark_dir}/ark ] && echo "$0: ark data is required" && exit 1;
+[ ! -d ${txt_dir}/txt ] && echo "$0: txt data is required" && exit 1;
+
+for n in $(seq $nj); do
+ echo "${ark_dir}/ark/feats.$n.ark ${txt_dir}/txt/text.$n.txt" || exit 1
+done > ${output_dir}/ark_txt.scp || exit 1
+
diff --git a/egs_modelscope/common_uniasr/utils/parse_options.sh b/egs_modelscope/common_uniasr/utils/parse_options.sh
new file mode 100755
index 0000000..71fb9e5
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/parse_options.sh
@@ -0,0 +1,97 @@
+#!/usr/bin/env bash
+
+# Copyright 2012 Johns Hopkins University (Author: Daniel Povey);
+# Arnab Ghoshal, Karel Vesely
+
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+# Parse command-line options.
+# To be sourced by another script (as in ". parse_options.sh").
+# Option format is: --option-name arg
+# and shell variable "option_name" gets set to value "arg."
+# The exception is --help, which takes no arguments, but prints the
+# $help_message variable (if defined).
+
+
+###
+### The --config file options have lower priority to command line
+### options, so we need to import them first...
+###
+
+# Now import all the configs specified by command-line, in left-to-right order
+for ((argpos=1; argpos<$#; argpos++)); do
+ if [ "${!argpos}" == "--config" ]; then
+ argpos_plus1=$((argpos+1))
+ config=${!argpos_plus1}
+ [ ! -r $config ] && echo "$0: missing config '$config'" && exit 1
+ . $config # source the config file.
+ fi
+done
+
+
+###
+### Now we process the command line options
+###
+while true; do
+ [ -z "${1:-}" ] && break; # break if there are no arguments
+ case "$1" in
+ # If the enclosing script is called with --help option, print the help
+ # message and exit. Scripts should put help messages in $help_message
+ --help|-h) if [ -z "$help_message" ]; then echo "No help found." 1>&2;
+ else printf "$help_message\n" 1>&2 ; fi;
+ exit 0 ;;
+ --*=*) echo "$0: options to scripts must be of the form --name value, got '$1'"
+ exit 1 ;;
+ # If the first command-line argument begins with "--" (e.g. --foo-bar),
+ # then work out the variable name as $name, which will equal "foo_bar".
+ --*) name=`echo "$1" | sed s/^--// | sed s/-/_/g`;
+ # Next we test whether the variable in question is undefned-- if so it's
+ # an invalid option and we die. Note: $0 evaluates to the name of the
+ # enclosing script.
+ # The test [ -z ${foo_bar+xxx} ] will return true if the variable foo_bar
+ # is undefined. We then have to wrap this test inside "eval" because
+ # foo_bar is itself inside a variable ($name).
+ eval '[ -z "${'$name'+xxx}" ]' && echo "$0: invalid option $1" 1>&2 && exit 1;
+
+ oldval="`eval echo \\$$name`";
+ # Work out whether we seem to be expecting a Boolean argument.
+ if [ "$oldval" == "true" ] || [ "$oldval" == "false" ]; then
+ was_bool=true;
+ else
+ was_bool=false;
+ fi
+
+ # Set the variable to the right value-- the escaped quotes make it work if
+ # the option had spaces, like --cmd "queue.pl -sync y"
+ eval $name=\"$2\";
+
+ # Check that Boolean-valued arguments are really Boolean.
+ if $was_bool && [[ "$2" != "true" && "$2" != "false" ]]; then
+ echo "$0: expected \"true\" or \"false\": $1 $2" 1>&2
+ exit 1;
+ fi
+ shift 2;
+ ;;
+ *) break;
+ esac
+done
+
+
+# Check for an empty argument to the --cmd option, which can easily occur as a
+# result of scripting errors.
+[ ! -z "${cmd+xxx}" ] && [ -z "$cmd" ] && echo "$0: empty argument to --cmd option" 1>&2 && exit 1;
+
+
+true; # so this script returns exit code 0.
diff --git a/egs_modelscope/common_uniasr/utils/print_args.py b/egs_modelscope/common_uniasr/utils/print_args.py
new file mode 100755
index 0000000..b0c61e5
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/print_args.py
@@ -0,0 +1,45 @@
+#!/usr/bin/env python
+import sys
+
+
+def get_commandline_args(no_executable=True):
+ extra_chars = [
+ " ",
+ ";",
+ "&",
+ "|",
+ "<",
+ ">",
+ "?",
+ "*",
+ "~",
+ "`",
+ '"',
+ "'",
+ "\\",
+ "{",
+ "}",
+ "(",
+ ")",
+ ]
+
+ # Escape the extra characters for shell
+ argv = [
+ arg.replace("'", "'\\''")
+ if all(char not in arg for char in extra_chars)
+ else "'" + arg.replace("'", "'\\''") + "'"
+ for arg in sys.argv
+ ]
+
+ if no_executable:
+ return " ".join(argv[1:])
+ else:
+ return sys.executable + " " + " ".join(argv)
+
+
+def main():
+ print(get_commandline_args())
+
+
+if __name__ == "__main__":
+ main()
diff --git a/egs_modelscope/common_uniasr/utils/proc_conf_oss.py b/egs_modelscope/common_uniasr/utils/proc_conf_oss.py
new file mode 100755
index 0000000..c4a90c5
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/proc_conf_oss.py
@@ -0,0 +1,35 @@
+from pathlib import Path
+
+import torch
+import yaml
+
+
+class NoAliasSafeDumper(yaml.SafeDumper):
+ # Disable anchor/alias in yaml because looks ugly
+ def ignore_aliases(self, data):
+ return True
+
+
+def yaml_no_alias_safe_dump(data, stream=None, **kwargs):
+ """Safe-dump in yaml with no anchor/alias"""
+ return yaml.dump(
+ data, stream, allow_unicode=True, Dumper=NoAliasSafeDumper, **kwargs
+ )
+
+
+def gen_conf(file, out_dir):
+ conf = torch.load(file)["config"]
+ conf["oss_bucket"] = "null"
+ print(conf)
+ output_dir = Path(out_dir)
+ output_dir.mkdir(parents=True, exist_ok=True)
+ with (output_dir / "config.yaml").open("w", encoding="utf-8") as f:
+ yaml_no_alias_safe_dump(conf, f, indent=4, sort_keys=False)
+
+
+if __name__ == "__main__":
+ import sys
+
+ in_f = sys.argv[1]
+ out_f = sys.argv[2]
+ gen_conf(in_f, out_f)
diff --git a/egs_modelscope/common_uniasr/utils/proce_text.py b/egs_modelscope/common_uniasr/utils/proce_text.py
new file mode 100755
index 0000000..9e517a4
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/proce_text.py
@@ -0,0 +1,31 @@
+
+import sys
+import re
+
+in_f = sys.argv[1]
+out_f = sys.argv[2]
+
+
+with open(in_f, "r", encoding="utf-8") as f:
+ lines = f.readlines()
+
+with open(out_f, "w", encoding="utf-8") as f:
+ for line in lines:
+ outs = line.strip().split(" ", 1)
+ if len(outs) == 2:
+ idx, text = outs
+ text = re.sub("</s>", "", text)
+ text = re.sub("<s>", "", text)
+ text = re.sub("@@", "", text)
+ text = re.sub("@", "", text)
+ text = re.sub("<unk>", "", text)
+ text = re.sub(" ", "", text)
+ text = text.lower()
+ else:
+ idx = outs[0]
+ text = " "
+
+ text = [x for x in text]
+ text = " ".join(text)
+ out = "{} {}\n".format(idx, text)
+ f.write(out)
diff --git a/egs_modelscope/common_uniasr/utils/run.pl b/egs_modelscope/common_uniasr/utils/run.pl
new file mode 100755
index 0000000..483f95b
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/run.pl
@@ -0,0 +1,356 @@
+#!/usr/bin/env perl
+use warnings; #sed replacement for -w perl parameter
+# In general, doing
+# run.pl some.log a b c is like running the command a b c in
+# the bash shell, and putting the standard error and output into some.log.
+# To run parallel jobs (backgrounded on the host machine), you can do (e.g.)
+# run.pl JOB=1:4 some.JOB.log a b c JOB is like running the command a b c JOB
+# and putting it in some.JOB.log, for each one. [Note: JOB can be any identifier].
+# If any of the jobs fails, this script will fail.
+
+# A typical example is:
+# run.pl some.log my-prog "--opt=foo bar" foo \| other-prog baz
+# and run.pl will run something like:
+# ( my-prog '--opt=foo bar' foo | other-prog baz ) >& some.log
+#
+# Basically it takes the command-line arguments, quotes them
+# as necessary to preserve spaces, and evaluates them with bash.
+# In addition it puts the command line at the top of the log, and
+# the start and end times of the command at the beginning and end.
+# The reason why this is useful is so that we can create a different
+# version of this program that uses a queueing system instead.
+
+#use Data::Dumper;
+
+@ARGV < 2 && die "usage: run.pl log-file command-line arguments...";
+
+#print STDERR "COMMAND-LINE: " . Dumper(\@ARGV) . "\n";
+$job_pick = 'all';
+$max_jobs_run = -1;
+$jobstart = 1;
+$jobend = 1;
+$ignored_opts = ""; # These will be ignored.
+
+# First parse an option like JOB=1:4, and any
+# options that would normally be given to
+# queue.pl, which we will just discard.
+
+for (my $x = 1; $x <= 2; $x++) { # This for-loop is to
+ # allow the JOB=1:n option to be interleaved with the
+ # options to qsub.
+ while (@ARGV >= 2 && $ARGV[0] =~ m:^-:) {
+ # parse any options that would normally go to qsub, but which will be ignored here.
+ my $switch = shift @ARGV;
+ if ($switch eq "-V") {
+ $ignored_opts .= "-V ";
+ } elsif ($switch eq "--max-jobs-run" || $switch eq "-tc") {
+ # we do support the option --max-jobs-run n, and its GridEngine form -tc n.
+ # if the command appears multiple times uses the smallest option.
+ if ( $max_jobs_run <= 0 ) {
+ $max_jobs_run = shift @ARGV;
+ } else {
+ my $new_constraint = shift @ARGV;
+ if ( ($new_constraint < $max_jobs_run) ) {
+ $max_jobs_run = $new_constraint;
+ }
+ }
+
+ if (! ($max_jobs_run > 0)) {
+ die "run.pl: invalid option --max-jobs-run $max_jobs_run";
+ }
+ } else {
+ my $argument = shift @ARGV;
+ if ($argument =~ m/^--/) {
+ print STDERR "run.pl: WARNING: suspicious argument '$argument' to $switch; starts with '-'\n";
+ }
+ if ($switch eq "-sync" && $argument =~ m/^[yY]/) {
+ $ignored_opts .= "-sync "; # Note: in the
+ # corresponding code in queue.pl it says instead, just "$sync = 1;".
+ } elsif ($switch eq "-pe") { # e.g. -pe smp 5
+ my $argument2 = shift @ARGV;
+ $ignored_opts .= "$switch $argument $argument2 ";
+ } elsif ($switch eq "--gpu") {
+ $using_gpu = $argument;
+ } elsif ($switch eq "--pick") {
+ if($argument =~ m/^(all|failed|incomplete)$/) {
+ $job_pick = $argument;
+ } else {
+ print STDERR "run.pl: ERROR: --pick argument must be one of 'all', 'failed' or 'incomplete'"
+ }
+ } else {
+ # Ignore option.
+ $ignored_opts .= "$switch $argument ";
+ }
+ }
+ }
+ if ($ARGV[0] =~ m/^([\w_][\w\d_]*)+=(\d+):(\d+)$/) { # e.g. JOB=1:20
+ $jobname = $1;
+ $jobstart = $2;
+ $jobend = $3;
+ if ($jobstart > $jobend) {
+ die "run.pl: invalid job range $ARGV[0]";
+ }
+ if ($jobstart <= 0) {
+ die "run.pl: invalid job range $ARGV[0], start must be strictly positive (this is required for GridEngine compatibility).";
+ }
+ shift;
+ } elsif ($ARGV[0] =~ m/^([\w_][\w\d_]*)+=(\d+)$/) { # e.g. JOB=1.
+ $jobname = $1;
+ $jobstart = $2;
+ $jobend = $2;
+ shift;
+ } elsif ($ARGV[0] =~ m/.+\=.*\:.*$/) {
+ print STDERR "run.pl: Warning: suspicious first argument to run.pl: $ARGV[0]\n";
+ }
+}
+
+# Users found this message confusing so we are removing it.
+# if ($ignored_opts ne "") {
+# print STDERR "run.pl: Warning: ignoring options \"$ignored_opts\"\n";
+# }
+
+if ($max_jobs_run == -1) { # If --max-jobs-run option not set,
+ # then work out the number of processors if possible,
+ # and set it based on that.
+ $max_jobs_run = 0;
+ if ($using_gpu) {
+ if (open(P, "nvidia-smi -L |")) {
+ $max_jobs_run++ while (<P>);
+ close(P);
+ }
+ if ($max_jobs_run == 0) {
+ $max_jobs_run = 1;
+ print STDERR "run.pl: Warning: failed to detect number of GPUs from nvidia-smi, using ${max_jobs_run}\n";
+ }
+ } elsif (open(P, "</proc/cpuinfo")) { # Linux
+ while (<P>) { if (m/^processor/) { $max_jobs_run++; } }
+ if ($max_jobs_run == 0) {
+ print STDERR "run.pl: Warning: failed to detect any processors from /proc/cpuinfo\n";
+ $max_jobs_run = 10; # reasonable default.
+ }
+ close(P);
+ } elsif (open(P, "sysctl -a |")) { # BSD/Darwin
+ while (<P>) {
+ if (m/hw\.ncpu\s*[:=]\s*(\d+)/) { # hw.ncpu = 4, or hw.ncpu: 4
+ $max_jobs_run = $1;
+ last;
+ }
+ }
+ close(P);
+ if ($max_jobs_run == 0) {
+ print STDERR "run.pl: Warning: failed to detect any processors from sysctl -a\n";
+ $max_jobs_run = 10; # reasonable default.
+ }
+ } else {
+ # allow at most 32 jobs at once, on non-UNIX systems; change this code
+ # if you need to change this default.
+ $max_jobs_run = 32;
+ }
+ # The just-computed value of $max_jobs_run is just the number of processors
+ # (or our best guess); and if it happens that the number of jobs we need to
+ # run is just slightly above $max_jobs_run, it will make sense to increase
+ # $max_jobs_run to equal the number of jobs, so we don't have a small number
+ # of leftover jobs.
+ $num_jobs = $jobend - $jobstart + 1;
+ if (!$using_gpu &&
+ $num_jobs > $max_jobs_run && $num_jobs < 1.4 * $max_jobs_run) {
+ $max_jobs_run = $num_jobs;
+ }
+}
+
+sub pick_or_exit {
+ # pick_or_exit ( $logfile )
+ # Invoked before each job is started helps to run jobs selectively.
+ #
+ # Given the name of the output logfile decides whether the job must be
+ # executed (by returning from the subroutine) or not (by terminating the
+ # process calling exit)
+ #
+ # PRE: $job_pick is a global variable set by command line switch --pick
+ # and indicates which class of jobs must be executed.
+ #
+ # 1) If a failed job is not executed the process exit code will indicate
+ # failure, just as if the task was just executed and failed.
+ #
+ # 2) If a task is incomplete it will be executed. Incomplete may be either
+ # a job whose log file does not contain the accounting notes in the end,
+ # or a job whose log file does not exist.
+ #
+ # 3) If the $job_pick is set to 'all' (default behavior) a task will be
+ # executed regardless of the result of previous attempts.
+ #
+ # This logic could have been implemented in the main execution loop
+ # but a subroutine to preserve the current level of readability of
+ # that part of the code.
+ #
+ # Alexandre Felipe, (o.alexandre.felipe@gmail.com) 14th of August of 2020
+ #
+ if($job_pick eq 'all'){
+ return; # no need to bother with the previous log
+ }
+ open my $fh, "<", $_[0] or return; # job not executed yet
+ my $log_line;
+ my $cur_line;
+ while ($cur_line = <$fh>) {
+ if( $cur_line =~ m/# Ended \(code .*/ ) {
+ $log_line = $cur_line;
+ }
+ }
+ close $fh;
+ if (! defined($log_line)){
+ return; # incomplete
+ }
+ if ( $log_line =~ m/# Ended \(code 0\).*/ ) {
+ exit(0); # complete
+ } elsif ( $log_line =~ m/# Ended \(code \d+(; signal \d+)?\).*/ ){
+ if ($job_pick !~ m/^(failed|all)$/) {
+ exit(1); # failed but not going to run
+ } else {
+ return; # failed
+ }
+ } elsif ( $log_line =~ m/.*\S.*/ ) {
+ return; # incomplete jobs are always run
+ }
+}
+
+
+$logfile = shift @ARGV;
+
+if (defined $jobname && $logfile !~ m/$jobname/ &&
+ $jobend > $jobstart) {
+ print STDERR "run.pl: you are trying to run a parallel job but "
+ . "you are putting the output into just one log file ($logfile)\n";
+ exit(1);
+}
+
+$cmd = "";
+
+foreach $x (@ARGV) {
+ if ($x =~ m/^\S+$/) { $cmd .= $x . " "; }
+ elsif ($x =~ m:\":) { $cmd .= "'$x' "; }
+ else { $cmd .= "\"$x\" "; }
+}
+
+#$Data::Dumper::Indent=0;
+$ret = 0;
+$numfail = 0;
+%active_pids=();
+
+use POSIX ":sys_wait_h";
+for ($jobid = $jobstart; $jobid <= $jobend; $jobid++) {
+ if (scalar(keys %active_pids) >= $max_jobs_run) {
+
+ # Lets wait for a change in any child's status
+ # Then we have to work out which child finished
+ $r = waitpid(-1, 0);
+ $code = $?;
+ if ($r < 0 ) { die "run.pl: Error waiting for child process"; } # should never happen.
+ if ( defined $active_pids{$r} ) {
+ $jid=$active_pids{$r};
+ $fail[$jid]=$code;
+ if ($code !=0) { $numfail++;}
+ delete $active_pids{$r};
+ # print STDERR "Finished: $r/$jid " . Dumper(\%active_pids) . "\n";
+ } else {
+ die "run.pl: Cannot find the PID of the child process that just finished.";
+ }
+
+ # In theory we could do a non-blocking waitpid over all jobs running just
+ # to find out if only one or more jobs finished during the previous waitpid()
+ # However, we just omit this and will reap the next one in the next pass
+ # through the for(;;) cycle
+ }
+ $childpid = fork();
+ if (!defined $childpid) { die "run.pl: Error forking in run.pl (writing to $logfile)"; }
+ if ($childpid == 0) { # We're in the child... this branch
+ # executes the job and returns (possibly with an error status).
+ if (defined $jobname) {
+ $cmd =~ s/$jobname/$jobid/g;
+ $logfile =~ s/$jobname/$jobid/g;
+ }
+ # exit if the job does not need to be executed
+ pick_or_exit( $logfile );
+
+ system("mkdir -p `dirname $logfile` 2>/dev/null");
+ open(F, ">$logfile") || die "run.pl: Error opening log file $logfile";
+ print F "# " . $cmd . "\n";
+ print F "# Started at " . `date`;
+ $starttime = `date +'%s'`;
+ print F "#\n";
+ close(F);
+
+ # Pipe into bash.. make sure we're not using any other shell.
+ open(B, "|bash") || die "run.pl: Error opening shell command";
+ print B "( " . $cmd . ") 2>>$logfile >> $logfile";
+ close(B); # If there was an error, exit status is in $?
+ $ret = $?;
+
+ $lowbits = $ret & 127;
+ $highbits = $ret >> 8;
+ if ($lowbits != 0) { $return_str = "code $highbits; signal $lowbits" }
+ else { $return_str = "code $highbits"; }
+
+ $endtime = `date +'%s'`;
+ open(F, ">>$logfile") || die "run.pl: Error opening log file $logfile (again)";
+ $enddate = `date`;
+ chop $enddate;
+ print F "# Accounting: time=" . ($endtime - $starttime) . " threads=1\n";
+ print F "# Ended ($return_str) at " . $enddate . ", elapsed time " . ($endtime-$starttime) . " seconds\n";
+ close(F);
+ exit($ret == 0 ? 0 : 1);
+ } else {
+ $pid[$jobid] = $childpid;
+ $active_pids{$childpid} = $jobid;
+ # print STDERR "Queued: " . Dumper(\%active_pids) . "\n";
+ }
+}
+
+# Now we have submitted all the jobs, lets wait until all the jobs finish
+foreach $child (keys %active_pids) {
+ $jobid=$active_pids{$child};
+ $r = waitpid($pid[$jobid], 0);
+ $code = $?;
+ if ($r == -1) { die "run.pl: Error waiting for child process"; } # should never happen.
+ if ($r != 0) { $fail[$jobid]=$code; $numfail++ if $code!=0; } # Completed successfully
+}
+
+# Some sanity checks:
+# The $fail array should not contain undefined codes
+# The number of non-zeros in that array should be equal to $numfail
+# We cannot do foreach() here, as the JOB ids do not start at zero
+$failed_jids=0;
+for ($jobid = $jobstart; $jobid <= $jobend; $jobid++) {
+ $job_return = $fail[$jobid];
+ if (not defined $job_return ) {
+ # print Dumper(\@fail);
+
+ die "run.pl: Sanity check failed: we have indication that some jobs are running " .
+ "even after we waited for all jobs to finish" ;
+ }
+ if ($job_return != 0 ){ $failed_jids++;}
+}
+if ($failed_jids != $numfail) {
+ die "run.pl: Sanity check failed: cannot find out how many jobs failed ($failed_jids x $numfail)."
+}
+if ($numfail > 0) { $ret = 1; }
+
+if ($ret != 0) {
+ $njobs = $jobend - $jobstart + 1;
+ if ($njobs == 1) {
+ if (defined $jobname) {
+ $logfile =~ s/$jobname/$jobstart/; # only one numbered job, so replace name with
+ # that job.
+ }
+ print STDERR "run.pl: job failed, log is in $logfile\n";
+ if ($logfile =~ m/JOB/) {
+ print STDERR "run.pl: probably you forgot to put JOB=1:\$nj in your script.";
+ }
+ }
+ else {
+ $logfile =~ s/$jobname/*/g;
+ print STDERR "run.pl: $numfail / $njobs failed, log is in $logfile\n";
+ }
+}
+
+
+exit ($ret);
diff --git a/egs_modelscope/common_uniasr/utils/shuffle_list.pl b/egs_modelscope/common_uniasr/utils/shuffle_list.pl
new file mode 100755
index 0000000..a116200
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/shuffle_list.pl
@@ -0,0 +1,44 @@
+#!/usr/bin/env perl
+
+# Copyright 2013 Johns Hopkins University (author: Daniel Povey)
+
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+if ($ARGV[0] eq "--srand") {
+ $n = $ARGV[1];
+ $n =~ m/\d+/ || die "Bad argument to --srand option: \"$n\"";
+ srand($ARGV[1]);
+ shift;
+ shift;
+} else {
+ srand(0); # Gives inconsistent behavior if we don't seed.
+}
+
+if (@ARGV > 1 || $ARGV[0] =~ m/^-.+/) { # >1 args, or an option we
+ # don't understand.
+ print "Usage: shuffle_list.pl [--srand N] [input file] > output\n";
+ print "randomizes the order of lines of input.\n";
+ exit(1);
+}
+
+@lines;
+while (<>) {
+ push @lines, [ (rand(), $_)] ;
+}
+
+@lines = sort { $a->[0] cmp $b->[0] } @lines;
+foreach $l (@lines) {
+ print $l->[1];
+}
\ No newline at end of file
diff --git a/egs_modelscope/common_uniasr/utils/split_data.py b/egs_modelscope/common_uniasr/utils/split_data.py
new file mode 100755
index 0000000..060eae6
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/split_data.py
@@ -0,0 +1,60 @@
+import os
+import sys
+import random
+
+
+in_dir = sys.argv[1]
+out_dir = sys.argv[2]
+num_split = sys.argv[3]
+
+
+def split_scp(scp, num):
+ assert len(scp) >= num
+ avg = len(scp) // num
+ out = []
+ begin = 0
+
+ for i in range(num):
+ if i == num - 1:
+ out.append(scp[begin:])
+ else:
+ out.append(scp[begin:begin+avg])
+ begin += avg
+
+ return out
+
+
+os.path.exists("{}/wav.scp".format(in_dir))
+os.path.exists("{}/text".format(in_dir))
+
+with open("{}/wav.scp".format(in_dir), 'r') as infile:
+ wav_list = infile.readlines()
+
+with open("{}/text".format(in_dir), 'r') as infile:
+ text_list = infile.readlines()
+
+assert len(wav_list) == len(text_list)
+
+x = list(zip(wav_list, text_list))
+random.shuffle(x)
+wav_shuffle_list, text_shuffle_list = zip(*x)
+
+num_split = int(num_split)
+wav_split_list = split_scp(wav_shuffle_list, num_split)
+text_split_list = split_scp(text_shuffle_list, num_split)
+
+for idx, wav_list in enumerate(wav_split_list, 1):
+ path = out_dir + "/split" + str(num_split) + "/" + str(idx)
+ if not os.path.exists(path):
+ os.makedirs(path)
+ with open("{}/wav.scp".format(path), 'w') as wav_writer:
+ for line in wav_list:
+ wav_writer.write(line)
+
+for idx, text_list in enumerate(text_split_list, 1):
+ path = out_dir + "/split" + str(num_split) + "/" + str(idx)
+ if not os.path.exists(path):
+ os.makedirs(path)
+ with open("{}/text".format(path), 'w') as text_writer:
+ for line in text_list:
+ text_writer.write(line)
diff --git a/egs_modelscope/common_uniasr/utils/split_scp.pl b/egs_modelscope/common_uniasr/utils/split_scp.pl
new file mode 100755
index 0000000..0876dcb
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/split_scp.pl
@@ -0,0 +1,246 @@
+#!/usr/bin/env perl
+
+# Copyright 2010-2011 Microsoft Corporation
+
+# See ../../COPYING for clarification regarding multiple authors
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+# This program splits up any kind of .scp or archive-type file.
+# If there is no utt2spk option it will work on any text file and
+# will split it up with an approximately equal number of lines in
+# each but.
+# With the --utt2spk option it will work on anything that has the
+# utterance-id as the first entry on each line; the utt2spk file is
+# of the form "utterance speaker" (on each line).
+# It splits it into equal size chunks as far as it can. If you use the utt2spk
+# option it will make sure these chunks coincide with speaker boundaries. In
+# this case, if there are more chunks than speakers (and in some other
+# circumstances), some of the resulting chunks will be empty and it will print
+# an error message and exit with nonzero status.
+# You will normally call this like:
+# split_scp.pl scp scp.1 scp.2 scp.3 ...
+# or
+# split_scp.pl --utt2spk=utt2spk scp scp.1 scp.2 scp.3 ...
+# Note that you can use this script to split the utt2spk file itself,
+# e.g. split_scp.pl --utt2spk=utt2spk utt2spk utt2spk.1 utt2spk.2 ...
+
+# You can also call the scripts like:
+# split_scp.pl -j 3 0 scp scp.0
+# [note: with this option, it assumes zero-based indexing of the split parts,
+# i.e. the second number must be 0 <= n < num-jobs.]
+
+use warnings;
+
+$num_jobs = 0;
+$job_id = 0;
+$utt2spk_file = "";
+$one_based = 0;
+
+for ($x = 1; $x <= 3 && @ARGV > 0; $x++) {
+ if ($ARGV[0] eq "-j") {
+ shift @ARGV;
+ $num_jobs = shift @ARGV;
+ $job_id = shift @ARGV;
+ }
+ if ($ARGV[0] =~ /--utt2spk=(.+)/) {
+ $utt2spk_file=$1;
+ shift;
+ }
+ if ($ARGV[0] eq '--one-based') {
+ $one_based = 1;
+ shift @ARGV;
+ }
+}
+
+if ($num_jobs != 0 && ($num_jobs < 0 || $job_id - $one_based < 0 ||
+ $job_id - $one_based >= $num_jobs)) {
+ die "$0: Invalid job number/index values for '-j $num_jobs $job_id" .
+ ($one_based ? " --one-based" : "") . "'\n"
+}
+
+$one_based
+ and $job_id--;
+
+if(($num_jobs == 0 && @ARGV < 2) || ($num_jobs > 0 && (@ARGV < 1 || @ARGV > 2))) {
+ die
+"Usage: split_scp.pl [--utt2spk=<utt2spk_file>] in.scp out1.scp out2.scp ...
+ or: split_scp.pl -j num-jobs job-id [--one-based] [--utt2spk=<utt2spk_file>] in.scp [out.scp]
+ ... where 0 <= job-id < num-jobs, or 1 <= job-id <- num-jobs if --one-based.\n";
+}
+
+$error = 0;
+$inscp = shift @ARGV;
+if ($num_jobs == 0) { # without -j option
+ @OUTPUTS = @ARGV;
+} else {
+ for ($j = 0; $j < $num_jobs; $j++) {
+ if ($j == $job_id) {
+ if (@ARGV > 0) { push @OUTPUTS, $ARGV[0]; }
+ else { push @OUTPUTS, "-"; }
+ } else {
+ push @OUTPUTS, "/dev/null";
+ }
+ }
+}
+
+if ($utt2spk_file ne "") { # We have the --utt2spk option...
+ open($u_fh, '<', $utt2spk_file) || die "$0: Error opening utt2spk file $utt2spk_file: $!\n";
+ while(<$u_fh>) {
+ @A = split;
+ @A == 2 || die "$0: Bad line $_ in utt2spk file $utt2spk_file\n";
+ ($u,$s) = @A;
+ $utt2spk{$u} = $s;
+ }
+ close $u_fh;
+ open($i_fh, '<', $inscp) || die "$0: Error opening input scp file $inscp: $!\n";
+ @spkrs = ();
+ while(<$i_fh>) {
+ @A = split;
+ if(@A == 0) { die "$0: Empty or space-only line in scp file $inscp\n"; }
+ $u = $A[0];
+ $s = $utt2spk{$u};
+ defined $s || die "$0: No utterance $u in utt2spk file $utt2spk_file\n";
+ if(!defined $spk_count{$s}) {
+ push @spkrs, $s;
+ $spk_count{$s} = 0;
+ $spk_data{$s} = []; # ref to new empty array.
+ }
+ $spk_count{$s}++;
+ push @{$spk_data{$s}}, $_;
+ }
+ # Now split as equally as possible ..
+ # First allocate spks to files by allocating an approximately
+ # equal number of speakers.
+ $numspks = @spkrs; # number of speakers.
+ $numscps = @OUTPUTS; # number of output files.
+ if ($numspks < $numscps) {
+ die "$0: Refusing to split data because number of speakers $numspks " .
+ "is less than the number of output .scp files $numscps\n";
+ }
+ for($scpidx = 0; $scpidx < $numscps; $scpidx++) {
+ $scparray[$scpidx] = []; # [] is array reference.
+ }
+ for ($spkidx = 0; $spkidx < $numspks; $spkidx++) {
+ $scpidx = int(($spkidx*$numscps) / $numspks);
+ $spk = $spkrs[$spkidx];
+ push @{$scparray[$scpidx]}, $spk;
+ $scpcount[$scpidx] += $spk_count{$spk};
+ }
+
+ # Now will try to reassign beginning + ending speakers
+ # to different scp's and see if it gets more balanced.
+ # Suppose objf we're minimizing is sum_i (num utts in scp[i] - average)^2.
+ # We can show that if considering changing just 2 scp's, we minimize
+ # this by minimizing the squared difference in sizes. This is
+ # equivalent to minimizing the absolute difference in sizes. This
+ # shows this method is bound to converge.
+
+ $changed = 1;
+ while($changed) {
+ $changed = 0;
+ for($scpidx = 0; $scpidx < $numscps; $scpidx++) {
+ # First try to reassign ending spk of this scp.
+ if($scpidx < $numscps-1) {
+ $sz = @{$scparray[$scpidx]};
+ if($sz > 0) {
+ $spk = $scparray[$scpidx]->[$sz-1];
+ $count = $spk_count{$spk};
+ $nutt1 = $scpcount[$scpidx];
+ $nutt2 = $scpcount[$scpidx+1];
+ if( abs( ($nutt2+$count) - ($nutt1-$count))
+ < abs($nutt2 - $nutt1)) { # Would decrease
+ # size-diff by reassigning spk...
+ $scpcount[$scpidx+1] += $count;
+ $scpcount[$scpidx] -= $count;
+ pop @{$scparray[$scpidx]};
+ unshift @{$scparray[$scpidx+1]}, $spk;
+ $changed = 1;
+ }
+ }
+ }
+ if($scpidx > 0 && @{$scparray[$scpidx]} > 0) {
+ $spk = $scparray[$scpidx]->[0];
+ $count = $spk_count{$spk};
+ $nutt1 = $scpcount[$scpidx-1];
+ $nutt2 = $scpcount[$scpidx];
+ if( abs( ($nutt2-$count) - ($nutt1+$count))
+ < abs($nutt2 - $nutt1)) { # Would decrease
+ # size-diff by reassigning spk...
+ $scpcount[$scpidx-1] += $count;
+ $scpcount[$scpidx] -= $count;
+ shift @{$scparray[$scpidx]};
+ push @{$scparray[$scpidx-1]}, $spk;
+ $changed = 1;
+ }
+ }
+ }
+ }
+ # Now print out the files...
+ for($scpidx = 0; $scpidx < $numscps; $scpidx++) {
+ $scpfile = $OUTPUTS[$scpidx];
+ ($scpfile ne '-' ? open($f_fh, '>', $scpfile)
+ : open($f_fh, '>&', \*STDOUT)) ||
+ die "$0: Could not open scp file $scpfile for writing: $!\n";
+ $count = 0;
+ if(@{$scparray[$scpidx]} == 0) {
+ print STDERR "$0: eError: split_scp.pl producing empty .scp file " .
+ "$scpfile (too many splits and too few speakers?)\n";
+ $error = 1;
+ } else {
+ foreach $spk ( @{$scparray[$scpidx]} ) {
+ print $f_fh @{$spk_data{$spk}};
+ $count += $spk_count{$spk};
+ }
+ $count == $scpcount[$scpidx] || die "Count mismatch [code error]";
+ }
+ close($f_fh);
+ }
+} else {
+ # This block is the "normal" case where there is no --utt2spk
+ # option and we just break into equal size chunks.
+
+ open($i_fh, '<', $inscp) || die "$0: Error opening input scp file $inscp: $!\n";
+
+ $numscps = @OUTPUTS; # size of array.
+ @F = ();
+ while(<$i_fh>) {
+ push @F, $_;
+ }
+ $numlines = @F;
+ if($numlines == 0) {
+ print STDERR "$0: error: empty input scp file $inscp\n";
+ $error = 1;
+ }
+ $linesperscp = int( $numlines / $numscps); # the "whole part"..
+ $linesperscp >= 1 || die "$0: You are splitting into too many pieces! [reduce \$nj ($numscps) to be smaller than the number of lines ($numlines) in $inscp]\n";
+ $remainder = $numlines - ($linesperscp * $numscps);
+ ($remainder >= 0 && $remainder < $numlines) || die "bad remainder $remainder";
+ # [just doing int() rounds down].
+ $n = 0;
+ for($scpidx = 0; $scpidx < @OUTPUTS; $scpidx++) {
+ $scpfile = $OUTPUTS[$scpidx];
+ ($scpfile ne '-' ? open($o_fh, '>', $scpfile)
+ : open($o_fh, '>&', \*STDOUT)) ||
+ die "$0: Could not open scp file $scpfile for writing: $!\n";
+ for($k = 0; $k < $linesperscp + ($scpidx < $remainder ? 1 : 0); $k++) {
+ print $o_fh $F[$n++];
+ }
+ close($o_fh) || die "$0: Eror closing scp file $scpfile: $!\n";
+ }
+ $n == $numlines || die "$n != $numlines [code error]";
+}
+
+exit ($error);
diff --git a/egs_modelscope/common_uniasr/utils/subset_data_dir_tr_cv.sh b/egs_modelscope/common_uniasr/utils/subset_data_dir_tr_cv.sh
new file mode 100755
index 0000000..e16cebd
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/subset_data_dir_tr_cv.sh
@@ -0,0 +1,30 @@
+#!/usr/bin/env bash
+
+dev_num_utt=1000
+
+echo "$0 $@"
+. utils/parse_options.sh || exit 1;
+
+train_data=$1
+out_dir=$2
+
+[ ! -f ${train_data}/wav.scp ] && echo "$0: no such file ${train_data}/wav.scp" && exit 1;
+[ ! -f ${train_data}/text ] && echo "$0: no such file ${train_data}/text" && exit 1;
+
+mkdir -p ${out_dir}/train && mkdir -p ${out_dir}/dev
+
+cp ${train_data}/wav.scp ${out_dir}/train/wav.scp.bak
+cp ${train_data}/text ${out_dir}/train/text.bak
+
+num_utt=$(wc -l <${out_dir}/train/wav.scp.bak)
+
+utils/shuffle_list.pl --srand 1 ${out_dir}/train/wav.scp.bak > ${out_dir}/train/wav.scp.shuf
+head -n ${dev_num_utt} ${out_dir}/train/wav.scp.shuf > ${out_dir}/dev/wav.scp
+tail -n $((${num_utt}-${dev_num_utt})) ${out_dir}/train/wav.scp.shuf > ${out_dir}/train/wav.scp
+
+utils/shuffle_list.pl --srand 1 ${out_dir}/train/text.bak > ${out_dir}/train/text.shuf
+head -n ${dev_num_utt} ${out_dir}/train/text.shuf > ${out_dir}/dev/text
+tail -n $((${num_utt}-${dev_num_utt})) ${out_dir}/train/text.shuf > ${out_dir}/train/text
+
+rm ${out_dir}/train/wav.scp.bak ${out_dir}/train/text.bak
+rm ${out_dir}/train/wav.scp.shuf ${out_dir}/train/text.shuf
diff --git a/egs_modelscope/common_uniasr/utils/text2token.py b/egs_modelscope/common_uniasr/utils/text2token.py
new file mode 100755
index 0000000..56c3913
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/text2token.py
@@ -0,0 +1,135 @@
+#!/usr/bin/env python3
+
+# Copyright 2017 Johns Hopkins University (Shinji Watanabe)
+# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
+
+
+import argparse
+import codecs
+import re
+import sys
+
+is_python2 = sys.version_info[0] == 2
+
+
+def exist_or_not(i, match_pos):
+ start_pos = None
+ end_pos = None
+ for pos in match_pos:
+ if pos[0] <= i < pos[1]:
+ start_pos = pos[0]
+ end_pos = pos[1]
+ break
+
+ return start_pos, end_pos
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="convert raw text to tokenized text",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--nchar",
+ "-n",
+ default=1,
+ type=int,
+ help="number of characters to split, i.e., \
+ aabb -> a a b b with -n 1 and aa bb with -n 2",
+ )
+ parser.add_argument(
+ "--skip-ncols", "-s", default=0, type=int, help="skip first n columns"
+ )
+ parser.add_argument("--space", default="<space>", type=str, help="space symbol")
+ parser.add_argument(
+ "--non-lang-syms",
+ "-l",
+ default=None,
+ type=str,
+ help="list of non-linguistic symobles, e.g., <NOISE> etc.",
+ )
+ parser.add_argument("text", type=str, default=False, nargs="?", help="input text")
+ parser.add_argument(
+ "--trans_type",
+ "-t",
+ type=str,
+ default="char",
+ choices=["char", "phn"],
+ help="""Transcript type. char/phn. e.g., for TIMIT FADG0_SI1279 -
+ If trans_type is char,
+ read from SI1279.WRD file -> "bricks are an alternative"
+ Else if trans_type is phn,
+ read from SI1279.PHN file -> "sil b r ih sil k s aa r er n aa l
+ sil t er n ih sil t ih v sil" """,
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ rs = []
+ if args.non_lang_syms is not None:
+ with codecs.open(args.non_lang_syms, "r", encoding="utf-8") as f:
+ nls = [x.rstrip() for x in f.readlines()]
+ rs = [re.compile(re.escape(x)) for x in nls]
+
+ if args.text:
+ f = codecs.open(args.text, encoding="utf-8")
+ else:
+ f = codecs.getreader("utf-8")(sys.stdin if is_python2 else sys.stdin.buffer)
+
+ sys.stdout = codecs.getwriter("utf-8")(
+ sys.stdout if is_python2 else sys.stdout.buffer
+ )
+ line = f.readline()
+ n = args.nchar
+ while line:
+ x = line.split()
+ print(" ".join(x[: args.skip_ncols]), end=" ")
+ a = " ".join(x[args.skip_ncols :])
+
+ # get all matched positions
+ match_pos = []
+ for r in rs:
+ i = 0
+ while i >= 0:
+ m = r.search(a, i)
+ if m:
+ match_pos.append([m.start(), m.end()])
+ i = m.end()
+ else:
+ break
+
+ if args.trans_type == "phn":
+ a = a.split(" ")
+ else:
+ if len(match_pos) > 0:
+ chars = []
+ i = 0
+ while i < len(a):
+ start_pos, end_pos = exist_or_not(i, match_pos)
+ if start_pos is not None:
+ chars.append(a[start_pos:end_pos])
+ i = end_pos
+ else:
+ chars.append(a[i])
+ i += 1
+ a = chars
+
+ a = [a[j : j + n] for j in range(0, len(a), n)]
+
+ a_flat = []
+ for z in a:
+ a_flat.append("".join(z))
+
+ a_chars = [z.replace(" ", args.space) for z in a_flat]
+ if args.trans_type == "phn":
+ a_chars = [z.replace("sil", args.space) for z in a_chars]
+ print(" ".join(a_chars))
+ line = f.readline()
+
+
+if __name__ == "__main__":
+ main()
diff --git a/egs_modelscope/common_uniasr/utils/text_tokenize.py b/egs_modelscope/common_uniasr/utils/text_tokenize.py
new file mode 100755
index 0000000..962ea11
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/text_tokenize.py
@@ -0,0 +1,106 @@
+import re
+import argparse
+
+
+def load_dict(seg_file):
+ seg_dict = {}
+ with open(seg_file, 'r') as infile:
+ for line in infile:
+ s = line.strip().split()
+ key = s[0]
+ value = s[1:]
+ seg_dict[key] = " ".join(value)
+ return seg_dict
+
+
+def forward_segment(text, dic):
+ word_list = []
+ i = 0
+ while i < len(text):
+ longest_word = text[i]
+ for j in range(i + 1, len(text) + 1):
+ word = text[i:j]
+ if word in dic:
+ if len(word) > len(longest_word):
+ longest_word = word
+ word_list.append(longest_word)
+ i += len(longest_word)
+ return word_list
+
+
+def tokenize(txt,
+ seg_dict):
+ out_txt = ""
+ pattern = re.compile(r"([\u4E00-\u9FA5A-Za-z0-9])")
+ for word in txt:
+ if pattern.match(word):
+ if word in seg_dict:
+ out_txt += seg_dict[word] + " "
+ else:
+ out_txt += "<unk>" + " "
+ else:
+ continue
+ return out_txt.strip()
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="text tokenize",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--text-file",
+ "-t",
+ default=False,
+ required=True,
+ type=str,
+ help="input text",
+ )
+ parser.add_argument(
+ "--seg-file",
+ "-s",
+ default=False,
+ required=True,
+ type=str,
+ help="seg file",
+ )
+ parser.add_argument(
+ "--txt-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="txt index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ txt_writer = open("{}/text.{}.txt".format(args.output_dir, args.txt_index), 'w')
+ shape_writer = open("{}/len.{}".format(args.output_dir, args.txt_index), 'w')
+ seg_dict = load_dict(args.seg_file)
+ with open(args.text_file, 'r') as infile:
+ for line in infile:
+ s = line.strip().split()
+ text_id = s[0]
+ text_list = forward_segment("".join(s[1:]).lower(), seg_dict)
+ text = tokenize(text_list, seg_dict)
+ lens = len(text.strip().split())
+ txt_writer.write(text_id + " " + text + '\n')
+ shape_writer.write(text_id + " " + str(lens) + '\n')
+
+
+if __name__ == '__main__':
+ main()
+
diff --git a/egs_modelscope/common_uniasr/utils/text_tokenize.sh b/egs_modelscope/common_uniasr/utils/text_tokenize.sh
new file mode 100755
index 0000000..6b74fef
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/text_tokenize.sh
@@ -0,0 +1,35 @@
+#!/usr/bin/env bash
+
+
+# Begin configuration section.
+nj=32
+cmd=utils/run.pl
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+# tokenize configuration
+text_dir=$1
+seg_file=$2
+logdir=$3
+output_dir=$4
+
+txt_dir=${output_dir}/txt; mkdir -p ${output_dir}/txt
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/text_tokenize.JOB.log \
+ python utils/text_tokenize.py -t ${text_dir}/txt/text.JOB.txt \
+ -s ${seg_file} -i JOB -o ${txt_dir} \
+ || exit 1;
+
+# concatenate the text files together.
+for n in $(seq $nj); do
+ cat ${txt_dir}/text.$n.txt || exit 1
+done > ${output_dir}/text || exit 1
+
+for n in $(seq $nj); do
+ cat ${txt_dir}/len.$n || exit 1
+done > ${output_dir}/text_shape || exit 1
+
+echo "$0: Succeeded text tokenize"
diff --git a/egs_modelscope/common_uniasr/utils/textnorm_zh.py b/egs_modelscope/common_uniasr/utils/textnorm_zh.py
new file mode 100755
index 0000000..79feb83
--- /dev/null
+++ b/egs_modelscope/common_uniasr/utils/textnorm_zh.py
@@ -0,0 +1,834 @@
+#!/usr/bin/env python3
+# coding=utf-8
+
+# Authors:
+# 2019.5 Zhiyang Zhou (https://github.com/Joee1995/chn_text_norm.git)
+# 2019.9 Jiayu DU
+#
+# requirements:
+# - python 3.X
+# notes: python 2.X WILL fail or produce misleading results
+
+import sys, os, argparse, codecs, string, re
+
+# ================================================================================ #
+# basic constant
+# ================================================================================ #
+CHINESE_DIGIS = u'闆朵竴浜屼笁鍥涗簲鍏竷鍏節'
+BIG_CHINESE_DIGIS_SIMPLIFIED = u'闆跺9璐板弫鑲嗕紞闄嗘煉鎹岀帠'
+BIG_CHINESE_DIGIS_TRADITIONAL = u'闆跺9璨冲弮鑲嗕紞闄告煉鎹岀帠'
+SMALLER_BIG_CHINESE_UNITS_SIMPLIFIED = u'鍗佺櫨鍗冧竾'
+SMALLER_BIG_CHINESE_UNITS_TRADITIONAL = u'鎷句桨浠熻惉'
+LARGER_CHINESE_NUMERING_UNITS_SIMPLIFIED = u'浜垮厗浜灀绉┌娌熸锭姝h浇'
+LARGER_CHINESE_NUMERING_UNITS_TRADITIONAL = u'鍎勫厗浜灀绉┌婧濇緱姝h級'
+SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED = u'鍗佺櫨鍗冧竾'
+SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL = u'鎷句桨浠熻惉'
+
+ZERO_ALT = u'銆�'
+ONE_ALT = u'骞�'
+TWO_ALTS = [u'涓�', u'鍏�']
+
+POSITIVE = [u'姝�', u'姝�']
+NEGATIVE = [u'璐�', u'璨�']
+POINT = [u'鐐�', u'榛�']
+# PLUS = [u'鍔�', u'鍔�']
+# SIL = [u'鏉�', u'妲�']
+
+FILLER_CHARS = ['鍛�', '鍟�']
+ER_WHITELIST = '(鍎垮コ|鍎垮瓙|鍎垮瓩|濂冲効|鍎垮|濡诲効|' \
+ '鑳庡効|濠村効|鏂扮敓鍎縷濠村辜鍎縷骞煎効|灏戝効|灏忓効|鍎挎瓕|鍎跨|鍎跨|鎵樺効鎵�|瀛ゅ効|' \
+ '鍎挎垙|鍎垮寲|鍙板効搴剕楣垮効宀泑姝e効鍏粡|鍚婂効閮庡綋|鐢熷効鑲插コ|鎵樺効甯﹀コ|鍏诲効闃茶�亅鐥村効鍛嗗コ|' \
+ '浣冲効浣冲|鍎挎�滃吔鎵皘鍎挎棤甯哥埗|鍎夸笉瀚屾瘝涓憒鍎胯鍗冮噷姣嶆媴蹇鍎垮ぇ涓嶇敱鐖穦鑻忎篂鍎�)'
+
+# 涓枃鏁板瓧绯荤粺绫诲瀷
+NUMBERING_TYPES = ['low', 'mid', 'high']
+
+CURRENCY_NAMES = '(浜烘皯甯亅缇庡厓|鏃ュ厓|鑻遍晳|娆у厓|椹厠|娉曢儙|鍔犳嬁澶у厓|婢冲厓|娓竵|鍏堜护|鑺叞椹厠|鐖卞皵鍏伴晳|' \
+ '閲屾媺|鑽峰叞鐩緗鍩冩柉搴撳|姣斿濉攟鍗板凹鐩緗鏋楀悏鐗箌鏂拌タ鍏板厓|姣旂储|鍗㈠竷|鏂板姞鍧″厓|闊╁厓|娉伴摙)'
+CURRENCY_UNITS = '((浜縷鍗冧竾|鐧句竾|涓噟鍗億鐧�)|(浜縷鍗冧竾|鐧句竾|涓噟鍗億鐧緗)鍏億(浜縷鍗冧竾|鐧句竾|涓噟鍗億鐧緗)鍧梶瑙抾姣泑鍒�)'
+COM_QUANTIFIERS = '(鍖箌寮爘搴鍥瀨鍦簗灏緗鏉涓獆棣東闃檤闃祙缃憒鐐畖椤秥涓榺妫祙鍙獆鏀瘄琚瓅杈唡鎸憒鎷厊棰梶澹硘绐爘鏇瞸澧檤缇鑵攟' \
+ '鐮搴瀹璐瘄鎵巪鎹唡鍒�|浠鎵搢鎵媩缃梶鍧灞眧宀瓅姹焲婧獆閽焲闃焲鍗晐鍙寍瀵箌鍑簗鍙澶磡鑴殀鏉縷璺硘鏋潀浠秥璐磡' \
+ '閽坾绾縷绠鍚峾浣峾韬珅鍫倈璇緗鏈瑋椤祙瀹秥鎴穦灞倈涓潀姣珅鍘榺鍒唡閽眧涓鏂鎷厊閾鐭硘閽閿眧蹇絴(鍗億姣珅寰�)鍏媩' \
+ '姣珅鍘榺鍒唡瀵竱灏簗涓坾閲寍瀵粅甯竱閾簗绋媩(鍗億鍒唡鍘榺姣珅寰�)绫硘鎾畖鍕簗鍚坾鍗噟鏂梶鐭硘鐩榺纰梶纰焲鍙爘妗秥绗紎鐩唡' \
+ '鐩抾鏉瘄閽焲鏂泑閿厊绨媩绡畖鐩榺妗秥缃恷鐡秥澹秥鍗畖鐩弢绠﹟绠眧鐓瞸鍟東琚媩閽祙骞磡鏈坾鏃瀛鍒粅鏃秥鍛▅澶﹟绉抾鍒唡鏃瑋' \
+ '绾獆宀亅涓東鏇磡澶渱鏄澶弢绉媩鍐瑋浠浼弢杈坾涓竱娉绮抾棰梶骞鍫唡鏉鏍箌鏀瘄閬搢闈鐗噟寮爘棰梶鍧�)'
+
+# punctuation information are based on Zhon project (https://github.com/tsroten/zhon.git)
+CHINESE_PUNC_STOP = '锛侊紵锝°��'
+CHINESE_PUNC_NON_STOP = '锛傦純锛勶紖锛嗭紘锛堬級锛婏紜锛岋紞锛忥細锛涳紲锛濓紴锛狅蓟锛硷冀锛撅伎锝�锝涳綔锝濓綖锝燂綘锝剑锝ゃ�併�冦�嬨�屻�嶃�庛�忋�愩�戙�斻�曘�栥�椼�樸�欍�氥�涖�溿�濄�炪�熴�般�俱�库�撯�斺�樷�欌�涒�溾�濃�炩�熲�︹�э箯'
+CHINESE_PUNC_LIST = CHINESE_PUNC_STOP + CHINESE_PUNC_NON_STOP
+
+# ================================================================================ #
+# basic class
+# ================================================================================ #
+class ChineseChar(object):
+ """
+ 涓枃瀛楃
+ 姣忎釜瀛楃瀵瑰簲绠�浣撳拰绻佷綋,
+ e.g. 绠�浣� = '璐�', 绻佷綋 = '璨�'
+ 杞崲鏃跺彲杞崲涓虹畝浣撴垨绻佷綋
+ """
+
+ def __init__(self, simplified, traditional):
+ self.simplified = simplified
+ self.traditional = traditional
+ #self.__repr__ = self.__str__
+
+ def __str__(self):
+ return self.simplified or self.traditional or None
+
+ def __repr__(self):
+ return self.__str__()
+
+
+class ChineseNumberUnit(ChineseChar):
+ """
+ 涓枃鏁板瓧/鏁颁綅瀛楃
+ 姣忎釜瀛楃闄ょ箒绠�浣撳杩樻湁涓�涓澶栫殑澶у啓瀛楃
+ e.g. '闄�' 鍜� '闄�'
+ """
+
+ def __init__(self, power, simplified, traditional, big_s, big_t):
+ super(ChineseNumberUnit, self).__init__(simplified, traditional)
+ self.power = power
+ self.big_s = big_s
+ self.big_t = big_t
+
+ def __str__(self):
+ return '10^{}'.format(self.power)
+
+ @classmethod
+ def create(cls, index, value, numbering_type=NUMBERING_TYPES[1], small_unit=False):
+
+ if small_unit:
+ return ChineseNumberUnit(power=index + 1,
+ simplified=value[0], traditional=value[1], big_s=value[1], big_t=value[1])
+ elif numbering_type == NUMBERING_TYPES[0]:
+ return ChineseNumberUnit(power=index + 8,
+ simplified=value[0], traditional=value[1], big_s=value[0], big_t=value[1])
+ elif numbering_type == NUMBERING_TYPES[1]:
+ return ChineseNumberUnit(power=(index + 2) * 4,
+ simplified=value[0], traditional=value[1], big_s=value[0], big_t=value[1])
+ elif numbering_type == NUMBERING_TYPES[2]:
+ return ChineseNumberUnit(power=pow(2, index + 3),
+ simplified=value[0], traditional=value[1], big_s=value[0], big_t=value[1])
+ else:
+ raise ValueError(
+ 'Counting type should be in {0} ({1} provided).'.format(NUMBERING_TYPES, numbering_type))
+
+
+class ChineseNumberDigit(ChineseChar):
+ """
+ 涓枃鏁板瓧瀛楃
+ """
+
+ def __init__(self, value, simplified, traditional, big_s, big_t, alt_s=None, alt_t=None):
+ super(ChineseNumberDigit, self).__init__(simplified, traditional)
+ self.value = value
+ self.big_s = big_s
+ self.big_t = big_t
+ self.alt_s = alt_s
+ self.alt_t = alt_t
+
+ def __str__(self):
+ return str(self.value)
+
+ @classmethod
+ def create(cls, i, v):
+ return ChineseNumberDigit(i, v[0], v[1], v[2], v[3])
+
+
+class ChineseMath(ChineseChar):
+ """
+ 涓枃鏁颁綅瀛楃
+ """
+
+ def __init__(self, simplified, traditional, symbol, expression=None):
+ super(ChineseMath, self).__init__(simplified, traditional)
+ self.symbol = symbol
+ self.expression = expression
+ self.big_s = simplified
+ self.big_t = traditional
+
+
+CC, CNU, CND, CM = ChineseChar, ChineseNumberUnit, ChineseNumberDigit, ChineseMath
+
+
+class NumberSystem(object):
+ """
+ 涓枃鏁板瓧绯荤粺
+ """
+ pass
+
+
+class MathSymbol(object):
+ """
+ 鐢ㄤ簬涓枃鏁板瓧绯荤粺鐨勬暟瀛︾鍙� (绻�/绠�浣�), e.g.
+ positive = ['姝�', '姝�']
+ negative = ['璐�', '璨�']
+ point = ['鐐�', '榛�']
+ """
+
+ def __init__(self, positive, negative, point):
+ self.positive = positive
+ self.negative = negative
+ self.point = point
+
+ def __iter__(self):
+ for v in self.__dict__.values():
+ yield v
+
+
+# class OtherSymbol(object):
+# """
+# 鍏朵粬绗﹀彿
+# """
+#
+# def __init__(self, sil):
+# self.sil = sil
+#
+# def __iter__(self):
+# for v in self.__dict__.values():
+# yield v
+
+
+# ================================================================================ #
+# basic utils
+# ================================================================================ #
+def create_system(numbering_type=NUMBERING_TYPES[1]):
+ """
+ 鏍规嵁鏁板瓧绯荤粺绫诲瀷杩斿洖鍒涘缓鐩稿簲鐨勬暟瀛楃郴缁燂紝榛樿涓� mid
+ NUMBERING_TYPES = ['low', 'mid', 'high']: 涓枃鏁板瓧绯荤粺绫诲瀷
+ low: '鍏�' = '浜�' * '鍗�' = $10^{9}$, '浜�' = '鍏�' * '鍗�', etc.
+ mid: '鍏�' = '浜�' * '涓�' = $10^{12}$, '浜�' = '鍏�' * '涓�', etc.
+ high: '鍏�' = '浜�' * '浜�' = $10^{16}$, '浜�' = '鍏�' * '鍏�', etc.
+ 杩斿洖瀵瑰簲鐨勬暟瀛楃郴缁�
+ """
+
+ # chinese number units of '浜�' and larger
+ all_larger_units = zip(
+ LARGER_CHINESE_NUMERING_UNITS_SIMPLIFIED, LARGER_CHINESE_NUMERING_UNITS_TRADITIONAL)
+ larger_units = [CNU.create(i, v, numbering_type, False)
+ for i, v in enumerate(all_larger_units)]
+ # chinese number units of '鍗�, 鐧�, 鍗�, 涓�'
+ all_smaller_units = zip(
+ SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED, SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL)
+ smaller_units = [CNU.create(i, v, small_unit=True)
+ for i, v in enumerate(all_smaller_units)]
+ # digis
+ chinese_digis = zip(CHINESE_DIGIS, CHINESE_DIGIS,
+ BIG_CHINESE_DIGIS_SIMPLIFIED, BIG_CHINESE_DIGIS_TRADITIONAL)
+ digits = [CND.create(i, v) for i, v in enumerate(chinese_digis)]
+ digits[0].alt_s, digits[0].alt_t = ZERO_ALT, ZERO_ALT
+ digits[1].alt_s, digits[1].alt_t = ONE_ALT, ONE_ALT
+ digits[2].alt_s, digits[2].alt_t = TWO_ALTS[0], TWO_ALTS[1]
+
+ # symbols
+ positive_cn = CM(POSITIVE[0], POSITIVE[1], '+', lambda x: x)
+ negative_cn = CM(NEGATIVE[0], NEGATIVE[1], '-', lambda x: -x)
+ point_cn = CM(POINT[0], POINT[1], '.', lambda x,
+ y: float(str(x) + '.' + str(y)))
+ # sil_cn = CM(SIL[0], SIL[1], '-', lambda x, y: float(str(x) + '-' + str(y)))
+ system = NumberSystem()
+ system.units = smaller_units + larger_units
+ system.digits = digits
+ system.math = MathSymbol(positive_cn, negative_cn, point_cn)
+ # system.symbols = OtherSymbol(sil_cn)
+ return system
+
+
+def chn2num(chinese_string, numbering_type=NUMBERING_TYPES[1]):
+
+ def get_symbol(char, system):
+ for u in system.units:
+ if char in [u.traditional, u.simplified, u.big_s, u.big_t]:
+ return u
+ for d in system.digits:
+ if char in [d.traditional, d.simplified, d.big_s, d.big_t, d.alt_s, d.alt_t]:
+ return d
+ for m in system.math:
+ if char in [m.traditional, m.simplified]:
+ return m
+
+ def string2symbols(chinese_string, system):
+ int_string, dec_string = chinese_string, ''
+ for p in [system.math.point.simplified, system.math.point.traditional]:
+ if p in chinese_string:
+ int_string, dec_string = chinese_string.split(p)
+ break
+ return [get_symbol(c, system) for c in int_string], \
+ [get_symbol(c, system) for c in dec_string]
+
+ def correct_symbols(integer_symbols, system):
+ """
+ 涓�鐧惧叓 to 涓�鐧惧叓鍗�
+ 涓�浜夸竴鍗冧笁鐧句竾 to 涓�浜� 涓�鍗冧竾 涓夌櫨涓�
+ """
+
+ if integer_symbols and isinstance(integer_symbols[0], CNU):
+ if integer_symbols[0].power == 1:
+ integer_symbols = [system.digits[1]] + integer_symbols
+
+ if len(integer_symbols) > 1:
+ if isinstance(integer_symbols[-1], CND) and isinstance(integer_symbols[-2], CNU):
+ integer_symbols.append(
+ CNU(integer_symbols[-2].power - 1, None, None, None, None))
+
+ result = []
+ unit_count = 0
+ for s in integer_symbols:
+ if isinstance(s, CND):
+ result.append(s)
+ unit_count = 0
+ elif isinstance(s, CNU):
+ current_unit = CNU(s.power, None, None, None, None)
+ unit_count += 1
+
+ if unit_count == 1:
+ result.append(current_unit)
+ elif unit_count > 1:
+ for i in range(len(result)):
+ if isinstance(result[-i - 1], CNU) and result[-i - 1].power < current_unit.power:
+ result[-i - 1] = CNU(result[-i - 1].power +
+ current_unit.power, None, None, None, None)
+ return result
+
+ def compute_value(integer_symbols):
+ """
+ Compute the value.
+ When current unit is larger than previous unit, current unit * all previous units will be used as all previous units.
+ e.g. '涓ゅ崈涓�' = 2000 * 10000 not 2000 + 10000
+ """
+ value = [0]
+ last_power = 0
+ for s in integer_symbols:
+ if isinstance(s, CND):
+ value[-1] = s.value
+ elif isinstance(s, CNU):
+ value[-1] *= pow(10, s.power)
+ if s.power > last_power:
+ value[:-1] = list(map(lambda v: v *
+ pow(10, s.power), value[:-1]))
+ last_power = s.power
+ value.append(0)
+ return sum(value)
+
+ system = create_system(numbering_type)
+ int_part, dec_part = string2symbols(chinese_string, system)
+ int_part = correct_symbols(int_part, system)
+ int_str = str(compute_value(int_part))
+ dec_str = ''.join([str(d.value) for d in dec_part])
+ if dec_part:
+ return '{0}.{1}'.format(int_str, dec_str)
+ else:
+ return int_str
+
+
+def num2chn(number_string, numbering_type=NUMBERING_TYPES[1], big=False,
+ traditional=False, alt_zero=False, alt_one=False, alt_two=True,
+ use_zeros=True, use_units=True):
+
+ def get_value(value_string, use_zeros=True):
+
+ striped_string = value_string.lstrip('0')
+
+ # record nothing if all zeros
+ if not striped_string:
+ return []
+
+ # record one digits
+ elif len(striped_string) == 1:
+ if use_zeros and len(value_string) != len(striped_string):
+ return [system.digits[0], system.digits[int(striped_string)]]
+ else:
+ return [system.digits[int(striped_string)]]
+
+ # recursively record multiple digits
+ else:
+ result_unit = next(u for u in reversed(
+ system.units) if u.power < len(striped_string))
+ result_string = value_string[:-result_unit.power]
+ return get_value(result_string) + [result_unit] + get_value(striped_string[-result_unit.power:])
+
+ system = create_system(numbering_type)
+
+ int_dec = number_string.split('.')
+ if len(int_dec) == 1:
+ int_string = int_dec[0]
+ dec_string = ""
+ elif len(int_dec) == 2:
+ int_string = int_dec[0]
+ dec_string = int_dec[1]
+ else:
+ raise ValueError(
+ "invalid input num string with more than one dot: {}".format(number_string))
+
+ if use_units and len(int_string) > 1:
+ result_symbols = get_value(int_string)
+ else:
+ result_symbols = [system.digits[int(c)] for c in int_string]
+ dec_symbols = [system.digits[int(c)] for c in dec_string]
+ if dec_string:
+ result_symbols += [system.math.point] + dec_symbols
+
+ if alt_two:
+ liang = CND(2, system.digits[2].alt_s, system.digits[2].alt_t,
+ system.digits[2].big_s, system.digits[2].big_t)
+ for i, v in enumerate(result_symbols):
+ if isinstance(v, CND) and v.value == 2:
+ next_symbol = result_symbols[i +
+ 1] if i < len(result_symbols) - 1 else None
+ previous_symbol = result_symbols[i - 1] if i > 0 else None
+ if isinstance(next_symbol, CNU) and isinstance(previous_symbol, (CNU, type(None))):
+ if next_symbol.power != 1 and ((previous_symbol is None) or (previous_symbol.power != 1)):
+ result_symbols[i] = liang
+
+ # if big is True, '涓�' will not be used and `alt_two` has no impact on output
+ if big:
+ attr_name = 'big_'
+ if traditional:
+ attr_name += 't'
+ else:
+ attr_name += 's'
+ else:
+ if traditional:
+ attr_name = 'traditional'
+ else:
+ attr_name = 'simplified'
+
+ result = ''.join([getattr(s, attr_name) for s in result_symbols])
+
+ # if not use_zeros:
+ # result = result.strip(getattr(system.digits[0], attr_name))
+
+ if alt_zero:
+ result = result.replace(
+ getattr(system.digits[0], attr_name), system.digits[0].alt_s)
+
+ if alt_one:
+ result = result.replace(
+ getattr(system.digits[1], attr_name), system.digits[1].alt_s)
+
+ for i, p in enumerate(POINT):
+ if result.startswith(p):
+ return CHINESE_DIGIS[0] + result
+
+ # ^10, 11, .., 19
+ if len(result) >= 2 and result[1] in [SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED[0],
+ SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL[0]] and \
+ result[0] in [CHINESE_DIGIS[1], BIG_CHINESE_DIGIS_SIMPLIFIED[1], BIG_CHINESE_DIGIS_TRADITIONAL[1]]:
+ result = result[1:]
+
+ return result
+
+
+# ================================================================================ #
+# different types of rewriters
+# ================================================================================ #
+class Cardinal:
+ """
+ CARDINAL绫�
+ """
+
+ def __init__(self, cardinal=None, chntext=None):
+ self.cardinal = cardinal
+ self.chntext = chntext
+
+ def chntext2cardinal(self):
+ return chn2num(self.chntext)
+
+ def cardinal2chntext(self):
+ return num2chn(self.cardinal)
+
+class Digit:
+ """
+ DIGIT绫�
+ """
+
+ def __init__(self, digit=None, chntext=None):
+ self.digit = digit
+ self.chntext = chntext
+
+ # def chntext2digit(self):
+ # return chn2num(self.chntext)
+
+ def digit2chntext(self):
+ return num2chn(self.digit, alt_two=False, use_units=False)
+
+
+class TelePhone:
+ """
+ TELEPHONE绫�
+ """
+
+ def __init__(self, telephone=None, raw_chntext=None, chntext=None):
+ self.telephone = telephone
+ self.raw_chntext = raw_chntext
+ self.chntext = chntext
+
+ # def chntext2telephone(self):
+ # sil_parts = self.raw_chntext.split('<SIL>')
+ # self.telephone = '-'.join([
+ # str(chn2num(p)) for p in sil_parts
+ # ])
+ # return self.telephone
+
+ def telephone2chntext(self, fixed=False):
+
+ if fixed:
+ sil_parts = self.telephone.split('-')
+ self.raw_chntext = '<SIL>'.join([
+ num2chn(part, alt_two=False, use_units=False) for part in sil_parts
+ ])
+ self.chntext = self.raw_chntext.replace('<SIL>', '')
+ else:
+ sp_parts = self.telephone.strip('+').split()
+ self.raw_chntext = '<SP>'.join([
+ num2chn(part, alt_two=False, use_units=False) for part in sp_parts
+ ])
+ self.chntext = self.raw_chntext.replace('<SP>', '')
+ return self.chntext
+
+
+class Fraction:
+ """
+ FRACTION绫�
+ """
+
+ def __init__(self, fraction=None, chntext=None):
+ self.fraction = fraction
+ self.chntext = chntext
+
+ def chntext2fraction(self):
+ denominator, numerator = self.chntext.split('鍒嗕箣')
+ return chn2num(numerator) + '/' + chn2num(denominator)
+
+ def fraction2chntext(self):
+ numerator, denominator = self.fraction.split('/')
+ return num2chn(denominator) + '鍒嗕箣' + num2chn(numerator)
+
+
+class Date:
+ """
+ DATE绫�
+ """
+
+ def __init__(self, date=None, chntext=None):
+ self.date = date
+ self.chntext = chntext
+
+ # def chntext2date(self):
+ # chntext = self.chntext
+ # try:
+ # year, other = chntext.strip().split('骞�', maxsplit=1)
+ # year = Digit(chntext=year).digit2chntext() + '骞�'
+ # except ValueError:
+ # other = chntext
+ # year = ''
+ # if other:
+ # try:
+ # month, day = other.strip().split('鏈�', maxsplit=1)
+ # month = Cardinal(chntext=month).chntext2cardinal() + '鏈�'
+ # except ValueError:
+ # day = chntext
+ # month = ''
+ # if day:
+ # day = Cardinal(chntext=day[:-1]).chntext2cardinal() + day[-1]
+ # else:
+ # month = ''
+ # day = ''
+ # date = year + month + day
+ # self.date = date
+ # return self.date
+
+ def date2chntext(self):
+ date = self.date
+ try:
+ year, other = date.strip().split('骞�', 1)
+ year = Digit(digit=year).digit2chntext() + '骞�'
+ except ValueError:
+ other = date
+ year = ''
+ if other:
+ try:
+ month, day = other.strip().split('鏈�', 1)
+ month = Cardinal(cardinal=month).cardinal2chntext() + '鏈�'
+ except ValueError:
+ day = date
+ month = ''
+ if day:
+ day = Cardinal(cardinal=day[:-1]).cardinal2chntext() + day[-1]
+ else:
+ month = ''
+ day = ''
+ chntext = year + month + day
+ self.chntext = chntext
+ return self.chntext
+
+
+class Money:
+ """
+ MONEY绫�
+ """
+
+ def __init__(self, money=None, chntext=None):
+ self.money = money
+ self.chntext = chntext
+
+ # def chntext2money(self):
+ # return self.money
+
+ def money2chntext(self):
+ money = self.money
+ pattern = re.compile(r'(\d+(\.\d+)?)')
+ matchers = pattern.findall(money)
+ if matchers:
+ for matcher in matchers:
+ money = money.replace(matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext())
+ self.chntext = money
+ return self.chntext
+
+
+class Percentage:
+ """
+ PERCENTAGE绫�
+ """
+
+ def __init__(self, percentage=None, chntext=None):
+ self.percentage = percentage
+ self.chntext = chntext
+
+ def chntext2percentage(self):
+ return chn2num(self.chntext.strip().strip('鐧惧垎涔�')) + '%'
+
+ def percentage2chntext(self):
+ return '鐧惧垎涔�' + num2chn(self.percentage.strip().strip('%'))
+
+
+def remove_erhua(text, er_whitelist):
+ """
+ 鍘婚櫎鍎垮寲闊宠瘝涓殑鍎�:
+ 浠栧コ鍎垮湪閭h竟鍎� -> 浠栧コ鍎垮湪閭h竟
+ """
+
+ er_pattern = re.compile(er_whitelist)
+ new_str=''
+ while re.search('鍎�',text):
+ a = re.search('鍎�',text).span()
+ remove_er_flag = 0
+
+ if er_pattern.search(text):
+ b = er_pattern.search(text).span()
+ if b[0] <= a[0]:
+ remove_er_flag = 1
+
+ if remove_er_flag == 0 :
+ new_str = new_str + text[0:a[0]]
+ text = text[a[1]:]
+ else:
+ new_str = new_str + text[0:b[1]]
+ text = text[b[1]:]
+
+ text = new_str + text
+ return text
+
+# ================================================================================ #
+# NSW Normalizer
+# ================================================================================ #
+class NSWNormalizer:
+ def __init__(self, raw_text):
+ self.raw_text = '^' + raw_text + '$'
+ self.norm_text = ''
+
+ def _particular(self):
+ text = self.norm_text
+ pattern = re.compile(r"(([a-zA-Z]+)浜�([a-zA-Z]+))")
+ matchers = pattern.findall(text)
+ if matchers:
+ # print('particular')
+ for matcher in matchers:
+ text = text.replace(matcher[0], matcher[1]+'2'+matcher[2], 1)
+ self.norm_text = text
+ return self.norm_text
+
+ def normalize(self):
+ text = self.raw_text
+
+ # 瑙勮寖鍖栨棩鏈�
+ pattern = re.compile(r"\D+((([089]\d|(19|20)\d{2})骞�)?(\d{1,2}鏈�(\d{1,2}[鏃ュ彿])?)?)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('date')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Date(date=matcher[0]).date2chntext(), 1)
+
+ # 瑙勮寖鍖栭噾閽�
+ pattern = re.compile(r"\D+((\d+(\.\d+)?)[澶氫綑鍑燷?" + CURRENCY_UNITS + r"(\d" + CURRENCY_UNITS + r"?)?)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('money')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Money(money=matcher[0]).money2chntext(), 1)
+
+ # 瑙勮寖鍖栧浐璇�/鎵嬫満鍙风爜
+ # 鎵嬫満
+ # http://www.jihaoba.com/news/show/13680
+ # 绉诲姩锛�139銆�138銆�137銆�136銆�135銆�134銆�159銆�158銆�157銆�150銆�151銆�152銆�188銆�187銆�182銆�183銆�184銆�178銆�198
+ # 鑱旈�氾細130銆�131銆�132銆�156銆�155銆�186銆�185銆�176
+ # 鐢典俊锛�133銆�153銆�189銆�180銆�181銆�177
+ pattern = re.compile(r"\D((\+?86 ?)?1([38]\d|5[0-35-9]|7[678]|9[89])\d{8})\D")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('telephone')
+ for matcher in matchers:
+ text = text.replace(matcher[0], TelePhone(telephone=matcher[0]).telephone2chntext(), 1)
+ # 鍥鸿瘽
+ pattern = re.compile(r"\D((0(10|2[1-3]|[3-9]\d{2})-?)?[1-9]\d{6,7})\D")
+ matchers = pattern.findall(text)
+ if matchers:
+ # print('fixed telephone')
+ for matcher in matchers:
+ text = text.replace(matcher[0], TelePhone(telephone=matcher[0]).telephone2chntext(fixed=True), 1)
+
+ # 瑙勮寖鍖栧垎鏁�
+ pattern = re.compile(r"(\d+/\d+)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('fraction')
+ for matcher in matchers:
+ text = text.replace(matcher, Fraction(fraction=matcher).fraction2chntext(), 1)
+
+ # 瑙勮寖鍖栫櫨鍒嗘暟
+ text = text.replace('锛�', '%')
+ pattern = re.compile(r"(\d+(\.\d+)?%)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('percentage')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Percentage(percentage=matcher[0]).percentage2chntext(), 1)
+
+ # 瑙勮寖鍖栫函鏁�+閲忚瘝
+ pattern = re.compile(r"(\d+(\.\d+)?)[澶氫綑鍑燷?" + COM_QUANTIFIERS)
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('cardinal+quantifier')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext(), 1)
+
+ # 瑙勮寖鍖栨暟瀛楃紪鍙�
+ pattern = re.compile(r"(\d{4,32})")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('digit')
+ for matcher in matchers:
+ text = text.replace(matcher, Digit(digit=matcher).digit2chntext(), 1)
+
+ # 瑙勮寖鍖栫函鏁�
+ pattern = re.compile(r"(\d+(\.\d+)?)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('cardinal')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext(), 1)
+
+ self.norm_text = text
+ self._particular()
+
+ return self.norm_text.lstrip('^').rstrip('$')
+
+
+def nsw_test_case(raw_text):
+ print('I:' + raw_text)
+ print('O:' + NSWNormalizer(raw_text).normalize())
+ print('')
+
+
+def nsw_test():
+ nsw_test_case('鍥鸿瘽锛�0595-23865596鎴�23880880銆�')
+ nsw_test_case('鍥鸿瘽锛�0595-23865596鎴�23880880銆�')
+ nsw_test_case('鎵嬫満锛�+86 19859213959鎴�15659451527銆�')
+ nsw_test_case('鍒嗘暟锛�32477/76391銆�')
+ nsw_test_case('鐧惧垎鏁帮細80.03%銆�')
+ nsw_test_case('缂栧彿锛�31520181154418銆�')
+ nsw_test_case('绾暟锛�2983.07鍏嬫垨12345.60绫炽��')
+ nsw_test_case('鏃ユ湡锛�1999骞�2鏈�20鏃ユ垨09骞�3鏈�15鍙枫��')
+ nsw_test_case('閲戦挶锛�12鍧�5锛�34.5鍏冿紝20.1涓�')
+ nsw_test_case('鐗规畩锛歄2O鎴朆2C銆�')
+ nsw_test_case('3456涓囧惃')
+ nsw_test_case('2938涓�')
+ nsw_test_case('938')
+ nsw_test_case('浠婂ぉ鍚冧簡115涓皬绗煎寘231涓澶�')
+ nsw_test_case('鏈�62锛呯殑姒傜巼')
+
+
+if __name__ == '__main__':
+ #nsw_test()
+
+ p = argparse.ArgumentParser()
+ p.add_argument('ifile', help='input filename, assume utf-8 encoding')
+ p.add_argument('ofile', help='output filename')
+ p.add_argument('--to_upper', action='store_true', help='convert to upper case')
+ p.add_argument('--to_lower', action='store_true', help='convert to lower case')
+ p.add_argument('--has_key', action='store_true', help="input text has Kaldi's key as first field.")
+ p.add_argument('--remove_fillers', type=bool, default=True, help='remove filler chars such as "鍛�, 鍟�"')
+ p.add_argument('--remove_erhua', type=bool, default=True, help='remove erhua chars such as "杩欏効"')
+ p.add_argument('--log_interval', type=int, default=10000, help='log interval in number of processed lines')
+ args = p.parse_args()
+
+ ifile = codecs.open(args.ifile, 'r', 'utf8')
+ ofile = codecs.open(args.ofile, 'w+', 'utf8')
+
+ n = 0
+ for l in ifile:
+ key = ''
+ text = ''
+ if args.has_key:
+ cols = l.split(maxsplit=1)
+ key = cols[0]
+ if len(cols) == 2:
+ text = cols[1].strip()
+ else:
+ text = ''
+ else:
+ text = l.strip()
+
+ # cases
+ if args.to_upper and args.to_lower:
+ sys.stderr.write('text norm: to_upper OR to_lower?')
+ exit(1)
+ if args.to_upper:
+ text = text.upper()
+ if args.to_lower:
+ text = text.lower()
+
+ # Filler chars removal
+ if args.remove_fillers:
+ for ch in FILLER_CHARS:
+ text = text.replace(ch, '')
+
+ if args.remove_erhua:
+ text = remove_erhua(text, ER_WHITELIST)
+
+ # NSW(Non-Standard-Word) normalization
+ text = NSWNormalizer(text).normalize()
+
+ # Punctuations removal
+ old_chars = CHINESE_PUNC_LIST + string.punctuation # includes all CN and EN punctuations
+ new_chars = ' ' * len(old_chars)
+ del_chars = ''
+ text = text.translate(str.maketrans(old_chars, new_chars, del_chars))
+
+ #
+ if args.has_key:
+ ofile.write(key + '\t' + text + '\n')
+ else:
+ ofile.write(text + '\n')
+
+ n += 1
+ if n % args.log_interval == 0:
+ sys.stderr.write("text norm: {} lines done.\n".format(n))
+
+ sys.stderr.write("text norm: {} lines done in total.\n".format(n))
+
+ ifile.close()
+ ofile.close()
diff --git a/egs_modelscope/speechio/paraformer/modelscope_utils b/egs_modelscope/speechio/paraformer/modelscope_utils
deleted file mode 120000
index fc97768..0000000
--- a/egs_modelscope/speechio/paraformer/modelscope_utils
+++ /dev/null
@@ -1 +0,0 @@
-../../common/modelscope_utils
\ No newline at end of file
diff --git a/egs_modelscope/speechio/paraformer/modelscope_utils/download_model.py b/egs_modelscope/speechio/paraformer/modelscope_utils/download_model.py
new file mode 100755
index 0000000..51ba6b8
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/modelscope_utils/download_model.py
@@ -0,0 +1,25 @@
+#!/usr/bin/env python3
+import argparse
+
+from modelscope.pipelines import pipeline
+from modelscope.utils.constant import Tasks
+
+if __name__ == '__main__':
+ parser = argparse.ArgumentParser(
+ description="download model configs",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument("--model_name",
+ type=str,
+ default="speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-pytorch",
+ help="model name in modelscope")
+ parser.add_argument("--model_revision",
+ type=str,
+ default="v1.0.3",
+ help="model revision in modelscope")
+ args = parser.parse_args()
+
+ inference_pipeline = pipeline(
+ task=Tasks.auto_speech_recognition,
+ model='damo/{}'.format(args.model_name),
+ model_revision=args.model_revision)
diff --git a/egs_modelscope/speechio/paraformer/modelscope_utils/modelscope_infer.sh b/egs_modelscope/speechio/paraformer/modelscope_utils/modelscope_infer.sh
new file mode 100755
index 0000000..a0c606f
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/modelscope_utils/modelscope_infer.sh
@@ -0,0 +1,88 @@
+#!/usr/bin/env bash
+
+set -e
+set -u
+set -o pipefail
+
+data_dir=
+exp_dir=
+model_name=
+model_revision=
+inference_nj=32
+gpuid_list="0,1,2,3"
+njob=32
+gpu_inference=true
+
+test_sets="dev test"
+decode_cmd=utils/run.pl
+
+# LM configs
+use_lm=false
+beam_size=1
+lm_weight=0.0
+
+. utils/parse_options.sh
+
+if ${gpu_inference}; then
+ _ngpu=1
+else
+ _ngpu=0
+fi
+
+# download model from modelscope
+python modelscope_utils/download_model.py \
+ --model_name ${model_name} --model_revision ${model_revision}
+
+modelscope_dir=${HOME}/.cache/modelscope/hub/damo/${model_name}
+
+
+for dset in ${test_sets}; do
+ _dir=${exp_dir}/${model_name}/decode_asr/${dset}
+ _logdir=${_dir}/logdir
+ _data=${data_dir}/${dset}
+ if [ -d ${_dir} ]; then
+ echo "${_dir} is already exists. if you want to decode again, please delete ${_dir} first."
+ exit 1
+ else
+ mkdir -p "${_dir}"
+ mkdir -p "${_logdir}"
+ fi
+
+ if "${use_lm}"; then
+ cp ${modelscope_dir}/decoding.yaml ${modelscope_dir}/decoding.yaml.back
+ sed -i "s#beam_size: [0-9]*#beam_size: `echo $beam_size`#g" ${modelscope_dir}/decoding.yaml
+ sed -i "s#lm_weight: 0.[0-9]*#lm_weight: `echo $lm_weight`#g" ${modelscope_dir}/decoding.yaml
+ fi
+
+ for n in $(seq "${inference_nj}"); do
+ split_scps+=" ${_logdir}/keys.${n}.scp"
+ done
+ # shellcheck disable=SC2086
+ utils/split_scp.pl "${data_dir}/${dset}/wav.scp" ${split_scps}
+
+ echo "Decoding started... log: '${_logdir}/asr_inference.*.log'"
+ # shellcheck disable=SC2086
+ ${decode_cmd} --max-jobs-run "${inference_nj}" JOB=1:"${inference_nj}" "${_logdir}"/asr_inference.JOB.log \
+ python -m funasr.bin.modelscope_infer \
+ --model_name ${model_name} \
+ --model_revision ${model_revision} \
+ --wav_list ${_logdir}/keys.JOB.scp \
+ --output_file ${_logdir}/text.JOB \
+ --gpuid_list ${gpuid_list} \
+ --njob ${njob} \
+ --ngpu ${_ngpu} \
+
+ for i in $(seq ${inference_nj}); do
+ cat ${_logdir}/text.${i}
+ done | sort -k1 >${_dir}/text
+
+ python utils/proce_text.py ${_dir}/text ${_dir}/text.proc
+ python utils/proce_text.py ${_data}/text ${_data}/text.proc
+ python utils/compute_wer.py ${_data}/text.proc ${_dir}/text.proc ${_dir}/text.cer
+ tail -n 3 ${_dir}/text.cer > ${_dir}/text.cer.txt
+ cat ${_dir}/text.cer.txt
+done
+
+if "${use_lm}"; then
+ mv ${modelscope_dir}/decoding.yaml.back ${modelscope_dir}/decoding.yaml
+fi
diff --git a/egs_modelscope/speechio/paraformer/modelscope_utils/update_config.py b/egs_modelscope/speechio/paraformer/modelscope_utils/update_config.py
new file mode 100644
index 0000000..88466ed
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/modelscope_utils/update_config.py
@@ -0,0 +1,41 @@
+import yaml
+import argparse
+
+def update_dct(fin_configs, root):
+ if root == {}:
+ return {}
+ for root_key, root_value in root.items():
+ if not isinstance(root[root_key],dict):
+ fin_configs[root_key] = root[root_key]
+ else:
+ result = update_dct(fin_configs[root_key], root[root_key])
+ fin_configs[root_key] = result
+ return fin_configs
+
+if __name__ == '__main__':
+ parser = argparse.ArgumentParser(
+ description="update configs",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument("--modelscope_config",
+ type=str,
+ help="modelscope config file")
+ parser.add_argument("--finetune_config",
+ type=str,
+ help="finetune config file")
+ parser.add_argument("--output_config",
+ type=str,
+ help="output config file")
+ args = parser.parse_args()
+
+ with open(args.modelscope_config) as f:
+ modelscope_configs = yaml.safe_load(f)
+
+ with open(args.finetune_config) as f:
+ finetune_configs = yaml.safe_load(f)
+
+ # update configs, e.g., lr, batch_size, ...
+ modelscope_configs = update_dct(modelscope_configs, finetune_configs)
+
+ with open(args.output_config, "w") as f:
+ yaml.dump(modelscope_configs, f, indent=4)
diff --git a/egs_modelscope/speechio/paraformer/paraformer_large_infer.sh b/egs_modelscope/speechio/paraformer/paraformer_large_infer.sh
index 8cce760..0988612 100755
--- a/egs_modelscope/speechio/paraformer/paraformer_large_infer.sh
+++ b/egs_modelscope/speechio/paraformer/paraformer_large_infer.sh
@@ -8,7 +8,7 @@
data_dir=
exp_dir=
model_name=speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-pytorch
-model_revision="v1.0.3" # please do not modify the model revision
+model_revision="v1.0.4" # please do not modify the model revision
inference_nj=32
gpuid_list="0,1" # set gpus, e.g., gpuid_list="0,1"
ngpu=$(echo $gpuid_list | awk -F "," '{print NF}')
diff --git a/egs_modelscope/speechio/paraformer/utils b/egs_modelscope/speechio/paraformer/utils
deleted file mode 120000
index 37d9761..0000000
--- a/egs_modelscope/speechio/paraformer/utils
+++ /dev/null
@@ -1 +0,0 @@
-../../../egs/aishell/tranformer/utils/
\ No newline at end of file
diff --git a/egs_modelscope/speechio/paraformer/utils/__init__.py b/egs_modelscope/speechio/paraformer/utils/__init__.py
new file mode 100644
index 0000000..e69de29
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/__init__.py
diff --git a/egs_modelscope/speechio/paraformer/utils/apply_cmvn.py b/egs_modelscope/speechio/paraformer/utils/apply_cmvn.py
new file mode 100755
index 0000000..b5c5086
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/apply_cmvn.py
@@ -0,0 +1,79 @@
+from kaldiio import ReadHelper
+from kaldiio import WriteHelper
+
+import argparse
+import json
+import math
+import numpy as np
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="apply cmvn",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--ark-file",
+ "-a",
+ default=False,
+ required=True,
+ type=str,
+ help="fbank ark file",
+ )
+ parser.add_argument(
+ "--cmvn-file",
+ "-c",
+ default=False,
+ required=True,
+ type=str,
+ help="cmvn file",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ ark_file = args.output_dir + "/feats." + str(args.ark_index) + ".ark"
+ scp_file = args.output_dir + "/feats." + str(args.ark_index) + ".scp"
+ ark_writer = WriteHelper('ark,scp:{},{}'.format(ark_file, scp_file))
+
+ with open(args.cmvn_file) as f:
+ cmvn_stats = json.load(f)
+
+ means = cmvn_stats['mean_stats']
+ vars = cmvn_stats['var_stats']
+ total_frames = cmvn_stats['total_frames']
+
+ for i in range(len(means)):
+ means[i] /= total_frames
+ vars[i] = vars[i] / total_frames - means[i] * means[i]
+ if vars[i] < 1.0e-20:
+ vars[i] = 1.0e-20
+ vars[i] = 1.0 / math.sqrt(vars[i])
+
+ with ReadHelper('ark:{}'.format(args.ark_file)) as ark_reader:
+ for key, mat in ark_reader:
+ mat = (mat - means) * vars
+ ark_writer(key, mat)
+
+
+if __name__ == '__main__':
+ main()
diff --git a/egs_modelscope/speechio/paraformer/utils/apply_cmvn.sh b/egs_modelscope/speechio/paraformer/utils/apply_cmvn.sh
new file mode 100755
index 0000000..f8fd1d1
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/apply_cmvn.sh
@@ -0,0 +1,29 @@
+#!/usr/bin/env bash
+
+. ./path.sh || exit 1;
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+fbankdir=$1
+cmvn_file=$2
+logdir=$3
+output_dir=$4
+
+dump_dir=${output_dir}/ark; mkdir -p ${dump_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/apply_cmvn.JOB.log \
+ python utils/apply_cmvn.py -a $fbankdir/ark/feats.JOB.ark \
+ -c $cmvn_file -i JOB -o ${dump_dir} \
+ || exit 1;
+
+for n in $(seq $nj); do
+ cat ${dump_dir}/feats.$n.scp || exit 1
+done > ${output_dir}/feats.scp || exit 1
+
+echo "$0: Succeeded apply cmvn"
diff --git a/egs_modelscope/speechio/paraformer/utils/apply_lfr_and_cmvn.py b/egs_modelscope/speechio/paraformer/utils/apply_lfr_and_cmvn.py
new file mode 100755
index 0000000..50d18d1
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/apply_lfr_and_cmvn.py
@@ -0,0 +1,143 @@
+from kaldiio import ReadHelper, WriteHelper
+
+import argparse
+import numpy as np
+
+
+def build_LFR_features(inputs, m=7, n=6):
+ LFR_inputs = []
+ T = inputs.shape[0]
+ T_lfr = int(np.ceil(T / n))
+ left_padding = np.tile(inputs[0], ((m - 1) // 2, 1))
+ inputs = np.vstack((left_padding, inputs))
+ T = T + (m - 1) // 2
+ for i in range(T_lfr):
+ if m <= T - i * n:
+ LFR_inputs.append(np.hstack(inputs[i * n:i * n + m]))
+ else:
+ num_padding = m - (T - i * n)
+ frame = np.hstack(inputs[i * n:])
+ for _ in range(num_padding):
+ frame = np.hstack((frame, inputs[-1]))
+ LFR_inputs.append(frame)
+ return np.vstack(LFR_inputs)
+
+
+def build_CMVN_features(inputs, mvn_file): # noqa
+ with open(mvn_file, 'r', encoding='utf-8') as f:
+ lines = f.readlines()
+
+ add_shift_list = []
+ rescale_list = []
+ for i in range(len(lines)):
+ line_item = lines[i].split()
+ if line_item[0] == '<AddShift>':
+ line_item = lines[i + 1].split()
+ if line_item[0] == '<LearnRateCoef>':
+ add_shift_line = line_item[3:(len(line_item) - 1)]
+ add_shift_list = list(add_shift_line)
+ continue
+ elif line_item[0] == '<Rescale>':
+ line_item = lines[i + 1].split()
+ if line_item[0] == '<LearnRateCoef>':
+ rescale_line = line_item[3:(len(line_item) - 1)]
+ rescale_list = list(rescale_line)
+ continue
+
+ for j in range(inputs.shape[0]):
+ for k in range(inputs.shape[1]):
+ add_shift_value = add_shift_list[k]
+ rescale_value = rescale_list[k]
+ inputs[j, k] = float(inputs[j, k]) + float(add_shift_value)
+ inputs[j, k] = float(inputs[j, k]) * float(rescale_value)
+
+ return inputs
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="apply low_frame_rate and cmvn",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--ark-file",
+ "-a",
+ default=False,
+ required=True,
+ type=str,
+ help="fbank ark file",
+ )
+ parser.add_argument(
+ "--lfr",
+ "-f",
+ default=True,
+ type=str,
+ help="low frame rate",
+ )
+ parser.add_argument(
+ "--lfr-m",
+ "-m",
+ default=7,
+ type=int,
+ help="number of frames to stack",
+ )
+ parser.add_argument(
+ "--lfr-n",
+ "-n",
+ default=6,
+ type=int,
+ help="number of frames to skip",
+ )
+ parser.add_argument(
+ "--cmvn-file",
+ "-c",
+ default=False,
+ required=True,
+ type=str,
+ help="global cmvn file",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ dump_ark_file = args.output_dir + "/feats." + str(args.ark_index) + ".ark"
+ dump_scp_file = args.output_dir + "/feats." + str(args.ark_index) + ".scp"
+ shape_file = args.output_dir + "/len." + str(args.ark_index)
+ ark_writer = WriteHelper('ark,scp:{},{}'.format(dump_ark_file, dump_scp_file))
+
+ shape_writer = open(shape_file, 'w')
+ with ReadHelper('ark:{}'.format(args.ark_file)) as ark_reader:
+ for key, mat in ark_reader:
+ if args.lfr:
+ lfr = build_LFR_features(mat, args.lfr_m, args.lfr_n)
+ else:
+ lfr = mat
+ cmvn = build_CMVN_features(lfr, args.cmvn_file)
+ dims = cmvn.shape[1]
+ lens = cmvn.shape[0]
+ shape_writer.write(key + " " + str(lens) + "," + str(dims) + '\n')
+ ark_writer(key, cmvn)
+
+
+if __name__ == '__main__':
+ main()
+
diff --git a/egs_modelscope/speechio/paraformer/utils/apply_lfr_and_cmvn.sh b/egs_modelscope/speechio/paraformer/utils/apply_lfr_and_cmvn.sh
new file mode 100755
index 0000000..3119fdb
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/apply_lfr_and_cmvn.sh
@@ -0,0 +1,38 @@
+#!/usr/bin/env bash
+
+
+# Begin configuration section.
+nj=32
+cmd=utils/run.pl
+
+# feature configuration
+lfr=True
+lfr_m=7
+lfr_n=6
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+fbankdir=$1
+cmvn_file=$2
+logdir=$3
+output_dir=$4
+
+dump_dir=${output_dir}/ark; mkdir -p ${dump_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/apply_lfr_and_cmvn.JOB.log \
+ python utils/apply_lfr_and_cmvn.py -a $fbankdir/ark/feats.JOB.ark \
+ -f $lfr -m $lfr_m -n $lfr_n -c $cmvn_file -i JOB -o ${dump_dir} \
+ || exit 1;
+
+for n in $(seq $nj); do
+ cat ${dump_dir}/feats.$n.scp || exit 1
+done > ${output_dir}/feats.scp || exit 1
+
+for n in $(seq $nj); do
+ cat ${dump_dir}/len.$n || exit 1
+done > ${output_dir}/speech_shape || exit 1
+
+echo "$0: Succeeded apply low frame rate and cmvn"
diff --git a/egs_modelscope/speechio/paraformer/utils/combine_cmvn_file.py b/egs_modelscope/speechio/paraformer/utils/combine_cmvn_file.py
new file mode 100755
index 0000000..b2974a4
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/combine_cmvn_file.py
@@ -0,0 +1,73 @@
+import argparse
+import json
+import numpy as np
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="combine cmvn file",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--dims",
+ "-d",
+ default=80,
+ type=int,
+ help="feature dims",
+ )
+ parser.add_argument(
+ "--cmvn-dir",
+ "-c",
+ default=False,
+ required=True,
+ type=str,
+ help="cmvn dir",
+ )
+
+ parser.add_argument(
+ "--nj",
+ "-n",
+ default=1,
+ required=True,
+ type=int,
+ help="num of cmvn file",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ total_means = np.zeros(args.dims)
+ total_vars = np.zeros(args.dims)
+ total_frames = 0
+
+ cmvn_file = args.output_dir + "/cmvn.json"
+
+ for i in range(1, args.nj+1):
+ with open(args.cmvn_dir + "/cmvn." + str(i) + ".json", "r") as fin:
+ cmvn_stats = json.load(fin)
+
+ total_means += np.array(cmvn_stats["mean_stats"])
+ total_vars += np.array(cmvn_stats["var_stats"])
+ total_frames += cmvn_stats["total_frames"]
+
+ cmvn_info = {
+ 'mean_stats': list(total_means.tolist()),
+ 'var_stats': list(total_vars.tolist()),
+ 'total_frames': total_frames
+ }
+ with open(cmvn_file, 'w') as fout:
+ fout.write(json.dumps(cmvn_info))
+
+
+if __name__ == '__main__':
+ main()
diff --git a/egs_modelscope/speechio/paraformer/utils/compute_cmvn.py b/egs_modelscope/speechio/paraformer/utils/compute_cmvn.py
new file mode 100755
index 0000000..2b96e26
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/compute_cmvn.py
@@ -0,0 +1,74 @@
+from kaldiio import ReadHelper
+
+import argparse
+import numpy as np
+import json
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="computer global cmvn",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--dims",
+ "-d",
+ default=80,
+ type=int,
+ help="feature dims",
+ )
+ parser.add_argument(
+ "--ark-file",
+ "-a",
+ default=False,
+ required=True,
+ type=str,
+ help="fbank ark file",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ ark_file = args.ark_file + "/feats." + str(args.ark_index) + ".ark"
+ cmvn_file = args.output_dir + "/cmvn." + str(args.ark_index) + ".json"
+
+ mean_stats = np.zeros(args.dims)
+ var_stats = np.zeros(args.dims)
+ total_frames = 0
+
+ with ReadHelper('ark:{}'.format(ark_file)) as ark_reader:
+ for key, mat in ark_reader:
+ mean_stats += np.sum(mat, axis=0)
+ var_stats += np.sum(np.square(mat), axis=0)
+ total_frames += mat.shape[0]
+
+ cmvn_info = {
+ 'mean_stats': list(mean_stats.tolist()),
+ 'var_stats': list(var_stats.tolist()),
+ 'total_frames': total_frames
+ }
+ with open(cmvn_file, 'w') as fout:
+ fout.write(json.dumps(cmvn_info))
+
+
+if __name__ == '__main__':
+ main()
diff --git a/egs_modelscope/speechio/paraformer/utils/compute_cmvn.sh b/egs_modelscope/speechio/paraformer/utils/compute_cmvn.sh
new file mode 100755
index 0000000..12173ee
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/compute_cmvn.sh
@@ -0,0 +1,25 @@
+#!/usr/bin/env bash
+
+. ./path.sh || exit 1;
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+feats_dim=80
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+fbankdir=$1
+logdir=$2
+
+output_dir=${fbankdir}/cmvn; mkdir -p ${output_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/cmvn.JOB.log \
+ python utils/compute_cmvn.py -d ${feats_dim} -a $fbankdir/ark -i JOB -o ${output_dir} \
+ || exit 1;
+
+python utils/combine_cmvn_file.py -d ${feats_dim} -c ${output_dir} -n $nj -o $fbankdir
+
+echo "$0: Succeeded compute global cmvn"
diff --git a/egs_modelscope/speechio/paraformer/utils/compute_fbank.py b/egs_modelscope/speechio/paraformer/utils/compute_fbank.py
new file mode 100755
index 0000000..d03b5a8
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/compute_fbank.py
@@ -0,0 +1,153 @@
+from kaldiio import WriteHelper
+
+import argparse
+import numpy as np
+import json
+import torch
+import torchaudio
+import torchaudio.compliance.kaldi as kaldi
+
+
+def compute_fbank(wav_file,
+ num_mel_bins=80,
+ frame_length=25,
+ frame_shift=10,
+ dither=0.0,
+ resample_rate=16000,
+ speed=1.0):
+
+ waveform, sample_rate = torchaudio.load(wav_file)
+ if resample_rate != sample_rate:
+ waveform = torchaudio.transforms.Resample(orig_freq=sample_rate,
+ new_freq=resample_rate)(waveform)
+ if speed != 1.0:
+ waveform, _ = torchaudio.sox_effects.apply_effects_tensor(
+ waveform, resample_rate,
+ [['speed', str(speed)], ['rate', str(resample_rate)]]
+ )
+
+ waveform = waveform * (1 << 15)
+ mat = kaldi.fbank(waveform,
+ num_mel_bins=num_mel_bins,
+ frame_length=frame_length,
+ frame_shift=frame_shift,
+ dither=dither,
+ energy_floor=0.0,
+ window_type='hamming',
+ sample_frequency=resample_rate)
+
+ return mat.numpy()
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="computer features",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--wav-lists",
+ "-w",
+ default=False,
+ required=True,
+ type=str,
+ help="input wav lists",
+ )
+ parser.add_argument(
+ "--text-files",
+ "-t",
+ default=False,
+ required=True,
+ type=str,
+ help="input text files",
+ )
+ parser.add_argument(
+ "--dims",
+ "-d",
+ default=80,
+ type=int,
+ help="feature dims",
+ )
+ parser.add_argument(
+ "--sample-frequency",
+ "-s",
+ default=16000,
+ type=int,
+ help="sample frequency",
+ )
+ parser.add_argument(
+ "--speed-perturb",
+ "-p",
+ default="1.0",
+ type=str,
+ help="speed perturb",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-a",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ ark_file = args.output_dir + "/ark/feats." + str(args.ark_index) + ".ark"
+ scp_file = args.output_dir + "/ark/feats." + str(args.ark_index) + ".scp"
+ text_file = args.output_dir + "/txt/text." + str(args.ark_index) + ".txt"
+ feats_shape_file = args.output_dir + "/ark/len." + str(args.ark_index)
+ text_shape_file = args.output_dir + "/txt/len." + str(args.ark_index)
+
+ ark_writer = WriteHelper('ark,scp:{},{}'.format(ark_file, scp_file))
+ text_writer = open(text_file, 'w')
+ feats_shape_writer = open(feats_shape_file, 'w')
+ text_shape_writer = open(text_shape_file, 'w')
+
+ speed_perturb_list = args.speed_perturb.split(',')
+
+ for speed in speed_perturb_list:
+ with open(args.wav_lists, 'r', encoding='utf-8') as wavfile:
+ with open(args.text_files, 'r', encoding='utf-8') as textfile:
+ for wav, text in zip(wavfile, textfile):
+ s_w = wav.strip().split()
+ wav_id = s_w[0]
+ wav_file = s_w[1]
+
+ s_t = text.strip().split()
+ text_id = s_t[0]
+ txt = s_t[1:]
+ fbank = compute_fbank(wav_file,
+ num_mel_bins=args.dims,
+ resample_rate=args.sample_frequency,
+ speed=float(speed)
+ )
+ feats_dims = fbank.shape[1]
+ feats_lens = fbank.shape[0]
+ txt_lens = len(txt)
+ if speed == "1.0":
+ wav_id_sp = wav_id
+ else:
+ wav_id_sp = wav_id + "_sp" + speed
+
+ feats_shape_writer.write(wav_id_sp + " " + str(feats_lens) + "," + str(feats_dims) + '\n')
+ text_shape_writer.write(wav_id_sp + " " + str(txt_lens) + '\n')
+
+ text_writer.write(wav_id_sp + " " + " ".join(txt) + '\n')
+ ark_writer(wav_id_sp, fbank)
+
+
+if __name__ == '__main__':
+ main()
+
diff --git a/egs_modelscope/speechio/paraformer/utils/compute_fbank.sh b/egs_modelscope/speechio/paraformer/utils/compute_fbank.sh
new file mode 100755
index 0000000..92a4fe6
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/compute_fbank.sh
@@ -0,0 +1,51 @@
+#!/usr/bin/env bash
+
+. ./path.sh || exit 1;
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+
+# feature configuration
+feats_dim=80
+sample_frequency=16000
+speed_perturb="1.0"
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+data=$1
+logdir=$2
+fbankdir=$3
+
+[ ! -f $data/wav.scp ] && echo "$0: no such file $data/wav.scp" && exit 1;
+[ ! -f $data/text ] && echo "$0: no such file $data/text" && exit 1;
+
+python utils/split_data.py $data $data $nj
+
+ark_dir=${fbankdir}/ark; mkdir -p ${ark_dir}
+text_dir=${fbankdir}/txt; mkdir -p ${text_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/make_fbank.JOB.log \
+ python utils/compute_fbank.py -w $data/split${nj}/JOB/wav.scp -t $data/split${nj}/JOB/text \
+ -d $feats_dim -s $sample_frequency -p ${speed_perturb} -a JOB -o ${fbankdir} \
+ || exit 1;
+
+for n in $(seq $nj); do
+ cat ${ark_dir}/feats.$n.scp || exit 1
+done > $fbankdir/feats.scp || exit 1
+
+for n in $(seq $nj); do
+ cat ${text_dir}/text.$n.txt || exit 1
+done > $fbankdir/text || exit 1
+
+for n in $(seq $nj); do
+ cat ${ark_dir}/len.$n || exit 1
+done > $fbankdir/speech_shape || exit 1
+
+for n in $(seq $nj); do
+ cat ${text_dir}/len.$n || exit 1
+done > $fbankdir/text_shape || exit 1
+
+echo "$0: Succeeded compute FBANK features"
diff --git a/egs_modelscope/speechio/paraformer/utils/compute_wer.py b/egs_modelscope/speechio/paraformer/utils/compute_wer.py
new file mode 100755
index 0000000..349a3f6
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/compute_wer.py
@@ -0,0 +1,157 @@
+import os
+import numpy as np
+import sys
+
+def compute_wer(ref_file,
+ hyp_file,
+ cer_detail_file):
+ rst = {
+ 'Wrd': 0,
+ 'Corr': 0,
+ 'Ins': 0,
+ 'Del': 0,
+ 'Sub': 0,
+ 'Snt': 0,
+ 'Err': 0.0,
+ 'S.Err': 0.0,
+ 'wrong_words': 0,
+ 'wrong_sentences': 0
+ }
+
+ hyp_dict = {}
+ ref_dict = {}
+ with open(hyp_file, 'r') as hyp_reader:
+ for line in hyp_reader:
+ key = line.strip().split()[0]
+ value = line.strip().split()[1:]
+ hyp_dict[key] = value
+ with open(ref_file, 'r') as ref_reader:
+ for line in ref_reader:
+ key = line.strip().split()[0]
+ value = line.strip().split()[1:]
+ ref_dict[key] = value
+
+ cer_detail_writer = open(cer_detail_file, 'w')
+ for hyp_key in hyp_dict:
+ if hyp_key in ref_dict:
+ out_item = compute_wer_by_line(hyp_dict[hyp_key], ref_dict[hyp_key])
+ rst['Wrd'] += out_item['nwords']
+ rst['Corr'] += out_item['cor']
+ rst['wrong_words'] += out_item['wrong']
+ rst['Ins'] += out_item['ins']
+ rst['Del'] += out_item['del']
+ rst['Sub'] += out_item['sub']
+ rst['Snt'] += 1
+ if out_item['wrong'] > 0:
+ rst['wrong_sentences'] += 1
+ cer_detail_writer.write(hyp_key + print_cer_detail(out_item) + '\n')
+ cer_detail_writer.write("ref:" + '\t' + "".join(ref_dict[hyp_key]) + '\n')
+ cer_detail_writer.write("hyp:" + '\t' + "".join(hyp_dict[hyp_key]) + '\n')
+
+ if rst['Wrd'] > 0:
+ rst['Err'] = round(rst['wrong_words'] * 100 / rst['Wrd'], 2)
+ if rst['Snt'] > 0:
+ rst['S.Err'] = round(rst['wrong_sentences'] * 100 / rst['Snt'], 2)
+
+ cer_detail_writer.write('\n')
+ cer_detail_writer.write("%WER " + str(rst['Err']) + " [ " + str(rst['wrong_words'])+ " / " + str(rst['Wrd']) +
+ ", " + str(rst['Ins']) + " ins, " + str(rst['Del']) + " del, " + str(rst['Sub']) + " sub ]" + '\n')
+ cer_detail_writer.write("%SER " + str(rst['S.Err']) + " [ " + str(rst['wrong_sentences']) + " / " + str(rst['Snt']) + " ]" + '\n')
+ cer_detail_writer.write("Scored " + str(len(hyp_dict)) + " sentences, " + str(len(hyp_dict) - rst['Snt']) + " not present in hyp." + '\n')
+
+
+def compute_wer_by_line(hyp,
+ ref):
+ hyp = list(map(lambda x: x.lower(), hyp))
+ ref = list(map(lambda x: x.lower(), ref))
+
+ len_hyp = len(hyp)
+ len_ref = len(ref)
+
+ cost_matrix = np.zeros((len_hyp + 1, len_ref + 1), dtype=np.int16)
+
+ ops_matrix = np.zeros((len_hyp + 1, len_ref + 1), dtype=np.int8)
+
+ for i in range(len_hyp + 1):
+ cost_matrix[i][0] = i
+ for j in range(len_ref + 1):
+ cost_matrix[0][j] = j
+
+ for i in range(1, len_hyp + 1):
+ for j in range(1, len_ref + 1):
+ if hyp[i - 1] == ref[j - 1]:
+ cost_matrix[i][j] = cost_matrix[i - 1][j - 1]
+ else:
+ substitution = cost_matrix[i - 1][j - 1] + 1
+ insertion = cost_matrix[i - 1][j] + 1
+ deletion = cost_matrix[i][j - 1] + 1
+
+ compare_val = [substitution, insertion, deletion]
+
+ min_val = min(compare_val)
+ operation_idx = compare_val.index(min_val) + 1
+ cost_matrix[i][j] = min_val
+ ops_matrix[i][j] = operation_idx
+
+ match_idx = []
+ i = len_hyp
+ j = len_ref
+ rst = {
+ 'nwords': len_ref,
+ 'cor': 0,
+ 'wrong': 0,
+ 'ins': 0,
+ 'del': 0,
+ 'sub': 0
+ }
+ while i >= 0 or j >= 0:
+ i_idx = max(0, i)
+ j_idx = max(0, j)
+
+ if ops_matrix[i_idx][j_idx] == 0: # correct
+ if i - 1 >= 0 and j - 1 >= 0:
+ match_idx.append((j - 1, i - 1))
+ rst['cor'] += 1
+
+ i -= 1
+ j -= 1
+
+ elif ops_matrix[i_idx][j_idx] == 2: # insert
+ i -= 1
+ rst['ins'] += 1
+
+ elif ops_matrix[i_idx][j_idx] == 3: # delete
+ j -= 1
+ rst['del'] += 1
+
+ elif ops_matrix[i_idx][j_idx] == 1: # substitute
+ i -= 1
+ j -= 1
+ rst['sub'] += 1
+
+ if i < 0 and j >= 0:
+ rst['del'] += 1
+ elif j < 0 and i >= 0:
+ rst['ins'] += 1
+
+ match_idx.reverse()
+ wrong_cnt = cost_matrix[len_hyp][len_ref]
+ rst['wrong'] = wrong_cnt
+
+ return rst
+
+def print_cer_detail(rst):
+ return ("(" + "nwords=" + str(rst['nwords']) + ",cor=" + str(rst['cor'])
+ + ",ins=" + str(rst['ins']) + ",del=" + str(rst['del']) + ",sub="
+ + str(rst['sub']) + ") corr:" + '{:.2%}'.format(rst['cor']/rst['nwords'])
+ + ",cer:" + '{:.2%}'.format(rst['wrong']/rst['nwords']))
+
+if __name__ == '__main__':
+ if len(sys.argv) != 4:
+ print("usage : python compute-wer.py test.ref test.hyp test.wer")
+ sys.exit(0)
+
+ ref_file = sys.argv[1]
+ hyp_file = sys.argv[2]
+ cer_detail_file = sys.argv[3]
+ compute_wer(ref_file, hyp_file, cer_detail_file)
diff --git a/egs_modelscope/speechio/paraformer/utils/error_rate_zh b/egs_modelscope/speechio/paraformer/utils/error_rate_zh
new file mode 100755
index 0000000..6871a07
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/error_rate_zh
@@ -0,0 +1,370 @@
+#!/usr/bin/env python3
+# coding=utf8
+
+# Copyright 2021 Jiayu DU
+
+import sys
+import argparse
+import json
+import logging
+logging.basicConfig(stream=sys.stderr, level=logging.INFO, format='[%(levelname)s] %(message)s')
+
+DEBUG = None
+
+def GetEditType(ref_token, hyp_token):
+ if ref_token == None and hyp_token != None:
+ return 'I'
+ elif ref_token != None and hyp_token == None:
+ return 'D'
+ elif ref_token == hyp_token:
+ return 'C'
+ elif ref_token != hyp_token:
+ return 'S'
+ else:
+ raise RuntimeError
+
+class AlignmentArc:
+ def __init__(self, src, dst, ref, hyp):
+ self.src = src
+ self.dst = dst
+ self.ref = ref
+ self.hyp = hyp
+ self.edit_type = GetEditType(ref, hyp)
+
+def similarity_score_function(ref_token, hyp_token):
+ return 0 if (ref_token == hyp_token) else -1.0
+
+def insertion_score_function(token):
+ return -1.0
+
+def deletion_score_function(token):
+ return -1.0
+
+def EditDistance(
+ ref,
+ hyp,
+ similarity_score_function = similarity_score_function,
+ insertion_score_function = insertion_score_function,
+ deletion_score_function = deletion_score_function):
+ assert(len(ref) != 0)
+ class DPState:
+ def __init__(self):
+ self.score = -float('inf')
+ # backpointer
+ self.prev_r = None
+ self.prev_h = None
+
+ def print_search_grid(S, R, H, fstream):
+ print(file=fstream)
+ for r in range(R):
+ for h in range(H):
+ print(F'[{r},{h}]:{S[r][h].score:4.3f}:({S[r][h].prev_r},{S[r][h].prev_h}) ', end='', file=fstream)
+ print(file=fstream)
+
+ R = len(ref) + 1
+ H = len(hyp) + 1
+
+ # Construct DP search space, a (R x H) grid
+ S = [ [] for r in range(R) ]
+ for r in range(R):
+ S[r] = [ DPState() for x in range(H) ]
+
+ # initialize DP search grid origin, S(r = 0, h = 0)
+ S[0][0].score = 0.0
+ S[0][0].prev_r = None
+ S[0][0].prev_h = None
+
+ # initialize REF axis
+ for r in range(1, R):
+ S[r][0].score = S[r-1][0].score + deletion_score_function(ref[r-1])
+ S[r][0].prev_r = r-1
+ S[r][0].prev_h = 0
+
+ # initialize HYP axis
+ for h in range(1, H):
+ S[0][h].score = S[0][h-1].score + insertion_score_function(hyp[h-1])
+ S[0][h].prev_r = 0
+ S[0][h].prev_h = h-1
+
+ best_score = S[0][0].score
+ best_state = (0, 0)
+
+ for r in range(1, R):
+ for h in range(1, H):
+ sub_or_cor_score = similarity_score_function(ref[r-1], hyp[h-1])
+ new_score = S[r-1][h-1].score + sub_or_cor_score
+ if new_score >= S[r][h].score:
+ S[r][h].score = new_score
+ S[r][h].prev_r = r-1
+ S[r][h].prev_h = h-1
+
+ del_score = deletion_score_function(ref[r-1])
+ new_score = S[r-1][h].score + del_score
+ if new_score >= S[r][h].score:
+ S[r][h].score = new_score
+ S[r][h].prev_r = r - 1
+ S[r][h].prev_h = h
+
+ ins_score = insertion_score_function(hyp[h-1])
+ new_score = S[r][h-1].score + ins_score
+ if new_score >= S[r][h].score:
+ S[r][h].score = new_score
+ S[r][h].prev_r = r
+ S[r][h].prev_h = h-1
+
+ best_score = S[R-1][H-1].score
+ best_state = (R-1, H-1)
+
+ if DEBUG:
+ print_search_grid(S, R, H, sys.stderr)
+
+ # Backtracing best alignment path, i.e. a list of arcs
+ # arc = (src, dst, ref, hyp, edit_type)
+ # src/dst = (r, h), where r/h refers to search grid state-id along Ref/Hyp axis
+ best_path = []
+ r, h = best_state[0], best_state[1]
+ prev_r, prev_h = S[r][h].prev_r, S[r][h].prev_h
+ score = S[r][h].score
+ # loop invariant:
+ # 1. (prev_r, prev_h) -> (r, h) is a "forward arc" on best alignment path
+ # 2. score is the value of point(r, h) on DP search grid
+ while prev_r != None or prev_h != None:
+ src = (prev_r, prev_h)
+ dst = (r, h)
+ if (r == prev_r + 1 and h == prev_h + 1): # Substitution or correct
+ arc = AlignmentArc(src, dst, ref[prev_r], hyp[prev_h])
+ elif (r == prev_r + 1 and h == prev_h): # Deletion
+ arc = AlignmentArc(src, dst, ref[prev_r], None)
+ elif (r == prev_r and h == prev_h + 1): # Insertion
+ arc = AlignmentArc(src, dst, None, hyp[prev_h])
+ else:
+ raise RuntimeError
+ best_path.append(arc)
+ r, h = prev_r, prev_h
+ prev_r, prev_h = S[r][h].prev_r, S[r][h].prev_h
+ score = S[r][h].score
+
+ best_path.reverse()
+ return (best_path, best_score)
+
+def PrettyPrintAlignment(alignment, stream = sys.stderr):
+ def get_token_str(token):
+ if token == None:
+ return "*"
+ return token
+
+ def is_double_width_char(ch):
+ if (ch >= '\u4e00') and (ch <= '\u9fa5'): # codepoint ranges for Chinese chars
+ return True
+ # TODO: support other double-width-char language such as Japanese, Korean
+ else:
+ return False
+
+ def display_width(token_str):
+ m = 0
+ for c in token_str:
+ if is_double_width_char(c):
+ m += 2
+ else:
+ m += 1
+ return m
+
+ R = ' REF : '
+ H = ' HYP : '
+ E = ' EDIT : '
+ for arc in alignment:
+ r = get_token_str(arc.ref)
+ h = get_token_str(arc.hyp)
+ e = arc.edit_type if arc.edit_type != 'C' else ''
+
+ nr, nh, ne = display_width(r), display_width(h), display_width(e)
+ n = max(nr, nh, ne) + 1
+
+ R += r + ' ' * (n-nr)
+ H += h + ' ' * (n-nh)
+ E += e + ' ' * (n-ne)
+
+ print(R, file=stream)
+ print(H, file=stream)
+ print(E, file=stream)
+
+def CountEdits(alignment):
+ c, s, i, d = 0, 0, 0, 0
+ for arc in alignment:
+ if arc.edit_type == 'C':
+ c += 1
+ elif arc.edit_type == 'S':
+ s += 1
+ elif arc.edit_type == 'I':
+ i += 1
+ elif arc.edit_type == 'D':
+ d += 1
+ else:
+ raise RuntimeError
+ return (c, s, i, d)
+
+def ComputeTokenErrorRate(c, s, i, d):
+ return 100.0 * (s + d + i) / (s + d + c)
+
+def ComputeSentenceErrorRate(num_err_utts, num_utts):
+ assert(num_utts != 0)
+ return 100.0 * num_err_utts / num_utts
+
+
+class EvaluationResult:
+ def __init__(self):
+ self.num_ref_utts = 0
+ self.num_hyp_utts = 0
+ self.num_eval_utts = 0 # seen in both ref & hyp
+ self.num_hyp_without_ref = 0
+
+ self.C = 0
+ self.S = 0
+ self.I = 0
+ self.D = 0
+ self.token_error_rate = 0.0
+
+ self.num_utts_with_error = 0
+ self.sentence_error_rate = 0.0
+
+ def to_json(self):
+ return json.dumps(self.__dict__)
+
+ def to_kaldi(self):
+ info = (
+ F'%WER {self.token_error_rate:.2f} [ {self.S + self.D + self.I} / {self.C + self.S + self.D}, {self.I} ins, {self.D} del, {self.S} sub ]\n'
+ F'%SER {self.sentence_error_rate:.2f} [ {self.num_utts_with_error} / {self.num_eval_utts} ]\n'
+ )
+ return info
+
+ def to_sclite(self):
+ return "TODO"
+
+ def to_espnet(self):
+ return "TODO"
+
+ def to_summary(self):
+ #return json.dumps(self.__dict__, indent=4)
+ summary = (
+ '==================== Overall Statistics ====================\n'
+ F'num_ref_utts: {self.num_ref_utts}\n'
+ F'num_hyp_utts: {self.num_hyp_utts}\n'
+ F'num_hyp_without_ref: {self.num_hyp_without_ref}\n'
+ F'num_eval_utts: {self.num_eval_utts}\n'
+ F'sentence_error_rate: {self.sentence_error_rate:.2f}%\n'
+ F'token_error_rate: {self.token_error_rate:.2f}%\n'
+ F'token_stats:\n'
+ F' - tokens:{self.C + self.S + self.D:>7}\n'
+ F' - edits: {self.S + self.I + self.D:>7}\n'
+ F' - cor: {self.C:>7}\n'
+ F' - sub: {self.S:>7}\n'
+ F' - ins: {self.I:>7}\n'
+ F' - del: {self.D:>7}\n'
+ '============================================================\n'
+ )
+ return summary
+
+
+class Utterance:
+ def __init__(self, uid, text):
+ self.uid = uid
+ self.text = text
+
+
+def LoadUtterances(filepath, format):
+ utts = {}
+ if format == 'text': # utt_id word1 word2 ...
+ with open(filepath, 'r', encoding='utf8') as f:
+ for line in f:
+ line = line.strip()
+ if line:
+ cols = line.split(maxsplit=1)
+ assert(len(cols) == 2 or len(cols) == 1)
+ uid = cols[0]
+ text = cols[1] if len(cols) == 2 else ''
+ if utts.get(uid) != None:
+ raise RuntimeError(F'Found duplicated utterence id {uid}')
+ utts[uid] = Utterance(uid, text)
+ else:
+ raise RuntimeError(F'Unsupported text format {format}')
+ return utts
+
+
+def tokenize_text(text, tokenizer):
+ if tokenizer == 'whitespace':
+ return text.split()
+ elif tokenizer == 'char':
+ return [ ch for ch in ''.join(text.split()) ]
+ else:
+ raise RuntimeError(F'ERROR: Unsupported tokenizer {tokenizer}')
+
+
+if __name__ == '__main__':
+ parser = argparse.ArgumentParser()
+ # optional
+ parser.add_argument('--tokenizer', choices=['whitespace', 'char'], default='whitespace', help='whitespace for WER, char for CER')
+ parser.add_argument('--ref-format', choices=['text'], default='text', help='reference format, first col is utt_id, the rest is text')
+ parser.add_argument('--hyp-format', choices=['text'], default='text', help='hypothesis format, first col is utt_id, the rest is text')
+ # required
+ parser.add_argument('--ref', type=str, required=True, help='input reference file')
+ parser.add_argument('--hyp', type=str, required=True, help='input hypothesis file')
+
+ parser.add_argument('result_file', type=str)
+ args = parser.parse_args()
+ logging.info(args)
+
+ ref_utts = LoadUtterances(args.ref, args.ref_format)
+ hyp_utts = LoadUtterances(args.hyp, args.hyp_format)
+
+ r = EvaluationResult()
+
+ # check valid utterances in hyp that have matched non-empty reference
+ eval_utts = []
+ r.num_hyp_without_ref = 0
+ for uid in sorted(hyp_utts.keys()):
+ if uid in ref_utts.keys(): # TODO: efficiency
+ if ref_utts[uid].text.strip(): # non-empty reference
+ eval_utts.append(uid)
+ else:
+ logging.warn(F'Found {uid} with empty reference, skipping...')
+ else:
+ logging.warn(F'Found {uid} without reference, skipping...')
+ r.num_hyp_without_ref += 1
+
+ r.num_hyp_utts = len(hyp_utts)
+ r.num_ref_utts = len(ref_utts)
+ r.num_eval_utts = len(eval_utts)
+
+ with open(args.result_file, 'w+', encoding='utf8') as fo:
+ for uid in eval_utts:
+ ref = ref_utts[uid]
+ hyp = hyp_utts[uid]
+
+ alignment, score = EditDistance(
+ tokenize_text(ref.text, args.tokenizer),
+ tokenize_text(hyp.text, args.tokenizer)
+ )
+
+ c, s, i, d = CountEdits(alignment)
+ utt_ter = ComputeTokenErrorRate(c, s, i, d)
+
+ # utt-level evaluation result
+ print(F'{{"uid":{uid}, "score":{score}, "ter":{utt_ter:.2f}, "cor":{c}, "sub":{s}, "ins":{i}, "del":{d}}}', file=fo)
+ PrettyPrintAlignment(alignment, fo)
+
+ r.C += c
+ r.S += s
+ r.I += i
+ r.D += d
+
+ if utt_ter > 0:
+ r.num_utts_with_error += 1
+
+ # corpus level evaluation result
+ r.sentence_error_rate = ComputeSentenceErrorRate(r.num_utts_with_error, r.num_eval_utts)
+ r.token_error_rate = ComputeTokenErrorRate(r.C, r.S, r.I, r.D)
+
+ print(r.to_summary(), file=fo)
+
+ print(r.to_json())
+ print(r.to_kaldi())
diff --git a/egs_modelscope/speechio/paraformer/utils/extract_embeds.py b/egs_modelscope/speechio/paraformer/utils/extract_embeds.py
new file mode 100755
index 0000000..7b817d8
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/extract_embeds.py
@@ -0,0 +1,47 @@
+from transformers import AutoTokenizer, AutoModel, pipeline
+import numpy as np
+import sys
+import os
+import torch
+from kaldiio import WriteHelper
+import re
+text_file_json = sys.argv[1]
+out_ark = sys.argv[2]
+out_scp = sys.argv[3]
+out_shape = sys.argv[4]
+device = int(sys.argv[5])
+model_path = sys.argv[6]
+
+model = AutoModel.from_pretrained(model_path)
+tokenizer = AutoTokenizer.from_pretrained(model_path)
+extractor = pipeline(task="feature-extraction", model=model, tokenizer=tokenizer, device=device)
+
+with open(text_file_json, 'r') as f:
+ js = f.readlines()
+
+
+f_shape = open(out_shape, "w")
+with WriteHelper('ark,scp:{},{}'.format(out_ark, out_scp)) as writer:
+ with torch.no_grad():
+ for idx, line in enumerate(js):
+ id, tokens = line.strip().split(" ", 1)
+ tokens = re.sub(" ", "", tokens.strip())
+ tokens = ' '.join([j for j in tokens])
+ token_num = len(tokens.split(" "))
+ outputs = extractor(tokens)
+ outputs = np.array(outputs)
+ embeds = outputs[0, 1:-1, :]
+
+ token_num_embeds, dim = embeds.shape
+ if token_num == token_num_embeds:
+ writer(id, embeds)
+ shape_line = "{} {},{}\n".format(id, token_num_embeds, dim)
+ f_shape.write(shape_line)
+ else:
+ print("{}, size has changed, {}, {}, {}".format(id, token_num, token_num_embeds, tokens))
+
+
+
+f_shape.close()
+
+
diff --git a/egs_modelscope/speechio/paraformer/utils/filter_scp.pl b/egs_modelscope/speechio/paraformer/utils/filter_scp.pl
new file mode 100755
index 0000000..003530d
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/filter_scp.pl
@@ -0,0 +1,87 @@
+#!/usr/bin/env perl
+# Copyright 2010-2012 Microsoft Corporation
+# Johns Hopkins University (author: Daniel Povey)
+
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+# This script takes a list of utterance-ids or any file whose first field
+# of each line is an utterance-id, and filters an scp
+# file (or any file whose "n-th" field is an utterance id), printing
+# out only those lines whose "n-th" field is in id_list. The index of
+# the "n-th" field is 1, by default, but can be changed by using
+# the -f <n> switch
+
+$exclude = 0;
+$field = 1;
+$shifted = 0;
+
+do {
+ $shifted=0;
+ if ($ARGV[0] eq "--exclude") {
+ $exclude = 1;
+ shift @ARGV;
+ $shifted=1;
+ }
+ if ($ARGV[0] eq "-f") {
+ $field = $ARGV[1];
+ shift @ARGV; shift @ARGV;
+ $shifted=1
+ }
+} while ($shifted);
+
+if(@ARGV < 1 || @ARGV > 2) {
+ die "Usage: filter_scp.pl [--exclude] [-f <field-to-filter-on>] id_list [in.scp] > out.scp \n" .
+ "Prints only the input lines whose f'th field (default: first) is in 'id_list'.\n" .
+ "Note: only the first field of each line in id_list matters. With --exclude, prints\n" .
+ "only the lines that were *not* in id_list.\n" .
+ "Caution: previously, the -f option was interpreted as a zero-based field index.\n" .
+ "If your older scripts (written before Oct 2014) stopped working and you used the\n" .
+ "-f option, add 1 to the argument.\n" .
+ "See also: scripts/filter_scp.pl .\n";
+}
+
+
+$idlist = shift @ARGV;
+open(F, "<$idlist") || die "Could not open id-list file $idlist";
+while(<F>) {
+ @A = split;
+ @A>=1 || die "Invalid id-list file line $_";
+ $seen{$A[0]} = 1;
+}
+
+if ($field == 1) { # Treat this as special case, since it is common.
+ while(<>) {
+ $_ =~ m/\s*(\S+)\s*/ || die "Bad line $_, could not get first field.";
+ # $1 is what we filter on.
+ if ((!$exclude && $seen{$1}) || ($exclude && !defined $seen{$1})) {
+ print $_;
+ }
+ }
+} else {
+ while(<>) {
+ @A = split;
+ @A > 0 || die "Invalid scp file line $_";
+ @A >= $field || die "Invalid scp file line $_";
+ if ((!$exclude && $seen{$A[$field-1]}) || ($exclude && !defined $seen{$A[$field-1]})) {
+ print $_;
+ }
+ }
+}
+
+# tests:
+# the following should print "foo 1"
+# ( echo foo 1; echo bar 2 ) | scripts/filter_scp.pl <(echo foo)
+# the following should print "bar 2".
+# ( echo foo 1; echo bar 2 ) | scripts/filter_scp.pl -f 2 <(echo 2)
diff --git a/egs_modelscope/speechio/paraformer/utils/fix_data.sh b/egs_modelscope/speechio/paraformer/utils/fix_data.sh
new file mode 100755
index 0000000..32cdde5
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/fix_data.sh
@@ -0,0 +1,35 @@
+#!/usr/bin/env bash
+
+echo "$0 $@"
+data_dir=$1
+
+if [ ! -f ${data_dir}/wav.scp ]; then
+ echo "$0: wav.scp is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/text ]; then
+ echo "$0: text is not found"
+ exit 1;
+fi
+
+
+
+mkdir -p ${data_dir}/.backup
+
+awk '{print $1}' ${data_dir}/wav.scp > ${data_dir}/.backup/wav_id
+awk '{print $1}' ${data_dir}/text > ${data_dir}/.backup/text_id
+
+sort ${data_dir}/.backup/wav_id ${data_dir}/.backup/text_id | uniq -d > ${data_dir}/.backup/id
+
+cp ${data_dir}/wav.scp ${data_dir}/.backup/wav.scp
+cp ${data_dir}/text ${data_dir}/.backup/text
+
+mv ${data_dir}/wav.scp ${data_dir}/wav.scp.bak
+mv ${data_dir}/text ${data_dir}/text.bak
+
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/wav.scp.bak > ${data_dir}/wav.scp
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/text.bak > ${data_dir}/text
+
+rm ${data_dir}/wav.scp.bak
+rm ${data_dir}/text.bak
diff --git a/egs_modelscope/speechio/paraformer/utils/fix_data_feat.sh b/egs_modelscope/speechio/paraformer/utils/fix_data_feat.sh
new file mode 100755
index 0000000..2c92d7f
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/fix_data_feat.sh
@@ -0,0 +1,52 @@
+#!/usr/bin/env bash
+
+echo "$0 $@"
+data_dir=$1
+
+if [ ! -f ${data_dir}/feats.scp ]; then
+ echo "$0: feats.scp is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/text ]; then
+ echo "$0: text is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/speech_shape ]; then
+ echo "$0: feature lengths is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/text_shape ]; then
+ echo "$0: text lengths is not found"
+ exit 1;
+fi
+
+mkdir -p ${data_dir}/.backup
+
+awk '{print $1}' ${data_dir}/feats.scp > ${data_dir}/.backup/wav_id
+awk '{print $1}' ${data_dir}/text > ${data_dir}/.backup/text_id
+
+sort ${data_dir}/.backup/wav_id ${data_dir}/.backup/text_id | uniq -d > ${data_dir}/.backup/id
+
+cp ${data_dir}/feats.scp ${data_dir}/.backup/feats.scp
+cp ${data_dir}/text ${data_dir}/.backup/text
+cp ${data_dir}/speech_shape ${data_dir}/.backup/speech_shape
+cp ${data_dir}/text_shape ${data_dir}/.backup/text_shape
+
+mv ${data_dir}/feats.scp ${data_dir}/feats.scp.bak
+mv ${data_dir}/text ${data_dir}/text.bak
+mv ${data_dir}/speech_shape ${data_dir}/speech_shape.bak
+mv ${data_dir}/text_shape ${data_dir}/text_shape.bak
+
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/feats.scp.bak > ${data_dir}/feats.scp
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/text.bak > ${data_dir}/text
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/speech_shape.bak > ${data_dir}/speech_shape
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/text_shape.bak > ${data_dir}/text_shape
+
+rm ${data_dir}/feats.scp.bak
+rm ${data_dir}/text.bak
+rm ${data_dir}/speech_shape.bak
+rm ${data_dir}/text_shape.bak
+
diff --git a/egs_modelscope/speechio/paraformer/utils/gen_ark_list.sh b/egs_modelscope/speechio/paraformer/utils/gen_ark_list.sh
new file mode 100755
index 0000000..aebf356
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/gen_ark_list.sh
@@ -0,0 +1,22 @@
+#!/usr/bin/env bash
+
+
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+ark_dir=$1
+txt_dir=$2
+output_dir=$3
+
+[ ! -d ${ark_dir}/ark ] && echo "$0: ark data is required" && exit 1;
+[ ! -d ${txt_dir}/txt ] && echo "$0: txt data is required" && exit 1;
+
+for n in $(seq $nj); do
+ echo "${ark_dir}/ark/feats.$n.ark ${txt_dir}/txt/text.$n.txt" || exit 1
+done > ${output_dir}/ark_txt.scp || exit 1
+
diff --git a/egs_modelscope/speechio/paraformer/utils/parse_options.sh b/egs_modelscope/speechio/paraformer/utils/parse_options.sh
new file mode 100755
index 0000000..71fb9e5
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/parse_options.sh
@@ -0,0 +1,97 @@
+#!/usr/bin/env bash
+
+# Copyright 2012 Johns Hopkins University (Author: Daniel Povey);
+# Arnab Ghoshal, Karel Vesely
+
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+# Parse command-line options.
+# To be sourced by another script (as in ". parse_options.sh").
+# Option format is: --option-name arg
+# and shell variable "option_name" gets set to value "arg."
+# The exception is --help, which takes no arguments, but prints the
+# $help_message variable (if defined).
+
+
+###
+### The --config file options have lower priority to command line
+### options, so we need to import them first...
+###
+
+# Now import all the configs specified by command-line, in left-to-right order
+for ((argpos=1; argpos<$#; argpos++)); do
+ if [ "${!argpos}" == "--config" ]; then
+ argpos_plus1=$((argpos+1))
+ config=${!argpos_plus1}
+ [ ! -r $config ] && echo "$0: missing config '$config'" && exit 1
+ . $config # source the config file.
+ fi
+done
+
+
+###
+### Now we process the command line options
+###
+while true; do
+ [ -z "${1:-}" ] && break; # break if there are no arguments
+ case "$1" in
+ # If the enclosing script is called with --help option, print the help
+ # message and exit. Scripts should put help messages in $help_message
+ --help|-h) if [ -z "$help_message" ]; then echo "No help found." 1>&2;
+ else printf "$help_message\n" 1>&2 ; fi;
+ exit 0 ;;
+ --*=*) echo "$0: options to scripts must be of the form --name value, got '$1'"
+ exit 1 ;;
+ # If the first command-line argument begins with "--" (e.g. --foo-bar),
+ # then work out the variable name as $name, which will equal "foo_bar".
+ --*) name=`echo "$1" | sed s/^--// | sed s/-/_/g`;
+ # Next we test whether the variable in question is undefned-- if so it's
+ # an invalid option and we die. Note: $0 evaluates to the name of the
+ # enclosing script.
+ # The test [ -z ${foo_bar+xxx} ] will return true if the variable foo_bar
+ # is undefined. We then have to wrap this test inside "eval" because
+ # foo_bar is itself inside a variable ($name).
+ eval '[ -z "${'$name'+xxx}" ]' && echo "$0: invalid option $1" 1>&2 && exit 1;
+
+ oldval="`eval echo \\$$name`";
+ # Work out whether we seem to be expecting a Boolean argument.
+ if [ "$oldval" == "true" ] || [ "$oldval" == "false" ]; then
+ was_bool=true;
+ else
+ was_bool=false;
+ fi
+
+ # Set the variable to the right value-- the escaped quotes make it work if
+ # the option had spaces, like --cmd "queue.pl -sync y"
+ eval $name=\"$2\";
+
+ # Check that Boolean-valued arguments are really Boolean.
+ if $was_bool && [[ "$2" != "true" && "$2" != "false" ]]; then
+ echo "$0: expected \"true\" or \"false\": $1 $2" 1>&2
+ exit 1;
+ fi
+ shift 2;
+ ;;
+ *) break;
+ esac
+done
+
+
+# Check for an empty argument to the --cmd option, which can easily occur as a
+# result of scripting errors.
+[ ! -z "${cmd+xxx}" ] && [ -z "$cmd" ] && echo "$0: empty argument to --cmd option" 1>&2 && exit 1;
+
+
+true; # so this script returns exit code 0.
diff --git a/egs_modelscope/speechio/paraformer/utils/print_args.py b/egs_modelscope/speechio/paraformer/utils/print_args.py
new file mode 100755
index 0000000..b0c61e5
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/print_args.py
@@ -0,0 +1,45 @@
+#!/usr/bin/env python
+import sys
+
+
+def get_commandline_args(no_executable=True):
+ extra_chars = [
+ " ",
+ ";",
+ "&",
+ "|",
+ "<",
+ ">",
+ "?",
+ "*",
+ "~",
+ "`",
+ '"',
+ "'",
+ "\\",
+ "{",
+ "}",
+ "(",
+ ")",
+ ]
+
+ # Escape the extra characters for shell
+ argv = [
+ arg.replace("'", "'\\''")
+ if all(char not in arg for char in extra_chars)
+ else "'" + arg.replace("'", "'\\''") + "'"
+ for arg in sys.argv
+ ]
+
+ if no_executable:
+ return " ".join(argv[1:])
+ else:
+ return sys.executable + " " + " ".join(argv)
+
+
+def main():
+ print(get_commandline_args())
+
+
+if __name__ == "__main__":
+ main()
diff --git a/egs_modelscope/speechio/paraformer/utils/proc_conf_oss.py b/egs_modelscope/speechio/paraformer/utils/proc_conf_oss.py
new file mode 100755
index 0000000..c4a90c5
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/proc_conf_oss.py
@@ -0,0 +1,35 @@
+from pathlib import Path
+
+import torch
+import yaml
+
+
+class NoAliasSafeDumper(yaml.SafeDumper):
+ # Disable anchor/alias in yaml because looks ugly
+ def ignore_aliases(self, data):
+ return True
+
+
+def yaml_no_alias_safe_dump(data, stream=None, **kwargs):
+ """Safe-dump in yaml with no anchor/alias"""
+ return yaml.dump(
+ data, stream, allow_unicode=True, Dumper=NoAliasSafeDumper, **kwargs
+ )
+
+
+def gen_conf(file, out_dir):
+ conf = torch.load(file)["config"]
+ conf["oss_bucket"] = "null"
+ print(conf)
+ output_dir = Path(out_dir)
+ output_dir.mkdir(parents=True, exist_ok=True)
+ with (output_dir / "config.yaml").open("w", encoding="utf-8") as f:
+ yaml_no_alias_safe_dump(conf, f, indent=4, sort_keys=False)
+
+
+if __name__ == "__main__":
+ import sys
+
+ in_f = sys.argv[1]
+ out_f = sys.argv[2]
+ gen_conf(in_f, out_f)
diff --git a/egs_modelscope/speechio/paraformer/utils/proce_text.py b/egs_modelscope/speechio/paraformer/utils/proce_text.py
new file mode 100755
index 0000000..9e517a4
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/proce_text.py
@@ -0,0 +1,31 @@
+
+import sys
+import re
+
+in_f = sys.argv[1]
+out_f = sys.argv[2]
+
+
+with open(in_f, "r", encoding="utf-8") as f:
+ lines = f.readlines()
+
+with open(out_f, "w", encoding="utf-8") as f:
+ for line in lines:
+ outs = line.strip().split(" ", 1)
+ if len(outs) == 2:
+ idx, text = outs
+ text = re.sub("</s>", "", text)
+ text = re.sub("<s>", "", text)
+ text = re.sub("@@", "", text)
+ text = re.sub("@", "", text)
+ text = re.sub("<unk>", "", text)
+ text = re.sub(" ", "", text)
+ text = text.lower()
+ else:
+ idx = outs[0]
+ text = " "
+
+ text = [x for x in text]
+ text = " ".join(text)
+ out = "{} {}\n".format(idx, text)
+ f.write(out)
diff --git a/egs_modelscope/speechio/paraformer/utils/run.pl b/egs_modelscope/speechio/paraformer/utils/run.pl
new file mode 100755
index 0000000..483f95b
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/run.pl
@@ -0,0 +1,356 @@
+#!/usr/bin/env perl
+use warnings; #sed replacement for -w perl parameter
+# In general, doing
+# run.pl some.log a b c is like running the command a b c in
+# the bash shell, and putting the standard error and output into some.log.
+# To run parallel jobs (backgrounded on the host machine), you can do (e.g.)
+# run.pl JOB=1:4 some.JOB.log a b c JOB is like running the command a b c JOB
+# and putting it in some.JOB.log, for each one. [Note: JOB can be any identifier].
+# If any of the jobs fails, this script will fail.
+
+# A typical example is:
+# run.pl some.log my-prog "--opt=foo bar" foo \| other-prog baz
+# and run.pl will run something like:
+# ( my-prog '--opt=foo bar' foo | other-prog baz ) >& some.log
+#
+# Basically it takes the command-line arguments, quotes them
+# as necessary to preserve spaces, and evaluates them with bash.
+# In addition it puts the command line at the top of the log, and
+# the start and end times of the command at the beginning and end.
+# The reason why this is useful is so that we can create a different
+# version of this program that uses a queueing system instead.
+
+#use Data::Dumper;
+
+@ARGV < 2 && die "usage: run.pl log-file command-line arguments...";
+
+#print STDERR "COMMAND-LINE: " . Dumper(\@ARGV) . "\n";
+$job_pick = 'all';
+$max_jobs_run = -1;
+$jobstart = 1;
+$jobend = 1;
+$ignored_opts = ""; # These will be ignored.
+
+# First parse an option like JOB=1:4, and any
+# options that would normally be given to
+# queue.pl, which we will just discard.
+
+for (my $x = 1; $x <= 2; $x++) { # This for-loop is to
+ # allow the JOB=1:n option to be interleaved with the
+ # options to qsub.
+ while (@ARGV >= 2 && $ARGV[0] =~ m:^-:) {
+ # parse any options that would normally go to qsub, but which will be ignored here.
+ my $switch = shift @ARGV;
+ if ($switch eq "-V") {
+ $ignored_opts .= "-V ";
+ } elsif ($switch eq "--max-jobs-run" || $switch eq "-tc") {
+ # we do support the option --max-jobs-run n, and its GridEngine form -tc n.
+ # if the command appears multiple times uses the smallest option.
+ if ( $max_jobs_run <= 0 ) {
+ $max_jobs_run = shift @ARGV;
+ } else {
+ my $new_constraint = shift @ARGV;
+ if ( ($new_constraint < $max_jobs_run) ) {
+ $max_jobs_run = $new_constraint;
+ }
+ }
+
+ if (! ($max_jobs_run > 0)) {
+ die "run.pl: invalid option --max-jobs-run $max_jobs_run";
+ }
+ } else {
+ my $argument = shift @ARGV;
+ if ($argument =~ m/^--/) {
+ print STDERR "run.pl: WARNING: suspicious argument '$argument' to $switch; starts with '-'\n";
+ }
+ if ($switch eq "-sync" && $argument =~ m/^[yY]/) {
+ $ignored_opts .= "-sync "; # Note: in the
+ # corresponding code in queue.pl it says instead, just "$sync = 1;".
+ } elsif ($switch eq "-pe") { # e.g. -pe smp 5
+ my $argument2 = shift @ARGV;
+ $ignored_opts .= "$switch $argument $argument2 ";
+ } elsif ($switch eq "--gpu") {
+ $using_gpu = $argument;
+ } elsif ($switch eq "--pick") {
+ if($argument =~ m/^(all|failed|incomplete)$/) {
+ $job_pick = $argument;
+ } else {
+ print STDERR "run.pl: ERROR: --pick argument must be one of 'all', 'failed' or 'incomplete'"
+ }
+ } else {
+ # Ignore option.
+ $ignored_opts .= "$switch $argument ";
+ }
+ }
+ }
+ if ($ARGV[0] =~ m/^([\w_][\w\d_]*)+=(\d+):(\d+)$/) { # e.g. JOB=1:20
+ $jobname = $1;
+ $jobstart = $2;
+ $jobend = $3;
+ if ($jobstart > $jobend) {
+ die "run.pl: invalid job range $ARGV[0]";
+ }
+ if ($jobstart <= 0) {
+ die "run.pl: invalid job range $ARGV[0], start must be strictly positive (this is required for GridEngine compatibility).";
+ }
+ shift;
+ } elsif ($ARGV[0] =~ m/^([\w_][\w\d_]*)+=(\d+)$/) { # e.g. JOB=1.
+ $jobname = $1;
+ $jobstart = $2;
+ $jobend = $2;
+ shift;
+ } elsif ($ARGV[0] =~ m/.+\=.*\:.*$/) {
+ print STDERR "run.pl: Warning: suspicious first argument to run.pl: $ARGV[0]\n";
+ }
+}
+
+# Users found this message confusing so we are removing it.
+# if ($ignored_opts ne "") {
+# print STDERR "run.pl: Warning: ignoring options \"$ignored_opts\"\n";
+# }
+
+if ($max_jobs_run == -1) { # If --max-jobs-run option not set,
+ # then work out the number of processors if possible,
+ # and set it based on that.
+ $max_jobs_run = 0;
+ if ($using_gpu) {
+ if (open(P, "nvidia-smi -L |")) {
+ $max_jobs_run++ while (<P>);
+ close(P);
+ }
+ if ($max_jobs_run == 0) {
+ $max_jobs_run = 1;
+ print STDERR "run.pl: Warning: failed to detect number of GPUs from nvidia-smi, using ${max_jobs_run}\n";
+ }
+ } elsif (open(P, "</proc/cpuinfo")) { # Linux
+ while (<P>) { if (m/^processor/) { $max_jobs_run++; } }
+ if ($max_jobs_run == 0) {
+ print STDERR "run.pl: Warning: failed to detect any processors from /proc/cpuinfo\n";
+ $max_jobs_run = 10; # reasonable default.
+ }
+ close(P);
+ } elsif (open(P, "sysctl -a |")) { # BSD/Darwin
+ while (<P>) {
+ if (m/hw\.ncpu\s*[:=]\s*(\d+)/) { # hw.ncpu = 4, or hw.ncpu: 4
+ $max_jobs_run = $1;
+ last;
+ }
+ }
+ close(P);
+ if ($max_jobs_run == 0) {
+ print STDERR "run.pl: Warning: failed to detect any processors from sysctl -a\n";
+ $max_jobs_run = 10; # reasonable default.
+ }
+ } else {
+ # allow at most 32 jobs at once, on non-UNIX systems; change this code
+ # if you need to change this default.
+ $max_jobs_run = 32;
+ }
+ # The just-computed value of $max_jobs_run is just the number of processors
+ # (or our best guess); and if it happens that the number of jobs we need to
+ # run is just slightly above $max_jobs_run, it will make sense to increase
+ # $max_jobs_run to equal the number of jobs, so we don't have a small number
+ # of leftover jobs.
+ $num_jobs = $jobend - $jobstart + 1;
+ if (!$using_gpu &&
+ $num_jobs > $max_jobs_run && $num_jobs < 1.4 * $max_jobs_run) {
+ $max_jobs_run = $num_jobs;
+ }
+}
+
+sub pick_or_exit {
+ # pick_or_exit ( $logfile )
+ # Invoked before each job is started helps to run jobs selectively.
+ #
+ # Given the name of the output logfile decides whether the job must be
+ # executed (by returning from the subroutine) or not (by terminating the
+ # process calling exit)
+ #
+ # PRE: $job_pick is a global variable set by command line switch --pick
+ # and indicates which class of jobs must be executed.
+ #
+ # 1) If a failed job is not executed the process exit code will indicate
+ # failure, just as if the task was just executed and failed.
+ #
+ # 2) If a task is incomplete it will be executed. Incomplete may be either
+ # a job whose log file does not contain the accounting notes in the end,
+ # or a job whose log file does not exist.
+ #
+ # 3) If the $job_pick is set to 'all' (default behavior) a task will be
+ # executed regardless of the result of previous attempts.
+ #
+ # This logic could have been implemented in the main execution loop
+ # but a subroutine to preserve the current level of readability of
+ # that part of the code.
+ #
+ # Alexandre Felipe, (o.alexandre.felipe@gmail.com) 14th of August of 2020
+ #
+ if($job_pick eq 'all'){
+ return; # no need to bother with the previous log
+ }
+ open my $fh, "<", $_[0] or return; # job not executed yet
+ my $log_line;
+ my $cur_line;
+ while ($cur_line = <$fh>) {
+ if( $cur_line =~ m/# Ended \(code .*/ ) {
+ $log_line = $cur_line;
+ }
+ }
+ close $fh;
+ if (! defined($log_line)){
+ return; # incomplete
+ }
+ if ( $log_line =~ m/# Ended \(code 0\).*/ ) {
+ exit(0); # complete
+ } elsif ( $log_line =~ m/# Ended \(code \d+(; signal \d+)?\).*/ ){
+ if ($job_pick !~ m/^(failed|all)$/) {
+ exit(1); # failed but not going to run
+ } else {
+ return; # failed
+ }
+ } elsif ( $log_line =~ m/.*\S.*/ ) {
+ return; # incomplete jobs are always run
+ }
+}
+
+
+$logfile = shift @ARGV;
+
+if (defined $jobname && $logfile !~ m/$jobname/ &&
+ $jobend > $jobstart) {
+ print STDERR "run.pl: you are trying to run a parallel job but "
+ . "you are putting the output into just one log file ($logfile)\n";
+ exit(1);
+}
+
+$cmd = "";
+
+foreach $x (@ARGV) {
+ if ($x =~ m/^\S+$/) { $cmd .= $x . " "; }
+ elsif ($x =~ m:\":) { $cmd .= "'$x' "; }
+ else { $cmd .= "\"$x\" "; }
+}
+
+#$Data::Dumper::Indent=0;
+$ret = 0;
+$numfail = 0;
+%active_pids=();
+
+use POSIX ":sys_wait_h";
+for ($jobid = $jobstart; $jobid <= $jobend; $jobid++) {
+ if (scalar(keys %active_pids) >= $max_jobs_run) {
+
+ # Lets wait for a change in any child's status
+ # Then we have to work out which child finished
+ $r = waitpid(-1, 0);
+ $code = $?;
+ if ($r < 0 ) { die "run.pl: Error waiting for child process"; } # should never happen.
+ if ( defined $active_pids{$r} ) {
+ $jid=$active_pids{$r};
+ $fail[$jid]=$code;
+ if ($code !=0) { $numfail++;}
+ delete $active_pids{$r};
+ # print STDERR "Finished: $r/$jid " . Dumper(\%active_pids) . "\n";
+ } else {
+ die "run.pl: Cannot find the PID of the child process that just finished.";
+ }
+
+ # In theory we could do a non-blocking waitpid over all jobs running just
+ # to find out if only one or more jobs finished during the previous waitpid()
+ # However, we just omit this and will reap the next one in the next pass
+ # through the for(;;) cycle
+ }
+ $childpid = fork();
+ if (!defined $childpid) { die "run.pl: Error forking in run.pl (writing to $logfile)"; }
+ if ($childpid == 0) { # We're in the child... this branch
+ # executes the job and returns (possibly with an error status).
+ if (defined $jobname) {
+ $cmd =~ s/$jobname/$jobid/g;
+ $logfile =~ s/$jobname/$jobid/g;
+ }
+ # exit if the job does not need to be executed
+ pick_or_exit( $logfile );
+
+ system("mkdir -p `dirname $logfile` 2>/dev/null");
+ open(F, ">$logfile") || die "run.pl: Error opening log file $logfile";
+ print F "# " . $cmd . "\n";
+ print F "# Started at " . `date`;
+ $starttime = `date +'%s'`;
+ print F "#\n";
+ close(F);
+
+ # Pipe into bash.. make sure we're not using any other shell.
+ open(B, "|bash") || die "run.pl: Error opening shell command";
+ print B "( " . $cmd . ") 2>>$logfile >> $logfile";
+ close(B); # If there was an error, exit status is in $?
+ $ret = $?;
+
+ $lowbits = $ret & 127;
+ $highbits = $ret >> 8;
+ if ($lowbits != 0) { $return_str = "code $highbits; signal $lowbits" }
+ else { $return_str = "code $highbits"; }
+
+ $endtime = `date +'%s'`;
+ open(F, ">>$logfile") || die "run.pl: Error opening log file $logfile (again)";
+ $enddate = `date`;
+ chop $enddate;
+ print F "# Accounting: time=" . ($endtime - $starttime) . " threads=1\n";
+ print F "# Ended ($return_str) at " . $enddate . ", elapsed time " . ($endtime-$starttime) . " seconds\n";
+ close(F);
+ exit($ret == 0 ? 0 : 1);
+ } else {
+ $pid[$jobid] = $childpid;
+ $active_pids{$childpid} = $jobid;
+ # print STDERR "Queued: " . Dumper(\%active_pids) . "\n";
+ }
+}
+
+# Now we have submitted all the jobs, lets wait until all the jobs finish
+foreach $child (keys %active_pids) {
+ $jobid=$active_pids{$child};
+ $r = waitpid($pid[$jobid], 0);
+ $code = $?;
+ if ($r == -1) { die "run.pl: Error waiting for child process"; } # should never happen.
+ if ($r != 0) { $fail[$jobid]=$code; $numfail++ if $code!=0; } # Completed successfully
+}
+
+# Some sanity checks:
+# The $fail array should not contain undefined codes
+# The number of non-zeros in that array should be equal to $numfail
+# We cannot do foreach() here, as the JOB ids do not start at zero
+$failed_jids=0;
+for ($jobid = $jobstart; $jobid <= $jobend; $jobid++) {
+ $job_return = $fail[$jobid];
+ if (not defined $job_return ) {
+ # print Dumper(\@fail);
+
+ die "run.pl: Sanity check failed: we have indication that some jobs are running " .
+ "even after we waited for all jobs to finish" ;
+ }
+ if ($job_return != 0 ){ $failed_jids++;}
+}
+if ($failed_jids != $numfail) {
+ die "run.pl: Sanity check failed: cannot find out how many jobs failed ($failed_jids x $numfail)."
+}
+if ($numfail > 0) { $ret = 1; }
+
+if ($ret != 0) {
+ $njobs = $jobend - $jobstart + 1;
+ if ($njobs == 1) {
+ if (defined $jobname) {
+ $logfile =~ s/$jobname/$jobstart/; # only one numbered job, so replace name with
+ # that job.
+ }
+ print STDERR "run.pl: job failed, log is in $logfile\n";
+ if ($logfile =~ m/JOB/) {
+ print STDERR "run.pl: probably you forgot to put JOB=1:\$nj in your script.";
+ }
+ }
+ else {
+ $logfile =~ s/$jobname/*/g;
+ print STDERR "run.pl: $numfail / $njobs failed, log is in $logfile\n";
+ }
+}
+
+
+exit ($ret);
diff --git a/egs_modelscope/speechio/paraformer/utils/shuffle_list.pl b/egs_modelscope/speechio/paraformer/utils/shuffle_list.pl
new file mode 100755
index 0000000..a116200
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/shuffle_list.pl
@@ -0,0 +1,44 @@
+#!/usr/bin/env perl
+
+# Copyright 2013 Johns Hopkins University (author: Daniel Povey)
+
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+if ($ARGV[0] eq "--srand") {
+ $n = $ARGV[1];
+ $n =~ m/\d+/ || die "Bad argument to --srand option: \"$n\"";
+ srand($ARGV[1]);
+ shift;
+ shift;
+} else {
+ srand(0); # Gives inconsistent behavior if we don't seed.
+}
+
+if (@ARGV > 1 || $ARGV[0] =~ m/^-.+/) { # >1 args, or an option we
+ # don't understand.
+ print "Usage: shuffle_list.pl [--srand N] [input file] > output\n";
+ print "randomizes the order of lines of input.\n";
+ exit(1);
+}
+
+@lines;
+while (<>) {
+ push @lines, [ (rand(), $_)] ;
+}
+
+@lines = sort { $a->[0] cmp $b->[0] } @lines;
+foreach $l (@lines) {
+ print $l->[1];
+}
\ No newline at end of file
diff --git a/egs_modelscope/speechio/paraformer/utils/split_data.py b/egs_modelscope/speechio/paraformer/utils/split_data.py
new file mode 100755
index 0000000..060eae6
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/split_data.py
@@ -0,0 +1,60 @@
+import os
+import sys
+import random
+
+
+in_dir = sys.argv[1]
+out_dir = sys.argv[2]
+num_split = sys.argv[3]
+
+
+def split_scp(scp, num):
+ assert len(scp) >= num
+ avg = len(scp) // num
+ out = []
+ begin = 0
+
+ for i in range(num):
+ if i == num - 1:
+ out.append(scp[begin:])
+ else:
+ out.append(scp[begin:begin+avg])
+ begin += avg
+
+ return out
+
+
+os.path.exists("{}/wav.scp".format(in_dir))
+os.path.exists("{}/text".format(in_dir))
+
+with open("{}/wav.scp".format(in_dir), 'r') as infile:
+ wav_list = infile.readlines()
+
+with open("{}/text".format(in_dir), 'r') as infile:
+ text_list = infile.readlines()
+
+assert len(wav_list) == len(text_list)
+
+x = list(zip(wav_list, text_list))
+random.shuffle(x)
+wav_shuffle_list, text_shuffle_list = zip(*x)
+
+num_split = int(num_split)
+wav_split_list = split_scp(wav_shuffle_list, num_split)
+text_split_list = split_scp(text_shuffle_list, num_split)
+
+for idx, wav_list in enumerate(wav_split_list, 1):
+ path = out_dir + "/split" + str(num_split) + "/" + str(idx)
+ if not os.path.exists(path):
+ os.makedirs(path)
+ with open("{}/wav.scp".format(path), 'w') as wav_writer:
+ for line in wav_list:
+ wav_writer.write(line)
+
+for idx, text_list in enumerate(text_split_list, 1):
+ path = out_dir + "/split" + str(num_split) + "/" + str(idx)
+ if not os.path.exists(path):
+ os.makedirs(path)
+ with open("{}/text".format(path), 'w') as text_writer:
+ for line in text_list:
+ text_writer.write(line)
diff --git a/egs_modelscope/speechio/paraformer/utils/split_scp.pl b/egs_modelscope/speechio/paraformer/utils/split_scp.pl
new file mode 100755
index 0000000..0876dcb
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/split_scp.pl
@@ -0,0 +1,246 @@
+#!/usr/bin/env perl
+
+# Copyright 2010-2011 Microsoft Corporation
+
+# See ../../COPYING for clarification regarding multiple authors
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+# This program splits up any kind of .scp or archive-type file.
+# If there is no utt2spk option it will work on any text file and
+# will split it up with an approximately equal number of lines in
+# each but.
+# With the --utt2spk option it will work on anything that has the
+# utterance-id as the first entry on each line; the utt2spk file is
+# of the form "utterance speaker" (on each line).
+# It splits it into equal size chunks as far as it can. If you use the utt2spk
+# option it will make sure these chunks coincide with speaker boundaries. In
+# this case, if there are more chunks than speakers (and in some other
+# circumstances), some of the resulting chunks will be empty and it will print
+# an error message and exit with nonzero status.
+# You will normally call this like:
+# split_scp.pl scp scp.1 scp.2 scp.3 ...
+# or
+# split_scp.pl --utt2spk=utt2spk scp scp.1 scp.2 scp.3 ...
+# Note that you can use this script to split the utt2spk file itself,
+# e.g. split_scp.pl --utt2spk=utt2spk utt2spk utt2spk.1 utt2spk.2 ...
+
+# You can also call the scripts like:
+# split_scp.pl -j 3 0 scp scp.0
+# [note: with this option, it assumes zero-based indexing of the split parts,
+# i.e. the second number must be 0 <= n < num-jobs.]
+
+use warnings;
+
+$num_jobs = 0;
+$job_id = 0;
+$utt2spk_file = "";
+$one_based = 0;
+
+for ($x = 1; $x <= 3 && @ARGV > 0; $x++) {
+ if ($ARGV[0] eq "-j") {
+ shift @ARGV;
+ $num_jobs = shift @ARGV;
+ $job_id = shift @ARGV;
+ }
+ if ($ARGV[0] =~ /--utt2spk=(.+)/) {
+ $utt2spk_file=$1;
+ shift;
+ }
+ if ($ARGV[0] eq '--one-based') {
+ $one_based = 1;
+ shift @ARGV;
+ }
+}
+
+if ($num_jobs != 0 && ($num_jobs < 0 || $job_id - $one_based < 0 ||
+ $job_id - $one_based >= $num_jobs)) {
+ die "$0: Invalid job number/index values for '-j $num_jobs $job_id" .
+ ($one_based ? " --one-based" : "") . "'\n"
+}
+
+$one_based
+ and $job_id--;
+
+if(($num_jobs == 0 && @ARGV < 2) || ($num_jobs > 0 && (@ARGV < 1 || @ARGV > 2))) {
+ die
+"Usage: split_scp.pl [--utt2spk=<utt2spk_file>] in.scp out1.scp out2.scp ...
+ or: split_scp.pl -j num-jobs job-id [--one-based] [--utt2spk=<utt2spk_file>] in.scp [out.scp]
+ ... where 0 <= job-id < num-jobs, or 1 <= job-id <- num-jobs if --one-based.\n";
+}
+
+$error = 0;
+$inscp = shift @ARGV;
+if ($num_jobs == 0) { # without -j option
+ @OUTPUTS = @ARGV;
+} else {
+ for ($j = 0; $j < $num_jobs; $j++) {
+ if ($j == $job_id) {
+ if (@ARGV > 0) { push @OUTPUTS, $ARGV[0]; }
+ else { push @OUTPUTS, "-"; }
+ } else {
+ push @OUTPUTS, "/dev/null";
+ }
+ }
+}
+
+if ($utt2spk_file ne "") { # We have the --utt2spk option...
+ open($u_fh, '<', $utt2spk_file) || die "$0: Error opening utt2spk file $utt2spk_file: $!\n";
+ while(<$u_fh>) {
+ @A = split;
+ @A == 2 || die "$0: Bad line $_ in utt2spk file $utt2spk_file\n";
+ ($u,$s) = @A;
+ $utt2spk{$u} = $s;
+ }
+ close $u_fh;
+ open($i_fh, '<', $inscp) || die "$0: Error opening input scp file $inscp: $!\n";
+ @spkrs = ();
+ while(<$i_fh>) {
+ @A = split;
+ if(@A == 0) { die "$0: Empty or space-only line in scp file $inscp\n"; }
+ $u = $A[0];
+ $s = $utt2spk{$u};
+ defined $s || die "$0: No utterance $u in utt2spk file $utt2spk_file\n";
+ if(!defined $spk_count{$s}) {
+ push @spkrs, $s;
+ $spk_count{$s} = 0;
+ $spk_data{$s} = []; # ref to new empty array.
+ }
+ $spk_count{$s}++;
+ push @{$spk_data{$s}}, $_;
+ }
+ # Now split as equally as possible ..
+ # First allocate spks to files by allocating an approximately
+ # equal number of speakers.
+ $numspks = @spkrs; # number of speakers.
+ $numscps = @OUTPUTS; # number of output files.
+ if ($numspks < $numscps) {
+ die "$0: Refusing to split data because number of speakers $numspks " .
+ "is less than the number of output .scp files $numscps\n";
+ }
+ for($scpidx = 0; $scpidx < $numscps; $scpidx++) {
+ $scparray[$scpidx] = []; # [] is array reference.
+ }
+ for ($spkidx = 0; $spkidx < $numspks; $spkidx++) {
+ $scpidx = int(($spkidx*$numscps) / $numspks);
+ $spk = $spkrs[$spkidx];
+ push @{$scparray[$scpidx]}, $spk;
+ $scpcount[$scpidx] += $spk_count{$spk};
+ }
+
+ # Now will try to reassign beginning + ending speakers
+ # to different scp's and see if it gets more balanced.
+ # Suppose objf we're minimizing is sum_i (num utts in scp[i] - average)^2.
+ # We can show that if considering changing just 2 scp's, we minimize
+ # this by minimizing the squared difference in sizes. This is
+ # equivalent to minimizing the absolute difference in sizes. This
+ # shows this method is bound to converge.
+
+ $changed = 1;
+ while($changed) {
+ $changed = 0;
+ for($scpidx = 0; $scpidx < $numscps; $scpidx++) {
+ # First try to reassign ending spk of this scp.
+ if($scpidx < $numscps-1) {
+ $sz = @{$scparray[$scpidx]};
+ if($sz > 0) {
+ $spk = $scparray[$scpidx]->[$sz-1];
+ $count = $spk_count{$spk};
+ $nutt1 = $scpcount[$scpidx];
+ $nutt2 = $scpcount[$scpidx+1];
+ if( abs( ($nutt2+$count) - ($nutt1-$count))
+ < abs($nutt2 - $nutt1)) { # Would decrease
+ # size-diff by reassigning spk...
+ $scpcount[$scpidx+1] += $count;
+ $scpcount[$scpidx] -= $count;
+ pop @{$scparray[$scpidx]};
+ unshift @{$scparray[$scpidx+1]}, $spk;
+ $changed = 1;
+ }
+ }
+ }
+ if($scpidx > 0 && @{$scparray[$scpidx]} > 0) {
+ $spk = $scparray[$scpidx]->[0];
+ $count = $spk_count{$spk};
+ $nutt1 = $scpcount[$scpidx-1];
+ $nutt2 = $scpcount[$scpidx];
+ if( abs( ($nutt2-$count) - ($nutt1+$count))
+ < abs($nutt2 - $nutt1)) { # Would decrease
+ # size-diff by reassigning spk...
+ $scpcount[$scpidx-1] += $count;
+ $scpcount[$scpidx] -= $count;
+ shift @{$scparray[$scpidx]};
+ push @{$scparray[$scpidx-1]}, $spk;
+ $changed = 1;
+ }
+ }
+ }
+ }
+ # Now print out the files...
+ for($scpidx = 0; $scpidx < $numscps; $scpidx++) {
+ $scpfile = $OUTPUTS[$scpidx];
+ ($scpfile ne '-' ? open($f_fh, '>', $scpfile)
+ : open($f_fh, '>&', \*STDOUT)) ||
+ die "$0: Could not open scp file $scpfile for writing: $!\n";
+ $count = 0;
+ if(@{$scparray[$scpidx]} == 0) {
+ print STDERR "$0: eError: split_scp.pl producing empty .scp file " .
+ "$scpfile (too many splits and too few speakers?)\n";
+ $error = 1;
+ } else {
+ foreach $spk ( @{$scparray[$scpidx]} ) {
+ print $f_fh @{$spk_data{$spk}};
+ $count += $spk_count{$spk};
+ }
+ $count == $scpcount[$scpidx] || die "Count mismatch [code error]";
+ }
+ close($f_fh);
+ }
+} else {
+ # This block is the "normal" case where there is no --utt2spk
+ # option and we just break into equal size chunks.
+
+ open($i_fh, '<', $inscp) || die "$0: Error opening input scp file $inscp: $!\n";
+
+ $numscps = @OUTPUTS; # size of array.
+ @F = ();
+ while(<$i_fh>) {
+ push @F, $_;
+ }
+ $numlines = @F;
+ if($numlines == 0) {
+ print STDERR "$0: error: empty input scp file $inscp\n";
+ $error = 1;
+ }
+ $linesperscp = int( $numlines / $numscps); # the "whole part"..
+ $linesperscp >= 1 || die "$0: You are splitting into too many pieces! [reduce \$nj ($numscps) to be smaller than the number of lines ($numlines) in $inscp]\n";
+ $remainder = $numlines - ($linesperscp * $numscps);
+ ($remainder >= 0 && $remainder < $numlines) || die "bad remainder $remainder";
+ # [just doing int() rounds down].
+ $n = 0;
+ for($scpidx = 0; $scpidx < @OUTPUTS; $scpidx++) {
+ $scpfile = $OUTPUTS[$scpidx];
+ ($scpfile ne '-' ? open($o_fh, '>', $scpfile)
+ : open($o_fh, '>&', \*STDOUT)) ||
+ die "$0: Could not open scp file $scpfile for writing: $!\n";
+ for($k = 0; $k < $linesperscp + ($scpidx < $remainder ? 1 : 0); $k++) {
+ print $o_fh $F[$n++];
+ }
+ close($o_fh) || die "$0: Eror closing scp file $scpfile: $!\n";
+ }
+ $n == $numlines || die "$n != $numlines [code error]";
+}
+
+exit ($error);
diff --git a/egs_modelscope/speechio/paraformer/utils/subset_data_dir_tr_cv.sh b/egs_modelscope/speechio/paraformer/utils/subset_data_dir_tr_cv.sh
new file mode 100755
index 0000000..e16cebd
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/subset_data_dir_tr_cv.sh
@@ -0,0 +1,30 @@
+#!/usr/bin/env bash
+
+dev_num_utt=1000
+
+echo "$0 $@"
+. utils/parse_options.sh || exit 1;
+
+train_data=$1
+out_dir=$2
+
+[ ! -f ${train_data}/wav.scp ] && echo "$0: no such file ${train_data}/wav.scp" && exit 1;
+[ ! -f ${train_data}/text ] && echo "$0: no such file ${train_data}/text" && exit 1;
+
+mkdir -p ${out_dir}/train && mkdir -p ${out_dir}/dev
+
+cp ${train_data}/wav.scp ${out_dir}/train/wav.scp.bak
+cp ${train_data}/text ${out_dir}/train/text.bak
+
+num_utt=$(wc -l <${out_dir}/train/wav.scp.bak)
+
+utils/shuffle_list.pl --srand 1 ${out_dir}/train/wav.scp.bak > ${out_dir}/train/wav.scp.shuf
+head -n ${dev_num_utt} ${out_dir}/train/wav.scp.shuf > ${out_dir}/dev/wav.scp
+tail -n $((${num_utt}-${dev_num_utt})) ${out_dir}/train/wav.scp.shuf > ${out_dir}/train/wav.scp
+
+utils/shuffle_list.pl --srand 1 ${out_dir}/train/text.bak > ${out_dir}/train/text.shuf
+head -n ${dev_num_utt} ${out_dir}/train/text.shuf > ${out_dir}/dev/text
+tail -n $((${num_utt}-${dev_num_utt})) ${out_dir}/train/text.shuf > ${out_dir}/train/text
+
+rm ${out_dir}/train/wav.scp.bak ${out_dir}/train/text.bak
+rm ${out_dir}/train/wav.scp.shuf ${out_dir}/train/text.shuf
diff --git a/egs_modelscope/speechio/paraformer/utils/text2token.py b/egs_modelscope/speechio/paraformer/utils/text2token.py
new file mode 100755
index 0000000..56c3913
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/text2token.py
@@ -0,0 +1,135 @@
+#!/usr/bin/env python3
+
+# Copyright 2017 Johns Hopkins University (Shinji Watanabe)
+# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
+
+
+import argparse
+import codecs
+import re
+import sys
+
+is_python2 = sys.version_info[0] == 2
+
+
+def exist_or_not(i, match_pos):
+ start_pos = None
+ end_pos = None
+ for pos in match_pos:
+ if pos[0] <= i < pos[1]:
+ start_pos = pos[0]
+ end_pos = pos[1]
+ break
+
+ return start_pos, end_pos
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="convert raw text to tokenized text",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--nchar",
+ "-n",
+ default=1,
+ type=int,
+ help="number of characters to split, i.e., \
+ aabb -> a a b b with -n 1 and aa bb with -n 2",
+ )
+ parser.add_argument(
+ "--skip-ncols", "-s", default=0, type=int, help="skip first n columns"
+ )
+ parser.add_argument("--space", default="<space>", type=str, help="space symbol")
+ parser.add_argument(
+ "--non-lang-syms",
+ "-l",
+ default=None,
+ type=str,
+ help="list of non-linguistic symobles, e.g., <NOISE> etc.",
+ )
+ parser.add_argument("text", type=str, default=False, nargs="?", help="input text")
+ parser.add_argument(
+ "--trans_type",
+ "-t",
+ type=str,
+ default="char",
+ choices=["char", "phn"],
+ help="""Transcript type. char/phn. e.g., for TIMIT FADG0_SI1279 -
+ If trans_type is char,
+ read from SI1279.WRD file -> "bricks are an alternative"
+ Else if trans_type is phn,
+ read from SI1279.PHN file -> "sil b r ih sil k s aa r er n aa l
+ sil t er n ih sil t ih v sil" """,
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ rs = []
+ if args.non_lang_syms is not None:
+ with codecs.open(args.non_lang_syms, "r", encoding="utf-8") as f:
+ nls = [x.rstrip() for x in f.readlines()]
+ rs = [re.compile(re.escape(x)) for x in nls]
+
+ if args.text:
+ f = codecs.open(args.text, encoding="utf-8")
+ else:
+ f = codecs.getreader("utf-8")(sys.stdin if is_python2 else sys.stdin.buffer)
+
+ sys.stdout = codecs.getwriter("utf-8")(
+ sys.stdout if is_python2 else sys.stdout.buffer
+ )
+ line = f.readline()
+ n = args.nchar
+ while line:
+ x = line.split()
+ print(" ".join(x[: args.skip_ncols]), end=" ")
+ a = " ".join(x[args.skip_ncols :])
+
+ # get all matched positions
+ match_pos = []
+ for r in rs:
+ i = 0
+ while i >= 0:
+ m = r.search(a, i)
+ if m:
+ match_pos.append([m.start(), m.end()])
+ i = m.end()
+ else:
+ break
+
+ if args.trans_type == "phn":
+ a = a.split(" ")
+ else:
+ if len(match_pos) > 0:
+ chars = []
+ i = 0
+ while i < len(a):
+ start_pos, end_pos = exist_or_not(i, match_pos)
+ if start_pos is not None:
+ chars.append(a[start_pos:end_pos])
+ i = end_pos
+ else:
+ chars.append(a[i])
+ i += 1
+ a = chars
+
+ a = [a[j : j + n] for j in range(0, len(a), n)]
+
+ a_flat = []
+ for z in a:
+ a_flat.append("".join(z))
+
+ a_chars = [z.replace(" ", args.space) for z in a_flat]
+ if args.trans_type == "phn":
+ a_chars = [z.replace("sil", args.space) for z in a_chars]
+ print(" ".join(a_chars))
+ line = f.readline()
+
+
+if __name__ == "__main__":
+ main()
diff --git a/egs_modelscope/speechio/paraformer/utils/text_tokenize.py b/egs_modelscope/speechio/paraformer/utils/text_tokenize.py
new file mode 100755
index 0000000..962ea11
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/text_tokenize.py
@@ -0,0 +1,106 @@
+import re
+import argparse
+
+
+def load_dict(seg_file):
+ seg_dict = {}
+ with open(seg_file, 'r') as infile:
+ for line in infile:
+ s = line.strip().split()
+ key = s[0]
+ value = s[1:]
+ seg_dict[key] = " ".join(value)
+ return seg_dict
+
+
+def forward_segment(text, dic):
+ word_list = []
+ i = 0
+ while i < len(text):
+ longest_word = text[i]
+ for j in range(i + 1, len(text) + 1):
+ word = text[i:j]
+ if word in dic:
+ if len(word) > len(longest_word):
+ longest_word = word
+ word_list.append(longest_word)
+ i += len(longest_word)
+ return word_list
+
+
+def tokenize(txt,
+ seg_dict):
+ out_txt = ""
+ pattern = re.compile(r"([\u4E00-\u9FA5A-Za-z0-9])")
+ for word in txt:
+ if pattern.match(word):
+ if word in seg_dict:
+ out_txt += seg_dict[word] + " "
+ else:
+ out_txt += "<unk>" + " "
+ else:
+ continue
+ return out_txt.strip()
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="text tokenize",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--text-file",
+ "-t",
+ default=False,
+ required=True,
+ type=str,
+ help="input text",
+ )
+ parser.add_argument(
+ "--seg-file",
+ "-s",
+ default=False,
+ required=True,
+ type=str,
+ help="seg file",
+ )
+ parser.add_argument(
+ "--txt-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="txt index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ txt_writer = open("{}/text.{}.txt".format(args.output_dir, args.txt_index), 'w')
+ shape_writer = open("{}/len.{}".format(args.output_dir, args.txt_index), 'w')
+ seg_dict = load_dict(args.seg_file)
+ with open(args.text_file, 'r') as infile:
+ for line in infile:
+ s = line.strip().split()
+ text_id = s[0]
+ text_list = forward_segment("".join(s[1:]).lower(), seg_dict)
+ text = tokenize(text_list, seg_dict)
+ lens = len(text.strip().split())
+ txt_writer.write(text_id + " " + text + '\n')
+ shape_writer.write(text_id + " " + str(lens) + '\n')
+
+
+if __name__ == '__main__':
+ main()
+
diff --git a/egs_modelscope/speechio/paraformer/utils/text_tokenize.sh b/egs_modelscope/speechio/paraformer/utils/text_tokenize.sh
new file mode 100755
index 0000000..6b74fef
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/text_tokenize.sh
@@ -0,0 +1,35 @@
+#!/usr/bin/env bash
+
+
+# Begin configuration section.
+nj=32
+cmd=utils/run.pl
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+# tokenize configuration
+text_dir=$1
+seg_file=$2
+logdir=$3
+output_dir=$4
+
+txt_dir=${output_dir}/txt; mkdir -p ${output_dir}/txt
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/text_tokenize.JOB.log \
+ python utils/text_tokenize.py -t ${text_dir}/txt/text.JOB.txt \
+ -s ${seg_file} -i JOB -o ${txt_dir} \
+ || exit 1;
+
+# concatenate the text files together.
+for n in $(seq $nj); do
+ cat ${txt_dir}/text.$n.txt || exit 1
+done > ${output_dir}/text || exit 1
+
+for n in $(seq $nj); do
+ cat ${txt_dir}/len.$n || exit 1
+done > ${output_dir}/text_shape || exit 1
+
+echo "$0: Succeeded text tokenize"
diff --git a/egs_modelscope/speechio/paraformer/utils/textnorm_zh.py b/egs_modelscope/speechio/paraformer/utils/textnorm_zh.py
new file mode 100755
index 0000000..79feb83
--- /dev/null
+++ b/egs_modelscope/speechio/paraformer/utils/textnorm_zh.py
@@ -0,0 +1,834 @@
+#!/usr/bin/env python3
+# coding=utf-8
+
+# Authors:
+# 2019.5 Zhiyang Zhou (https://github.com/Joee1995/chn_text_norm.git)
+# 2019.9 Jiayu DU
+#
+# requirements:
+# - python 3.X
+# notes: python 2.X WILL fail or produce misleading results
+
+import sys, os, argparse, codecs, string, re
+
+# ================================================================================ #
+# basic constant
+# ================================================================================ #
+CHINESE_DIGIS = u'闆朵竴浜屼笁鍥涗簲鍏竷鍏節'
+BIG_CHINESE_DIGIS_SIMPLIFIED = u'闆跺9璐板弫鑲嗕紞闄嗘煉鎹岀帠'
+BIG_CHINESE_DIGIS_TRADITIONAL = u'闆跺9璨冲弮鑲嗕紞闄告煉鎹岀帠'
+SMALLER_BIG_CHINESE_UNITS_SIMPLIFIED = u'鍗佺櫨鍗冧竾'
+SMALLER_BIG_CHINESE_UNITS_TRADITIONAL = u'鎷句桨浠熻惉'
+LARGER_CHINESE_NUMERING_UNITS_SIMPLIFIED = u'浜垮厗浜灀绉┌娌熸锭姝h浇'
+LARGER_CHINESE_NUMERING_UNITS_TRADITIONAL = u'鍎勫厗浜灀绉┌婧濇緱姝h級'
+SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED = u'鍗佺櫨鍗冧竾'
+SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL = u'鎷句桨浠熻惉'
+
+ZERO_ALT = u'銆�'
+ONE_ALT = u'骞�'
+TWO_ALTS = [u'涓�', u'鍏�']
+
+POSITIVE = [u'姝�', u'姝�']
+NEGATIVE = [u'璐�', u'璨�']
+POINT = [u'鐐�', u'榛�']
+# PLUS = [u'鍔�', u'鍔�']
+# SIL = [u'鏉�', u'妲�']
+
+FILLER_CHARS = ['鍛�', '鍟�']
+ER_WHITELIST = '(鍎垮コ|鍎垮瓙|鍎垮瓩|濂冲効|鍎垮|濡诲効|' \
+ '鑳庡効|濠村効|鏂扮敓鍎縷濠村辜鍎縷骞煎効|灏戝効|灏忓効|鍎挎瓕|鍎跨|鍎跨|鎵樺効鎵�|瀛ゅ効|' \
+ '鍎挎垙|鍎垮寲|鍙板効搴剕楣垮効宀泑姝e効鍏粡|鍚婂効閮庡綋|鐢熷効鑲插コ|鎵樺効甯﹀コ|鍏诲効闃茶�亅鐥村効鍛嗗コ|' \
+ '浣冲効浣冲|鍎挎�滃吔鎵皘鍎挎棤甯哥埗|鍎夸笉瀚屾瘝涓憒鍎胯鍗冮噷姣嶆媴蹇鍎垮ぇ涓嶇敱鐖穦鑻忎篂鍎�)'
+
+# 涓枃鏁板瓧绯荤粺绫诲瀷
+NUMBERING_TYPES = ['low', 'mid', 'high']
+
+CURRENCY_NAMES = '(浜烘皯甯亅缇庡厓|鏃ュ厓|鑻遍晳|娆у厓|椹厠|娉曢儙|鍔犳嬁澶у厓|婢冲厓|娓竵|鍏堜护|鑺叞椹厠|鐖卞皵鍏伴晳|' \
+ '閲屾媺|鑽峰叞鐩緗鍩冩柉搴撳|姣斿濉攟鍗板凹鐩緗鏋楀悏鐗箌鏂拌タ鍏板厓|姣旂储|鍗㈠竷|鏂板姞鍧″厓|闊╁厓|娉伴摙)'
+CURRENCY_UNITS = '((浜縷鍗冧竾|鐧句竾|涓噟鍗億鐧�)|(浜縷鍗冧竾|鐧句竾|涓噟鍗億鐧緗)鍏億(浜縷鍗冧竾|鐧句竾|涓噟鍗億鐧緗)鍧梶瑙抾姣泑鍒�)'
+COM_QUANTIFIERS = '(鍖箌寮爘搴鍥瀨鍦簗灏緗鏉涓獆棣東闃檤闃祙缃憒鐐畖椤秥涓榺妫祙鍙獆鏀瘄琚瓅杈唡鎸憒鎷厊棰梶澹硘绐爘鏇瞸澧檤缇鑵攟' \
+ '鐮搴瀹璐瘄鎵巪鎹唡鍒�|浠鎵搢鎵媩缃梶鍧灞眧宀瓅姹焲婧獆閽焲闃焲鍗晐鍙寍瀵箌鍑簗鍙澶磡鑴殀鏉縷璺硘鏋潀浠秥璐磡' \
+ '閽坾绾縷绠鍚峾浣峾韬珅鍫倈璇緗鏈瑋椤祙瀹秥鎴穦灞倈涓潀姣珅鍘榺鍒唡閽眧涓鏂鎷厊閾鐭硘閽閿眧蹇絴(鍗億姣珅寰�)鍏媩' \
+ '姣珅鍘榺鍒唡瀵竱灏簗涓坾閲寍瀵粅甯竱閾簗绋媩(鍗億鍒唡鍘榺姣珅寰�)绫硘鎾畖鍕簗鍚坾鍗噟鏂梶鐭硘鐩榺纰梶纰焲鍙爘妗秥绗紎鐩唡' \
+ '鐩抾鏉瘄閽焲鏂泑閿厊绨媩绡畖鐩榺妗秥缃恷鐡秥澹秥鍗畖鐩弢绠﹟绠眧鐓瞸鍟東琚媩閽祙骞磡鏈坾鏃瀛鍒粅鏃秥鍛▅澶﹟绉抾鍒唡鏃瑋' \
+ '绾獆宀亅涓東鏇磡澶渱鏄澶弢绉媩鍐瑋浠浼弢杈坾涓竱娉绮抾棰梶骞鍫唡鏉鏍箌鏀瘄閬搢闈鐗噟寮爘棰梶鍧�)'
+
+# punctuation information are based on Zhon project (https://github.com/tsroten/zhon.git)
+CHINESE_PUNC_STOP = '锛侊紵锝°��'
+CHINESE_PUNC_NON_STOP = '锛傦純锛勶紖锛嗭紘锛堬級锛婏紜锛岋紞锛忥細锛涳紲锛濓紴锛狅蓟锛硷冀锛撅伎锝�锝涳綔锝濓綖锝燂綘锝剑锝ゃ�併�冦�嬨�屻�嶃�庛�忋�愩�戙�斻�曘�栥�椼�樸�欍�氥�涖�溿�濄�炪�熴�般�俱�库�撯�斺�樷�欌�涒�溾�濃�炩�熲�︹�э箯'
+CHINESE_PUNC_LIST = CHINESE_PUNC_STOP + CHINESE_PUNC_NON_STOP
+
+# ================================================================================ #
+# basic class
+# ================================================================================ #
+class ChineseChar(object):
+ """
+ 涓枃瀛楃
+ 姣忎釜瀛楃瀵瑰簲绠�浣撳拰绻佷綋,
+ e.g. 绠�浣� = '璐�', 绻佷綋 = '璨�'
+ 杞崲鏃跺彲杞崲涓虹畝浣撴垨绻佷綋
+ """
+
+ def __init__(self, simplified, traditional):
+ self.simplified = simplified
+ self.traditional = traditional
+ #self.__repr__ = self.__str__
+
+ def __str__(self):
+ return self.simplified or self.traditional or None
+
+ def __repr__(self):
+ return self.__str__()
+
+
+class ChineseNumberUnit(ChineseChar):
+ """
+ 涓枃鏁板瓧/鏁颁綅瀛楃
+ 姣忎釜瀛楃闄ょ箒绠�浣撳杩樻湁涓�涓澶栫殑澶у啓瀛楃
+ e.g. '闄�' 鍜� '闄�'
+ """
+
+ def __init__(self, power, simplified, traditional, big_s, big_t):
+ super(ChineseNumberUnit, self).__init__(simplified, traditional)
+ self.power = power
+ self.big_s = big_s
+ self.big_t = big_t
+
+ def __str__(self):
+ return '10^{}'.format(self.power)
+
+ @classmethod
+ def create(cls, index, value, numbering_type=NUMBERING_TYPES[1], small_unit=False):
+
+ if small_unit:
+ return ChineseNumberUnit(power=index + 1,
+ simplified=value[0], traditional=value[1], big_s=value[1], big_t=value[1])
+ elif numbering_type == NUMBERING_TYPES[0]:
+ return ChineseNumberUnit(power=index + 8,
+ simplified=value[0], traditional=value[1], big_s=value[0], big_t=value[1])
+ elif numbering_type == NUMBERING_TYPES[1]:
+ return ChineseNumberUnit(power=(index + 2) * 4,
+ simplified=value[0], traditional=value[1], big_s=value[0], big_t=value[1])
+ elif numbering_type == NUMBERING_TYPES[2]:
+ return ChineseNumberUnit(power=pow(2, index + 3),
+ simplified=value[0], traditional=value[1], big_s=value[0], big_t=value[1])
+ else:
+ raise ValueError(
+ 'Counting type should be in {0} ({1} provided).'.format(NUMBERING_TYPES, numbering_type))
+
+
+class ChineseNumberDigit(ChineseChar):
+ """
+ 涓枃鏁板瓧瀛楃
+ """
+
+ def __init__(self, value, simplified, traditional, big_s, big_t, alt_s=None, alt_t=None):
+ super(ChineseNumberDigit, self).__init__(simplified, traditional)
+ self.value = value
+ self.big_s = big_s
+ self.big_t = big_t
+ self.alt_s = alt_s
+ self.alt_t = alt_t
+
+ def __str__(self):
+ return str(self.value)
+
+ @classmethod
+ def create(cls, i, v):
+ return ChineseNumberDigit(i, v[0], v[1], v[2], v[3])
+
+
+class ChineseMath(ChineseChar):
+ """
+ 涓枃鏁颁綅瀛楃
+ """
+
+ def __init__(self, simplified, traditional, symbol, expression=None):
+ super(ChineseMath, self).__init__(simplified, traditional)
+ self.symbol = symbol
+ self.expression = expression
+ self.big_s = simplified
+ self.big_t = traditional
+
+
+CC, CNU, CND, CM = ChineseChar, ChineseNumberUnit, ChineseNumberDigit, ChineseMath
+
+
+class NumberSystem(object):
+ """
+ 涓枃鏁板瓧绯荤粺
+ """
+ pass
+
+
+class MathSymbol(object):
+ """
+ 鐢ㄤ簬涓枃鏁板瓧绯荤粺鐨勬暟瀛︾鍙� (绻�/绠�浣�), e.g.
+ positive = ['姝�', '姝�']
+ negative = ['璐�', '璨�']
+ point = ['鐐�', '榛�']
+ """
+
+ def __init__(self, positive, negative, point):
+ self.positive = positive
+ self.negative = negative
+ self.point = point
+
+ def __iter__(self):
+ for v in self.__dict__.values():
+ yield v
+
+
+# class OtherSymbol(object):
+# """
+# 鍏朵粬绗﹀彿
+# """
+#
+# def __init__(self, sil):
+# self.sil = sil
+#
+# def __iter__(self):
+# for v in self.__dict__.values():
+# yield v
+
+
+# ================================================================================ #
+# basic utils
+# ================================================================================ #
+def create_system(numbering_type=NUMBERING_TYPES[1]):
+ """
+ 鏍规嵁鏁板瓧绯荤粺绫诲瀷杩斿洖鍒涘缓鐩稿簲鐨勬暟瀛楃郴缁燂紝榛樿涓� mid
+ NUMBERING_TYPES = ['low', 'mid', 'high']: 涓枃鏁板瓧绯荤粺绫诲瀷
+ low: '鍏�' = '浜�' * '鍗�' = $10^{9}$, '浜�' = '鍏�' * '鍗�', etc.
+ mid: '鍏�' = '浜�' * '涓�' = $10^{12}$, '浜�' = '鍏�' * '涓�', etc.
+ high: '鍏�' = '浜�' * '浜�' = $10^{16}$, '浜�' = '鍏�' * '鍏�', etc.
+ 杩斿洖瀵瑰簲鐨勬暟瀛楃郴缁�
+ """
+
+ # chinese number units of '浜�' and larger
+ all_larger_units = zip(
+ LARGER_CHINESE_NUMERING_UNITS_SIMPLIFIED, LARGER_CHINESE_NUMERING_UNITS_TRADITIONAL)
+ larger_units = [CNU.create(i, v, numbering_type, False)
+ for i, v in enumerate(all_larger_units)]
+ # chinese number units of '鍗�, 鐧�, 鍗�, 涓�'
+ all_smaller_units = zip(
+ SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED, SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL)
+ smaller_units = [CNU.create(i, v, small_unit=True)
+ for i, v in enumerate(all_smaller_units)]
+ # digis
+ chinese_digis = zip(CHINESE_DIGIS, CHINESE_DIGIS,
+ BIG_CHINESE_DIGIS_SIMPLIFIED, BIG_CHINESE_DIGIS_TRADITIONAL)
+ digits = [CND.create(i, v) for i, v in enumerate(chinese_digis)]
+ digits[0].alt_s, digits[0].alt_t = ZERO_ALT, ZERO_ALT
+ digits[1].alt_s, digits[1].alt_t = ONE_ALT, ONE_ALT
+ digits[2].alt_s, digits[2].alt_t = TWO_ALTS[0], TWO_ALTS[1]
+
+ # symbols
+ positive_cn = CM(POSITIVE[0], POSITIVE[1], '+', lambda x: x)
+ negative_cn = CM(NEGATIVE[0], NEGATIVE[1], '-', lambda x: -x)
+ point_cn = CM(POINT[0], POINT[1], '.', lambda x,
+ y: float(str(x) + '.' + str(y)))
+ # sil_cn = CM(SIL[0], SIL[1], '-', lambda x, y: float(str(x) + '-' + str(y)))
+ system = NumberSystem()
+ system.units = smaller_units + larger_units
+ system.digits = digits
+ system.math = MathSymbol(positive_cn, negative_cn, point_cn)
+ # system.symbols = OtherSymbol(sil_cn)
+ return system
+
+
+def chn2num(chinese_string, numbering_type=NUMBERING_TYPES[1]):
+
+ def get_symbol(char, system):
+ for u in system.units:
+ if char in [u.traditional, u.simplified, u.big_s, u.big_t]:
+ return u
+ for d in system.digits:
+ if char in [d.traditional, d.simplified, d.big_s, d.big_t, d.alt_s, d.alt_t]:
+ return d
+ for m in system.math:
+ if char in [m.traditional, m.simplified]:
+ return m
+
+ def string2symbols(chinese_string, system):
+ int_string, dec_string = chinese_string, ''
+ for p in [system.math.point.simplified, system.math.point.traditional]:
+ if p in chinese_string:
+ int_string, dec_string = chinese_string.split(p)
+ break
+ return [get_symbol(c, system) for c in int_string], \
+ [get_symbol(c, system) for c in dec_string]
+
+ def correct_symbols(integer_symbols, system):
+ """
+ 涓�鐧惧叓 to 涓�鐧惧叓鍗�
+ 涓�浜夸竴鍗冧笁鐧句竾 to 涓�浜� 涓�鍗冧竾 涓夌櫨涓�
+ """
+
+ if integer_symbols and isinstance(integer_symbols[0], CNU):
+ if integer_symbols[0].power == 1:
+ integer_symbols = [system.digits[1]] + integer_symbols
+
+ if len(integer_symbols) > 1:
+ if isinstance(integer_symbols[-1], CND) and isinstance(integer_symbols[-2], CNU):
+ integer_symbols.append(
+ CNU(integer_symbols[-2].power - 1, None, None, None, None))
+
+ result = []
+ unit_count = 0
+ for s in integer_symbols:
+ if isinstance(s, CND):
+ result.append(s)
+ unit_count = 0
+ elif isinstance(s, CNU):
+ current_unit = CNU(s.power, None, None, None, None)
+ unit_count += 1
+
+ if unit_count == 1:
+ result.append(current_unit)
+ elif unit_count > 1:
+ for i in range(len(result)):
+ if isinstance(result[-i - 1], CNU) and result[-i - 1].power < current_unit.power:
+ result[-i - 1] = CNU(result[-i - 1].power +
+ current_unit.power, None, None, None, None)
+ return result
+
+ def compute_value(integer_symbols):
+ """
+ Compute the value.
+ When current unit is larger than previous unit, current unit * all previous units will be used as all previous units.
+ e.g. '涓ゅ崈涓�' = 2000 * 10000 not 2000 + 10000
+ """
+ value = [0]
+ last_power = 0
+ for s in integer_symbols:
+ if isinstance(s, CND):
+ value[-1] = s.value
+ elif isinstance(s, CNU):
+ value[-1] *= pow(10, s.power)
+ if s.power > last_power:
+ value[:-1] = list(map(lambda v: v *
+ pow(10, s.power), value[:-1]))
+ last_power = s.power
+ value.append(0)
+ return sum(value)
+
+ system = create_system(numbering_type)
+ int_part, dec_part = string2symbols(chinese_string, system)
+ int_part = correct_symbols(int_part, system)
+ int_str = str(compute_value(int_part))
+ dec_str = ''.join([str(d.value) for d in dec_part])
+ if dec_part:
+ return '{0}.{1}'.format(int_str, dec_str)
+ else:
+ return int_str
+
+
+def num2chn(number_string, numbering_type=NUMBERING_TYPES[1], big=False,
+ traditional=False, alt_zero=False, alt_one=False, alt_two=True,
+ use_zeros=True, use_units=True):
+
+ def get_value(value_string, use_zeros=True):
+
+ striped_string = value_string.lstrip('0')
+
+ # record nothing if all zeros
+ if not striped_string:
+ return []
+
+ # record one digits
+ elif len(striped_string) == 1:
+ if use_zeros and len(value_string) != len(striped_string):
+ return [system.digits[0], system.digits[int(striped_string)]]
+ else:
+ return [system.digits[int(striped_string)]]
+
+ # recursively record multiple digits
+ else:
+ result_unit = next(u for u in reversed(
+ system.units) if u.power < len(striped_string))
+ result_string = value_string[:-result_unit.power]
+ return get_value(result_string) + [result_unit] + get_value(striped_string[-result_unit.power:])
+
+ system = create_system(numbering_type)
+
+ int_dec = number_string.split('.')
+ if len(int_dec) == 1:
+ int_string = int_dec[0]
+ dec_string = ""
+ elif len(int_dec) == 2:
+ int_string = int_dec[0]
+ dec_string = int_dec[1]
+ else:
+ raise ValueError(
+ "invalid input num string with more than one dot: {}".format(number_string))
+
+ if use_units and len(int_string) > 1:
+ result_symbols = get_value(int_string)
+ else:
+ result_symbols = [system.digits[int(c)] for c in int_string]
+ dec_symbols = [system.digits[int(c)] for c in dec_string]
+ if dec_string:
+ result_symbols += [system.math.point] + dec_symbols
+
+ if alt_two:
+ liang = CND(2, system.digits[2].alt_s, system.digits[2].alt_t,
+ system.digits[2].big_s, system.digits[2].big_t)
+ for i, v in enumerate(result_symbols):
+ if isinstance(v, CND) and v.value == 2:
+ next_symbol = result_symbols[i +
+ 1] if i < len(result_symbols) - 1 else None
+ previous_symbol = result_symbols[i - 1] if i > 0 else None
+ if isinstance(next_symbol, CNU) and isinstance(previous_symbol, (CNU, type(None))):
+ if next_symbol.power != 1 and ((previous_symbol is None) or (previous_symbol.power != 1)):
+ result_symbols[i] = liang
+
+ # if big is True, '涓�' will not be used and `alt_two` has no impact on output
+ if big:
+ attr_name = 'big_'
+ if traditional:
+ attr_name += 't'
+ else:
+ attr_name += 's'
+ else:
+ if traditional:
+ attr_name = 'traditional'
+ else:
+ attr_name = 'simplified'
+
+ result = ''.join([getattr(s, attr_name) for s in result_symbols])
+
+ # if not use_zeros:
+ # result = result.strip(getattr(system.digits[0], attr_name))
+
+ if alt_zero:
+ result = result.replace(
+ getattr(system.digits[0], attr_name), system.digits[0].alt_s)
+
+ if alt_one:
+ result = result.replace(
+ getattr(system.digits[1], attr_name), system.digits[1].alt_s)
+
+ for i, p in enumerate(POINT):
+ if result.startswith(p):
+ return CHINESE_DIGIS[0] + result
+
+ # ^10, 11, .., 19
+ if len(result) >= 2 and result[1] in [SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED[0],
+ SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL[0]] and \
+ result[0] in [CHINESE_DIGIS[1], BIG_CHINESE_DIGIS_SIMPLIFIED[1], BIG_CHINESE_DIGIS_TRADITIONAL[1]]:
+ result = result[1:]
+
+ return result
+
+
+# ================================================================================ #
+# different types of rewriters
+# ================================================================================ #
+class Cardinal:
+ """
+ CARDINAL绫�
+ """
+
+ def __init__(self, cardinal=None, chntext=None):
+ self.cardinal = cardinal
+ self.chntext = chntext
+
+ def chntext2cardinal(self):
+ return chn2num(self.chntext)
+
+ def cardinal2chntext(self):
+ return num2chn(self.cardinal)
+
+class Digit:
+ """
+ DIGIT绫�
+ """
+
+ def __init__(self, digit=None, chntext=None):
+ self.digit = digit
+ self.chntext = chntext
+
+ # def chntext2digit(self):
+ # return chn2num(self.chntext)
+
+ def digit2chntext(self):
+ return num2chn(self.digit, alt_two=False, use_units=False)
+
+
+class TelePhone:
+ """
+ TELEPHONE绫�
+ """
+
+ def __init__(self, telephone=None, raw_chntext=None, chntext=None):
+ self.telephone = telephone
+ self.raw_chntext = raw_chntext
+ self.chntext = chntext
+
+ # def chntext2telephone(self):
+ # sil_parts = self.raw_chntext.split('<SIL>')
+ # self.telephone = '-'.join([
+ # str(chn2num(p)) for p in sil_parts
+ # ])
+ # return self.telephone
+
+ def telephone2chntext(self, fixed=False):
+
+ if fixed:
+ sil_parts = self.telephone.split('-')
+ self.raw_chntext = '<SIL>'.join([
+ num2chn(part, alt_two=False, use_units=False) for part in sil_parts
+ ])
+ self.chntext = self.raw_chntext.replace('<SIL>', '')
+ else:
+ sp_parts = self.telephone.strip('+').split()
+ self.raw_chntext = '<SP>'.join([
+ num2chn(part, alt_two=False, use_units=False) for part in sp_parts
+ ])
+ self.chntext = self.raw_chntext.replace('<SP>', '')
+ return self.chntext
+
+
+class Fraction:
+ """
+ FRACTION绫�
+ """
+
+ def __init__(self, fraction=None, chntext=None):
+ self.fraction = fraction
+ self.chntext = chntext
+
+ def chntext2fraction(self):
+ denominator, numerator = self.chntext.split('鍒嗕箣')
+ return chn2num(numerator) + '/' + chn2num(denominator)
+
+ def fraction2chntext(self):
+ numerator, denominator = self.fraction.split('/')
+ return num2chn(denominator) + '鍒嗕箣' + num2chn(numerator)
+
+
+class Date:
+ """
+ DATE绫�
+ """
+
+ def __init__(self, date=None, chntext=None):
+ self.date = date
+ self.chntext = chntext
+
+ # def chntext2date(self):
+ # chntext = self.chntext
+ # try:
+ # year, other = chntext.strip().split('骞�', maxsplit=1)
+ # year = Digit(chntext=year).digit2chntext() + '骞�'
+ # except ValueError:
+ # other = chntext
+ # year = ''
+ # if other:
+ # try:
+ # month, day = other.strip().split('鏈�', maxsplit=1)
+ # month = Cardinal(chntext=month).chntext2cardinal() + '鏈�'
+ # except ValueError:
+ # day = chntext
+ # month = ''
+ # if day:
+ # day = Cardinal(chntext=day[:-1]).chntext2cardinal() + day[-1]
+ # else:
+ # month = ''
+ # day = ''
+ # date = year + month + day
+ # self.date = date
+ # return self.date
+
+ def date2chntext(self):
+ date = self.date
+ try:
+ year, other = date.strip().split('骞�', 1)
+ year = Digit(digit=year).digit2chntext() + '骞�'
+ except ValueError:
+ other = date
+ year = ''
+ if other:
+ try:
+ month, day = other.strip().split('鏈�', 1)
+ month = Cardinal(cardinal=month).cardinal2chntext() + '鏈�'
+ except ValueError:
+ day = date
+ month = ''
+ if day:
+ day = Cardinal(cardinal=day[:-1]).cardinal2chntext() + day[-1]
+ else:
+ month = ''
+ day = ''
+ chntext = year + month + day
+ self.chntext = chntext
+ return self.chntext
+
+
+class Money:
+ """
+ MONEY绫�
+ """
+
+ def __init__(self, money=None, chntext=None):
+ self.money = money
+ self.chntext = chntext
+
+ # def chntext2money(self):
+ # return self.money
+
+ def money2chntext(self):
+ money = self.money
+ pattern = re.compile(r'(\d+(\.\d+)?)')
+ matchers = pattern.findall(money)
+ if matchers:
+ for matcher in matchers:
+ money = money.replace(matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext())
+ self.chntext = money
+ return self.chntext
+
+
+class Percentage:
+ """
+ PERCENTAGE绫�
+ """
+
+ def __init__(self, percentage=None, chntext=None):
+ self.percentage = percentage
+ self.chntext = chntext
+
+ def chntext2percentage(self):
+ return chn2num(self.chntext.strip().strip('鐧惧垎涔�')) + '%'
+
+ def percentage2chntext(self):
+ return '鐧惧垎涔�' + num2chn(self.percentage.strip().strip('%'))
+
+
+def remove_erhua(text, er_whitelist):
+ """
+ 鍘婚櫎鍎垮寲闊宠瘝涓殑鍎�:
+ 浠栧コ鍎垮湪閭h竟鍎� -> 浠栧コ鍎垮湪閭h竟
+ """
+
+ er_pattern = re.compile(er_whitelist)
+ new_str=''
+ while re.search('鍎�',text):
+ a = re.search('鍎�',text).span()
+ remove_er_flag = 0
+
+ if er_pattern.search(text):
+ b = er_pattern.search(text).span()
+ if b[0] <= a[0]:
+ remove_er_flag = 1
+
+ if remove_er_flag == 0 :
+ new_str = new_str + text[0:a[0]]
+ text = text[a[1]:]
+ else:
+ new_str = new_str + text[0:b[1]]
+ text = text[b[1]:]
+
+ text = new_str + text
+ return text
+
+# ================================================================================ #
+# NSW Normalizer
+# ================================================================================ #
+class NSWNormalizer:
+ def __init__(self, raw_text):
+ self.raw_text = '^' + raw_text + '$'
+ self.norm_text = ''
+
+ def _particular(self):
+ text = self.norm_text
+ pattern = re.compile(r"(([a-zA-Z]+)浜�([a-zA-Z]+))")
+ matchers = pattern.findall(text)
+ if matchers:
+ # print('particular')
+ for matcher in matchers:
+ text = text.replace(matcher[0], matcher[1]+'2'+matcher[2], 1)
+ self.norm_text = text
+ return self.norm_text
+
+ def normalize(self):
+ text = self.raw_text
+
+ # 瑙勮寖鍖栨棩鏈�
+ pattern = re.compile(r"\D+((([089]\d|(19|20)\d{2})骞�)?(\d{1,2}鏈�(\d{1,2}[鏃ュ彿])?)?)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('date')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Date(date=matcher[0]).date2chntext(), 1)
+
+ # 瑙勮寖鍖栭噾閽�
+ pattern = re.compile(r"\D+((\d+(\.\d+)?)[澶氫綑鍑燷?" + CURRENCY_UNITS + r"(\d" + CURRENCY_UNITS + r"?)?)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('money')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Money(money=matcher[0]).money2chntext(), 1)
+
+ # 瑙勮寖鍖栧浐璇�/鎵嬫満鍙风爜
+ # 鎵嬫満
+ # http://www.jihaoba.com/news/show/13680
+ # 绉诲姩锛�139銆�138銆�137銆�136銆�135銆�134銆�159銆�158銆�157銆�150銆�151銆�152銆�188銆�187銆�182銆�183銆�184銆�178銆�198
+ # 鑱旈�氾細130銆�131銆�132銆�156銆�155銆�186銆�185銆�176
+ # 鐢典俊锛�133銆�153銆�189銆�180銆�181銆�177
+ pattern = re.compile(r"\D((\+?86 ?)?1([38]\d|5[0-35-9]|7[678]|9[89])\d{8})\D")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('telephone')
+ for matcher in matchers:
+ text = text.replace(matcher[0], TelePhone(telephone=matcher[0]).telephone2chntext(), 1)
+ # 鍥鸿瘽
+ pattern = re.compile(r"\D((0(10|2[1-3]|[3-9]\d{2})-?)?[1-9]\d{6,7})\D")
+ matchers = pattern.findall(text)
+ if matchers:
+ # print('fixed telephone')
+ for matcher in matchers:
+ text = text.replace(matcher[0], TelePhone(telephone=matcher[0]).telephone2chntext(fixed=True), 1)
+
+ # 瑙勮寖鍖栧垎鏁�
+ pattern = re.compile(r"(\d+/\d+)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('fraction')
+ for matcher in matchers:
+ text = text.replace(matcher, Fraction(fraction=matcher).fraction2chntext(), 1)
+
+ # 瑙勮寖鍖栫櫨鍒嗘暟
+ text = text.replace('锛�', '%')
+ pattern = re.compile(r"(\d+(\.\d+)?%)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('percentage')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Percentage(percentage=matcher[0]).percentage2chntext(), 1)
+
+ # 瑙勮寖鍖栫函鏁�+閲忚瘝
+ pattern = re.compile(r"(\d+(\.\d+)?)[澶氫綑鍑燷?" + COM_QUANTIFIERS)
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('cardinal+quantifier')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext(), 1)
+
+ # 瑙勮寖鍖栨暟瀛楃紪鍙�
+ pattern = re.compile(r"(\d{4,32})")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('digit')
+ for matcher in matchers:
+ text = text.replace(matcher, Digit(digit=matcher).digit2chntext(), 1)
+
+ # 瑙勮寖鍖栫函鏁�
+ pattern = re.compile(r"(\d+(\.\d+)?)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('cardinal')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext(), 1)
+
+ self.norm_text = text
+ self._particular()
+
+ return self.norm_text.lstrip('^').rstrip('$')
+
+
+def nsw_test_case(raw_text):
+ print('I:' + raw_text)
+ print('O:' + NSWNormalizer(raw_text).normalize())
+ print('')
+
+
+def nsw_test():
+ nsw_test_case('鍥鸿瘽锛�0595-23865596鎴�23880880銆�')
+ nsw_test_case('鍥鸿瘽锛�0595-23865596鎴�23880880銆�')
+ nsw_test_case('鎵嬫満锛�+86 19859213959鎴�15659451527銆�')
+ nsw_test_case('鍒嗘暟锛�32477/76391銆�')
+ nsw_test_case('鐧惧垎鏁帮細80.03%銆�')
+ nsw_test_case('缂栧彿锛�31520181154418銆�')
+ nsw_test_case('绾暟锛�2983.07鍏嬫垨12345.60绫炽��')
+ nsw_test_case('鏃ユ湡锛�1999骞�2鏈�20鏃ユ垨09骞�3鏈�15鍙枫��')
+ nsw_test_case('閲戦挶锛�12鍧�5锛�34.5鍏冿紝20.1涓�')
+ nsw_test_case('鐗规畩锛歄2O鎴朆2C銆�')
+ nsw_test_case('3456涓囧惃')
+ nsw_test_case('2938涓�')
+ nsw_test_case('938')
+ nsw_test_case('浠婂ぉ鍚冧簡115涓皬绗煎寘231涓澶�')
+ nsw_test_case('鏈�62锛呯殑姒傜巼')
+
+
+if __name__ == '__main__':
+ #nsw_test()
+
+ p = argparse.ArgumentParser()
+ p.add_argument('ifile', help='input filename, assume utf-8 encoding')
+ p.add_argument('ofile', help='output filename')
+ p.add_argument('--to_upper', action='store_true', help='convert to upper case')
+ p.add_argument('--to_lower', action='store_true', help='convert to lower case')
+ p.add_argument('--has_key', action='store_true', help="input text has Kaldi's key as first field.")
+ p.add_argument('--remove_fillers', type=bool, default=True, help='remove filler chars such as "鍛�, 鍟�"')
+ p.add_argument('--remove_erhua', type=bool, default=True, help='remove erhua chars such as "杩欏効"')
+ p.add_argument('--log_interval', type=int, default=10000, help='log interval in number of processed lines')
+ args = p.parse_args()
+
+ ifile = codecs.open(args.ifile, 'r', 'utf8')
+ ofile = codecs.open(args.ofile, 'w+', 'utf8')
+
+ n = 0
+ for l in ifile:
+ key = ''
+ text = ''
+ if args.has_key:
+ cols = l.split(maxsplit=1)
+ key = cols[0]
+ if len(cols) == 2:
+ text = cols[1].strip()
+ else:
+ text = ''
+ else:
+ text = l.strip()
+
+ # cases
+ if args.to_upper and args.to_lower:
+ sys.stderr.write('text norm: to_upper OR to_lower?')
+ exit(1)
+ if args.to_upper:
+ text = text.upper()
+ if args.to_lower:
+ text = text.lower()
+
+ # Filler chars removal
+ if args.remove_fillers:
+ for ch in FILLER_CHARS:
+ text = text.replace(ch, '')
+
+ if args.remove_erhua:
+ text = remove_erhua(text, ER_WHITELIST)
+
+ # NSW(Non-Standard-Word) normalization
+ text = NSWNormalizer(text).normalize()
+
+ # Punctuations removal
+ old_chars = CHINESE_PUNC_LIST + string.punctuation # includes all CN and EN punctuations
+ new_chars = ' ' * len(old_chars)
+ del_chars = ''
+ text = text.translate(str.maketrans(old_chars, new_chars, del_chars))
+
+ #
+ if args.has_key:
+ ofile.write(key + '\t' + text + '\n')
+ else:
+ ofile.write(text + '\n')
+
+ n += 1
+ if n % args.log_interval == 0:
+ sys.stderr.write("text norm: {} lines done.\n".format(n))
+
+ sys.stderr.write("text norm: {} lines done in total.\n".format(n))
+
+ ifile.close()
+ ofile.close()
diff --git a/egs_modelscope/wenetspeech/paraformer/modelscope_utils b/egs_modelscope/wenetspeech/paraformer/modelscope_utils
deleted file mode 120000
index fc97768..0000000
--- a/egs_modelscope/wenetspeech/paraformer/modelscope_utils
+++ /dev/null
@@ -1 +0,0 @@
-../../common/modelscope_utils
\ No newline at end of file
diff --git a/egs_modelscope/wenetspeech/paraformer/modelscope_utils/download_model.py b/egs_modelscope/wenetspeech/paraformer/modelscope_utils/download_model.py
new file mode 100755
index 0000000..51ba6b8
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/modelscope_utils/download_model.py
@@ -0,0 +1,25 @@
+#!/usr/bin/env python3
+import argparse
+
+from modelscope.pipelines import pipeline
+from modelscope.utils.constant import Tasks
+
+if __name__ == '__main__':
+ parser = argparse.ArgumentParser(
+ description="download model configs",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument("--model_name",
+ type=str,
+ default="speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-pytorch",
+ help="model name in modelscope")
+ parser.add_argument("--model_revision",
+ type=str,
+ default="v1.0.3",
+ help="model revision in modelscope")
+ args = parser.parse_args()
+
+ inference_pipeline = pipeline(
+ task=Tasks.auto_speech_recognition,
+ model='damo/{}'.format(args.model_name),
+ model_revision=args.model_revision)
diff --git a/egs_modelscope/wenetspeech/paraformer/modelscope_utils/modelscope_infer.sh b/egs_modelscope/wenetspeech/paraformer/modelscope_utils/modelscope_infer.sh
new file mode 100755
index 0000000..a0c606f
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/modelscope_utils/modelscope_infer.sh
@@ -0,0 +1,88 @@
+#!/usr/bin/env bash
+
+set -e
+set -u
+set -o pipefail
+
+data_dir=
+exp_dir=
+model_name=
+model_revision=
+inference_nj=32
+gpuid_list="0,1,2,3"
+njob=32
+gpu_inference=true
+
+test_sets="dev test"
+decode_cmd=utils/run.pl
+
+# LM configs
+use_lm=false
+beam_size=1
+lm_weight=0.0
+
+. utils/parse_options.sh
+
+if ${gpu_inference}; then
+ _ngpu=1
+else
+ _ngpu=0
+fi
+
+# download model from modelscope
+python modelscope_utils/download_model.py \
+ --model_name ${model_name} --model_revision ${model_revision}
+
+modelscope_dir=${HOME}/.cache/modelscope/hub/damo/${model_name}
+
+
+for dset in ${test_sets}; do
+ _dir=${exp_dir}/${model_name}/decode_asr/${dset}
+ _logdir=${_dir}/logdir
+ _data=${data_dir}/${dset}
+ if [ -d ${_dir} ]; then
+ echo "${_dir} is already exists. if you want to decode again, please delete ${_dir} first."
+ exit 1
+ else
+ mkdir -p "${_dir}"
+ mkdir -p "${_logdir}"
+ fi
+
+ if "${use_lm}"; then
+ cp ${modelscope_dir}/decoding.yaml ${modelscope_dir}/decoding.yaml.back
+ sed -i "s#beam_size: [0-9]*#beam_size: `echo $beam_size`#g" ${modelscope_dir}/decoding.yaml
+ sed -i "s#lm_weight: 0.[0-9]*#lm_weight: `echo $lm_weight`#g" ${modelscope_dir}/decoding.yaml
+ fi
+
+ for n in $(seq "${inference_nj}"); do
+ split_scps+=" ${_logdir}/keys.${n}.scp"
+ done
+ # shellcheck disable=SC2086
+ utils/split_scp.pl "${data_dir}/${dset}/wav.scp" ${split_scps}
+
+ echo "Decoding started... log: '${_logdir}/asr_inference.*.log'"
+ # shellcheck disable=SC2086
+ ${decode_cmd} --max-jobs-run "${inference_nj}" JOB=1:"${inference_nj}" "${_logdir}"/asr_inference.JOB.log \
+ python -m funasr.bin.modelscope_infer \
+ --model_name ${model_name} \
+ --model_revision ${model_revision} \
+ --wav_list ${_logdir}/keys.JOB.scp \
+ --output_file ${_logdir}/text.JOB \
+ --gpuid_list ${gpuid_list} \
+ --njob ${njob} \
+ --ngpu ${_ngpu} \
+
+ for i in $(seq ${inference_nj}); do
+ cat ${_logdir}/text.${i}
+ done | sort -k1 >${_dir}/text
+
+ python utils/proce_text.py ${_dir}/text ${_dir}/text.proc
+ python utils/proce_text.py ${_data}/text ${_data}/text.proc
+ python utils/compute_wer.py ${_data}/text.proc ${_dir}/text.proc ${_dir}/text.cer
+ tail -n 3 ${_dir}/text.cer > ${_dir}/text.cer.txt
+ cat ${_dir}/text.cer.txt
+done
+
+if "${use_lm}"; then
+ mv ${modelscope_dir}/decoding.yaml.back ${modelscope_dir}/decoding.yaml
+fi
diff --git a/egs_modelscope/wenetspeech/paraformer/modelscope_utils/update_config.py b/egs_modelscope/wenetspeech/paraformer/modelscope_utils/update_config.py
new file mode 100644
index 0000000..88466ed
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/modelscope_utils/update_config.py
@@ -0,0 +1,41 @@
+import yaml
+import argparse
+
+def update_dct(fin_configs, root):
+ if root == {}:
+ return {}
+ for root_key, root_value in root.items():
+ if not isinstance(root[root_key],dict):
+ fin_configs[root_key] = root[root_key]
+ else:
+ result = update_dct(fin_configs[root_key], root[root_key])
+ fin_configs[root_key] = result
+ return fin_configs
+
+if __name__ == '__main__':
+ parser = argparse.ArgumentParser(
+ description="update configs",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument("--modelscope_config",
+ type=str,
+ help="modelscope config file")
+ parser.add_argument("--finetune_config",
+ type=str,
+ help="finetune config file")
+ parser.add_argument("--output_config",
+ type=str,
+ help="output config file")
+ args = parser.parse_args()
+
+ with open(args.modelscope_config) as f:
+ modelscope_configs = yaml.safe_load(f)
+
+ with open(args.finetune_config) as f:
+ finetune_configs = yaml.safe_load(f)
+
+ # update configs, e.g., lr, batch_size, ...
+ modelscope_configs = update_dct(modelscope_configs, finetune_configs)
+
+ with open(args.output_config, "w") as f:
+ yaml.dump(modelscope_configs, f, indent=4)
diff --git a/egs_modelscope/wenetspeech/paraformer/paraformer_large_infer.sh b/egs_modelscope/wenetspeech/paraformer/paraformer_large_infer.sh
index 88e0990..e55e47d 100755
--- a/egs_modelscope/wenetspeech/paraformer/paraformer_large_infer.sh
+++ b/egs_modelscope/wenetspeech/paraformer/paraformer_large_infer.sh
@@ -8,7 +8,7 @@
data_dir=
exp_dir=
model_name=speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-pytorch
-model_revision="v1.0.3" # please do not modify the model revision
+model_revision="v1.0.4" # please do not modify the model revision
inference_nj=32
gpuid_list="0,1" # set gpus, e.g., gpuid_list="0,1"
ngpu=$(echo $gpuid_list | awk -F "," '{print NF}')
diff --git a/egs_modelscope/wenetspeech/paraformer/utils b/egs_modelscope/wenetspeech/paraformer/utils
deleted file mode 120000
index 37d9761..0000000
--- a/egs_modelscope/wenetspeech/paraformer/utils
+++ /dev/null
@@ -1 +0,0 @@
-../../../egs/aishell/tranformer/utils/
\ No newline at end of file
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/__init__.py b/egs_modelscope/wenetspeech/paraformer/utils/__init__.py
new file mode 100644
index 0000000..e69de29
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/__init__.py
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/apply_cmvn.py b/egs_modelscope/wenetspeech/paraformer/utils/apply_cmvn.py
new file mode 100755
index 0000000..b5c5086
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/apply_cmvn.py
@@ -0,0 +1,79 @@
+from kaldiio import ReadHelper
+from kaldiio import WriteHelper
+
+import argparse
+import json
+import math
+import numpy as np
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="apply cmvn",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--ark-file",
+ "-a",
+ default=False,
+ required=True,
+ type=str,
+ help="fbank ark file",
+ )
+ parser.add_argument(
+ "--cmvn-file",
+ "-c",
+ default=False,
+ required=True,
+ type=str,
+ help="cmvn file",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ ark_file = args.output_dir + "/feats." + str(args.ark_index) + ".ark"
+ scp_file = args.output_dir + "/feats." + str(args.ark_index) + ".scp"
+ ark_writer = WriteHelper('ark,scp:{},{}'.format(ark_file, scp_file))
+
+ with open(args.cmvn_file) as f:
+ cmvn_stats = json.load(f)
+
+ means = cmvn_stats['mean_stats']
+ vars = cmvn_stats['var_stats']
+ total_frames = cmvn_stats['total_frames']
+
+ for i in range(len(means)):
+ means[i] /= total_frames
+ vars[i] = vars[i] / total_frames - means[i] * means[i]
+ if vars[i] < 1.0e-20:
+ vars[i] = 1.0e-20
+ vars[i] = 1.0 / math.sqrt(vars[i])
+
+ with ReadHelper('ark:{}'.format(args.ark_file)) as ark_reader:
+ for key, mat in ark_reader:
+ mat = (mat - means) * vars
+ ark_writer(key, mat)
+
+
+if __name__ == '__main__':
+ main()
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/apply_cmvn.sh b/egs_modelscope/wenetspeech/paraformer/utils/apply_cmvn.sh
new file mode 100755
index 0000000..f8fd1d1
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/apply_cmvn.sh
@@ -0,0 +1,29 @@
+#!/usr/bin/env bash
+
+. ./path.sh || exit 1;
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+fbankdir=$1
+cmvn_file=$2
+logdir=$3
+output_dir=$4
+
+dump_dir=${output_dir}/ark; mkdir -p ${dump_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/apply_cmvn.JOB.log \
+ python utils/apply_cmvn.py -a $fbankdir/ark/feats.JOB.ark \
+ -c $cmvn_file -i JOB -o ${dump_dir} \
+ || exit 1;
+
+for n in $(seq $nj); do
+ cat ${dump_dir}/feats.$n.scp || exit 1
+done > ${output_dir}/feats.scp || exit 1
+
+echo "$0: Succeeded apply cmvn"
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/apply_lfr_and_cmvn.py b/egs_modelscope/wenetspeech/paraformer/utils/apply_lfr_and_cmvn.py
new file mode 100755
index 0000000..50d18d1
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/apply_lfr_and_cmvn.py
@@ -0,0 +1,143 @@
+from kaldiio import ReadHelper, WriteHelper
+
+import argparse
+import numpy as np
+
+
+def build_LFR_features(inputs, m=7, n=6):
+ LFR_inputs = []
+ T = inputs.shape[0]
+ T_lfr = int(np.ceil(T / n))
+ left_padding = np.tile(inputs[0], ((m - 1) // 2, 1))
+ inputs = np.vstack((left_padding, inputs))
+ T = T + (m - 1) // 2
+ for i in range(T_lfr):
+ if m <= T - i * n:
+ LFR_inputs.append(np.hstack(inputs[i * n:i * n + m]))
+ else:
+ num_padding = m - (T - i * n)
+ frame = np.hstack(inputs[i * n:])
+ for _ in range(num_padding):
+ frame = np.hstack((frame, inputs[-1]))
+ LFR_inputs.append(frame)
+ return np.vstack(LFR_inputs)
+
+
+def build_CMVN_features(inputs, mvn_file): # noqa
+ with open(mvn_file, 'r', encoding='utf-8') as f:
+ lines = f.readlines()
+
+ add_shift_list = []
+ rescale_list = []
+ for i in range(len(lines)):
+ line_item = lines[i].split()
+ if line_item[0] == '<AddShift>':
+ line_item = lines[i + 1].split()
+ if line_item[0] == '<LearnRateCoef>':
+ add_shift_line = line_item[3:(len(line_item) - 1)]
+ add_shift_list = list(add_shift_line)
+ continue
+ elif line_item[0] == '<Rescale>':
+ line_item = lines[i + 1].split()
+ if line_item[0] == '<LearnRateCoef>':
+ rescale_line = line_item[3:(len(line_item) - 1)]
+ rescale_list = list(rescale_line)
+ continue
+
+ for j in range(inputs.shape[0]):
+ for k in range(inputs.shape[1]):
+ add_shift_value = add_shift_list[k]
+ rescale_value = rescale_list[k]
+ inputs[j, k] = float(inputs[j, k]) + float(add_shift_value)
+ inputs[j, k] = float(inputs[j, k]) * float(rescale_value)
+
+ return inputs
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="apply low_frame_rate and cmvn",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--ark-file",
+ "-a",
+ default=False,
+ required=True,
+ type=str,
+ help="fbank ark file",
+ )
+ parser.add_argument(
+ "--lfr",
+ "-f",
+ default=True,
+ type=str,
+ help="low frame rate",
+ )
+ parser.add_argument(
+ "--lfr-m",
+ "-m",
+ default=7,
+ type=int,
+ help="number of frames to stack",
+ )
+ parser.add_argument(
+ "--lfr-n",
+ "-n",
+ default=6,
+ type=int,
+ help="number of frames to skip",
+ )
+ parser.add_argument(
+ "--cmvn-file",
+ "-c",
+ default=False,
+ required=True,
+ type=str,
+ help="global cmvn file",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ dump_ark_file = args.output_dir + "/feats." + str(args.ark_index) + ".ark"
+ dump_scp_file = args.output_dir + "/feats." + str(args.ark_index) + ".scp"
+ shape_file = args.output_dir + "/len." + str(args.ark_index)
+ ark_writer = WriteHelper('ark,scp:{},{}'.format(dump_ark_file, dump_scp_file))
+
+ shape_writer = open(shape_file, 'w')
+ with ReadHelper('ark:{}'.format(args.ark_file)) as ark_reader:
+ for key, mat in ark_reader:
+ if args.lfr:
+ lfr = build_LFR_features(mat, args.lfr_m, args.lfr_n)
+ else:
+ lfr = mat
+ cmvn = build_CMVN_features(lfr, args.cmvn_file)
+ dims = cmvn.shape[1]
+ lens = cmvn.shape[0]
+ shape_writer.write(key + " " + str(lens) + "," + str(dims) + '\n')
+ ark_writer(key, cmvn)
+
+
+if __name__ == '__main__':
+ main()
+
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/apply_lfr_and_cmvn.sh b/egs_modelscope/wenetspeech/paraformer/utils/apply_lfr_and_cmvn.sh
new file mode 100755
index 0000000..3119fdb
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/apply_lfr_and_cmvn.sh
@@ -0,0 +1,38 @@
+#!/usr/bin/env bash
+
+
+# Begin configuration section.
+nj=32
+cmd=utils/run.pl
+
+# feature configuration
+lfr=True
+lfr_m=7
+lfr_n=6
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+fbankdir=$1
+cmvn_file=$2
+logdir=$3
+output_dir=$4
+
+dump_dir=${output_dir}/ark; mkdir -p ${dump_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/apply_lfr_and_cmvn.JOB.log \
+ python utils/apply_lfr_and_cmvn.py -a $fbankdir/ark/feats.JOB.ark \
+ -f $lfr -m $lfr_m -n $lfr_n -c $cmvn_file -i JOB -o ${dump_dir} \
+ || exit 1;
+
+for n in $(seq $nj); do
+ cat ${dump_dir}/feats.$n.scp || exit 1
+done > ${output_dir}/feats.scp || exit 1
+
+for n in $(seq $nj); do
+ cat ${dump_dir}/len.$n || exit 1
+done > ${output_dir}/speech_shape || exit 1
+
+echo "$0: Succeeded apply low frame rate and cmvn"
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/combine_cmvn_file.py b/egs_modelscope/wenetspeech/paraformer/utils/combine_cmvn_file.py
new file mode 100755
index 0000000..b2974a4
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/combine_cmvn_file.py
@@ -0,0 +1,73 @@
+import argparse
+import json
+import numpy as np
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="combine cmvn file",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--dims",
+ "-d",
+ default=80,
+ type=int,
+ help="feature dims",
+ )
+ parser.add_argument(
+ "--cmvn-dir",
+ "-c",
+ default=False,
+ required=True,
+ type=str,
+ help="cmvn dir",
+ )
+
+ parser.add_argument(
+ "--nj",
+ "-n",
+ default=1,
+ required=True,
+ type=int,
+ help="num of cmvn file",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ total_means = np.zeros(args.dims)
+ total_vars = np.zeros(args.dims)
+ total_frames = 0
+
+ cmvn_file = args.output_dir + "/cmvn.json"
+
+ for i in range(1, args.nj+1):
+ with open(args.cmvn_dir + "/cmvn." + str(i) + ".json", "r") as fin:
+ cmvn_stats = json.load(fin)
+
+ total_means += np.array(cmvn_stats["mean_stats"])
+ total_vars += np.array(cmvn_stats["var_stats"])
+ total_frames += cmvn_stats["total_frames"]
+
+ cmvn_info = {
+ 'mean_stats': list(total_means.tolist()),
+ 'var_stats': list(total_vars.tolist()),
+ 'total_frames': total_frames
+ }
+ with open(cmvn_file, 'w') as fout:
+ fout.write(json.dumps(cmvn_info))
+
+
+if __name__ == '__main__':
+ main()
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/compute_cmvn.py b/egs_modelscope/wenetspeech/paraformer/utils/compute_cmvn.py
new file mode 100755
index 0000000..2b96e26
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/compute_cmvn.py
@@ -0,0 +1,74 @@
+from kaldiio import ReadHelper
+
+import argparse
+import numpy as np
+import json
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="computer global cmvn",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--dims",
+ "-d",
+ default=80,
+ type=int,
+ help="feature dims",
+ )
+ parser.add_argument(
+ "--ark-file",
+ "-a",
+ default=False,
+ required=True,
+ type=str,
+ help="fbank ark file",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ ark_file = args.ark_file + "/feats." + str(args.ark_index) + ".ark"
+ cmvn_file = args.output_dir + "/cmvn." + str(args.ark_index) + ".json"
+
+ mean_stats = np.zeros(args.dims)
+ var_stats = np.zeros(args.dims)
+ total_frames = 0
+
+ with ReadHelper('ark:{}'.format(ark_file)) as ark_reader:
+ for key, mat in ark_reader:
+ mean_stats += np.sum(mat, axis=0)
+ var_stats += np.sum(np.square(mat), axis=0)
+ total_frames += mat.shape[0]
+
+ cmvn_info = {
+ 'mean_stats': list(mean_stats.tolist()),
+ 'var_stats': list(var_stats.tolist()),
+ 'total_frames': total_frames
+ }
+ with open(cmvn_file, 'w') as fout:
+ fout.write(json.dumps(cmvn_info))
+
+
+if __name__ == '__main__':
+ main()
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/compute_cmvn.sh b/egs_modelscope/wenetspeech/paraformer/utils/compute_cmvn.sh
new file mode 100755
index 0000000..12173ee
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/compute_cmvn.sh
@@ -0,0 +1,25 @@
+#!/usr/bin/env bash
+
+. ./path.sh || exit 1;
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+feats_dim=80
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+fbankdir=$1
+logdir=$2
+
+output_dir=${fbankdir}/cmvn; mkdir -p ${output_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/cmvn.JOB.log \
+ python utils/compute_cmvn.py -d ${feats_dim} -a $fbankdir/ark -i JOB -o ${output_dir} \
+ || exit 1;
+
+python utils/combine_cmvn_file.py -d ${feats_dim} -c ${output_dir} -n $nj -o $fbankdir
+
+echo "$0: Succeeded compute global cmvn"
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/compute_fbank.py b/egs_modelscope/wenetspeech/paraformer/utils/compute_fbank.py
new file mode 100755
index 0000000..d03b5a8
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/compute_fbank.py
@@ -0,0 +1,153 @@
+from kaldiio import WriteHelper
+
+import argparse
+import numpy as np
+import json
+import torch
+import torchaudio
+import torchaudio.compliance.kaldi as kaldi
+
+
+def compute_fbank(wav_file,
+ num_mel_bins=80,
+ frame_length=25,
+ frame_shift=10,
+ dither=0.0,
+ resample_rate=16000,
+ speed=1.0):
+
+ waveform, sample_rate = torchaudio.load(wav_file)
+ if resample_rate != sample_rate:
+ waveform = torchaudio.transforms.Resample(orig_freq=sample_rate,
+ new_freq=resample_rate)(waveform)
+ if speed != 1.0:
+ waveform, _ = torchaudio.sox_effects.apply_effects_tensor(
+ waveform, resample_rate,
+ [['speed', str(speed)], ['rate', str(resample_rate)]]
+ )
+
+ waveform = waveform * (1 << 15)
+ mat = kaldi.fbank(waveform,
+ num_mel_bins=num_mel_bins,
+ frame_length=frame_length,
+ frame_shift=frame_shift,
+ dither=dither,
+ energy_floor=0.0,
+ window_type='hamming',
+ sample_frequency=resample_rate)
+
+ return mat.numpy()
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="computer features",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--wav-lists",
+ "-w",
+ default=False,
+ required=True,
+ type=str,
+ help="input wav lists",
+ )
+ parser.add_argument(
+ "--text-files",
+ "-t",
+ default=False,
+ required=True,
+ type=str,
+ help="input text files",
+ )
+ parser.add_argument(
+ "--dims",
+ "-d",
+ default=80,
+ type=int,
+ help="feature dims",
+ )
+ parser.add_argument(
+ "--sample-frequency",
+ "-s",
+ default=16000,
+ type=int,
+ help="sample frequency",
+ )
+ parser.add_argument(
+ "--speed-perturb",
+ "-p",
+ default="1.0",
+ type=str,
+ help="speed perturb",
+ )
+ parser.add_argument(
+ "--ark-index",
+ "-a",
+ default=1,
+ required=True,
+ type=int,
+ help="ark index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ ark_file = args.output_dir + "/ark/feats." + str(args.ark_index) + ".ark"
+ scp_file = args.output_dir + "/ark/feats." + str(args.ark_index) + ".scp"
+ text_file = args.output_dir + "/txt/text." + str(args.ark_index) + ".txt"
+ feats_shape_file = args.output_dir + "/ark/len." + str(args.ark_index)
+ text_shape_file = args.output_dir + "/txt/len." + str(args.ark_index)
+
+ ark_writer = WriteHelper('ark,scp:{},{}'.format(ark_file, scp_file))
+ text_writer = open(text_file, 'w')
+ feats_shape_writer = open(feats_shape_file, 'w')
+ text_shape_writer = open(text_shape_file, 'w')
+
+ speed_perturb_list = args.speed_perturb.split(',')
+
+ for speed in speed_perturb_list:
+ with open(args.wav_lists, 'r', encoding='utf-8') as wavfile:
+ with open(args.text_files, 'r', encoding='utf-8') as textfile:
+ for wav, text in zip(wavfile, textfile):
+ s_w = wav.strip().split()
+ wav_id = s_w[0]
+ wav_file = s_w[1]
+
+ s_t = text.strip().split()
+ text_id = s_t[0]
+ txt = s_t[1:]
+ fbank = compute_fbank(wav_file,
+ num_mel_bins=args.dims,
+ resample_rate=args.sample_frequency,
+ speed=float(speed)
+ )
+ feats_dims = fbank.shape[1]
+ feats_lens = fbank.shape[0]
+ txt_lens = len(txt)
+ if speed == "1.0":
+ wav_id_sp = wav_id
+ else:
+ wav_id_sp = wav_id + "_sp" + speed
+
+ feats_shape_writer.write(wav_id_sp + " " + str(feats_lens) + "," + str(feats_dims) + '\n')
+ text_shape_writer.write(wav_id_sp + " " + str(txt_lens) + '\n')
+
+ text_writer.write(wav_id_sp + " " + " ".join(txt) + '\n')
+ ark_writer(wav_id_sp, fbank)
+
+
+if __name__ == '__main__':
+ main()
+
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/compute_fbank.sh b/egs_modelscope/wenetspeech/paraformer/utils/compute_fbank.sh
new file mode 100755
index 0000000..92a4fe6
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/compute_fbank.sh
@@ -0,0 +1,51 @@
+#!/usr/bin/env bash
+
+. ./path.sh || exit 1;
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+
+# feature configuration
+feats_dim=80
+sample_frequency=16000
+speed_perturb="1.0"
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+data=$1
+logdir=$2
+fbankdir=$3
+
+[ ! -f $data/wav.scp ] && echo "$0: no such file $data/wav.scp" && exit 1;
+[ ! -f $data/text ] && echo "$0: no such file $data/text" && exit 1;
+
+python utils/split_data.py $data $data $nj
+
+ark_dir=${fbankdir}/ark; mkdir -p ${ark_dir}
+text_dir=${fbankdir}/txt; mkdir -p ${text_dir}
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/make_fbank.JOB.log \
+ python utils/compute_fbank.py -w $data/split${nj}/JOB/wav.scp -t $data/split${nj}/JOB/text \
+ -d $feats_dim -s $sample_frequency -p ${speed_perturb} -a JOB -o ${fbankdir} \
+ || exit 1;
+
+for n in $(seq $nj); do
+ cat ${ark_dir}/feats.$n.scp || exit 1
+done > $fbankdir/feats.scp || exit 1
+
+for n in $(seq $nj); do
+ cat ${text_dir}/text.$n.txt || exit 1
+done > $fbankdir/text || exit 1
+
+for n in $(seq $nj); do
+ cat ${ark_dir}/len.$n || exit 1
+done > $fbankdir/speech_shape || exit 1
+
+for n in $(seq $nj); do
+ cat ${text_dir}/len.$n || exit 1
+done > $fbankdir/text_shape || exit 1
+
+echo "$0: Succeeded compute FBANK features"
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/compute_wer.py b/egs_modelscope/wenetspeech/paraformer/utils/compute_wer.py
new file mode 100755
index 0000000..349a3f6
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/compute_wer.py
@@ -0,0 +1,157 @@
+import os
+import numpy as np
+import sys
+
+def compute_wer(ref_file,
+ hyp_file,
+ cer_detail_file):
+ rst = {
+ 'Wrd': 0,
+ 'Corr': 0,
+ 'Ins': 0,
+ 'Del': 0,
+ 'Sub': 0,
+ 'Snt': 0,
+ 'Err': 0.0,
+ 'S.Err': 0.0,
+ 'wrong_words': 0,
+ 'wrong_sentences': 0
+ }
+
+ hyp_dict = {}
+ ref_dict = {}
+ with open(hyp_file, 'r') as hyp_reader:
+ for line in hyp_reader:
+ key = line.strip().split()[0]
+ value = line.strip().split()[1:]
+ hyp_dict[key] = value
+ with open(ref_file, 'r') as ref_reader:
+ for line in ref_reader:
+ key = line.strip().split()[0]
+ value = line.strip().split()[1:]
+ ref_dict[key] = value
+
+ cer_detail_writer = open(cer_detail_file, 'w')
+ for hyp_key in hyp_dict:
+ if hyp_key in ref_dict:
+ out_item = compute_wer_by_line(hyp_dict[hyp_key], ref_dict[hyp_key])
+ rst['Wrd'] += out_item['nwords']
+ rst['Corr'] += out_item['cor']
+ rst['wrong_words'] += out_item['wrong']
+ rst['Ins'] += out_item['ins']
+ rst['Del'] += out_item['del']
+ rst['Sub'] += out_item['sub']
+ rst['Snt'] += 1
+ if out_item['wrong'] > 0:
+ rst['wrong_sentences'] += 1
+ cer_detail_writer.write(hyp_key + print_cer_detail(out_item) + '\n')
+ cer_detail_writer.write("ref:" + '\t' + "".join(ref_dict[hyp_key]) + '\n')
+ cer_detail_writer.write("hyp:" + '\t' + "".join(hyp_dict[hyp_key]) + '\n')
+
+ if rst['Wrd'] > 0:
+ rst['Err'] = round(rst['wrong_words'] * 100 / rst['Wrd'], 2)
+ if rst['Snt'] > 0:
+ rst['S.Err'] = round(rst['wrong_sentences'] * 100 / rst['Snt'], 2)
+
+ cer_detail_writer.write('\n')
+ cer_detail_writer.write("%WER " + str(rst['Err']) + " [ " + str(rst['wrong_words'])+ " / " + str(rst['Wrd']) +
+ ", " + str(rst['Ins']) + " ins, " + str(rst['Del']) + " del, " + str(rst['Sub']) + " sub ]" + '\n')
+ cer_detail_writer.write("%SER " + str(rst['S.Err']) + " [ " + str(rst['wrong_sentences']) + " / " + str(rst['Snt']) + " ]" + '\n')
+ cer_detail_writer.write("Scored " + str(len(hyp_dict)) + " sentences, " + str(len(hyp_dict) - rst['Snt']) + " not present in hyp." + '\n')
+
+
+def compute_wer_by_line(hyp,
+ ref):
+ hyp = list(map(lambda x: x.lower(), hyp))
+ ref = list(map(lambda x: x.lower(), ref))
+
+ len_hyp = len(hyp)
+ len_ref = len(ref)
+
+ cost_matrix = np.zeros((len_hyp + 1, len_ref + 1), dtype=np.int16)
+
+ ops_matrix = np.zeros((len_hyp + 1, len_ref + 1), dtype=np.int8)
+
+ for i in range(len_hyp + 1):
+ cost_matrix[i][0] = i
+ for j in range(len_ref + 1):
+ cost_matrix[0][j] = j
+
+ for i in range(1, len_hyp + 1):
+ for j in range(1, len_ref + 1):
+ if hyp[i - 1] == ref[j - 1]:
+ cost_matrix[i][j] = cost_matrix[i - 1][j - 1]
+ else:
+ substitution = cost_matrix[i - 1][j - 1] + 1
+ insertion = cost_matrix[i - 1][j] + 1
+ deletion = cost_matrix[i][j - 1] + 1
+
+ compare_val = [substitution, insertion, deletion]
+
+ min_val = min(compare_val)
+ operation_idx = compare_val.index(min_val) + 1
+ cost_matrix[i][j] = min_val
+ ops_matrix[i][j] = operation_idx
+
+ match_idx = []
+ i = len_hyp
+ j = len_ref
+ rst = {
+ 'nwords': len_ref,
+ 'cor': 0,
+ 'wrong': 0,
+ 'ins': 0,
+ 'del': 0,
+ 'sub': 0
+ }
+ while i >= 0 or j >= 0:
+ i_idx = max(0, i)
+ j_idx = max(0, j)
+
+ if ops_matrix[i_idx][j_idx] == 0: # correct
+ if i - 1 >= 0 and j - 1 >= 0:
+ match_idx.append((j - 1, i - 1))
+ rst['cor'] += 1
+
+ i -= 1
+ j -= 1
+
+ elif ops_matrix[i_idx][j_idx] == 2: # insert
+ i -= 1
+ rst['ins'] += 1
+
+ elif ops_matrix[i_idx][j_idx] == 3: # delete
+ j -= 1
+ rst['del'] += 1
+
+ elif ops_matrix[i_idx][j_idx] == 1: # substitute
+ i -= 1
+ j -= 1
+ rst['sub'] += 1
+
+ if i < 0 and j >= 0:
+ rst['del'] += 1
+ elif j < 0 and i >= 0:
+ rst['ins'] += 1
+
+ match_idx.reverse()
+ wrong_cnt = cost_matrix[len_hyp][len_ref]
+ rst['wrong'] = wrong_cnt
+
+ return rst
+
+def print_cer_detail(rst):
+ return ("(" + "nwords=" + str(rst['nwords']) + ",cor=" + str(rst['cor'])
+ + ",ins=" + str(rst['ins']) + ",del=" + str(rst['del']) + ",sub="
+ + str(rst['sub']) + ") corr:" + '{:.2%}'.format(rst['cor']/rst['nwords'])
+ + ",cer:" + '{:.2%}'.format(rst['wrong']/rst['nwords']))
+
+if __name__ == '__main__':
+ if len(sys.argv) != 4:
+ print("usage : python compute-wer.py test.ref test.hyp test.wer")
+ sys.exit(0)
+
+ ref_file = sys.argv[1]
+ hyp_file = sys.argv[2]
+ cer_detail_file = sys.argv[3]
+ compute_wer(ref_file, hyp_file, cer_detail_file)
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/error_rate_zh b/egs_modelscope/wenetspeech/paraformer/utils/error_rate_zh
new file mode 100755
index 0000000..6871a07
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/error_rate_zh
@@ -0,0 +1,370 @@
+#!/usr/bin/env python3
+# coding=utf8
+
+# Copyright 2021 Jiayu DU
+
+import sys
+import argparse
+import json
+import logging
+logging.basicConfig(stream=sys.stderr, level=logging.INFO, format='[%(levelname)s] %(message)s')
+
+DEBUG = None
+
+def GetEditType(ref_token, hyp_token):
+ if ref_token == None and hyp_token != None:
+ return 'I'
+ elif ref_token != None and hyp_token == None:
+ return 'D'
+ elif ref_token == hyp_token:
+ return 'C'
+ elif ref_token != hyp_token:
+ return 'S'
+ else:
+ raise RuntimeError
+
+class AlignmentArc:
+ def __init__(self, src, dst, ref, hyp):
+ self.src = src
+ self.dst = dst
+ self.ref = ref
+ self.hyp = hyp
+ self.edit_type = GetEditType(ref, hyp)
+
+def similarity_score_function(ref_token, hyp_token):
+ return 0 if (ref_token == hyp_token) else -1.0
+
+def insertion_score_function(token):
+ return -1.0
+
+def deletion_score_function(token):
+ return -1.0
+
+def EditDistance(
+ ref,
+ hyp,
+ similarity_score_function = similarity_score_function,
+ insertion_score_function = insertion_score_function,
+ deletion_score_function = deletion_score_function):
+ assert(len(ref) != 0)
+ class DPState:
+ def __init__(self):
+ self.score = -float('inf')
+ # backpointer
+ self.prev_r = None
+ self.prev_h = None
+
+ def print_search_grid(S, R, H, fstream):
+ print(file=fstream)
+ for r in range(R):
+ for h in range(H):
+ print(F'[{r},{h}]:{S[r][h].score:4.3f}:({S[r][h].prev_r},{S[r][h].prev_h}) ', end='', file=fstream)
+ print(file=fstream)
+
+ R = len(ref) + 1
+ H = len(hyp) + 1
+
+ # Construct DP search space, a (R x H) grid
+ S = [ [] for r in range(R) ]
+ for r in range(R):
+ S[r] = [ DPState() for x in range(H) ]
+
+ # initialize DP search grid origin, S(r = 0, h = 0)
+ S[0][0].score = 0.0
+ S[0][0].prev_r = None
+ S[0][0].prev_h = None
+
+ # initialize REF axis
+ for r in range(1, R):
+ S[r][0].score = S[r-1][0].score + deletion_score_function(ref[r-1])
+ S[r][0].prev_r = r-1
+ S[r][0].prev_h = 0
+
+ # initialize HYP axis
+ for h in range(1, H):
+ S[0][h].score = S[0][h-1].score + insertion_score_function(hyp[h-1])
+ S[0][h].prev_r = 0
+ S[0][h].prev_h = h-1
+
+ best_score = S[0][0].score
+ best_state = (0, 0)
+
+ for r in range(1, R):
+ for h in range(1, H):
+ sub_or_cor_score = similarity_score_function(ref[r-1], hyp[h-1])
+ new_score = S[r-1][h-1].score + sub_or_cor_score
+ if new_score >= S[r][h].score:
+ S[r][h].score = new_score
+ S[r][h].prev_r = r-1
+ S[r][h].prev_h = h-1
+
+ del_score = deletion_score_function(ref[r-1])
+ new_score = S[r-1][h].score + del_score
+ if new_score >= S[r][h].score:
+ S[r][h].score = new_score
+ S[r][h].prev_r = r - 1
+ S[r][h].prev_h = h
+
+ ins_score = insertion_score_function(hyp[h-1])
+ new_score = S[r][h-1].score + ins_score
+ if new_score >= S[r][h].score:
+ S[r][h].score = new_score
+ S[r][h].prev_r = r
+ S[r][h].prev_h = h-1
+
+ best_score = S[R-1][H-1].score
+ best_state = (R-1, H-1)
+
+ if DEBUG:
+ print_search_grid(S, R, H, sys.stderr)
+
+ # Backtracing best alignment path, i.e. a list of arcs
+ # arc = (src, dst, ref, hyp, edit_type)
+ # src/dst = (r, h), where r/h refers to search grid state-id along Ref/Hyp axis
+ best_path = []
+ r, h = best_state[0], best_state[1]
+ prev_r, prev_h = S[r][h].prev_r, S[r][h].prev_h
+ score = S[r][h].score
+ # loop invariant:
+ # 1. (prev_r, prev_h) -> (r, h) is a "forward arc" on best alignment path
+ # 2. score is the value of point(r, h) on DP search grid
+ while prev_r != None or prev_h != None:
+ src = (prev_r, prev_h)
+ dst = (r, h)
+ if (r == prev_r + 1 and h == prev_h + 1): # Substitution or correct
+ arc = AlignmentArc(src, dst, ref[prev_r], hyp[prev_h])
+ elif (r == prev_r + 1 and h == prev_h): # Deletion
+ arc = AlignmentArc(src, dst, ref[prev_r], None)
+ elif (r == prev_r and h == prev_h + 1): # Insertion
+ arc = AlignmentArc(src, dst, None, hyp[prev_h])
+ else:
+ raise RuntimeError
+ best_path.append(arc)
+ r, h = prev_r, prev_h
+ prev_r, prev_h = S[r][h].prev_r, S[r][h].prev_h
+ score = S[r][h].score
+
+ best_path.reverse()
+ return (best_path, best_score)
+
+def PrettyPrintAlignment(alignment, stream = sys.stderr):
+ def get_token_str(token):
+ if token == None:
+ return "*"
+ return token
+
+ def is_double_width_char(ch):
+ if (ch >= '\u4e00') and (ch <= '\u9fa5'): # codepoint ranges for Chinese chars
+ return True
+ # TODO: support other double-width-char language such as Japanese, Korean
+ else:
+ return False
+
+ def display_width(token_str):
+ m = 0
+ for c in token_str:
+ if is_double_width_char(c):
+ m += 2
+ else:
+ m += 1
+ return m
+
+ R = ' REF : '
+ H = ' HYP : '
+ E = ' EDIT : '
+ for arc in alignment:
+ r = get_token_str(arc.ref)
+ h = get_token_str(arc.hyp)
+ e = arc.edit_type if arc.edit_type != 'C' else ''
+
+ nr, nh, ne = display_width(r), display_width(h), display_width(e)
+ n = max(nr, nh, ne) + 1
+
+ R += r + ' ' * (n-nr)
+ H += h + ' ' * (n-nh)
+ E += e + ' ' * (n-ne)
+
+ print(R, file=stream)
+ print(H, file=stream)
+ print(E, file=stream)
+
+def CountEdits(alignment):
+ c, s, i, d = 0, 0, 0, 0
+ for arc in alignment:
+ if arc.edit_type == 'C':
+ c += 1
+ elif arc.edit_type == 'S':
+ s += 1
+ elif arc.edit_type == 'I':
+ i += 1
+ elif arc.edit_type == 'D':
+ d += 1
+ else:
+ raise RuntimeError
+ return (c, s, i, d)
+
+def ComputeTokenErrorRate(c, s, i, d):
+ return 100.0 * (s + d + i) / (s + d + c)
+
+def ComputeSentenceErrorRate(num_err_utts, num_utts):
+ assert(num_utts != 0)
+ return 100.0 * num_err_utts / num_utts
+
+
+class EvaluationResult:
+ def __init__(self):
+ self.num_ref_utts = 0
+ self.num_hyp_utts = 0
+ self.num_eval_utts = 0 # seen in both ref & hyp
+ self.num_hyp_without_ref = 0
+
+ self.C = 0
+ self.S = 0
+ self.I = 0
+ self.D = 0
+ self.token_error_rate = 0.0
+
+ self.num_utts_with_error = 0
+ self.sentence_error_rate = 0.0
+
+ def to_json(self):
+ return json.dumps(self.__dict__)
+
+ def to_kaldi(self):
+ info = (
+ F'%WER {self.token_error_rate:.2f} [ {self.S + self.D + self.I} / {self.C + self.S + self.D}, {self.I} ins, {self.D} del, {self.S} sub ]\n'
+ F'%SER {self.sentence_error_rate:.2f} [ {self.num_utts_with_error} / {self.num_eval_utts} ]\n'
+ )
+ return info
+
+ def to_sclite(self):
+ return "TODO"
+
+ def to_espnet(self):
+ return "TODO"
+
+ def to_summary(self):
+ #return json.dumps(self.__dict__, indent=4)
+ summary = (
+ '==================== Overall Statistics ====================\n'
+ F'num_ref_utts: {self.num_ref_utts}\n'
+ F'num_hyp_utts: {self.num_hyp_utts}\n'
+ F'num_hyp_without_ref: {self.num_hyp_without_ref}\n'
+ F'num_eval_utts: {self.num_eval_utts}\n'
+ F'sentence_error_rate: {self.sentence_error_rate:.2f}%\n'
+ F'token_error_rate: {self.token_error_rate:.2f}%\n'
+ F'token_stats:\n'
+ F' - tokens:{self.C + self.S + self.D:>7}\n'
+ F' - edits: {self.S + self.I + self.D:>7}\n'
+ F' - cor: {self.C:>7}\n'
+ F' - sub: {self.S:>7}\n'
+ F' - ins: {self.I:>7}\n'
+ F' - del: {self.D:>7}\n'
+ '============================================================\n'
+ )
+ return summary
+
+
+class Utterance:
+ def __init__(self, uid, text):
+ self.uid = uid
+ self.text = text
+
+
+def LoadUtterances(filepath, format):
+ utts = {}
+ if format == 'text': # utt_id word1 word2 ...
+ with open(filepath, 'r', encoding='utf8') as f:
+ for line in f:
+ line = line.strip()
+ if line:
+ cols = line.split(maxsplit=1)
+ assert(len(cols) == 2 or len(cols) == 1)
+ uid = cols[0]
+ text = cols[1] if len(cols) == 2 else ''
+ if utts.get(uid) != None:
+ raise RuntimeError(F'Found duplicated utterence id {uid}')
+ utts[uid] = Utterance(uid, text)
+ else:
+ raise RuntimeError(F'Unsupported text format {format}')
+ return utts
+
+
+def tokenize_text(text, tokenizer):
+ if tokenizer == 'whitespace':
+ return text.split()
+ elif tokenizer == 'char':
+ return [ ch for ch in ''.join(text.split()) ]
+ else:
+ raise RuntimeError(F'ERROR: Unsupported tokenizer {tokenizer}')
+
+
+if __name__ == '__main__':
+ parser = argparse.ArgumentParser()
+ # optional
+ parser.add_argument('--tokenizer', choices=['whitespace', 'char'], default='whitespace', help='whitespace for WER, char for CER')
+ parser.add_argument('--ref-format', choices=['text'], default='text', help='reference format, first col is utt_id, the rest is text')
+ parser.add_argument('--hyp-format', choices=['text'], default='text', help='hypothesis format, first col is utt_id, the rest is text')
+ # required
+ parser.add_argument('--ref', type=str, required=True, help='input reference file')
+ parser.add_argument('--hyp', type=str, required=True, help='input hypothesis file')
+
+ parser.add_argument('result_file', type=str)
+ args = parser.parse_args()
+ logging.info(args)
+
+ ref_utts = LoadUtterances(args.ref, args.ref_format)
+ hyp_utts = LoadUtterances(args.hyp, args.hyp_format)
+
+ r = EvaluationResult()
+
+ # check valid utterances in hyp that have matched non-empty reference
+ eval_utts = []
+ r.num_hyp_without_ref = 0
+ for uid in sorted(hyp_utts.keys()):
+ if uid in ref_utts.keys(): # TODO: efficiency
+ if ref_utts[uid].text.strip(): # non-empty reference
+ eval_utts.append(uid)
+ else:
+ logging.warn(F'Found {uid} with empty reference, skipping...')
+ else:
+ logging.warn(F'Found {uid} without reference, skipping...')
+ r.num_hyp_without_ref += 1
+
+ r.num_hyp_utts = len(hyp_utts)
+ r.num_ref_utts = len(ref_utts)
+ r.num_eval_utts = len(eval_utts)
+
+ with open(args.result_file, 'w+', encoding='utf8') as fo:
+ for uid in eval_utts:
+ ref = ref_utts[uid]
+ hyp = hyp_utts[uid]
+
+ alignment, score = EditDistance(
+ tokenize_text(ref.text, args.tokenizer),
+ tokenize_text(hyp.text, args.tokenizer)
+ )
+
+ c, s, i, d = CountEdits(alignment)
+ utt_ter = ComputeTokenErrorRate(c, s, i, d)
+
+ # utt-level evaluation result
+ print(F'{{"uid":{uid}, "score":{score}, "ter":{utt_ter:.2f}, "cor":{c}, "sub":{s}, "ins":{i}, "del":{d}}}', file=fo)
+ PrettyPrintAlignment(alignment, fo)
+
+ r.C += c
+ r.S += s
+ r.I += i
+ r.D += d
+
+ if utt_ter > 0:
+ r.num_utts_with_error += 1
+
+ # corpus level evaluation result
+ r.sentence_error_rate = ComputeSentenceErrorRate(r.num_utts_with_error, r.num_eval_utts)
+ r.token_error_rate = ComputeTokenErrorRate(r.C, r.S, r.I, r.D)
+
+ print(r.to_summary(), file=fo)
+
+ print(r.to_json())
+ print(r.to_kaldi())
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/extract_embeds.py b/egs_modelscope/wenetspeech/paraformer/utils/extract_embeds.py
new file mode 100755
index 0000000..7b817d8
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/extract_embeds.py
@@ -0,0 +1,47 @@
+from transformers import AutoTokenizer, AutoModel, pipeline
+import numpy as np
+import sys
+import os
+import torch
+from kaldiio import WriteHelper
+import re
+text_file_json = sys.argv[1]
+out_ark = sys.argv[2]
+out_scp = sys.argv[3]
+out_shape = sys.argv[4]
+device = int(sys.argv[5])
+model_path = sys.argv[6]
+
+model = AutoModel.from_pretrained(model_path)
+tokenizer = AutoTokenizer.from_pretrained(model_path)
+extractor = pipeline(task="feature-extraction", model=model, tokenizer=tokenizer, device=device)
+
+with open(text_file_json, 'r') as f:
+ js = f.readlines()
+
+
+f_shape = open(out_shape, "w")
+with WriteHelper('ark,scp:{},{}'.format(out_ark, out_scp)) as writer:
+ with torch.no_grad():
+ for idx, line in enumerate(js):
+ id, tokens = line.strip().split(" ", 1)
+ tokens = re.sub(" ", "", tokens.strip())
+ tokens = ' '.join([j for j in tokens])
+ token_num = len(tokens.split(" "))
+ outputs = extractor(tokens)
+ outputs = np.array(outputs)
+ embeds = outputs[0, 1:-1, :]
+
+ token_num_embeds, dim = embeds.shape
+ if token_num == token_num_embeds:
+ writer(id, embeds)
+ shape_line = "{} {},{}\n".format(id, token_num_embeds, dim)
+ f_shape.write(shape_line)
+ else:
+ print("{}, size has changed, {}, {}, {}".format(id, token_num, token_num_embeds, tokens))
+
+
+
+f_shape.close()
+
+
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/filter_scp.pl b/egs_modelscope/wenetspeech/paraformer/utils/filter_scp.pl
new file mode 100755
index 0000000..003530d
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/filter_scp.pl
@@ -0,0 +1,87 @@
+#!/usr/bin/env perl
+# Copyright 2010-2012 Microsoft Corporation
+# Johns Hopkins University (author: Daniel Povey)
+
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+# This script takes a list of utterance-ids or any file whose first field
+# of each line is an utterance-id, and filters an scp
+# file (or any file whose "n-th" field is an utterance id), printing
+# out only those lines whose "n-th" field is in id_list. The index of
+# the "n-th" field is 1, by default, but can be changed by using
+# the -f <n> switch
+
+$exclude = 0;
+$field = 1;
+$shifted = 0;
+
+do {
+ $shifted=0;
+ if ($ARGV[0] eq "--exclude") {
+ $exclude = 1;
+ shift @ARGV;
+ $shifted=1;
+ }
+ if ($ARGV[0] eq "-f") {
+ $field = $ARGV[1];
+ shift @ARGV; shift @ARGV;
+ $shifted=1
+ }
+} while ($shifted);
+
+if(@ARGV < 1 || @ARGV > 2) {
+ die "Usage: filter_scp.pl [--exclude] [-f <field-to-filter-on>] id_list [in.scp] > out.scp \n" .
+ "Prints only the input lines whose f'th field (default: first) is in 'id_list'.\n" .
+ "Note: only the first field of each line in id_list matters. With --exclude, prints\n" .
+ "only the lines that were *not* in id_list.\n" .
+ "Caution: previously, the -f option was interpreted as a zero-based field index.\n" .
+ "If your older scripts (written before Oct 2014) stopped working and you used the\n" .
+ "-f option, add 1 to the argument.\n" .
+ "See also: scripts/filter_scp.pl .\n";
+}
+
+
+$idlist = shift @ARGV;
+open(F, "<$idlist") || die "Could not open id-list file $idlist";
+while(<F>) {
+ @A = split;
+ @A>=1 || die "Invalid id-list file line $_";
+ $seen{$A[0]} = 1;
+}
+
+if ($field == 1) { # Treat this as special case, since it is common.
+ while(<>) {
+ $_ =~ m/\s*(\S+)\s*/ || die "Bad line $_, could not get first field.";
+ # $1 is what we filter on.
+ if ((!$exclude && $seen{$1}) || ($exclude && !defined $seen{$1})) {
+ print $_;
+ }
+ }
+} else {
+ while(<>) {
+ @A = split;
+ @A > 0 || die "Invalid scp file line $_";
+ @A >= $field || die "Invalid scp file line $_";
+ if ((!$exclude && $seen{$A[$field-1]}) || ($exclude && !defined $seen{$A[$field-1]})) {
+ print $_;
+ }
+ }
+}
+
+# tests:
+# the following should print "foo 1"
+# ( echo foo 1; echo bar 2 ) | scripts/filter_scp.pl <(echo foo)
+# the following should print "bar 2".
+# ( echo foo 1; echo bar 2 ) | scripts/filter_scp.pl -f 2 <(echo 2)
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/fix_data.sh b/egs_modelscope/wenetspeech/paraformer/utils/fix_data.sh
new file mode 100755
index 0000000..32cdde5
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/fix_data.sh
@@ -0,0 +1,35 @@
+#!/usr/bin/env bash
+
+echo "$0 $@"
+data_dir=$1
+
+if [ ! -f ${data_dir}/wav.scp ]; then
+ echo "$0: wav.scp is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/text ]; then
+ echo "$0: text is not found"
+ exit 1;
+fi
+
+
+
+mkdir -p ${data_dir}/.backup
+
+awk '{print $1}' ${data_dir}/wav.scp > ${data_dir}/.backup/wav_id
+awk '{print $1}' ${data_dir}/text > ${data_dir}/.backup/text_id
+
+sort ${data_dir}/.backup/wav_id ${data_dir}/.backup/text_id | uniq -d > ${data_dir}/.backup/id
+
+cp ${data_dir}/wav.scp ${data_dir}/.backup/wav.scp
+cp ${data_dir}/text ${data_dir}/.backup/text
+
+mv ${data_dir}/wav.scp ${data_dir}/wav.scp.bak
+mv ${data_dir}/text ${data_dir}/text.bak
+
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/wav.scp.bak > ${data_dir}/wav.scp
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/text.bak > ${data_dir}/text
+
+rm ${data_dir}/wav.scp.bak
+rm ${data_dir}/text.bak
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/fix_data_feat.sh b/egs_modelscope/wenetspeech/paraformer/utils/fix_data_feat.sh
new file mode 100755
index 0000000..2c92d7f
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/fix_data_feat.sh
@@ -0,0 +1,52 @@
+#!/usr/bin/env bash
+
+echo "$0 $@"
+data_dir=$1
+
+if [ ! -f ${data_dir}/feats.scp ]; then
+ echo "$0: feats.scp is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/text ]; then
+ echo "$0: text is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/speech_shape ]; then
+ echo "$0: feature lengths is not found"
+ exit 1;
+fi
+
+if [ ! -f ${data_dir}/text_shape ]; then
+ echo "$0: text lengths is not found"
+ exit 1;
+fi
+
+mkdir -p ${data_dir}/.backup
+
+awk '{print $1}' ${data_dir}/feats.scp > ${data_dir}/.backup/wav_id
+awk '{print $1}' ${data_dir}/text > ${data_dir}/.backup/text_id
+
+sort ${data_dir}/.backup/wav_id ${data_dir}/.backup/text_id | uniq -d > ${data_dir}/.backup/id
+
+cp ${data_dir}/feats.scp ${data_dir}/.backup/feats.scp
+cp ${data_dir}/text ${data_dir}/.backup/text
+cp ${data_dir}/speech_shape ${data_dir}/.backup/speech_shape
+cp ${data_dir}/text_shape ${data_dir}/.backup/text_shape
+
+mv ${data_dir}/feats.scp ${data_dir}/feats.scp.bak
+mv ${data_dir}/text ${data_dir}/text.bak
+mv ${data_dir}/speech_shape ${data_dir}/speech_shape.bak
+mv ${data_dir}/text_shape ${data_dir}/text_shape.bak
+
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/feats.scp.bak > ${data_dir}/feats.scp
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/text.bak > ${data_dir}/text
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/speech_shape.bak > ${data_dir}/speech_shape
+utils/filter_scp.pl -f 1 ${data_dir}/.backup/id ${data_dir}/text_shape.bak > ${data_dir}/text_shape
+
+rm ${data_dir}/feats.scp.bak
+rm ${data_dir}/text.bak
+rm ${data_dir}/speech_shape.bak
+rm ${data_dir}/text_shape.bak
+
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/gen_ark_list.sh b/egs_modelscope/wenetspeech/paraformer/utils/gen_ark_list.sh
new file mode 100755
index 0000000..aebf356
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/gen_ark_list.sh
@@ -0,0 +1,22 @@
+#!/usr/bin/env bash
+
+
+# Begin configuration section.
+nj=32
+cmd=./utils/run.pl
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+ark_dir=$1
+txt_dir=$2
+output_dir=$3
+
+[ ! -d ${ark_dir}/ark ] && echo "$0: ark data is required" && exit 1;
+[ ! -d ${txt_dir}/txt ] && echo "$0: txt data is required" && exit 1;
+
+for n in $(seq $nj); do
+ echo "${ark_dir}/ark/feats.$n.ark ${txt_dir}/txt/text.$n.txt" || exit 1
+done > ${output_dir}/ark_txt.scp || exit 1
+
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/parse_options.sh b/egs_modelscope/wenetspeech/paraformer/utils/parse_options.sh
new file mode 100755
index 0000000..71fb9e5
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/parse_options.sh
@@ -0,0 +1,97 @@
+#!/usr/bin/env bash
+
+# Copyright 2012 Johns Hopkins University (Author: Daniel Povey);
+# Arnab Ghoshal, Karel Vesely
+
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+# Parse command-line options.
+# To be sourced by another script (as in ". parse_options.sh").
+# Option format is: --option-name arg
+# and shell variable "option_name" gets set to value "arg."
+# The exception is --help, which takes no arguments, but prints the
+# $help_message variable (if defined).
+
+
+###
+### The --config file options have lower priority to command line
+### options, so we need to import them first...
+###
+
+# Now import all the configs specified by command-line, in left-to-right order
+for ((argpos=1; argpos<$#; argpos++)); do
+ if [ "${!argpos}" == "--config" ]; then
+ argpos_plus1=$((argpos+1))
+ config=${!argpos_plus1}
+ [ ! -r $config ] && echo "$0: missing config '$config'" && exit 1
+ . $config # source the config file.
+ fi
+done
+
+
+###
+### Now we process the command line options
+###
+while true; do
+ [ -z "${1:-}" ] && break; # break if there are no arguments
+ case "$1" in
+ # If the enclosing script is called with --help option, print the help
+ # message and exit. Scripts should put help messages in $help_message
+ --help|-h) if [ -z "$help_message" ]; then echo "No help found." 1>&2;
+ else printf "$help_message\n" 1>&2 ; fi;
+ exit 0 ;;
+ --*=*) echo "$0: options to scripts must be of the form --name value, got '$1'"
+ exit 1 ;;
+ # If the first command-line argument begins with "--" (e.g. --foo-bar),
+ # then work out the variable name as $name, which will equal "foo_bar".
+ --*) name=`echo "$1" | sed s/^--// | sed s/-/_/g`;
+ # Next we test whether the variable in question is undefned-- if so it's
+ # an invalid option and we die. Note: $0 evaluates to the name of the
+ # enclosing script.
+ # The test [ -z ${foo_bar+xxx} ] will return true if the variable foo_bar
+ # is undefined. We then have to wrap this test inside "eval" because
+ # foo_bar is itself inside a variable ($name).
+ eval '[ -z "${'$name'+xxx}" ]' && echo "$0: invalid option $1" 1>&2 && exit 1;
+
+ oldval="`eval echo \\$$name`";
+ # Work out whether we seem to be expecting a Boolean argument.
+ if [ "$oldval" == "true" ] || [ "$oldval" == "false" ]; then
+ was_bool=true;
+ else
+ was_bool=false;
+ fi
+
+ # Set the variable to the right value-- the escaped quotes make it work if
+ # the option had spaces, like --cmd "queue.pl -sync y"
+ eval $name=\"$2\";
+
+ # Check that Boolean-valued arguments are really Boolean.
+ if $was_bool && [[ "$2" != "true" && "$2" != "false" ]]; then
+ echo "$0: expected \"true\" or \"false\": $1 $2" 1>&2
+ exit 1;
+ fi
+ shift 2;
+ ;;
+ *) break;
+ esac
+done
+
+
+# Check for an empty argument to the --cmd option, which can easily occur as a
+# result of scripting errors.
+[ ! -z "${cmd+xxx}" ] && [ -z "$cmd" ] && echo "$0: empty argument to --cmd option" 1>&2 && exit 1;
+
+
+true; # so this script returns exit code 0.
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/print_args.py b/egs_modelscope/wenetspeech/paraformer/utils/print_args.py
new file mode 100755
index 0000000..b0c61e5
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/print_args.py
@@ -0,0 +1,45 @@
+#!/usr/bin/env python
+import sys
+
+
+def get_commandline_args(no_executable=True):
+ extra_chars = [
+ " ",
+ ";",
+ "&",
+ "|",
+ "<",
+ ">",
+ "?",
+ "*",
+ "~",
+ "`",
+ '"',
+ "'",
+ "\\",
+ "{",
+ "}",
+ "(",
+ ")",
+ ]
+
+ # Escape the extra characters for shell
+ argv = [
+ arg.replace("'", "'\\''")
+ if all(char not in arg for char in extra_chars)
+ else "'" + arg.replace("'", "'\\''") + "'"
+ for arg in sys.argv
+ ]
+
+ if no_executable:
+ return " ".join(argv[1:])
+ else:
+ return sys.executable + " " + " ".join(argv)
+
+
+def main():
+ print(get_commandline_args())
+
+
+if __name__ == "__main__":
+ main()
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/proc_conf_oss.py b/egs_modelscope/wenetspeech/paraformer/utils/proc_conf_oss.py
new file mode 100755
index 0000000..c4a90c5
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/proc_conf_oss.py
@@ -0,0 +1,35 @@
+from pathlib import Path
+
+import torch
+import yaml
+
+
+class NoAliasSafeDumper(yaml.SafeDumper):
+ # Disable anchor/alias in yaml because looks ugly
+ def ignore_aliases(self, data):
+ return True
+
+
+def yaml_no_alias_safe_dump(data, stream=None, **kwargs):
+ """Safe-dump in yaml with no anchor/alias"""
+ return yaml.dump(
+ data, stream, allow_unicode=True, Dumper=NoAliasSafeDumper, **kwargs
+ )
+
+
+def gen_conf(file, out_dir):
+ conf = torch.load(file)["config"]
+ conf["oss_bucket"] = "null"
+ print(conf)
+ output_dir = Path(out_dir)
+ output_dir.mkdir(parents=True, exist_ok=True)
+ with (output_dir / "config.yaml").open("w", encoding="utf-8") as f:
+ yaml_no_alias_safe_dump(conf, f, indent=4, sort_keys=False)
+
+
+if __name__ == "__main__":
+ import sys
+
+ in_f = sys.argv[1]
+ out_f = sys.argv[2]
+ gen_conf(in_f, out_f)
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/proce_text.py b/egs_modelscope/wenetspeech/paraformer/utils/proce_text.py
new file mode 100755
index 0000000..9e517a4
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/proce_text.py
@@ -0,0 +1,31 @@
+
+import sys
+import re
+
+in_f = sys.argv[1]
+out_f = sys.argv[2]
+
+
+with open(in_f, "r", encoding="utf-8") as f:
+ lines = f.readlines()
+
+with open(out_f, "w", encoding="utf-8") as f:
+ for line in lines:
+ outs = line.strip().split(" ", 1)
+ if len(outs) == 2:
+ idx, text = outs
+ text = re.sub("</s>", "", text)
+ text = re.sub("<s>", "", text)
+ text = re.sub("@@", "", text)
+ text = re.sub("@", "", text)
+ text = re.sub("<unk>", "", text)
+ text = re.sub(" ", "", text)
+ text = text.lower()
+ else:
+ idx = outs[0]
+ text = " "
+
+ text = [x for x in text]
+ text = " ".join(text)
+ out = "{} {}\n".format(idx, text)
+ f.write(out)
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/run.pl b/egs_modelscope/wenetspeech/paraformer/utils/run.pl
new file mode 100755
index 0000000..483f95b
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/run.pl
@@ -0,0 +1,356 @@
+#!/usr/bin/env perl
+use warnings; #sed replacement for -w perl parameter
+# In general, doing
+# run.pl some.log a b c is like running the command a b c in
+# the bash shell, and putting the standard error and output into some.log.
+# To run parallel jobs (backgrounded on the host machine), you can do (e.g.)
+# run.pl JOB=1:4 some.JOB.log a b c JOB is like running the command a b c JOB
+# and putting it in some.JOB.log, for each one. [Note: JOB can be any identifier].
+# If any of the jobs fails, this script will fail.
+
+# A typical example is:
+# run.pl some.log my-prog "--opt=foo bar" foo \| other-prog baz
+# and run.pl will run something like:
+# ( my-prog '--opt=foo bar' foo | other-prog baz ) >& some.log
+#
+# Basically it takes the command-line arguments, quotes them
+# as necessary to preserve spaces, and evaluates them with bash.
+# In addition it puts the command line at the top of the log, and
+# the start and end times of the command at the beginning and end.
+# The reason why this is useful is so that we can create a different
+# version of this program that uses a queueing system instead.
+
+#use Data::Dumper;
+
+@ARGV < 2 && die "usage: run.pl log-file command-line arguments...";
+
+#print STDERR "COMMAND-LINE: " . Dumper(\@ARGV) . "\n";
+$job_pick = 'all';
+$max_jobs_run = -1;
+$jobstart = 1;
+$jobend = 1;
+$ignored_opts = ""; # These will be ignored.
+
+# First parse an option like JOB=1:4, and any
+# options that would normally be given to
+# queue.pl, which we will just discard.
+
+for (my $x = 1; $x <= 2; $x++) { # This for-loop is to
+ # allow the JOB=1:n option to be interleaved with the
+ # options to qsub.
+ while (@ARGV >= 2 && $ARGV[0] =~ m:^-:) {
+ # parse any options that would normally go to qsub, but which will be ignored here.
+ my $switch = shift @ARGV;
+ if ($switch eq "-V") {
+ $ignored_opts .= "-V ";
+ } elsif ($switch eq "--max-jobs-run" || $switch eq "-tc") {
+ # we do support the option --max-jobs-run n, and its GridEngine form -tc n.
+ # if the command appears multiple times uses the smallest option.
+ if ( $max_jobs_run <= 0 ) {
+ $max_jobs_run = shift @ARGV;
+ } else {
+ my $new_constraint = shift @ARGV;
+ if ( ($new_constraint < $max_jobs_run) ) {
+ $max_jobs_run = $new_constraint;
+ }
+ }
+
+ if (! ($max_jobs_run > 0)) {
+ die "run.pl: invalid option --max-jobs-run $max_jobs_run";
+ }
+ } else {
+ my $argument = shift @ARGV;
+ if ($argument =~ m/^--/) {
+ print STDERR "run.pl: WARNING: suspicious argument '$argument' to $switch; starts with '-'\n";
+ }
+ if ($switch eq "-sync" && $argument =~ m/^[yY]/) {
+ $ignored_opts .= "-sync "; # Note: in the
+ # corresponding code in queue.pl it says instead, just "$sync = 1;".
+ } elsif ($switch eq "-pe") { # e.g. -pe smp 5
+ my $argument2 = shift @ARGV;
+ $ignored_opts .= "$switch $argument $argument2 ";
+ } elsif ($switch eq "--gpu") {
+ $using_gpu = $argument;
+ } elsif ($switch eq "--pick") {
+ if($argument =~ m/^(all|failed|incomplete)$/) {
+ $job_pick = $argument;
+ } else {
+ print STDERR "run.pl: ERROR: --pick argument must be one of 'all', 'failed' or 'incomplete'"
+ }
+ } else {
+ # Ignore option.
+ $ignored_opts .= "$switch $argument ";
+ }
+ }
+ }
+ if ($ARGV[0] =~ m/^([\w_][\w\d_]*)+=(\d+):(\d+)$/) { # e.g. JOB=1:20
+ $jobname = $1;
+ $jobstart = $2;
+ $jobend = $3;
+ if ($jobstart > $jobend) {
+ die "run.pl: invalid job range $ARGV[0]";
+ }
+ if ($jobstart <= 0) {
+ die "run.pl: invalid job range $ARGV[0], start must be strictly positive (this is required for GridEngine compatibility).";
+ }
+ shift;
+ } elsif ($ARGV[0] =~ m/^([\w_][\w\d_]*)+=(\d+)$/) { # e.g. JOB=1.
+ $jobname = $1;
+ $jobstart = $2;
+ $jobend = $2;
+ shift;
+ } elsif ($ARGV[0] =~ m/.+\=.*\:.*$/) {
+ print STDERR "run.pl: Warning: suspicious first argument to run.pl: $ARGV[0]\n";
+ }
+}
+
+# Users found this message confusing so we are removing it.
+# if ($ignored_opts ne "") {
+# print STDERR "run.pl: Warning: ignoring options \"$ignored_opts\"\n";
+# }
+
+if ($max_jobs_run == -1) { # If --max-jobs-run option not set,
+ # then work out the number of processors if possible,
+ # and set it based on that.
+ $max_jobs_run = 0;
+ if ($using_gpu) {
+ if (open(P, "nvidia-smi -L |")) {
+ $max_jobs_run++ while (<P>);
+ close(P);
+ }
+ if ($max_jobs_run == 0) {
+ $max_jobs_run = 1;
+ print STDERR "run.pl: Warning: failed to detect number of GPUs from nvidia-smi, using ${max_jobs_run}\n";
+ }
+ } elsif (open(P, "</proc/cpuinfo")) { # Linux
+ while (<P>) { if (m/^processor/) { $max_jobs_run++; } }
+ if ($max_jobs_run == 0) {
+ print STDERR "run.pl: Warning: failed to detect any processors from /proc/cpuinfo\n";
+ $max_jobs_run = 10; # reasonable default.
+ }
+ close(P);
+ } elsif (open(P, "sysctl -a |")) { # BSD/Darwin
+ while (<P>) {
+ if (m/hw\.ncpu\s*[:=]\s*(\d+)/) { # hw.ncpu = 4, or hw.ncpu: 4
+ $max_jobs_run = $1;
+ last;
+ }
+ }
+ close(P);
+ if ($max_jobs_run == 0) {
+ print STDERR "run.pl: Warning: failed to detect any processors from sysctl -a\n";
+ $max_jobs_run = 10; # reasonable default.
+ }
+ } else {
+ # allow at most 32 jobs at once, on non-UNIX systems; change this code
+ # if you need to change this default.
+ $max_jobs_run = 32;
+ }
+ # The just-computed value of $max_jobs_run is just the number of processors
+ # (or our best guess); and if it happens that the number of jobs we need to
+ # run is just slightly above $max_jobs_run, it will make sense to increase
+ # $max_jobs_run to equal the number of jobs, so we don't have a small number
+ # of leftover jobs.
+ $num_jobs = $jobend - $jobstart + 1;
+ if (!$using_gpu &&
+ $num_jobs > $max_jobs_run && $num_jobs < 1.4 * $max_jobs_run) {
+ $max_jobs_run = $num_jobs;
+ }
+}
+
+sub pick_or_exit {
+ # pick_or_exit ( $logfile )
+ # Invoked before each job is started helps to run jobs selectively.
+ #
+ # Given the name of the output logfile decides whether the job must be
+ # executed (by returning from the subroutine) or not (by terminating the
+ # process calling exit)
+ #
+ # PRE: $job_pick is a global variable set by command line switch --pick
+ # and indicates which class of jobs must be executed.
+ #
+ # 1) If a failed job is not executed the process exit code will indicate
+ # failure, just as if the task was just executed and failed.
+ #
+ # 2) If a task is incomplete it will be executed. Incomplete may be either
+ # a job whose log file does not contain the accounting notes in the end,
+ # or a job whose log file does not exist.
+ #
+ # 3) If the $job_pick is set to 'all' (default behavior) a task will be
+ # executed regardless of the result of previous attempts.
+ #
+ # This logic could have been implemented in the main execution loop
+ # but a subroutine to preserve the current level of readability of
+ # that part of the code.
+ #
+ # Alexandre Felipe, (o.alexandre.felipe@gmail.com) 14th of August of 2020
+ #
+ if($job_pick eq 'all'){
+ return; # no need to bother with the previous log
+ }
+ open my $fh, "<", $_[0] or return; # job not executed yet
+ my $log_line;
+ my $cur_line;
+ while ($cur_line = <$fh>) {
+ if( $cur_line =~ m/# Ended \(code .*/ ) {
+ $log_line = $cur_line;
+ }
+ }
+ close $fh;
+ if (! defined($log_line)){
+ return; # incomplete
+ }
+ if ( $log_line =~ m/# Ended \(code 0\).*/ ) {
+ exit(0); # complete
+ } elsif ( $log_line =~ m/# Ended \(code \d+(; signal \d+)?\).*/ ){
+ if ($job_pick !~ m/^(failed|all)$/) {
+ exit(1); # failed but not going to run
+ } else {
+ return; # failed
+ }
+ } elsif ( $log_line =~ m/.*\S.*/ ) {
+ return; # incomplete jobs are always run
+ }
+}
+
+
+$logfile = shift @ARGV;
+
+if (defined $jobname && $logfile !~ m/$jobname/ &&
+ $jobend > $jobstart) {
+ print STDERR "run.pl: you are trying to run a parallel job but "
+ . "you are putting the output into just one log file ($logfile)\n";
+ exit(1);
+}
+
+$cmd = "";
+
+foreach $x (@ARGV) {
+ if ($x =~ m/^\S+$/) { $cmd .= $x . " "; }
+ elsif ($x =~ m:\":) { $cmd .= "'$x' "; }
+ else { $cmd .= "\"$x\" "; }
+}
+
+#$Data::Dumper::Indent=0;
+$ret = 0;
+$numfail = 0;
+%active_pids=();
+
+use POSIX ":sys_wait_h";
+for ($jobid = $jobstart; $jobid <= $jobend; $jobid++) {
+ if (scalar(keys %active_pids) >= $max_jobs_run) {
+
+ # Lets wait for a change in any child's status
+ # Then we have to work out which child finished
+ $r = waitpid(-1, 0);
+ $code = $?;
+ if ($r < 0 ) { die "run.pl: Error waiting for child process"; } # should never happen.
+ if ( defined $active_pids{$r} ) {
+ $jid=$active_pids{$r};
+ $fail[$jid]=$code;
+ if ($code !=0) { $numfail++;}
+ delete $active_pids{$r};
+ # print STDERR "Finished: $r/$jid " . Dumper(\%active_pids) . "\n";
+ } else {
+ die "run.pl: Cannot find the PID of the child process that just finished.";
+ }
+
+ # In theory we could do a non-blocking waitpid over all jobs running just
+ # to find out if only one or more jobs finished during the previous waitpid()
+ # However, we just omit this and will reap the next one in the next pass
+ # through the for(;;) cycle
+ }
+ $childpid = fork();
+ if (!defined $childpid) { die "run.pl: Error forking in run.pl (writing to $logfile)"; }
+ if ($childpid == 0) { # We're in the child... this branch
+ # executes the job and returns (possibly with an error status).
+ if (defined $jobname) {
+ $cmd =~ s/$jobname/$jobid/g;
+ $logfile =~ s/$jobname/$jobid/g;
+ }
+ # exit if the job does not need to be executed
+ pick_or_exit( $logfile );
+
+ system("mkdir -p `dirname $logfile` 2>/dev/null");
+ open(F, ">$logfile") || die "run.pl: Error opening log file $logfile";
+ print F "# " . $cmd . "\n";
+ print F "# Started at " . `date`;
+ $starttime = `date +'%s'`;
+ print F "#\n";
+ close(F);
+
+ # Pipe into bash.. make sure we're not using any other shell.
+ open(B, "|bash") || die "run.pl: Error opening shell command";
+ print B "( " . $cmd . ") 2>>$logfile >> $logfile";
+ close(B); # If there was an error, exit status is in $?
+ $ret = $?;
+
+ $lowbits = $ret & 127;
+ $highbits = $ret >> 8;
+ if ($lowbits != 0) { $return_str = "code $highbits; signal $lowbits" }
+ else { $return_str = "code $highbits"; }
+
+ $endtime = `date +'%s'`;
+ open(F, ">>$logfile") || die "run.pl: Error opening log file $logfile (again)";
+ $enddate = `date`;
+ chop $enddate;
+ print F "# Accounting: time=" . ($endtime - $starttime) . " threads=1\n";
+ print F "# Ended ($return_str) at " . $enddate . ", elapsed time " . ($endtime-$starttime) . " seconds\n";
+ close(F);
+ exit($ret == 0 ? 0 : 1);
+ } else {
+ $pid[$jobid] = $childpid;
+ $active_pids{$childpid} = $jobid;
+ # print STDERR "Queued: " . Dumper(\%active_pids) . "\n";
+ }
+}
+
+# Now we have submitted all the jobs, lets wait until all the jobs finish
+foreach $child (keys %active_pids) {
+ $jobid=$active_pids{$child};
+ $r = waitpid($pid[$jobid], 0);
+ $code = $?;
+ if ($r == -1) { die "run.pl: Error waiting for child process"; } # should never happen.
+ if ($r != 0) { $fail[$jobid]=$code; $numfail++ if $code!=0; } # Completed successfully
+}
+
+# Some sanity checks:
+# The $fail array should not contain undefined codes
+# The number of non-zeros in that array should be equal to $numfail
+# We cannot do foreach() here, as the JOB ids do not start at zero
+$failed_jids=0;
+for ($jobid = $jobstart; $jobid <= $jobend; $jobid++) {
+ $job_return = $fail[$jobid];
+ if (not defined $job_return ) {
+ # print Dumper(\@fail);
+
+ die "run.pl: Sanity check failed: we have indication that some jobs are running " .
+ "even after we waited for all jobs to finish" ;
+ }
+ if ($job_return != 0 ){ $failed_jids++;}
+}
+if ($failed_jids != $numfail) {
+ die "run.pl: Sanity check failed: cannot find out how many jobs failed ($failed_jids x $numfail)."
+}
+if ($numfail > 0) { $ret = 1; }
+
+if ($ret != 0) {
+ $njobs = $jobend - $jobstart + 1;
+ if ($njobs == 1) {
+ if (defined $jobname) {
+ $logfile =~ s/$jobname/$jobstart/; # only one numbered job, so replace name with
+ # that job.
+ }
+ print STDERR "run.pl: job failed, log is in $logfile\n";
+ if ($logfile =~ m/JOB/) {
+ print STDERR "run.pl: probably you forgot to put JOB=1:\$nj in your script.";
+ }
+ }
+ else {
+ $logfile =~ s/$jobname/*/g;
+ print STDERR "run.pl: $numfail / $njobs failed, log is in $logfile\n";
+ }
+}
+
+
+exit ($ret);
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/shuffle_list.pl b/egs_modelscope/wenetspeech/paraformer/utils/shuffle_list.pl
new file mode 100755
index 0000000..a116200
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/shuffle_list.pl
@@ -0,0 +1,44 @@
+#!/usr/bin/env perl
+
+# Copyright 2013 Johns Hopkins University (author: Daniel Povey)
+
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+if ($ARGV[0] eq "--srand") {
+ $n = $ARGV[1];
+ $n =~ m/\d+/ || die "Bad argument to --srand option: \"$n\"";
+ srand($ARGV[1]);
+ shift;
+ shift;
+} else {
+ srand(0); # Gives inconsistent behavior if we don't seed.
+}
+
+if (@ARGV > 1 || $ARGV[0] =~ m/^-.+/) { # >1 args, or an option we
+ # don't understand.
+ print "Usage: shuffle_list.pl [--srand N] [input file] > output\n";
+ print "randomizes the order of lines of input.\n";
+ exit(1);
+}
+
+@lines;
+while (<>) {
+ push @lines, [ (rand(), $_)] ;
+}
+
+@lines = sort { $a->[0] cmp $b->[0] } @lines;
+foreach $l (@lines) {
+ print $l->[1];
+}
\ No newline at end of file
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/split_data.py b/egs_modelscope/wenetspeech/paraformer/utils/split_data.py
new file mode 100755
index 0000000..060eae6
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/split_data.py
@@ -0,0 +1,60 @@
+import os
+import sys
+import random
+
+
+in_dir = sys.argv[1]
+out_dir = sys.argv[2]
+num_split = sys.argv[3]
+
+
+def split_scp(scp, num):
+ assert len(scp) >= num
+ avg = len(scp) // num
+ out = []
+ begin = 0
+
+ for i in range(num):
+ if i == num - 1:
+ out.append(scp[begin:])
+ else:
+ out.append(scp[begin:begin+avg])
+ begin += avg
+
+ return out
+
+
+os.path.exists("{}/wav.scp".format(in_dir))
+os.path.exists("{}/text".format(in_dir))
+
+with open("{}/wav.scp".format(in_dir), 'r') as infile:
+ wav_list = infile.readlines()
+
+with open("{}/text".format(in_dir), 'r') as infile:
+ text_list = infile.readlines()
+
+assert len(wav_list) == len(text_list)
+
+x = list(zip(wav_list, text_list))
+random.shuffle(x)
+wav_shuffle_list, text_shuffle_list = zip(*x)
+
+num_split = int(num_split)
+wav_split_list = split_scp(wav_shuffle_list, num_split)
+text_split_list = split_scp(text_shuffle_list, num_split)
+
+for idx, wav_list in enumerate(wav_split_list, 1):
+ path = out_dir + "/split" + str(num_split) + "/" + str(idx)
+ if not os.path.exists(path):
+ os.makedirs(path)
+ with open("{}/wav.scp".format(path), 'w') as wav_writer:
+ for line in wav_list:
+ wav_writer.write(line)
+
+for idx, text_list in enumerate(text_split_list, 1):
+ path = out_dir + "/split" + str(num_split) + "/" + str(idx)
+ if not os.path.exists(path):
+ os.makedirs(path)
+ with open("{}/text".format(path), 'w') as text_writer:
+ for line in text_list:
+ text_writer.write(line)
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/split_scp.pl b/egs_modelscope/wenetspeech/paraformer/utils/split_scp.pl
new file mode 100755
index 0000000..0876dcb
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/split_scp.pl
@@ -0,0 +1,246 @@
+#!/usr/bin/env perl
+
+# Copyright 2010-2011 Microsoft Corporation
+
+# See ../../COPYING for clarification regarding multiple authors
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# THIS CODE IS PROVIDED *AS IS* BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION ANY IMPLIED
+# WARRANTIES OR CONDITIONS OF TITLE, FITNESS FOR A PARTICULAR PURPOSE,
+# MERCHANTABLITY OR NON-INFRINGEMENT.
+# See the Apache 2 License for the specific language governing permissions and
+# limitations under the License.
+
+
+# This program splits up any kind of .scp or archive-type file.
+# If there is no utt2spk option it will work on any text file and
+# will split it up with an approximately equal number of lines in
+# each but.
+# With the --utt2spk option it will work on anything that has the
+# utterance-id as the first entry on each line; the utt2spk file is
+# of the form "utterance speaker" (on each line).
+# It splits it into equal size chunks as far as it can. If you use the utt2spk
+# option it will make sure these chunks coincide with speaker boundaries. In
+# this case, if there are more chunks than speakers (and in some other
+# circumstances), some of the resulting chunks will be empty and it will print
+# an error message and exit with nonzero status.
+# You will normally call this like:
+# split_scp.pl scp scp.1 scp.2 scp.3 ...
+# or
+# split_scp.pl --utt2spk=utt2spk scp scp.1 scp.2 scp.3 ...
+# Note that you can use this script to split the utt2spk file itself,
+# e.g. split_scp.pl --utt2spk=utt2spk utt2spk utt2spk.1 utt2spk.2 ...
+
+# You can also call the scripts like:
+# split_scp.pl -j 3 0 scp scp.0
+# [note: with this option, it assumes zero-based indexing of the split parts,
+# i.e. the second number must be 0 <= n < num-jobs.]
+
+use warnings;
+
+$num_jobs = 0;
+$job_id = 0;
+$utt2spk_file = "";
+$one_based = 0;
+
+for ($x = 1; $x <= 3 && @ARGV > 0; $x++) {
+ if ($ARGV[0] eq "-j") {
+ shift @ARGV;
+ $num_jobs = shift @ARGV;
+ $job_id = shift @ARGV;
+ }
+ if ($ARGV[0] =~ /--utt2spk=(.+)/) {
+ $utt2spk_file=$1;
+ shift;
+ }
+ if ($ARGV[0] eq '--one-based') {
+ $one_based = 1;
+ shift @ARGV;
+ }
+}
+
+if ($num_jobs != 0 && ($num_jobs < 0 || $job_id - $one_based < 0 ||
+ $job_id - $one_based >= $num_jobs)) {
+ die "$0: Invalid job number/index values for '-j $num_jobs $job_id" .
+ ($one_based ? " --one-based" : "") . "'\n"
+}
+
+$one_based
+ and $job_id--;
+
+if(($num_jobs == 0 && @ARGV < 2) || ($num_jobs > 0 && (@ARGV < 1 || @ARGV > 2))) {
+ die
+"Usage: split_scp.pl [--utt2spk=<utt2spk_file>] in.scp out1.scp out2.scp ...
+ or: split_scp.pl -j num-jobs job-id [--one-based] [--utt2spk=<utt2spk_file>] in.scp [out.scp]
+ ... where 0 <= job-id < num-jobs, or 1 <= job-id <- num-jobs if --one-based.\n";
+}
+
+$error = 0;
+$inscp = shift @ARGV;
+if ($num_jobs == 0) { # without -j option
+ @OUTPUTS = @ARGV;
+} else {
+ for ($j = 0; $j < $num_jobs; $j++) {
+ if ($j == $job_id) {
+ if (@ARGV > 0) { push @OUTPUTS, $ARGV[0]; }
+ else { push @OUTPUTS, "-"; }
+ } else {
+ push @OUTPUTS, "/dev/null";
+ }
+ }
+}
+
+if ($utt2spk_file ne "") { # We have the --utt2spk option...
+ open($u_fh, '<', $utt2spk_file) || die "$0: Error opening utt2spk file $utt2spk_file: $!\n";
+ while(<$u_fh>) {
+ @A = split;
+ @A == 2 || die "$0: Bad line $_ in utt2spk file $utt2spk_file\n";
+ ($u,$s) = @A;
+ $utt2spk{$u} = $s;
+ }
+ close $u_fh;
+ open($i_fh, '<', $inscp) || die "$0: Error opening input scp file $inscp: $!\n";
+ @spkrs = ();
+ while(<$i_fh>) {
+ @A = split;
+ if(@A == 0) { die "$0: Empty or space-only line in scp file $inscp\n"; }
+ $u = $A[0];
+ $s = $utt2spk{$u};
+ defined $s || die "$0: No utterance $u in utt2spk file $utt2spk_file\n";
+ if(!defined $spk_count{$s}) {
+ push @spkrs, $s;
+ $spk_count{$s} = 0;
+ $spk_data{$s} = []; # ref to new empty array.
+ }
+ $spk_count{$s}++;
+ push @{$spk_data{$s}}, $_;
+ }
+ # Now split as equally as possible ..
+ # First allocate spks to files by allocating an approximately
+ # equal number of speakers.
+ $numspks = @spkrs; # number of speakers.
+ $numscps = @OUTPUTS; # number of output files.
+ if ($numspks < $numscps) {
+ die "$0: Refusing to split data because number of speakers $numspks " .
+ "is less than the number of output .scp files $numscps\n";
+ }
+ for($scpidx = 0; $scpidx < $numscps; $scpidx++) {
+ $scparray[$scpidx] = []; # [] is array reference.
+ }
+ for ($spkidx = 0; $spkidx < $numspks; $spkidx++) {
+ $scpidx = int(($spkidx*$numscps) / $numspks);
+ $spk = $spkrs[$spkidx];
+ push @{$scparray[$scpidx]}, $spk;
+ $scpcount[$scpidx] += $spk_count{$spk};
+ }
+
+ # Now will try to reassign beginning + ending speakers
+ # to different scp's and see if it gets more balanced.
+ # Suppose objf we're minimizing is sum_i (num utts in scp[i] - average)^2.
+ # We can show that if considering changing just 2 scp's, we minimize
+ # this by minimizing the squared difference in sizes. This is
+ # equivalent to minimizing the absolute difference in sizes. This
+ # shows this method is bound to converge.
+
+ $changed = 1;
+ while($changed) {
+ $changed = 0;
+ for($scpidx = 0; $scpidx < $numscps; $scpidx++) {
+ # First try to reassign ending spk of this scp.
+ if($scpidx < $numscps-1) {
+ $sz = @{$scparray[$scpidx]};
+ if($sz > 0) {
+ $spk = $scparray[$scpidx]->[$sz-1];
+ $count = $spk_count{$spk};
+ $nutt1 = $scpcount[$scpidx];
+ $nutt2 = $scpcount[$scpidx+1];
+ if( abs( ($nutt2+$count) - ($nutt1-$count))
+ < abs($nutt2 - $nutt1)) { # Would decrease
+ # size-diff by reassigning spk...
+ $scpcount[$scpidx+1] += $count;
+ $scpcount[$scpidx] -= $count;
+ pop @{$scparray[$scpidx]};
+ unshift @{$scparray[$scpidx+1]}, $spk;
+ $changed = 1;
+ }
+ }
+ }
+ if($scpidx > 0 && @{$scparray[$scpidx]} > 0) {
+ $spk = $scparray[$scpidx]->[0];
+ $count = $spk_count{$spk};
+ $nutt1 = $scpcount[$scpidx-1];
+ $nutt2 = $scpcount[$scpidx];
+ if( abs( ($nutt2-$count) - ($nutt1+$count))
+ < abs($nutt2 - $nutt1)) { # Would decrease
+ # size-diff by reassigning spk...
+ $scpcount[$scpidx-1] += $count;
+ $scpcount[$scpidx] -= $count;
+ shift @{$scparray[$scpidx]};
+ push @{$scparray[$scpidx-1]}, $spk;
+ $changed = 1;
+ }
+ }
+ }
+ }
+ # Now print out the files...
+ for($scpidx = 0; $scpidx < $numscps; $scpidx++) {
+ $scpfile = $OUTPUTS[$scpidx];
+ ($scpfile ne '-' ? open($f_fh, '>', $scpfile)
+ : open($f_fh, '>&', \*STDOUT)) ||
+ die "$0: Could not open scp file $scpfile for writing: $!\n";
+ $count = 0;
+ if(@{$scparray[$scpidx]} == 0) {
+ print STDERR "$0: eError: split_scp.pl producing empty .scp file " .
+ "$scpfile (too many splits and too few speakers?)\n";
+ $error = 1;
+ } else {
+ foreach $spk ( @{$scparray[$scpidx]} ) {
+ print $f_fh @{$spk_data{$spk}};
+ $count += $spk_count{$spk};
+ }
+ $count == $scpcount[$scpidx] || die "Count mismatch [code error]";
+ }
+ close($f_fh);
+ }
+} else {
+ # This block is the "normal" case where there is no --utt2spk
+ # option and we just break into equal size chunks.
+
+ open($i_fh, '<', $inscp) || die "$0: Error opening input scp file $inscp: $!\n";
+
+ $numscps = @OUTPUTS; # size of array.
+ @F = ();
+ while(<$i_fh>) {
+ push @F, $_;
+ }
+ $numlines = @F;
+ if($numlines == 0) {
+ print STDERR "$0: error: empty input scp file $inscp\n";
+ $error = 1;
+ }
+ $linesperscp = int( $numlines / $numscps); # the "whole part"..
+ $linesperscp >= 1 || die "$0: You are splitting into too many pieces! [reduce \$nj ($numscps) to be smaller than the number of lines ($numlines) in $inscp]\n";
+ $remainder = $numlines - ($linesperscp * $numscps);
+ ($remainder >= 0 && $remainder < $numlines) || die "bad remainder $remainder";
+ # [just doing int() rounds down].
+ $n = 0;
+ for($scpidx = 0; $scpidx < @OUTPUTS; $scpidx++) {
+ $scpfile = $OUTPUTS[$scpidx];
+ ($scpfile ne '-' ? open($o_fh, '>', $scpfile)
+ : open($o_fh, '>&', \*STDOUT)) ||
+ die "$0: Could not open scp file $scpfile for writing: $!\n";
+ for($k = 0; $k < $linesperscp + ($scpidx < $remainder ? 1 : 0); $k++) {
+ print $o_fh $F[$n++];
+ }
+ close($o_fh) || die "$0: Eror closing scp file $scpfile: $!\n";
+ }
+ $n == $numlines || die "$n != $numlines [code error]";
+}
+
+exit ($error);
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/subset_data_dir_tr_cv.sh b/egs_modelscope/wenetspeech/paraformer/utils/subset_data_dir_tr_cv.sh
new file mode 100755
index 0000000..e16cebd
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/subset_data_dir_tr_cv.sh
@@ -0,0 +1,30 @@
+#!/usr/bin/env bash
+
+dev_num_utt=1000
+
+echo "$0 $@"
+. utils/parse_options.sh || exit 1;
+
+train_data=$1
+out_dir=$2
+
+[ ! -f ${train_data}/wav.scp ] && echo "$0: no such file ${train_data}/wav.scp" && exit 1;
+[ ! -f ${train_data}/text ] && echo "$0: no such file ${train_data}/text" && exit 1;
+
+mkdir -p ${out_dir}/train && mkdir -p ${out_dir}/dev
+
+cp ${train_data}/wav.scp ${out_dir}/train/wav.scp.bak
+cp ${train_data}/text ${out_dir}/train/text.bak
+
+num_utt=$(wc -l <${out_dir}/train/wav.scp.bak)
+
+utils/shuffle_list.pl --srand 1 ${out_dir}/train/wav.scp.bak > ${out_dir}/train/wav.scp.shuf
+head -n ${dev_num_utt} ${out_dir}/train/wav.scp.shuf > ${out_dir}/dev/wav.scp
+tail -n $((${num_utt}-${dev_num_utt})) ${out_dir}/train/wav.scp.shuf > ${out_dir}/train/wav.scp
+
+utils/shuffle_list.pl --srand 1 ${out_dir}/train/text.bak > ${out_dir}/train/text.shuf
+head -n ${dev_num_utt} ${out_dir}/train/text.shuf > ${out_dir}/dev/text
+tail -n $((${num_utt}-${dev_num_utt})) ${out_dir}/train/text.shuf > ${out_dir}/train/text
+
+rm ${out_dir}/train/wav.scp.bak ${out_dir}/train/text.bak
+rm ${out_dir}/train/wav.scp.shuf ${out_dir}/train/text.shuf
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/text2token.py b/egs_modelscope/wenetspeech/paraformer/utils/text2token.py
new file mode 100755
index 0000000..56c3913
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/text2token.py
@@ -0,0 +1,135 @@
+#!/usr/bin/env python3
+
+# Copyright 2017 Johns Hopkins University (Shinji Watanabe)
+# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
+
+
+import argparse
+import codecs
+import re
+import sys
+
+is_python2 = sys.version_info[0] == 2
+
+
+def exist_or_not(i, match_pos):
+ start_pos = None
+ end_pos = None
+ for pos in match_pos:
+ if pos[0] <= i < pos[1]:
+ start_pos = pos[0]
+ end_pos = pos[1]
+ break
+
+ return start_pos, end_pos
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="convert raw text to tokenized text",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--nchar",
+ "-n",
+ default=1,
+ type=int,
+ help="number of characters to split, i.e., \
+ aabb -> a a b b with -n 1 and aa bb with -n 2",
+ )
+ parser.add_argument(
+ "--skip-ncols", "-s", default=0, type=int, help="skip first n columns"
+ )
+ parser.add_argument("--space", default="<space>", type=str, help="space symbol")
+ parser.add_argument(
+ "--non-lang-syms",
+ "-l",
+ default=None,
+ type=str,
+ help="list of non-linguistic symobles, e.g., <NOISE> etc.",
+ )
+ parser.add_argument("text", type=str, default=False, nargs="?", help="input text")
+ parser.add_argument(
+ "--trans_type",
+ "-t",
+ type=str,
+ default="char",
+ choices=["char", "phn"],
+ help="""Transcript type. char/phn. e.g., for TIMIT FADG0_SI1279 -
+ If trans_type is char,
+ read from SI1279.WRD file -> "bricks are an alternative"
+ Else if trans_type is phn,
+ read from SI1279.PHN file -> "sil b r ih sil k s aa r er n aa l
+ sil t er n ih sil t ih v sil" """,
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ rs = []
+ if args.non_lang_syms is not None:
+ with codecs.open(args.non_lang_syms, "r", encoding="utf-8") as f:
+ nls = [x.rstrip() for x in f.readlines()]
+ rs = [re.compile(re.escape(x)) for x in nls]
+
+ if args.text:
+ f = codecs.open(args.text, encoding="utf-8")
+ else:
+ f = codecs.getreader("utf-8")(sys.stdin if is_python2 else sys.stdin.buffer)
+
+ sys.stdout = codecs.getwriter("utf-8")(
+ sys.stdout if is_python2 else sys.stdout.buffer
+ )
+ line = f.readline()
+ n = args.nchar
+ while line:
+ x = line.split()
+ print(" ".join(x[: args.skip_ncols]), end=" ")
+ a = " ".join(x[args.skip_ncols :])
+
+ # get all matched positions
+ match_pos = []
+ for r in rs:
+ i = 0
+ while i >= 0:
+ m = r.search(a, i)
+ if m:
+ match_pos.append([m.start(), m.end()])
+ i = m.end()
+ else:
+ break
+
+ if args.trans_type == "phn":
+ a = a.split(" ")
+ else:
+ if len(match_pos) > 0:
+ chars = []
+ i = 0
+ while i < len(a):
+ start_pos, end_pos = exist_or_not(i, match_pos)
+ if start_pos is not None:
+ chars.append(a[start_pos:end_pos])
+ i = end_pos
+ else:
+ chars.append(a[i])
+ i += 1
+ a = chars
+
+ a = [a[j : j + n] for j in range(0, len(a), n)]
+
+ a_flat = []
+ for z in a:
+ a_flat.append("".join(z))
+
+ a_chars = [z.replace(" ", args.space) for z in a_flat]
+ if args.trans_type == "phn":
+ a_chars = [z.replace("sil", args.space) for z in a_chars]
+ print(" ".join(a_chars))
+ line = f.readline()
+
+
+if __name__ == "__main__":
+ main()
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/text_tokenize.py b/egs_modelscope/wenetspeech/paraformer/utils/text_tokenize.py
new file mode 100755
index 0000000..962ea11
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/text_tokenize.py
@@ -0,0 +1,106 @@
+import re
+import argparse
+
+
+def load_dict(seg_file):
+ seg_dict = {}
+ with open(seg_file, 'r') as infile:
+ for line in infile:
+ s = line.strip().split()
+ key = s[0]
+ value = s[1:]
+ seg_dict[key] = " ".join(value)
+ return seg_dict
+
+
+def forward_segment(text, dic):
+ word_list = []
+ i = 0
+ while i < len(text):
+ longest_word = text[i]
+ for j in range(i + 1, len(text) + 1):
+ word = text[i:j]
+ if word in dic:
+ if len(word) > len(longest_word):
+ longest_word = word
+ word_list.append(longest_word)
+ i += len(longest_word)
+ return word_list
+
+
+def tokenize(txt,
+ seg_dict):
+ out_txt = ""
+ pattern = re.compile(r"([\u4E00-\u9FA5A-Za-z0-9])")
+ for word in txt:
+ if pattern.match(word):
+ if word in seg_dict:
+ out_txt += seg_dict[word] + " "
+ else:
+ out_txt += "<unk>" + " "
+ else:
+ continue
+ return out_txt.strip()
+
+
+def get_parser():
+ parser = argparse.ArgumentParser(
+ description="text tokenize",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument(
+ "--text-file",
+ "-t",
+ default=False,
+ required=True,
+ type=str,
+ help="input text",
+ )
+ parser.add_argument(
+ "--seg-file",
+ "-s",
+ default=False,
+ required=True,
+ type=str,
+ help="seg file",
+ )
+ parser.add_argument(
+ "--txt-index",
+ "-i",
+ default=1,
+ required=True,
+ type=int,
+ help="txt index",
+ )
+ parser.add_argument(
+ "--output-dir",
+ "-o",
+ default=False,
+ required=True,
+ type=str,
+ help="output dir",
+ )
+ return parser
+
+
+def main():
+ parser = get_parser()
+ args = parser.parse_args()
+
+ txt_writer = open("{}/text.{}.txt".format(args.output_dir, args.txt_index), 'w')
+ shape_writer = open("{}/len.{}".format(args.output_dir, args.txt_index), 'w')
+ seg_dict = load_dict(args.seg_file)
+ with open(args.text_file, 'r') as infile:
+ for line in infile:
+ s = line.strip().split()
+ text_id = s[0]
+ text_list = forward_segment("".join(s[1:]).lower(), seg_dict)
+ text = tokenize(text_list, seg_dict)
+ lens = len(text.strip().split())
+ txt_writer.write(text_id + " " + text + '\n')
+ shape_writer.write(text_id + " " + str(lens) + '\n')
+
+
+if __name__ == '__main__':
+ main()
+
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/text_tokenize.sh b/egs_modelscope/wenetspeech/paraformer/utils/text_tokenize.sh
new file mode 100755
index 0000000..6b74fef
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/text_tokenize.sh
@@ -0,0 +1,35 @@
+#!/usr/bin/env bash
+
+
+# Begin configuration section.
+nj=32
+cmd=utils/run.pl
+
+echo "$0 $@"
+
+. utils/parse_options.sh || exit 1;
+
+# tokenize configuration
+text_dir=$1
+seg_file=$2
+logdir=$3
+output_dir=$4
+
+txt_dir=${output_dir}/txt; mkdir -p ${output_dir}/txt
+mkdir -p ${logdir}
+
+$cmd JOB=1:$nj $logdir/text_tokenize.JOB.log \
+ python utils/text_tokenize.py -t ${text_dir}/txt/text.JOB.txt \
+ -s ${seg_file} -i JOB -o ${txt_dir} \
+ || exit 1;
+
+# concatenate the text files together.
+for n in $(seq $nj); do
+ cat ${txt_dir}/text.$n.txt || exit 1
+done > ${output_dir}/text || exit 1
+
+for n in $(seq $nj); do
+ cat ${txt_dir}/len.$n || exit 1
+done > ${output_dir}/text_shape || exit 1
+
+echo "$0: Succeeded text tokenize"
diff --git a/egs_modelscope/wenetspeech/paraformer/utils/textnorm_zh.py b/egs_modelscope/wenetspeech/paraformer/utils/textnorm_zh.py
new file mode 100755
index 0000000..79feb83
--- /dev/null
+++ b/egs_modelscope/wenetspeech/paraformer/utils/textnorm_zh.py
@@ -0,0 +1,834 @@
+#!/usr/bin/env python3
+# coding=utf-8
+
+# Authors:
+# 2019.5 Zhiyang Zhou (https://github.com/Joee1995/chn_text_norm.git)
+# 2019.9 Jiayu DU
+#
+# requirements:
+# - python 3.X
+# notes: python 2.X WILL fail or produce misleading results
+
+import sys, os, argparse, codecs, string, re
+
+# ================================================================================ #
+# basic constant
+# ================================================================================ #
+CHINESE_DIGIS = u'闆朵竴浜屼笁鍥涗簲鍏竷鍏節'
+BIG_CHINESE_DIGIS_SIMPLIFIED = u'闆跺9璐板弫鑲嗕紞闄嗘煉鎹岀帠'
+BIG_CHINESE_DIGIS_TRADITIONAL = u'闆跺9璨冲弮鑲嗕紞闄告煉鎹岀帠'
+SMALLER_BIG_CHINESE_UNITS_SIMPLIFIED = u'鍗佺櫨鍗冧竾'
+SMALLER_BIG_CHINESE_UNITS_TRADITIONAL = u'鎷句桨浠熻惉'
+LARGER_CHINESE_NUMERING_UNITS_SIMPLIFIED = u'浜垮厗浜灀绉┌娌熸锭姝h浇'
+LARGER_CHINESE_NUMERING_UNITS_TRADITIONAL = u'鍎勫厗浜灀绉┌婧濇緱姝h級'
+SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED = u'鍗佺櫨鍗冧竾'
+SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL = u'鎷句桨浠熻惉'
+
+ZERO_ALT = u'銆�'
+ONE_ALT = u'骞�'
+TWO_ALTS = [u'涓�', u'鍏�']
+
+POSITIVE = [u'姝�', u'姝�']
+NEGATIVE = [u'璐�', u'璨�']
+POINT = [u'鐐�', u'榛�']
+# PLUS = [u'鍔�', u'鍔�']
+# SIL = [u'鏉�', u'妲�']
+
+FILLER_CHARS = ['鍛�', '鍟�']
+ER_WHITELIST = '(鍎垮コ|鍎垮瓙|鍎垮瓩|濂冲効|鍎垮|濡诲効|' \
+ '鑳庡効|濠村効|鏂扮敓鍎縷濠村辜鍎縷骞煎効|灏戝効|灏忓効|鍎挎瓕|鍎跨|鍎跨|鎵樺効鎵�|瀛ゅ効|' \
+ '鍎挎垙|鍎垮寲|鍙板効搴剕楣垮効宀泑姝e効鍏粡|鍚婂効閮庡綋|鐢熷効鑲插コ|鎵樺効甯﹀コ|鍏诲効闃茶�亅鐥村効鍛嗗コ|' \
+ '浣冲効浣冲|鍎挎�滃吔鎵皘鍎挎棤甯哥埗|鍎夸笉瀚屾瘝涓憒鍎胯鍗冮噷姣嶆媴蹇鍎垮ぇ涓嶇敱鐖穦鑻忎篂鍎�)'
+
+# 涓枃鏁板瓧绯荤粺绫诲瀷
+NUMBERING_TYPES = ['low', 'mid', 'high']
+
+CURRENCY_NAMES = '(浜烘皯甯亅缇庡厓|鏃ュ厓|鑻遍晳|娆у厓|椹厠|娉曢儙|鍔犳嬁澶у厓|婢冲厓|娓竵|鍏堜护|鑺叞椹厠|鐖卞皵鍏伴晳|' \
+ '閲屾媺|鑽峰叞鐩緗鍩冩柉搴撳|姣斿濉攟鍗板凹鐩緗鏋楀悏鐗箌鏂拌タ鍏板厓|姣旂储|鍗㈠竷|鏂板姞鍧″厓|闊╁厓|娉伴摙)'
+CURRENCY_UNITS = '((浜縷鍗冧竾|鐧句竾|涓噟鍗億鐧�)|(浜縷鍗冧竾|鐧句竾|涓噟鍗億鐧緗)鍏億(浜縷鍗冧竾|鐧句竾|涓噟鍗億鐧緗)鍧梶瑙抾姣泑鍒�)'
+COM_QUANTIFIERS = '(鍖箌寮爘搴鍥瀨鍦簗灏緗鏉涓獆棣東闃檤闃祙缃憒鐐畖椤秥涓榺妫祙鍙獆鏀瘄琚瓅杈唡鎸憒鎷厊棰梶澹硘绐爘鏇瞸澧檤缇鑵攟' \
+ '鐮搴瀹璐瘄鎵巪鎹唡鍒�|浠鎵搢鎵媩缃梶鍧灞眧宀瓅姹焲婧獆閽焲闃焲鍗晐鍙寍瀵箌鍑簗鍙澶磡鑴殀鏉縷璺硘鏋潀浠秥璐磡' \
+ '閽坾绾縷绠鍚峾浣峾韬珅鍫倈璇緗鏈瑋椤祙瀹秥鎴穦灞倈涓潀姣珅鍘榺鍒唡閽眧涓鏂鎷厊閾鐭硘閽閿眧蹇絴(鍗億姣珅寰�)鍏媩' \
+ '姣珅鍘榺鍒唡瀵竱灏簗涓坾閲寍瀵粅甯竱閾簗绋媩(鍗億鍒唡鍘榺姣珅寰�)绫硘鎾畖鍕簗鍚坾鍗噟鏂梶鐭硘鐩榺纰梶纰焲鍙爘妗秥绗紎鐩唡' \
+ '鐩抾鏉瘄閽焲鏂泑閿厊绨媩绡畖鐩榺妗秥缃恷鐡秥澹秥鍗畖鐩弢绠﹟绠眧鐓瞸鍟東琚媩閽祙骞磡鏈坾鏃瀛鍒粅鏃秥鍛▅澶﹟绉抾鍒唡鏃瑋' \
+ '绾獆宀亅涓東鏇磡澶渱鏄澶弢绉媩鍐瑋浠浼弢杈坾涓竱娉绮抾棰梶骞鍫唡鏉鏍箌鏀瘄閬搢闈鐗噟寮爘棰梶鍧�)'
+
+# punctuation information are based on Zhon project (https://github.com/tsroten/zhon.git)
+CHINESE_PUNC_STOP = '锛侊紵锝°��'
+CHINESE_PUNC_NON_STOP = '锛傦純锛勶紖锛嗭紘锛堬級锛婏紜锛岋紞锛忥細锛涳紲锛濓紴锛狅蓟锛硷冀锛撅伎锝�锝涳綔锝濓綖锝燂綘锝剑锝ゃ�併�冦�嬨�屻�嶃�庛�忋�愩�戙�斻�曘�栥�椼�樸�欍�氥�涖�溿�濄�炪�熴�般�俱�库�撯�斺�樷�欌�涒�溾�濃�炩�熲�︹�э箯'
+CHINESE_PUNC_LIST = CHINESE_PUNC_STOP + CHINESE_PUNC_NON_STOP
+
+# ================================================================================ #
+# basic class
+# ================================================================================ #
+class ChineseChar(object):
+ """
+ 涓枃瀛楃
+ 姣忎釜瀛楃瀵瑰簲绠�浣撳拰绻佷綋,
+ e.g. 绠�浣� = '璐�', 绻佷綋 = '璨�'
+ 杞崲鏃跺彲杞崲涓虹畝浣撴垨绻佷綋
+ """
+
+ def __init__(self, simplified, traditional):
+ self.simplified = simplified
+ self.traditional = traditional
+ #self.__repr__ = self.__str__
+
+ def __str__(self):
+ return self.simplified or self.traditional or None
+
+ def __repr__(self):
+ return self.__str__()
+
+
+class ChineseNumberUnit(ChineseChar):
+ """
+ 涓枃鏁板瓧/鏁颁綅瀛楃
+ 姣忎釜瀛楃闄ょ箒绠�浣撳杩樻湁涓�涓澶栫殑澶у啓瀛楃
+ e.g. '闄�' 鍜� '闄�'
+ """
+
+ def __init__(self, power, simplified, traditional, big_s, big_t):
+ super(ChineseNumberUnit, self).__init__(simplified, traditional)
+ self.power = power
+ self.big_s = big_s
+ self.big_t = big_t
+
+ def __str__(self):
+ return '10^{}'.format(self.power)
+
+ @classmethod
+ def create(cls, index, value, numbering_type=NUMBERING_TYPES[1], small_unit=False):
+
+ if small_unit:
+ return ChineseNumberUnit(power=index + 1,
+ simplified=value[0], traditional=value[1], big_s=value[1], big_t=value[1])
+ elif numbering_type == NUMBERING_TYPES[0]:
+ return ChineseNumberUnit(power=index + 8,
+ simplified=value[0], traditional=value[1], big_s=value[0], big_t=value[1])
+ elif numbering_type == NUMBERING_TYPES[1]:
+ return ChineseNumberUnit(power=(index + 2) * 4,
+ simplified=value[0], traditional=value[1], big_s=value[0], big_t=value[1])
+ elif numbering_type == NUMBERING_TYPES[2]:
+ return ChineseNumberUnit(power=pow(2, index + 3),
+ simplified=value[0], traditional=value[1], big_s=value[0], big_t=value[1])
+ else:
+ raise ValueError(
+ 'Counting type should be in {0} ({1} provided).'.format(NUMBERING_TYPES, numbering_type))
+
+
+class ChineseNumberDigit(ChineseChar):
+ """
+ 涓枃鏁板瓧瀛楃
+ """
+
+ def __init__(self, value, simplified, traditional, big_s, big_t, alt_s=None, alt_t=None):
+ super(ChineseNumberDigit, self).__init__(simplified, traditional)
+ self.value = value
+ self.big_s = big_s
+ self.big_t = big_t
+ self.alt_s = alt_s
+ self.alt_t = alt_t
+
+ def __str__(self):
+ return str(self.value)
+
+ @classmethod
+ def create(cls, i, v):
+ return ChineseNumberDigit(i, v[0], v[1], v[2], v[3])
+
+
+class ChineseMath(ChineseChar):
+ """
+ 涓枃鏁颁綅瀛楃
+ """
+
+ def __init__(self, simplified, traditional, symbol, expression=None):
+ super(ChineseMath, self).__init__(simplified, traditional)
+ self.symbol = symbol
+ self.expression = expression
+ self.big_s = simplified
+ self.big_t = traditional
+
+
+CC, CNU, CND, CM = ChineseChar, ChineseNumberUnit, ChineseNumberDigit, ChineseMath
+
+
+class NumberSystem(object):
+ """
+ 涓枃鏁板瓧绯荤粺
+ """
+ pass
+
+
+class MathSymbol(object):
+ """
+ 鐢ㄤ簬涓枃鏁板瓧绯荤粺鐨勬暟瀛︾鍙� (绻�/绠�浣�), e.g.
+ positive = ['姝�', '姝�']
+ negative = ['璐�', '璨�']
+ point = ['鐐�', '榛�']
+ """
+
+ def __init__(self, positive, negative, point):
+ self.positive = positive
+ self.negative = negative
+ self.point = point
+
+ def __iter__(self):
+ for v in self.__dict__.values():
+ yield v
+
+
+# class OtherSymbol(object):
+# """
+# 鍏朵粬绗﹀彿
+# """
+#
+# def __init__(self, sil):
+# self.sil = sil
+#
+# def __iter__(self):
+# for v in self.__dict__.values():
+# yield v
+
+
+# ================================================================================ #
+# basic utils
+# ================================================================================ #
+def create_system(numbering_type=NUMBERING_TYPES[1]):
+ """
+ 鏍规嵁鏁板瓧绯荤粺绫诲瀷杩斿洖鍒涘缓鐩稿簲鐨勬暟瀛楃郴缁燂紝榛樿涓� mid
+ NUMBERING_TYPES = ['low', 'mid', 'high']: 涓枃鏁板瓧绯荤粺绫诲瀷
+ low: '鍏�' = '浜�' * '鍗�' = $10^{9}$, '浜�' = '鍏�' * '鍗�', etc.
+ mid: '鍏�' = '浜�' * '涓�' = $10^{12}$, '浜�' = '鍏�' * '涓�', etc.
+ high: '鍏�' = '浜�' * '浜�' = $10^{16}$, '浜�' = '鍏�' * '鍏�', etc.
+ 杩斿洖瀵瑰簲鐨勬暟瀛楃郴缁�
+ """
+
+ # chinese number units of '浜�' and larger
+ all_larger_units = zip(
+ LARGER_CHINESE_NUMERING_UNITS_SIMPLIFIED, LARGER_CHINESE_NUMERING_UNITS_TRADITIONAL)
+ larger_units = [CNU.create(i, v, numbering_type, False)
+ for i, v in enumerate(all_larger_units)]
+ # chinese number units of '鍗�, 鐧�, 鍗�, 涓�'
+ all_smaller_units = zip(
+ SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED, SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL)
+ smaller_units = [CNU.create(i, v, small_unit=True)
+ for i, v in enumerate(all_smaller_units)]
+ # digis
+ chinese_digis = zip(CHINESE_DIGIS, CHINESE_DIGIS,
+ BIG_CHINESE_DIGIS_SIMPLIFIED, BIG_CHINESE_DIGIS_TRADITIONAL)
+ digits = [CND.create(i, v) for i, v in enumerate(chinese_digis)]
+ digits[0].alt_s, digits[0].alt_t = ZERO_ALT, ZERO_ALT
+ digits[1].alt_s, digits[1].alt_t = ONE_ALT, ONE_ALT
+ digits[2].alt_s, digits[2].alt_t = TWO_ALTS[0], TWO_ALTS[1]
+
+ # symbols
+ positive_cn = CM(POSITIVE[0], POSITIVE[1], '+', lambda x: x)
+ negative_cn = CM(NEGATIVE[0], NEGATIVE[1], '-', lambda x: -x)
+ point_cn = CM(POINT[0], POINT[1], '.', lambda x,
+ y: float(str(x) + '.' + str(y)))
+ # sil_cn = CM(SIL[0], SIL[1], '-', lambda x, y: float(str(x) + '-' + str(y)))
+ system = NumberSystem()
+ system.units = smaller_units + larger_units
+ system.digits = digits
+ system.math = MathSymbol(positive_cn, negative_cn, point_cn)
+ # system.symbols = OtherSymbol(sil_cn)
+ return system
+
+
+def chn2num(chinese_string, numbering_type=NUMBERING_TYPES[1]):
+
+ def get_symbol(char, system):
+ for u in system.units:
+ if char in [u.traditional, u.simplified, u.big_s, u.big_t]:
+ return u
+ for d in system.digits:
+ if char in [d.traditional, d.simplified, d.big_s, d.big_t, d.alt_s, d.alt_t]:
+ return d
+ for m in system.math:
+ if char in [m.traditional, m.simplified]:
+ return m
+
+ def string2symbols(chinese_string, system):
+ int_string, dec_string = chinese_string, ''
+ for p in [system.math.point.simplified, system.math.point.traditional]:
+ if p in chinese_string:
+ int_string, dec_string = chinese_string.split(p)
+ break
+ return [get_symbol(c, system) for c in int_string], \
+ [get_symbol(c, system) for c in dec_string]
+
+ def correct_symbols(integer_symbols, system):
+ """
+ 涓�鐧惧叓 to 涓�鐧惧叓鍗�
+ 涓�浜夸竴鍗冧笁鐧句竾 to 涓�浜� 涓�鍗冧竾 涓夌櫨涓�
+ """
+
+ if integer_symbols and isinstance(integer_symbols[0], CNU):
+ if integer_symbols[0].power == 1:
+ integer_symbols = [system.digits[1]] + integer_symbols
+
+ if len(integer_symbols) > 1:
+ if isinstance(integer_symbols[-1], CND) and isinstance(integer_symbols[-2], CNU):
+ integer_symbols.append(
+ CNU(integer_symbols[-2].power - 1, None, None, None, None))
+
+ result = []
+ unit_count = 0
+ for s in integer_symbols:
+ if isinstance(s, CND):
+ result.append(s)
+ unit_count = 0
+ elif isinstance(s, CNU):
+ current_unit = CNU(s.power, None, None, None, None)
+ unit_count += 1
+
+ if unit_count == 1:
+ result.append(current_unit)
+ elif unit_count > 1:
+ for i in range(len(result)):
+ if isinstance(result[-i - 1], CNU) and result[-i - 1].power < current_unit.power:
+ result[-i - 1] = CNU(result[-i - 1].power +
+ current_unit.power, None, None, None, None)
+ return result
+
+ def compute_value(integer_symbols):
+ """
+ Compute the value.
+ When current unit is larger than previous unit, current unit * all previous units will be used as all previous units.
+ e.g. '涓ゅ崈涓�' = 2000 * 10000 not 2000 + 10000
+ """
+ value = [0]
+ last_power = 0
+ for s in integer_symbols:
+ if isinstance(s, CND):
+ value[-1] = s.value
+ elif isinstance(s, CNU):
+ value[-1] *= pow(10, s.power)
+ if s.power > last_power:
+ value[:-1] = list(map(lambda v: v *
+ pow(10, s.power), value[:-1]))
+ last_power = s.power
+ value.append(0)
+ return sum(value)
+
+ system = create_system(numbering_type)
+ int_part, dec_part = string2symbols(chinese_string, system)
+ int_part = correct_symbols(int_part, system)
+ int_str = str(compute_value(int_part))
+ dec_str = ''.join([str(d.value) for d in dec_part])
+ if dec_part:
+ return '{0}.{1}'.format(int_str, dec_str)
+ else:
+ return int_str
+
+
+def num2chn(number_string, numbering_type=NUMBERING_TYPES[1], big=False,
+ traditional=False, alt_zero=False, alt_one=False, alt_two=True,
+ use_zeros=True, use_units=True):
+
+ def get_value(value_string, use_zeros=True):
+
+ striped_string = value_string.lstrip('0')
+
+ # record nothing if all zeros
+ if not striped_string:
+ return []
+
+ # record one digits
+ elif len(striped_string) == 1:
+ if use_zeros and len(value_string) != len(striped_string):
+ return [system.digits[0], system.digits[int(striped_string)]]
+ else:
+ return [system.digits[int(striped_string)]]
+
+ # recursively record multiple digits
+ else:
+ result_unit = next(u for u in reversed(
+ system.units) if u.power < len(striped_string))
+ result_string = value_string[:-result_unit.power]
+ return get_value(result_string) + [result_unit] + get_value(striped_string[-result_unit.power:])
+
+ system = create_system(numbering_type)
+
+ int_dec = number_string.split('.')
+ if len(int_dec) == 1:
+ int_string = int_dec[0]
+ dec_string = ""
+ elif len(int_dec) == 2:
+ int_string = int_dec[0]
+ dec_string = int_dec[1]
+ else:
+ raise ValueError(
+ "invalid input num string with more than one dot: {}".format(number_string))
+
+ if use_units and len(int_string) > 1:
+ result_symbols = get_value(int_string)
+ else:
+ result_symbols = [system.digits[int(c)] for c in int_string]
+ dec_symbols = [system.digits[int(c)] for c in dec_string]
+ if dec_string:
+ result_symbols += [system.math.point] + dec_symbols
+
+ if alt_two:
+ liang = CND(2, system.digits[2].alt_s, system.digits[2].alt_t,
+ system.digits[2].big_s, system.digits[2].big_t)
+ for i, v in enumerate(result_symbols):
+ if isinstance(v, CND) and v.value == 2:
+ next_symbol = result_symbols[i +
+ 1] if i < len(result_symbols) - 1 else None
+ previous_symbol = result_symbols[i - 1] if i > 0 else None
+ if isinstance(next_symbol, CNU) and isinstance(previous_symbol, (CNU, type(None))):
+ if next_symbol.power != 1 and ((previous_symbol is None) or (previous_symbol.power != 1)):
+ result_symbols[i] = liang
+
+ # if big is True, '涓�' will not be used and `alt_two` has no impact on output
+ if big:
+ attr_name = 'big_'
+ if traditional:
+ attr_name += 't'
+ else:
+ attr_name += 's'
+ else:
+ if traditional:
+ attr_name = 'traditional'
+ else:
+ attr_name = 'simplified'
+
+ result = ''.join([getattr(s, attr_name) for s in result_symbols])
+
+ # if not use_zeros:
+ # result = result.strip(getattr(system.digits[0], attr_name))
+
+ if alt_zero:
+ result = result.replace(
+ getattr(system.digits[0], attr_name), system.digits[0].alt_s)
+
+ if alt_one:
+ result = result.replace(
+ getattr(system.digits[1], attr_name), system.digits[1].alt_s)
+
+ for i, p in enumerate(POINT):
+ if result.startswith(p):
+ return CHINESE_DIGIS[0] + result
+
+ # ^10, 11, .., 19
+ if len(result) >= 2 and result[1] in [SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED[0],
+ SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL[0]] and \
+ result[0] in [CHINESE_DIGIS[1], BIG_CHINESE_DIGIS_SIMPLIFIED[1], BIG_CHINESE_DIGIS_TRADITIONAL[1]]:
+ result = result[1:]
+
+ return result
+
+
+# ================================================================================ #
+# different types of rewriters
+# ================================================================================ #
+class Cardinal:
+ """
+ CARDINAL绫�
+ """
+
+ def __init__(self, cardinal=None, chntext=None):
+ self.cardinal = cardinal
+ self.chntext = chntext
+
+ def chntext2cardinal(self):
+ return chn2num(self.chntext)
+
+ def cardinal2chntext(self):
+ return num2chn(self.cardinal)
+
+class Digit:
+ """
+ DIGIT绫�
+ """
+
+ def __init__(self, digit=None, chntext=None):
+ self.digit = digit
+ self.chntext = chntext
+
+ # def chntext2digit(self):
+ # return chn2num(self.chntext)
+
+ def digit2chntext(self):
+ return num2chn(self.digit, alt_two=False, use_units=False)
+
+
+class TelePhone:
+ """
+ TELEPHONE绫�
+ """
+
+ def __init__(self, telephone=None, raw_chntext=None, chntext=None):
+ self.telephone = telephone
+ self.raw_chntext = raw_chntext
+ self.chntext = chntext
+
+ # def chntext2telephone(self):
+ # sil_parts = self.raw_chntext.split('<SIL>')
+ # self.telephone = '-'.join([
+ # str(chn2num(p)) for p in sil_parts
+ # ])
+ # return self.telephone
+
+ def telephone2chntext(self, fixed=False):
+
+ if fixed:
+ sil_parts = self.telephone.split('-')
+ self.raw_chntext = '<SIL>'.join([
+ num2chn(part, alt_two=False, use_units=False) for part in sil_parts
+ ])
+ self.chntext = self.raw_chntext.replace('<SIL>', '')
+ else:
+ sp_parts = self.telephone.strip('+').split()
+ self.raw_chntext = '<SP>'.join([
+ num2chn(part, alt_two=False, use_units=False) for part in sp_parts
+ ])
+ self.chntext = self.raw_chntext.replace('<SP>', '')
+ return self.chntext
+
+
+class Fraction:
+ """
+ FRACTION绫�
+ """
+
+ def __init__(self, fraction=None, chntext=None):
+ self.fraction = fraction
+ self.chntext = chntext
+
+ def chntext2fraction(self):
+ denominator, numerator = self.chntext.split('鍒嗕箣')
+ return chn2num(numerator) + '/' + chn2num(denominator)
+
+ def fraction2chntext(self):
+ numerator, denominator = self.fraction.split('/')
+ return num2chn(denominator) + '鍒嗕箣' + num2chn(numerator)
+
+
+class Date:
+ """
+ DATE绫�
+ """
+
+ def __init__(self, date=None, chntext=None):
+ self.date = date
+ self.chntext = chntext
+
+ # def chntext2date(self):
+ # chntext = self.chntext
+ # try:
+ # year, other = chntext.strip().split('骞�', maxsplit=1)
+ # year = Digit(chntext=year).digit2chntext() + '骞�'
+ # except ValueError:
+ # other = chntext
+ # year = ''
+ # if other:
+ # try:
+ # month, day = other.strip().split('鏈�', maxsplit=1)
+ # month = Cardinal(chntext=month).chntext2cardinal() + '鏈�'
+ # except ValueError:
+ # day = chntext
+ # month = ''
+ # if day:
+ # day = Cardinal(chntext=day[:-1]).chntext2cardinal() + day[-1]
+ # else:
+ # month = ''
+ # day = ''
+ # date = year + month + day
+ # self.date = date
+ # return self.date
+
+ def date2chntext(self):
+ date = self.date
+ try:
+ year, other = date.strip().split('骞�', 1)
+ year = Digit(digit=year).digit2chntext() + '骞�'
+ except ValueError:
+ other = date
+ year = ''
+ if other:
+ try:
+ month, day = other.strip().split('鏈�', 1)
+ month = Cardinal(cardinal=month).cardinal2chntext() + '鏈�'
+ except ValueError:
+ day = date
+ month = ''
+ if day:
+ day = Cardinal(cardinal=day[:-1]).cardinal2chntext() + day[-1]
+ else:
+ month = ''
+ day = ''
+ chntext = year + month + day
+ self.chntext = chntext
+ return self.chntext
+
+
+class Money:
+ """
+ MONEY绫�
+ """
+
+ def __init__(self, money=None, chntext=None):
+ self.money = money
+ self.chntext = chntext
+
+ # def chntext2money(self):
+ # return self.money
+
+ def money2chntext(self):
+ money = self.money
+ pattern = re.compile(r'(\d+(\.\d+)?)')
+ matchers = pattern.findall(money)
+ if matchers:
+ for matcher in matchers:
+ money = money.replace(matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext())
+ self.chntext = money
+ return self.chntext
+
+
+class Percentage:
+ """
+ PERCENTAGE绫�
+ """
+
+ def __init__(self, percentage=None, chntext=None):
+ self.percentage = percentage
+ self.chntext = chntext
+
+ def chntext2percentage(self):
+ return chn2num(self.chntext.strip().strip('鐧惧垎涔�')) + '%'
+
+ def percentage2chntext(self):
+ return '鐧惧垎涔�' + num2chn(self.percentage.strip().strip('%'))
+
+
+def remove_erhua(text, er_whitelist):
+ """
+ 鍘婚櫎鍎垮寲闊宠瘝涓殑鍎�:
+ 浠栧コ鍎垮湪閭h竟鍎� -> 浠栧コ鍎垮湪閭h竟
+ """
+
+ er_pattern = re.compile(er_whitelist)
+ new_str=''
+ while re.search('鍎�',text):
+ a = re.search('鍎�',text).span()
+ remove_er_flag = 0
+
+ if er_pattern.search(text):
+ b = er_pattern.search(text).span()
+ if b[0] <= a[0]:
+ remove_er_flag = 1
+
+ if remove_er_flag == 0 :
+ new_str = new_str + text[0:a[0]]
+ text = text[a[1]:]
+ else:
+ new_str = new_str + text[0:b[1]]
+ text = text[b[1]:]
+
+ text = new_str + text
+ return text
+
+# ================================================================================ #
+# NSW Normalizer
+# ================================================================================ #
+class NSWNormalizer:
+ def __init__(self, raw_text):
+ self.raw_text = '^' + raw_text + '$'
+ self.norm_text = ''
+
+ def _particular(self):
+ text = self.norm_text
+ pattern = re.compile(r"(([a-zA-Z]+)浜�([a-zA-Z]+))")
+ matchers = pattern.findall(text)
+ if matchers:
+ # print('particular')
+ for matcher in matchers:
+ text = text.replace(matcher[0], matcher[1]+'2'+matcher[2], 1)
+ self.norm_text = text
+ return self.norm_text
+
+ def normalize(self):
+ text = self.raw_text
+
+ # 瑙勮寖鍖栨棩鏈�
+ pattern = re.compile(r"\D+((([089]\d|(19|20)\d{2})骞�)?(\d{1,2}鏈�(\d{1,2}[鏃ュ彿])?)?)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('date')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Date(date=matcher[0]).date2chntext(), 1)
+
+ # 瑙勮寖鍖栭噾閽�
+ pattern = re.compile(r"\D+((\d+(\.\d+)?)[澶氫綑鍑燷?" + CURRENCY_UNITS + r"(\d" + CURRENCY_UNITS + r"?)?)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('money')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Money(money=matcher[0]).money2chntext(), 1)
+
+ # 瑙勮寖鍖栧浐璇�/鎵嬫満鍙风爜
+ # 鎵嬫満
+ # http://www.jihaoba.com/news/show/13680
+ # 绉诲姩锛�139銆�138銆�137銆�136銆�135銆�134銆�159銆�158銆�157銆�150銆�151銆�152銆�188銆�187銆�182銆�183銆�184銆�178銆�198
+ # 鑱旈�氾細130銆�131銆�132銆�156銆�155銆�186銆�185銆�176
+ # 鐢典俊锛�133銆�153銆�189銆�180銆�181銆�177
+ pattern = re.compile(r"\D((\+?86 ?)?1([38]\d|5[0-35-9]|7[678]|9[89])\d{8})\D")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('telephone')
+ for matcher in matchers:
+ text = text.replace(matcher[0], TelePhone(telephone=matcher[0]).telephone2chntext(), 1)
+ # 鍥鸿瘽
+ pattern = re.compile(r"\D((0(10|2[1-3]|[3-9]\d{2})-?)?[1-9]\d{6,7})\D")
+ matchers = pattern.findall(text)
+ if matchers:
+ # print('fixed telephone')
+ for matcher in matchers:
+ text = text.replace(matcher[0], TelePhone(telephone=matcher[0]).telephone2chntext(fixed=True), 1)
+
+ # 瑙勮寖鍖栧垎鏁�
+ pattern = re.compile(r"(\d+/\d+)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('fraction')
+ for matcher in matchers:
+ text = text.replace(matcher, Fraction(fraction=matcher).fraction2chntext(), 1)
+
+ # 瑙勮寖鍖栫櫨鍒嗘暟
+ text = text.replace('锛�', '%')
+ pattern = re.compile(r"(\d+(\.\d+)?%)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('percentage')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Percentage(percentage=matcher[0]).percentage2chntext(), 1)
+
+ # 瑙勮寖鍖栫函鏁�+閲忚瘝
+ pattern = re.compile(r"(\d+(\.\d+)?)[澶氫綑鍑燷?" + COM_QUANTIFIERS)
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('cardinal+quantifier')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext(), 1)
+
+ # 瑙勮寖鍖栨暟瀛楃紪鍙�
+ pattern = re.compile(r"(\d{4,32})")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('digit')
+ for matcher in matchers:
+ text = text.replace(matcher, Digit(digit=matcher).digit2chntext(), 1)
+
+ # 瑙勮寖鍖栫函鏁�
+ pattern = re.compile(r"(\d+(\.\d+)?)")
+ matchers = pattern.findall(text)
+ if matchers:
+ #print('cardinal')
+ for matcher in matchers:
+ text = text.replace(matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext(), 1)
+
+ self.norm_text = text
+ self._particular()
+
+ return self.norm_text.lstrip('^').rstrip('$')
+
+
+def nsw_test_case(raw_text):
+ print('I:' + raw_text)
+ print('O:' + NSWNormalizer(raw_text).normalize())
+ print('')
+
+
+def nsw_test():
+ nsw_test_case('鍥鸿瘽锛�0595-23865596鎴�23880880銆�')
+ nsw_test_case('鍥鸿瘽锛�0595-23865596鎴�23880880銆�')
+ nsw_test_case('鎵嬫満锛�+86 19859213959鎴�15659451527銆�')
+ nsw_test_case('鍒嗘暟锛�32477/76391銆�')
+ nsw_test_case('鐧惧垎鏁帮細80.03%銆�')
+ nsw_test_case('缂栧彿锛�31520181154418銆�')
+ nsw_test_case('绾暟锛�2983.07鍏嬫垨12345.60绫炽��')
+ nsw_test_case('鏃ユ湡锛�1999骞�2鏈�20鏃ユ垨09骞�3鏈�15鍙枫��')
+ nsw_test_case('閲戦挶锛�12鍧�5锛�34.5鍏冿紝20.1涓�')
+ nsw_test_case('鐗规畩锛歄2O鎴朆2C銆�')
+ nsw_test_case('3456涓囧惃')
+ nsw_test_case('2938涓�')
+ nsw_test_case('938')
+ nsw_test_case('浠婂ぉ鍚冧簡115涓皬绗煎寘231涓澶�')
+ nsw_test_case('鏈�62锛呯殑姒傜巼')
+
+
+if __name__ == '__main__':
+ #nsw_test()
+
+ p = argparse.ArgumentParser()
+ p.add_argument('ifile', help='input filename, assume utf-8 encoding')
+ p.add_argument('ofile', help='output filename')
+ p.add_argument('--to_upper', action='store_true', help='convert to upper case')
+ p.add_argument('--to_lower', action='store_true', help='convert to lower case')
+ p.add_argument('--has_key', action='store_true', help="input text has Kaldi's key as first field.")
+ p.add_argument('--remove_fillers', type=bool, default=True, help='remove filler chars such as "鍛�, 鍟�"')
+ p.add_argument('--remove_erhua', type=bool, default=True, help='remove erhua chars such as "杩欏効"')
+ p.add_argument('--log_interval', type=int, default=10000, help='log interval in number of processed lines')
+ args = p.parse_args()
+
+ ifile = codecs.open(args.ifile, 'r', 'utf8')
+ ofile = codecs.open(args.ofile, 'w+', 'utf8')
+
+ n = 0
+ for l in ifile:
+ key = ''
+ text = ''
+ if args.has_key:
+ cols = l.split(maxsplit=1)
+ key = cols[0]
+ if len(cols) == 2:
+ text = cols[1].strip()
+ else:
+ text = ''
+ else:
+ text = l.strip()
+
+ # cases
+ if args.to_upper and args.to_lower:
+ sys.stderr.write('text norm: to_upper OR to_lower?')
+ exit(1)
+ if args.to_upper:
+ text = text.upper()
+ if args.to_lower:
+ text = text.lower()
+
+ # Filler chars removal
+ if args.remove_fillers:
+ for ch in FILLER_CHARS:
+ text = text.replace(ch, '')
+
+ if args.remove_erhua:
+ text = remove_erhua(text, ER_WHITELIST)
+
+ # NSW(Non-Standard-Word) normalization
+ text = NSWNormalizer(text).normalize()
+
+ # Punctuations removal
+ old_chars = CHINESE_PUNC_LIST + string.punctuation # includes all CN and EN punctuations
+ new_chars = ' ' * len(old_chars)
+ del_chars = ''
+ text = text.translate(str.maketrans(old_chars, new_chars, del_chars))
+
+ #
+ if args.has_key:
+ ofile.write(key + '\t' + text + '\n')
+ else:
+ ofile.write(text + '\n')
+
+ n += 1
+ if n % args.log_interval == 0:
+ sys.stderr.write("text norm: {} lines done.\n".format(n))
+
+ sys.stderr.write("text norm: {} lines done in total.\n".format(n))
+
+ ifile.close()
+ ofile.close()
diff --git a/funasr/bin/asr_inference.py b/funasr/bin/asr_inference.py
index 6ee0ffe..bd5d7f4 100755
--- a/funasr/bin/asr_inference.py
+++ b/funasr/bin/asr_inference.py
@@ -12,6 +12,7 @@
from typing import Sequence
from typing import Tuple
from typing import Union
+from typing import Dict
import numpy as np
import torch
@@ -38,7 +39,21 @@
from funasr.utils.types import str2bool
from funasr.utils.types import str2triple_str
from funasr.utils.types import str_or_none
+from funasr.utils import asr_utils, wav_utils, postprocess_utils
+from funasr.models.frontend.wav_frontend import WavFrontend
+from modelscope.utils.logger import get_logger
+
+logger = get_logger()
+
+header_colors = '\033[95m'
+end_colors = '\033[0m'
+
+global_asr_language: str = 'zh-cn'
+global_sample_rate: Union[int, Dict[Any, int]] = {
+ 'audio_fs': 16000,
+ 'model_fs': 16000
+}
class Speech2Text:
"""Speech2Text class
@@ -72,6 +87,7 @@
penalty: float = 0.0,
nbest: int = 1,
streaming: bool = False,
+ frontend_conf: dict = None,
**kwargs,
):
assert check_argument_types()
@@ -81,6 +97,9 @@
asr_model, asr_train_args = ASRTask.build_model_from_file(
asr_train_config, asr_model_file, device
)
+ if asr_model.frontend is None and frontend_conf is not None:
+ frontend = WavFrontend(**frontend_conf)
+ asr_model.frontend = frontend
logging.info("asr_model: {}".format(asr_model))
logging.info("asr_train_args: {}".format(asr_train_args))
asr_model.to(dtype=getattr(torch, dtype)).eval()
@@ -129,36 +148,6 @@
pre_beam_score_key=None if ctc_weight == 1.0 else "full",
)
- # TODO(karita): make all scorers batchfied
- if batch_size == 1:
- non_batch = [
- k
- for k, v in beam_search.full_scorers.items()
- if not isinstance(v, BatchScorerInterface)
- ]
- if len(non_batch) == 0:
- if streaming:
- beam_search.__class__ = BatchBeamSearchOnlineSim
- beam_search.set_streaming_config(asr_train_config)
- logging.info(
- "BatchBeamSearchOnlineSim implementation is selected."
- )
- else:
- beam_search.__class__ = BatchBeamSearch
- logging.info("BatchBeamSearch implementation is selected.")
- else:
- logging.warning(
- f"As non-batch scorers {non_batch} are found, "
- f"fall back to non-batch implementation."
- )
-
- beam_search.to(device=device, dtype=getattr(torch, dtype)).eval()
- for scorer in scorers.values():
- if isinstance(scorer, torch.nn.Module):
- scorer.to(device=device, dtype=getattr(torch, dtype)).eval()
- logging.info(f"Beam_search: {beam_search}")
- logging.info(f"Decoding device={device}, dtype={dtype}")
-
# 5. [Optional] Build Text converter: e.g. bpe-sym -> Text
if token_type is None:
token_type = asr_train_args.token_type
@@ -203,7 +192,7 @@
"""Inference
Args:
- data: Input speech data
+ speech: Input speech data
Returns:
text, token, token_int, hyp
@@ -216,6 +205,7 @@
# data: (Nsamples,) -> (1, Nsamples)
speech = speech.unsqueeze(0).to(getattr(torch, self.dtype))
+ lfr_factor = max(1, (speech.size()[-1] // 80) - 1)
# lengths: (1,)
lengths = speech.new_full([1], dtype=torch.long, fill_value=speech.size(1))
batch = {"speech": speech, "speech_lengths": lengths}
@@ -264,32 +254,36 @@
def inference(
- output_dir: str,
maxlenratio: float,
minlenratio: float,
batch_size: int,
- dtype: str,
beam_size: int,
ngpu: int,
- seed: int,
ctc_weight: float,
lm_weight: float,
- ngram_weight: float,
penalty: float,
- nbest: int,
- num_workers: int,
log_level: Union[int, str],
- data_path_and_name_and_type: Sequence[Tuple[str, str, str]],
- key_file: Optional[str],
+ data_path_and_name_and_type,
asr_train_config: Optional[str],
asr_model_file: Optional[str],
- lm_train_config: Optional[str],
- lm_file: Optional[str],
- word_lm_train_config: Optional[str],
- token_type: Optional[str],
- bpemodel: Optional[str],
- allow_variable_data_keys: bool,
- streaming: bool,
+ audio_lists: Union[List[Any], bytes] = None,
+ lm_train_config: Optional[str] = None,
+ lm_file: Optional[str] = None,
+ token_type: Optional[str] = None,
+ key_file: Optional[str] = None,
+ word_lm_train_config: Optional[str] = None,
+ bpemodel: Optional[str] = None,
+ allow_variable_data_keys: bool = False,
+ streaming: bool = False,
+ output_dir: Optional[str] = None,
+ dtype: str = "float32",
+ seed: int = 0,
+ ngram_weight: float = 0.9,
+ nbest: int = 1,
+ num_workers: int = 1,
+ frontend_conf: dict = None,
+ fs: Union[dict, int] = 16000,
+ lang: Optional[str] = None,
**kwargs,
):
assert check_argument_types()
@@ -309,7 +303,46 @@
device = "cuda"
else:
device = "cpu"
+ hop_length: int = 160
+ sr: int = 16000
+ if isinstance(fs, int):
+ sr = fs
+ else:
+ if 'model_fs' in fs and fs['model_fs'] is not None:
+ sr = fs['model_fs']
+ # data_path_and_name_and_type for modelscope: (data from audio_lists)
+ # ['speech', 'sound', 'am.mvn']
+ # data_path_and_name_and_type for funasr:
+ # [('/mnt/data/jiangyu.xzy/exp/maas/mvn.1.scp', 'speech', 'kaldi_ark')]
+ if isinstance(data_path_and_name_and_type[0], Tuple):
+ features_type: str = data_path_and_name_and_type[0][1]
+ elif isinstance(data_path_and_name_and_type[0], str):
+ features_type: str = data_path_and_name_and_type[1]
+ else:
+ raise NotImplementedError("unknown features type:{0}".format(data_path_and_name_and_type))
+ if features_type != 'sound':
+ frontend_conf = None
+ flag_modelscope = False
+ else:
+ flag_modelscope = True
+ if frontend_conf is not None:
+ if 'hop_length' in frontend_conf:
+ hop_length = frontend_conf['hop_length']
+ finish_count = 0
+ file_count = 1
+ if flag_modelscope and not isinstance(data_path_and_name_and_type[0], Tuple):
+ data_path_and_name_and_type_new = [
+ audio_lists, data_path_and_name_and_type[0], data_path_and_name_and_type[1]
+ ]
+ if isinstance(audio_lists, bytes):
+ file_count = 1
+ else:
+ file_count = len(audio_lists)
+ if len(data_path_and_name_and_type) >= 3 and frontend_conf is not None:
+ mvn_file = data_path_and_name_and_type[2]
+ mvn_data = wav_utils.extract_CMVN_featrures(mvn_file)
+ frontend_conf['mvn_data'] = mvn_data
# 1. Set random-seed
set_all_random_seed(seed)
@@ -332,45 +365,66 @@
penalty=penalty,
nbest=nbest,
streaming=streaming,
+ frontend_conf=frontend_conf,
)
logging.info("speech2text_kwargs: {}".format(speech2text_kwargs))
speech2text = Speech2Text(**speech2text_kwargs)
# 3. Build data-iterator
- loader = ASRTask.build_streaming_iterator(
- data_path_and_name_and_type,
- dtype=dtype,
- batch_size=batch_size,
- key_file=key_file,
- num_workers=num_workers,
- preprocess_fn=ASRTask.build_preprocess_fn(speech2text.asr_train_args, False),
- collate_fn=ASRTask.build_collate_fn(speech2text.asr_train_args, False),
- allow_variable_data_keys=allow_variable_data_keys,
- inference=True,
- )
+ if flag_modelscope:
+ loader = ASRTask.build_streaming_iterator_modelscope(
+ data_path_and_name_and_type_new,
+ dtype=dtype,
+ batch_size=batch_size,
+ key_file=key_file,
+ num_workers=num_workers,
+ preprocess_fn=ASRTask.build_preprocess_fn(speech2text.asr_train_args, False),
+ collate_fn=ASRTask.build_collate_fn(speech2text.asr_train_args, False),
+ allow_variable_data_keys=allow_variable_data_keys,
+ inference=True,
+ sample_rate=fs
+ )
+ else:
+ loader = ASRTask.build_streaming_iterator(
+ data_path_and_name_and_type,
+ dtype=dtype,
+ batch_size=batch_size,
+ key_file=key_file,
+ num_workers=num_workers,
+ preprocess_fn=ASRTask.build_preprocess_fn(speech2text.asr_train_args, False),
+ collate_fn=ASRTask.build_collate_fn(speech2text.asr_train_args, False),
+ allow_variable_data_keys=allow_variable_data_keys,
+ inference=True,
+ )
# 7 .Start for-loop
# FIXME(kamo): The output format should be discussed about
- with DatadirWriter(output_dir) as writer:
- for keys, batch in loader:
- assert isinstance(batch, dict), type(batch)
- assert all(isinstance(s, str) for s in keys), keys
- _bs = len(next(iter(batch.values())))
- assert len(keys) == _bs, f"{len(keys)} != {_bs}"
- batch = {k: v[0] for k, v in batch.items() if not k.endswith("_lengths")}
+ asr_result_list = []
+ if output_dir is not None:
+ writer = DatadirWriter(output_dir)
+ else:
+ writer = None
- # N-best list of (text, token, token_int, hyp_object)
- try:
- results = speech2text(**batch)
- except TooShortUttError as e:
- logging.warning(f"Utterance {keys} {e}")
- hyp = Hypothesis(score=0.0, scores={}, states={}, yseq=[])
- results = [[" ", ["<space>"], [2], hyp]] * nbest
+ for keys, batch in loader:
+ assert isinstance(batch, dict), type(batch)
+ assert all(isinstance(s, str) for s in keys), keys
+ _bs = len(next(iter(batch.values())))
+ assert len(keys) == _bs, f"{len(keys)} != {_bs}"
+ batch = {k: v[0] for k, v in batch.items() if not k.endswith("_lengths")}
- # Only supporting batch_size==1
- key = keys[0]
- for n, (text, token, token_int, hyp) in zip(range(1, nbest + 1), results):
- # Create a directory: outdir/{n}best_recog
+ # N-best list of (text, token, token_int, hyp_object)
+ try:
+ results = speech2text(**batch)
+ except TooShortUttError as e:
+ logging.warning(f"Utterance {keys} {e}")
+ hyp = Hypothesis(score=0.0, scores={}, states={}, yseq=[])
+ results = [[" ", ["<space>"], [2], hyp]] * nbest
+
+ # Only supporting batch_size==1
+ key = keys[0]
+ for n, (text, token, token_int, hyp) in zip(range(1, nbest + 1), results):
+ # Create a directory: outdir/{n}best_recog
+ if writer is not None:
ibest_writer = writer[f"{n}best_recog"]
# Write the result to each file
@@ -378,8 +432,25 @@
ibest_writer["token_int"][key] = " ".join(map(str, token_int))
ibest_writer["score"][key] = str(hyp.score)
- if text is not None:
+ if text is not None:
+ text_postprocessed = postprocess_utils.sentence_postprocess(token)
+ item = {'key': key, 'value': text_postprocessed}
+ asr_result_list.append(item)
+ finish_count += 1
+ asr_utils.print_progress(finish_count / file_count)
+ if writer is not None:
ibest_writer["text"][key] = text
+ return asr_result_list
+
+
+def set_parameters(language: str = None,
+ sample_rate: Union[int, Dict[Any, int]] = None):
+ if language is not None:
+ global global_asr_language
+ global_asr_language = language
+ if sample_rate is not None:
+ global global_sample_rate
+ global_sample_rate = sample_rate
def get_parser():
@@ -432,6 +503,8 @@
required=True,
action="append",
)
+ group.add_argument("--audio_lists", type=list, default=None)
+ # example=[{'key':'EdevDEWdIYQ_0021','file':'/mnt/data/jiangyu.xzy/test_data/speech_io/SPEECHIO_ASR_ZH00007_zhibodaihuo/wav/EdevDEWdIYQ_0021.wav'}])
group.add_argument("--key_file", type=str_or_none)
group.add_argument("--allow_variable_data_keys", type=str2bool, default=False)
diff --git a/funasr/bin/asr_inference_launch.py b/funasr/bin/asr_inference_launch.py
index 9d328ad..84e1422 100755
--- a/funasr/bin/asr_inference_launch.py
+++ b/funasr/bin/asr_inference_launch.py
@@ -6,6 +6,7 @@
import logging
import os
import sys
+from typing import Union, Dict, Any
from funasr.utils import config_argparse
from funasr.utils.cli_utils import get_commandline_args
@@ -181,6 +182,31 @@
return parser
+def set_parameters(language: str = None,
+ sample_rate: Union[int, Dict[Any, int]] = None):
+ if language is not None:
+ global global_asr_language
+ global_asr_language = language
+ if sample_rate is not None:
+ global global_sample_rate
+ global_sample_rate = sample_rate
+
+
+def inference_launch(mode, **kwargs):
+ if mode == "asr":
+ from funasr.bin.asr_inference import inference
+ return inference(**kwargs)
+ elif mode == "uniasr":
+ from funasr.bin.asr_inference_uniasr import inference
+ return inference(**kwargs)
+ elif mode == "paraformer":
+ from funasr.bin.asr_inference_paraformer import inference
+ return inference(**kwargs)
+ else:
+ logging.info("Unknown decoding mode: {}".format(mode))
+ return None
+
+
def main(cmd=None):
print(get_commandline_args(), file=sys.stderr)
parser = get_parser()
@@ -208,17 +234,7 @@
os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID"
os.environ["CUDA_VISIBLE_DEVICES"] = gpuid
- if args.mode == "asr":
- from funasr.bin.asr_inference import inference
- inference(**kwargs)
- elif args.mode == "uniasr":
- from funasr.bin.asr_inference_uniasr import inference
- inference(**kwargs)
- elif args.mode == "paraformer":
- from funasr.bin.asr_inference_paraformer import inference
- inference(**kwargs)
- else:
- logging.info("Unknown decoding mode: {}".format(args.mode))
+ inference_launch(**kwargs)
if __name__ == "__main__":
diff --git a/funasr/bin/asr_inference_modelscope.py b/funasr/bin/asr_inference_modelscope.py
deleted file mode 100755
index fd9bd66..0000000
--- a/funasr/bin/asr_inference_modelscope.py
+++ /dev/null
@@ -1,687 +0,0 @@
-#!/usr/bin/env python3
-# Copyright ESPnet (https://github.com/espnet/espnet). All Rights Reserved.
-# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
-
-import argparse
-import logging
-import sys
-from pathlib import Path
-from typing import Any
-from typing import List
-from typing import Optional
-from typing import Sequence
-from typing import Tuple
-from typing import Union
-from typing import Dict
-
-import numpy as np
-import torch
-from typeguard import check_argument_types
-from typeguard import check_return_type
-
-from funasr.fileio.datadir_writer import DatadirWriter
-from funasr.modules.beam_search.batch_beam_search import BatchBeamSearch
-from funasr.modules.beam_search.batch_beam_search_online_sim import BatchBeamSearchOnlineSim
-from funasr.modules.beam_search.beam_search import BeamSearch
-from funasr.modules.beam_search.beam_search import Hypothesis
-from funasr.modules.scorers.ctc import CTCPrefixScorer
-from funasr.modules.scorers.length_bonus import LengthBonus
-from funasr.modules.scorers.scorer_interface import BatchScorerInterface
-from funasr.modules.subsampling import TooShortUttError
-from funasr.tasks.asr import ASRTask
-from funasr.tasks.lm import LMTask
-from funasr.text.build_tokenizer import build_tokenizer
-from funasr.text.token_id_converter import TokenIDConverter
-from funasr.torch_utils.device_funcs import to_device
-from funasr.torch_utils.set_all_random_seed import set_all_random_seed
-from funasr.utils import config_argparse
-from funasr.utils.cli_utils import get_commandline_args
-from funasr.utils.types import str2bool
-from funasr.utils.types import str2triple_str
-from funasr.utils.types import str_or_none
-from funasr.utils import asr_utils, wav_utils, postprocess_utils
-from funasr.models.frontend.wav_frontend import WavFrontend
-
-from modelscope.utils.logger import get_logger
-
-logger = get_logger()
-
-header_colors = '\033[95m'
-end_colors = '\033[0m'
-
-global_asr_language: str = 'zh-cn'
-global_sample_rate: Union[int, Dict[Any, int]] = {
- 'audio_fs': 16000,
- 'model_fs': 16000
-}
-
-class Speech2Text:
- """Speech2Text class
-
- Examples:
- >>> import soundfile
- >>> speech2text = Speech2Text("asr_config.yml", "asr.pth")
- >>> audio, rate = soundfile.read("speech.wav")
- >>> speech2text(audio)
- [(text, token, token_int, hypothesis object), ...]
-
- """
-
- def __init__(
- self,
- asr_train_config: Union[Path, str] = None,
- asr_model_file: Union[Path, str] = None,
- lm_train_config: Union[Path, str] = None,
- lm_file: Union[Path, str] = None,
- token_type: str = None,
- bpemodel: str = None,
- device: str = "cpu",
- maxlenratio: float = 0.0,
- minlenratio: float = 0.0,
- batch_size: int = 1,
- dtype: str = "float32",
- beam_size: int = 20,
- ctc_weight: float = 0.5,
- lm_weight: float = 1.0,
- ngram_weight: float = 0.9,
- penalty: float = 0.0,
- nbest: int = 1,
- streaming: bool = False,
- frontend_conf: dict = None,
- **kwargs,
- ):
- assert check_argument_types()
-
- # 1. Build ASR model
- scorers = {}
- asr_model, asr_train_args = ASRTask.build_model_from_file(
- asr_train_config, asr_model_file, device
- )
- if asr_model.frontend is None and frontend_conf is not None:
- frontend = WavFrontend(**frontend_conf)
- asr_model.frontend = frontend
- asr_model.to(dtype=getattr(torch, dtype)).eval()
-
- decoder = asr_model.decoder
-
- ctc = CTCPrefixScorer(ctc=asr_model.ctc, eos=asr_model.eos)
- token_list = asr_model.token_list
- scorers.update(
- decoder=decoder,
- ctc=ctc,
- length_bonus=LengthBonus(len(token_list)),
- )
-
- # 2. Build Language model
- if lm_train_config is not None:
- lm, lm_train_args = LMTask.build_model_from_file(
- lm_train_config, lm_file, device
- )
- scorers["lm"] = lm.lm
-
- # 3. Build ngram model
- # ngram is not supported now
- ngram = None
- scorers["ngram"] = ngram
-
- # 4. Build BeamSearch object
- # transducer is not supported now
- beam_search_transducer = None
-
- weights = dict(
- decoder=1.0 - ctc_weight,
- ctc=ctc_weight,
- lm=lm_weight,
- ngram=ngram_weight,
- length_bonus=penalty,
- )
- beam_search = BeamSearch(
- beam_size=beam_size,
- weights=weights,
- scorers=scorers,
- sos=asr_model.sos,
- eos=asr_model.eos,
- vocab_size=len(token_list),
- token_list=token_list,
- pre_beam_score_key=None if ctc_weight == 1.0 else "full",
- )
-
- # TODO(karita): make all scorers batchfied
- if batch_size == 1:
- non_batch = [
- k
- for k, v in beam_search.full_scorers.items()
- if not isinstance(v, BatchScorerInterface)
- ]
- if len(non_batch) == 0:
- if streaming:
- beam_search.__class__ = BatchBeamSearchOnlineSim
- beam_search.set_streaming_config(asr_train_config)
- logging.info(
- "BatchBeamSearchOnlineSim implementation is selected."
- )
- else:
- beam_search.__class__ = BatchBeamSearch
- logging.info("BatchBeamSearch implementation is selected.")
- else:
- logging.warning(
- f"As non-batch scorers {non_batch} are found, "
- f"fall back to non-batch implementation."
- )
-
- beam_search.to(device=device, dtype=getattr(torch, dtype)).eval()
- for scorer in scorers.values():
- if isinstance(scorer, torch.nn.Module):
- scorer.to(device=device, dtype=getattr(torch, dtype)).eval()
- logging.info(f"Beam_search: {beam_search}")
- logging.info(f"Decoding device={device}, dtype={dtype}")
-
- # 5. [Optional] Build Text converter: e.g. bpe-sym -> Text
- if token_type is None:
- token_type = asr_train_args.token_type
- if bpemodel is None:
- bpemodel = asr_train_args.bpemodel
-
- if token_type is None:
- tokenizer = None
- elif token_type == "bpe":
- if bpemodel is not None:
- tokenizer = build_tokenizer(token_type=token_type, bpemodel=bpemodel)
- else:
- tokenizer = None
- else:
- tokenizer = build_tokenizer(token_type=token_type)
- converter = TokenIDConverter(token_list=token_list)
- logging.info(f"Text tokenizer: {tokenizer}")
-
- self.asr_model = asr_model
- self.asr_train_args = asr_train_args
- self.converter = converter
- self.tokenizer = tokenizer
- self.beam_search = beam_search
- self.beam_search_transducer = beam_search_transducer
- self.maxlenratio = maxlenratio
- self.minlenratio = minlenratio
- self.device = device
- self.dtype = dtype
- self.nbest = nbest
-
- @torch.no_grad()
- def __call__(
- self, speech: Union[torch.Tensor, np.ndarray]
- ) -> List[
- Tuple[
- Optional[str],
- List[str],
- List[int],
- Union[Hypothesis],
- ]
- ]:
- """Inference
-
- Args:
- speech: Input speech data
- Returns:
- text, token, token_int, hyp
-
- """
- assert check_argument_types()
-
- # Input as audio signal
- if isinstance(speech, np.ndarray):
- speech = torch.tensor(speech)
-
- # data: (Nsamples,) -> (1, Nsamples)
- speech = speech.unsqueeze(0).to(getattr(torch, self.dtype))
- lfr_factor = max(1, (speech.size()[-1] // 80) - 1)
- # lengths: (1,)
- lengths = speech.new_full([1], dtype=torch.long, fill_value=speech.size(1))
- batch = {"speech": speech, "speech_lengths": lengths}
-
- # a. To device
- batch = to_device(batch, device=self.device)
-
- # b. Forward Encoder
- enc, _ = self.asr_model.encode(**batch)
- if isinstance(enc, tuple):
- enc = enc[0]
- assert len(enc) == 1, len(enc)
-
- # c. Passed the encoder result and the beam search
- nbest_hyps = self.beam_search(
- x=enc[0], maxlenratio=self.maxlenratio, minlenratio=self.minlenratio
- )
-
- nbest_hyps = nbest_hyps[: self.nbest]
-
- results = []
- for hyp in nbest_hyps:
- assert isinstance(hyp, (Hypothesis)), type(hyp)
-
- # remove sos/eos and get results
- last_pos = -1
- if isinstance(hyp.yseq, list):
- token_int = hyp.yseq[1:last_pos]
- else:
- token_int = hyp.yseq[1:last_pos].tolist()
-
- # remove blank symbol id, which is assumed to be 0
- token_int = list(filter(lambda x: x != 0, token_int))
-
- # Change integer-ids to tokens
- token = self.converter.ids2tokens(token_int)
-
- if self.tokenizer is not None:
- text = self.tokenizer.tokens2text(token)
- else:
- text = None
- results.append((text, token, token_int, hyp))
-
- assert check_return_type(results)
- return results
-
-
-def inference(
- maxlenratio: float,
- minlenratio: float,
- batch_size: int,
- dtype: str,
- beam_size: int,
- ngpu: int,
- seed: int,
- ctc_weight: float,
- lm_weight: float,
- ngram_weight: float,
- penalty: float,
- nbest: int,
- num_workers: int,
- log_level: Union[int, str],
- data_path_and_name_and_type: list,
- audio_lists: Union[List[Any], bytes],
- key_file: Optional[str],
- asr_train_config: Optional[str],
- asr_model_file: Optional[str],
- lm_train_config: Optional[str],
- lm_file: Optional[str],
- word_lm_train_config: Optional[str],
- token_type: Optional[str],
- bpemodel: Optional[str],
- output_dir: Optional[str],
- allow_variable_data_keys: bool,
- streaming: bool,
- frontend_conf: dict = None,
- fs: Union[dict, int] = 16000,
- **kwargs,
-) -> List[Any]:
- assert check_argument_types()
- if batch_size > 1:
- raise NotImplementedError("batch decoding is not implemented")
- if word_lm_train_config is not None:
- raise NotImplementedError("Word LM is not implemented")
- if ngpu > 1:
- raise NotImplementedError("only single GPU decoding is supported")
-
- logging.basicConfig(
- level=log_level,
- format="%(asctime)s (%(module)s:%(lineno)d) %(levelname)s: %(message)s",
- )
-
- if ngpu >= 1:
- device = "cuda"
- else:
- device = "cpu"
- features_type: str = data_path_and_name_and_type[1]
- hop_length: int = 160
- sr: int = 16000
- if isinstance(fs, int):
- sr = fs
- else:
- if 'model_fs' in fs and fs['model_fs'] is not None:
- sr = fs['model_fs']
- if features_type != 'sound':
- frontend_conf = None
- if frontend_conf is not None:
- if 'hop_length' in frontend_conf:
- hop_length = frontend_conf['hop_length']
-
- finish_count = 0
- file_count = 1
- if isinstance(audio_lists, bytes):
- file_count = 1
- else:
- file_count = len(audio_lists)
- if len(data_path_and_name_and_type) >= 3 and frontend_conf is not None:
- mvn_file = data_path_and_name_and_type[2]
- mvn_data = wav_utils.extract_CMVN_featrures(mvn_file)
- frontend_conf['mvn_data'] = mvn_data
- # 1. Set random-seed
- set_all_random_seed(seed)
-
- # 2. Build speech2text
- speech2text_kwargs = dict(
- asr_train_config=asr_train_config,
- asr_model_file=asr_model_file,
- lm_train_config=lm_train_config,
- lm_file=lm_file,
- token_type=token_type,
- bpemodel=bpemodel,
- device=device,
- maxlenratio=maxlenratio,
- minlenratio=minlenratio,
- dtype=dtype,
- beam_size=beam_size,
- ctc_weight=ctc_weight,
- lm_weight=lm_weight,
- ngram_weight=ngram_weight,
- penalty=penalty,
- nbest=nbest,
- streaming=streaming,
- frontend_conf=frontend_conf,
- )
- speech2text = Speech2Text(**speech2text_kwargs)
- data_path_and_name_and_type_new = [
- audio_lists, data_path_and_name_and_type[0], data_path_and_name_and_type[1]
- ]
- # 3. Build data-iterator
- loader = ASRTask.build_streaming_iterator_modelscope(
- data_path_and_name_and_type_new,
- dtype=dtype,
- batch_size=batch_size,
- key_file=key_file,
- num_workers=num_workers,
- preprocess_fn=ASRTask.build_preprocess_fn(speech2text.asr_train_args, False),
- collate_fn=ASRTask.build_collate_fn(speech2text.asr_train_args, False),
- allow_variable_data_keys=allow_variable_data_keys,
- inference=True,
- sample_rate=fs
- )
-
- # 7 .Start for-loop
- # FIXME(kamo): The output format should be discussed about
- asr_result_list = []
- for keys, batch in loader:
- assert isinstance(batch, dict), type(batch)
- assert all(isinstance(s, str) for s in keys), keys
- _bs = len(next(iter(batch.values())))
- assert len(keys) == _bs, f"{len(keys)} != {_bs}"
- batch = {k: v[0] for k, v in batch.items() if not k.endswith("_lengths")}
-
- # N-best list of (text, token, token_int, hyp_object)
- try:
- results = speech2text(**batch)
- except TooShortUttError as e:
- logging.warning(f"Utterance {keys} {e}")
- hyp = Hypothesis(score=0.0, scores={}, states={}, yseq=[])
- results = [[" ", ["<space>"], [2], hyp]] * nbest
-
- # Only supporting batch_size==1
- key = keys[0]
- for n, (text, token, token_int, hyp) in zip(range(1, nbest + 1), results):
- if text is not None:
- text_postprocessed = postprocess_utils.sentence_postprocess(token)
- item = {'key': key, 'value': text_postprocessed}
- asr_result_list.append(item)
- finish_count += 1
- asr_utils.print_progress(finish_count / file_count)
-
- return asr_result_list
-
-
-
-def set_parameters(language: str = None,
- sample_rate: Union[int, Dict[Any, int]] = None):
- if language is not None:
- global global_asr_language
- global_asr_language = language
- if sample_rate is not None:
- global global_sample_rate
- global_sample_rate = sample_rate
-
-
-def asr_inference(maxlenratio: float,
- minlenratio: float,
- beam_size: int,
- ngpu: int,
- ctc_weight: float,
- lm_weight: float,
- penalty: float,
- name_and_type: list,
- audio_lists: Union[List[Any], bytes],
- asr_train_config: Optional[str],
- asr_model_file: Optional[str],
- nbest: int = 1,
- num_workers: int = 1,
- log_level: Union[int, str] = 'INFO',
- batch_size: int = 1,
- dtype: str = 'float32',
- seed: int = 0,
- key_file: Optional[str] = None,
- lm_train_config: Optional[str] = None,
- lm_file: Optional[str] = None,
- word_lm_train_config: Optional[str] = None,
- word_lm_file: Optional[str] = None,
- ngram_file: Optional[str] = None,
- ngram_weight: float = 0.9,
- model_tag: Optional[str] = None,
- token_type: Optional[str] = None,
- bpemodel: Optional[str] = None,
- allow_variable_data_keys: bool = False,
- transducer_conf: Optional[dict] = None,
- streaming: bool = False,
- frontend_conf: dict = None,
- fs: Union[dict, int] = None,
- lang: Optional[str] = None,
- outputdir: Optional[str] = None):
- if lang is not None:
- global global_asr_language
- global_asr_language = lang
- if fs is not None:
- global global_sample_rate
- global_sample_rate = fs
-
- # force use CPU if data type is bytes
- if isinstance(audio_lists, bytes):
- num_workers = 0
- ngpu = 0
-
- return inference(output_dir=outputdir,
- maxlenratio=maxlenratio,
- minlenratio=minlenratio,
- batch_size=batch_size,
- dtype=dtype,
- beam_size=beam_size,
- ngpu=ngpu,
- seed=seed,
- ctc_weight=ctc_weight,
- lm_weight=lm_weight,
- ngram_weight=ngram_weight,
- penalty=penalty,
- nbest=nbest,
- num_workers=num_workers,
- log_level=log_level,
- data_path_and_name_and_type=name_and_type,
- audio_lists=audio_lists,
- key_file=key_file,
- asr_train_config=asr_train_config,
- asr_model_file=asr_model_file,
- lm_train_config=lm_train_config,
- lm_file=lm_file,
- word_lm_train_config=word_lm_train_config,
- word_lm_file=word_lm_file,
- ngram_file=ngram_file,
- model_tag=model_tag,
- token_type=token_type,
- bpemodel=bpemodel,
- allow_variable_data_keys=allow_variable_data_keys,
- transducer_conf=transducer_conf,
- streaming=streaming,
- frontend_conf=frontend_conf)
-
-
-def get_parser():
- parser = config_argparse.ArgumentParser(
- description="ASR Decoding",
- formatter_class=argparse.ArgumentDefaultsHelpFormatter,
- )
-
- # Note(kamo): Use '_' instead of '-' as separator.
- # '-' is confusing if written in yaml.
- parser.add_argument(
- "--log_level",
- type=lambda x: x.upper(),
- default="INFO",
- choices=("CRITICAL", "ERROR", "WARNING", "INFO", "DEBUG", "NOTSET"),
- help="The verbose level of logging",
- )
-
- parser.add_argument("--output_dir", type=str, required=True)
- parser.add_argument(
- "--ngpu",
- type=int,
- default=0,
- help="The number of gpus. 0 indicates CPU mode",
- )
- parser.add_argument(
- "--gpuid_list",
- type=str,
- default="",
- help="The visible gpus",
- )
- parser.add_argument("--seed", type=int, default=0, help="Random seed")
- parser.add_argument(
- "--dtype",
- default="float32",
- choices=["float16", "float32", "float64"],
- help="Data type",
- )
- parser.add_argument(
- "--num_workers",
- type=int,
- default=1,
- help="The number of workers used for DataLoader",
- )
-
- group = parser.add_argument_group("Input data related")
- group.add_argument(
- "--data_path_and_name_and_type",
- type=str2triple_str,
- required=True,
- action="append",
- )
- group.add_argument("--audio_lists", type=list,
- default=[{'key':'EdevDEWdIYQ_0021',
- 'file':'/mnt/data/jiangyu.xzy/test_data/speech_io/SPEECHIO_ASR_ZH00007_zhibodaihuo/wav/EdevDEWdIYQ_0021.wav'}])
- group.add_argument("--key_file", type=str_or_none)
- group.add_argument("--allow_variable_data_keys", type=str2bool, default=False)
-
- group = parser.add_argument_group("The model configuration related")
- group.add_argument(
- "--asr_train_config",
- type=str,
- help="ASR training configuration",
- )
- group.add_argument(
- "--asr_model_file",
- type=str,
- help="ASR model parameter file",
- )
- group.add_argument(
- "--lm_train_config",
- type=str,
- help="LM training configuration",
- )
- group.add_argument(
- "--lm_file",
- type=str,
- help="LM parameter file",
- )
- group.add_argument(
- "--word_lm_train_config",
- type=str,
- help="Word LM training configuration",
- )
- group.add_argument(
- "--word_lm_file",
- type=str,
- help="Word LM parameter file",
- )
- group.add_argument(
- "--ngram_file",
- type=str,
- help="N-gram parameter file",
- )
- group.add_argument(
- "--model_tag",
- type=str,
- help="Pretrained model tag. If specify this option, *_train_config and "
- "*_file will be overwritten",
- )
-
- group = parser.add_argument_group("Beam-search related")
- group.add_argument(
- "--batch_size",
- type=int,
- default=1,
- help="The batch size for inference",
- )
- group.add_argument("--nbest", type=int, default=1, help="Output N-best hypotheses")
- group.add_argument("--beam_size", type=int, default=20, help="Beam size")
- group.add_argument("--penalty", type=float, default=0.0, help="Insertion penalty")
- group.add_argument(
- "--maxlenratio",
- type=float,
- default=0.0,
- help="Input length ratio to obtain max output length. "
- "If maxlenratio=0.0 (default), it uses a end-detect "
- "function "
- "to automatically find maximum hypothesis lengths."
- "If maxlenratio<0.0, its absolute value is interpreted"
- "as a constant max output length",
- )
- group.add_argument(
- "--minlenratio",
- type=float,
- default=0.0,
- help="Input length ratio to obtain min output length",
- )
- group.add_argument(
- "--ctc_weight",
- type=float,
- default=0.5,
- help="CTC weight in joint decoding",
- )
- group.add_argument("--lm_weight", type=float, default=1.0, help="RNNLM weight")
- group.add_argument("--ngram_weight", type=float, default=0.9, help="ngram weight")
- group.add_argument("--streaming", type=str2bool, default=False)
-
- group = parser.add_argument_group("Text converter related")
- group.add_argument(
- "--token_type",
- type=str_or_none,
- default=None,
- choices=["char", "bpe", None],
- help="The token type for ASR model. "
- "If not given, refers from the training args",
- )
- group.add_argument(
- "--bpemodel",
- type=str_or_none,
- default=None,
- help="The model path of sentencepiece. "
- "If not given, refers from the training args",
- )
-
- return parser
-
-
-def main(cmd=None):
- print(get_commandline_args(), file=sys.stderr)
- parser = get_parser()
- args = parser.parse_args(cmd)
- kwargs = vars(args)
- kwargs.pop("config", None)
- inference(**kwargs)
-
-
-if __name__ == "__main__":
- main()
diff --git a/funasr/bin/asr_inference_paraformer.py b/funasr/bin/asr_inference_paraformer.py
index 179a62b..15a37f7 100755
--- a/funasr/bin/asr_inference_paraformer.py
+++ b/funasr/bin/asr_inference_paraformer.py
@@ -8,6 +8,9 @@
from typing import Sequence
from typing import Tuple
from typing import Union
+from typing import Dict
+from typing import Any
+from typing import List
import numpy as np
import torch
@@ -30,7 +33,21 @@
from funasr.utils.types import str2bool
from funasr.utils.types import str2triple_str
from funasr.utils.types import str_or_none
+from funasr.utils import asr_utils, wav_utils, postprocess_utils
+from funasr.models.frontend.wav_frontend import WavFrontend
+from modelscope.utils.logger import get_logger
+
+logger = get_logger()
+
+header_colors = '\033[95m'
+end_colors = '\033[0m'
+
+global_asr_language: str = 'zh-cn'
+global_sample_rate: Union[int, Dict[Any, int]] = {
+ 'audio_fs': 16000,
+ 'model_fs': 16000
+}
class Speech2Text:
"""Speech2Text class
@@ -62,6 +79,7 @@
ngram_weight: float = 0.9,
penalty: float = 0.0,
nbest: int = 1,
+ frontend_conf: dict = None,
**kwargs,
):
assert check_argument_types()
@@ -71,6 +89,9 @@
asr_model, asr_train_args = ASRTask.build_model_from_file(
asr_train_config, asr_model_file, device
)
+ if asr_model.frontend is None and frontend_conf is not None:
+ frontend = WavFrontend(**frontend_conf)
+ asr_model.frontend = frontend
logging.info("asr_model: {}".format(asr_model))
logging.info("asr_train_args: {}".format(asr_train_args))
asr_model.to(dtype=getattr(torch, dtype)).eval()
@@ -145,6 +166,9 @@
self.asr_train_args = asr_train_args
self.converter = converter
self.tokenizer = tokenizer
+ has_lm = lm_weight == 0.0 or lm_file is None
+ if ctc_weight == 0.0 and has_lm:
+ beam_search = None
self.beam_search = beam_search
self.beam_search_transducer = beam_search_transducer
self.maxlenratio = maxlenratio
@@ -155,12 +179,12 @@
@torch.no_grad()
def __call__(
- self, speech: Union[torch.Tensor, np.ndarray]
+ self, speech: Union[torch.Tensor, np.ndarray], speech_lengths: Union[torch.Tensor, np.ndarray] = None
):
"""Inference
Args:
- data: Input speech data
+ speech: Input speech data
Returns:
text, token, token_int, hyp
@@ -172,11 +196,13 @@
speech = torch.tensor(speech)
# data: (Nsamples,) -> (1, Nsamples)
- speech = speech.unsqueeze(0).to(getattr(torch, self.dtype))
- lfr_factor = max(1, (speech.size()[-1]//80)-1)
# lengths: (1,)
- lengths = speech.new_full([1], dtype=torch.long, fill_value=speech.size(1))
- batch = {"speech": speech, "speech_lengths": lengths}
+ if len(speech.size()) < 3:
+ speech = speech.unsqueeze(0).to(getattr(torch, self.dtype))
+ speech_lengths = speech.new_full([1], dtype=torch.long, fill_value=speech.size(1))
+ lfr_factor = max(1, (speech.size()[-1]//80)-1)
+
+ batch = {"speech": speech, "speech_lengths": speech_lengths}
# a. To device
batch = to_device(batch, device=self.device)
@@ -185,78 +211,98 @@
enc, enc_len = self.asr_model.encode(**batch)
if isinstance(enc, tuple):
enc = enc[0]
- assert len(enc) == 1, len(enc)
+ # assert len(enc) == 1, len(enc)
+ enc_len_batch_total = torch.sum(enc_len).item()
predictor_outs = self.asr_model.calc_predictor(enc, enc_len)
pre_acoustic_embeds, pre_token_length = predictor_outs[0], predictor_outs[1]
- pre_token_length = pre_token_length.long()
+ pre_token_length = pre_token_length.round().long()
decoder_outs = self.asr_model.cal_decoder_with_predictor(enc, enc_len, pre_acoustic_embeds, pre_token_length)
decoder_out, ys_pad_lens = decoder_outs[0], decoder_outs[1]
- nbest_hyps = self.beam_search(
- x=enc[0], am_scores=decoder_out[0], maxlenratio=self.maxlenratio, minlenratio=self.minlenratio
- )
-
- nbest_hyps = nbest_hyps[: self.nbest]
results = []
- for hyp in nbest_hyps:
- assert isinstance(hyp, (Hypothesis)), type(hyp)
-
- # remove sos/eos and get results
- last_pos = -1
- if isinstance(hyp.yseq, list):
- token_int = hyp.yseq[1:last_pos]
+ b, n, d = decoder_out.size()
+ for i in range(b):
+ x = enc[i, :enc_len[i], :]
+ am_scores = decoder_out[i, :pre_token_length[i], :]
+ if self.beam_search is not None:
+ nbest_hyps = self.beam_search(
+ x=x, am_scores=am_scores, maxlenratio=self.maxlenratio, minlenratio=self.minlenratio
+ )
+
+ nbest_hyps = nbest_hyps[: self.nbest]
else:
- token_int = hyp.yseq[1:last_pos].tolist()
-
- # remove blank symbol id, which is assumed to be 0
- token_int = list(filter(lambda x: x != 0, token_int))
-
- # Change integer-ids to tokens
- token = self.converter.ids2tokens(token_int)
-
- if self.tokenizer is not None:
- text = self.tokenizer.tokens2text(token)
- else:
- text = None
-
- results.append((text, token, token_int, hyp, speech.size(1), lfr_factor))
+ yseq = am_scores.argmax(dim=-1)
+ score = am_scores.max(dim=-1)[0]
+ score = torch.sum(score, dim=-1)
+ # pad with mask tokens to ensure compatibility with sos/eos tokens
+ yseq = torch.tensor(
+ [self.asr_model.sos] + yseq.tolist() + [self.asr_model.eos], device=yseq.device
+ )
+ nbest_hyps = [Hypothesis(yseq=yseq, score=score)]
+
+ for hyp in nbest_hyps:
+ assert isinstance(hyp, (Hypothesis)), type(hyp)
+
+ # remove sos/eos and get results
+ last_pos = -1
+ if isinstance(hyp.yseq, list):
+ token_int = hyp.yseq[1:last_pos]
+ else:
+ token_int = hyp.yseq[1:last_pos].tolist()
+
+ # remove blank symbol id, which is assumed to be 0
+ token_int = list(filter(lambda x: x != 0, token_int))
+
+ # Change integer-ids to tokens
+ token = self.converter.ids2tokens(token_int)
+
+ if self.tokenizer is not None:
+ text = self.tokenizer.tokens2text(token)
+ else:
+ text = None
+
+ results.append((text, token, token_int, hyp, enc_len_batch_total, lfr_factor))
# assert check_return_type(results)
return results
def inference(
- output_dir: str,
maxlenratio: float,
minlenratio: float,
batch_size: int,
- dtype: str,
beam_size: int,
ngpu: int,
- seed: int,
ctc_weight: float,
lm_weight: float,
- ngram_weight: float,
penalty: float,
- nbest: int,
- num_workers: int,
log_level: Union[int, str],
- data_path_and_name_and_type: Sequence[Tuple[str, str, str]],
- key_file: Optional[str],
+ data_path_and_name_and_type,
asr_train_config: Optional[str],
asr_model_file: Optional[str],
- lm_train_config: Optional[str],
- lm_file: Optional[str],
- word_lm_train_config: Optional[str],
- token_type: Optional[str],
- bpemodel: Optional[str],
- allow_variable_data_keys: bool,
+ audio_lists: Union[List[Any], bytes] = None,
+ lm_train_config: Optional[str] = None,
+ lm_file: Optional[str] = None,
+ token_type: Optional[str] = None,
+ key_file: Optional[str] = None,
+ word_lm_train_config: Optional[str] = None,
+ bpemodel: Optional[str] = None,
+ allow_variable_data_keys: bool = False,
+ streaming: bool = False,
+ output_dir: Optional[str] = None,
+ dtype: str = "float32",
+ seed: int = 0,
+ ngram_weight: float = 0.9,
+ nbest: int = 1,
+ num_workers: int = 1,
+ frontend_conf: dict = None,
+ fs: Union[dict, int] = 16000,
+ lang: Optional[str] = None,
**kwargs,
):
assert check_argument_types()
- if batch_size > 1:
- raise NotImplementedError("batch decoding is not implemented")
+
if word_lm_train_config is not None:
raise NotImplementedError("Word LM is not implemented")
if ngpu > 1:
@@ -271,7 +317,46 @@
device = "cuda"
else:
device = "cpu"
+ hop_length: int = 160
+ sr: int = 16000
+ if isinstance(fs, int):
+ sr = fs
+ else:
+ if 'model_fs' in fs and fs['model_fs'] is not None:
+ sr = fs['model_fs']
+ # data_path_and_name_and_type for modelscope: (data from audio_lists)
+ # ['speech', 'sound', 'am.mvn']
+ # data_path_and_name_and_type for funasr:
+ # [('/mnt/data/jiangyu.xzy/exp/maas/mvn.1.scp', 'speech', 'kaldi_ark')]
+ if isinstance(data_path_and_name_and_type[0], Tuple):
+ features_type: str = data_path_and_name_and_type[0][1]
+ elif isinstance(data_path_and_name_and_type[0], str):
+ features_type: str = data_path_and_name_and_type[1]
+ else:
+ raise NotImplementedError("unknown features type:{0}".format(data_path_and_name_and_type))
+ if features_type != 'sound':
+ frontend_conf = None
+ flag_modelscope = False
+ else:
+ flag_modelscope = True
+ if frontend_conf is not None:
+ if 'hop_length' in frontend_conf:
+ hop_length = frontend_conf['hop_length']
+ finish_count = 0
+ file_count = 1
+ if flag_modelscope and not isinstance(data_path_and_name_and_type[0], Tuple):
+ data_path_and_name_and_type_new = [
+ audio_lists, data_path_and_name_and_type[0], data_path_and_name_and_type[1]
+ ]
+ if isinstance(audio_lists, bytes):
+ file_count = 1
+ else:
+ file_count = len(audio_lists)
+ if len(data_path_and_name_and_type) >= 3 and frontend_conf is not None:
+ mvn_file = data_path_and_name_and_type[2]
+ mvn_data = wav_utils.extract_CMVN_featrures(mvn_file)
+ frontend_conf['mvn_data'] = mvn_data
# 1. Set random-seed
set_all_random_seed(seed)
@@ -293,73 +378,107 @@
ngram_weight=ngram_weight,
penalty=penalty,
nbest=nbest,
+ frontend_conf=frontend_conf,
)
speech2text = Speech2Text(**speech2text_kwargs)
# 3. Build data-iterator
- loader = ASRTask.build_streaming_iterator(
- data_path_and_name_and_type,
- dtype=dtype,
- batch_size=batch_size,
- key_file=key_file,
- num_workers=num_workers,
- preprocess_fn=ASRTask.build_preprocess_fn(speech2text.asr_train_args, False),
- collate_fn=ASRTask.build_collate_fn(speech2text.asr_train_args, False),
- allow_variable_data_keys=allow_variable_data_keys,
- inference=True,
- )
+ if flag_modelscope:
+ loader = ASRTask.build_streaming_iterator_modelscope(
+ data_path_and_name_and_type_new,
+ dtype=dtype,
+ batch_size=batch_size,
+ key_file=key_file,
+ num_workers=num_workers,
+ preprocess_fn=ASRTask.build_preprocess_fn(speech2text.asr_train_args, False),
+ collate_fn=ASRTask.build_collate_fn(speech2text.asr_train_args, False),
+ allow_variable_data_keys=allow_variable_data_keys,
+ inference=True,
+ sample_rate=fs
+ )
+ else:
+ loader = ASRTask.build_streaming_iterator(
+ data_path_and_name_and_type,
+ dtype=dtype,
+ batch_size=batch_size,
+ key_file=key_file,
+ num_workers=num_workers,
+ preprocess_fn=ASRTask.build_preprocess_fn(speech2text.asr_train_args, False),
+ collate_fn=ASRTask.build_collate_fn(speech2text.asr_train_args, False),
+ allow_variable_data_keys=allow_variable_data_keys,
+ inference=True,
+ )
forward_time_total = 0.0
length_total = 0.0
# 7 .Start for-loop
# FIXME(kamo): The output format should be discussed about
- with DatadirWriter(output_dir) as writer:
- for keys, batch in loader:
- assert isinstance(batch, dict), type(batch)
- assert all(isinstance(s, str) for s in keys), keys
- _bs = len(next(iter(batch.values())))
- assert len(keys) == _bs, f"{len(keys)} != {_bs}"
- batch = {k: v[0] for k, v in batch.items() if not k.endswith("_lengths")}
+ asr_result_list = []
+ if output_dir is not None:
+ writer = DatadirWriter(output_dir)
+ else:
+ writer = None
- logging.info("decoding, utt_id: {}".format(keys))
- # N-best list of (text, token, token_int, hyp_object)
+ for keys, batch in loader:
+ assert isinstance(batch, dict), type(batch)
+ assert all(isinstance(s, str) for s in keys), keys
+ _bs = len(next(iter(batch.values())))
+ assert len(keys) == _bs, f"{len(keys)} != {_bs}"
+ # batch = {k: v for k, v in batch.items() if not k.endswith("_lengths")}
- try:
- time_beg = time.time()
- results = speech2text(**batch)
- time_end = time.time()
- forward_time = time_end - time_beg
- lfr_factor = results[0][-1]
- length = results[0][-2]
- results = [results[0][:-2]]
- forward_time_total += forward_time
- length_total += length
- logging.info(
- "decoding, feature length: {}, forward_time: {:.4f}, rtf: {:.4f}".
- format(length, forward_time, 100 * forward_time / (length*lfr_factor)))
- except TooShortUttError as e:
- logging.warning(f"Utterance {keys} {e}")
- hyp = Hypothesis(score=0.0, scores={}, states={}, yseq=[])
- results = [[" ", ["<space>"], [2], hyp]] * nbest
+ logging.info("decoding, utt_id: {}".format(keys))
+ # N-best list of (text, token, token_int, hyp_object)
- # Only supporting batch_size==1
- key = keys[0]
- for n, (text, token, token_int, hyp) in zip(range(1, nbest + 1), results):
+ time_beg = time.time()
+ results = speech2text(**batch)
+ time_end = time.time()
+ forward_time = time_end - time_beg
+ lfr_factor = results[0][-1]
+ length = results[0][-2]
+ forward_time_total += forward_time
+ length_total += length
+ logging.info(
+ "decoding, feature length: {}, forward_time: {:.4f}, rtf: {:.4f}".
+ format(length, forward_time, 100 * forward_time / (length*lfr_factor)))
+
+ for batch_id in range(len(results)):
+ result = [results[batch_id][:-2]]
+
+ key = keys[batch_id]
+ for n, (text, token, token_int, hyp) in zip(range(1, nbest + 1), result):
# Create a directory: outdir/{n}best_recog
- ibest_writer = writer[f"{n}best_recog"]
-
- # Write the result to each file
- ibest_writer["token"][key] = " ".join(token)
- ibest_writer["token_int"][key] = " ".join(map(str, token_int))
- ibest_writer["score"][key] = str(hyp.score)
-
+ if writer is not None:
+ ibest_writer = writer[f"{n}best_recog"]
+
+ # Write the result to each file
+ ibest_writer["token"][key] = " ".join(token)
+ ibest_writer["token_int"][key] = " ".join(map(str, token_int))
+ ibest_writer["score"][key] = str(hyp.score)
+
if text is not None:
- ibest_writer["text"][key] = text
-
- logging.info("decoding, predictions: {}".format(text))
+ text_postprocessed = postprocess_utils.sentence_postprocess(token)
+ item = {'key': key, 'value': text_postprocessed}
+ asr_result_list.append(item)
+ finish_count += 1
+ asr_utils.print_progress(finish_count / file_count)
+ if writer is not None:
+ ibest_writer["text"][key] = text
+
+ logging.info("decoding, utt: {}, predictions: {}".format(key, text))
logging.info("decoding, feature length total: {}, forward_time total: {:.4f}, rtf avg: {:.4f}".
format(length_total, forward_time_total, 100 * forward_time_total / (length_total*lfr_factor)))
+ return asr_result_list
+
+
+def set_parameters(language: str = None,
+ sample_rate: Union[int, Dict[Any, int]] = None):
+ if language is not None:
+ global global_asr_language
+ global_asr_language = language
+ if sample_rate is not None:
+ global global_sample_rate
+ global_sample_rate = sample_rate
def get_parser():
@@ -494,6 +613,8 @@
default=None,
help="",
)
+ group.add_argument("--audio_lists", type=list, default=None)
+ # example=[{'key':'EdevDEWdIYQ_0021','file':'/mnt/data/jiangyu.xzy/test_data/speech_io/SPEECHIO_ASR_ZH00007_zhibodaihuo/wav/EdevDEWdIYQ_0021.wav'}])
group = parser.add_argument_group("Text converter related")
group.add_argument(
diff --git a/funasr/bin/asr_inference_paraformer_modelscope.py b/funasr/bin/asr_inference_paraformer_modelscope.py
deleted file mode 100755
index d64fe2b..0000000
--- a/funasr/bin/asr_inference_paraformer_modelscope.py
+++ /dev/null
@@ -1,686 +0,0 @@
-#!/usr/bin/env python3
-import argparse
-import logging
-import sys
-import time
-from pathlib import Path
-from typing import Any
-from typing import Optional
-from typing import Sequence
-from typing import Tuple
-from typing import Union
-from typing import List
-from typing import Dict
-
-import numpy as np
-import torch
-from typeguard import check_argument_types
-
-from funasr.modules.beam_search.beam_search import BeamSearchPara as BeamSearch
-from funasr.modules.beam_search.beam_search import Hypothesis
-from funasr.modules.scorers.ctc import CTCPrefixScorer
-from funasr.modules.scorers.length_bonus import LengthBonus
-from funasr.modules.subsampling import TooShortUttError
-from funasr.tasks.asr import ASRTaskParaformer as ASRTask
-from funasr.tasks.lm import LMTask
-from funasr.text.build_tokenizer import build_tokenizer
-from funasr.text.token_id_converter import TokenIDConverter
-from funasr.torch_utils.device_funcs import to_device
-from funasr.torch_utils.set_all_random_seed import set_all_random_seed
-from funasr.utils import config_argparse
-from funasr.utils.cli_utils import get_commandline_args
-from funasr.utils.types import str2bool
-from funasr.utils.types import str2triple_str
-from funasr.utils.types import str_or_none
-from funasr.utils import asr_utils, wav_utils, postprocess_utils
-from funasr.models.frontend.wav_frontend import WavFrontend
-
-from modelscope.utils.logger import get_logger
-
-logger = get_logger()
-
-header_colors = '\033[95m'
-end_colors = '\033[0m'
-
-global_asr_language: str = 'zh-cn'
-global_sample_rate: Union[int, Dict[Any, int]] = {
- 'audio_fs': 16000,
- 'model_fs': 16000
-}
-
-
-class Speech2Text:
- """Speech2Text class
-
- Examples:
- >>> import soundfile
- >>> speech2text = Speech2Text("asr_config.yml", "asr.pth")
- >>> audio, rate = soundfile.read("speech.wav")
- >>> speech2text(audio)
- [(text, token, token_int, hypothesis object), ...]
-
- """
-
- def __init__(
- self,
- asr_train_config: Union[Path, str] = None,
- asr_model_file: Union[Path, str] = None,
- lm_train_config: Union[Path, str] = None,
- lm_file: Union[Path, str] = None,
- token_type: str = None,
- bpemodel: str = None,
- device: str = "cpu",
- maxlenratio: float = 0.0,
- minlenratio: float = 0.0,
- dtype: str = "float32",
- beam_size: int = 20,
- ctc_weight: float = 0.5,
- lm_weight: float = 1.0,
- ngram_weight: float = 0.9,
- penalty: float = 0.0,
- nbest: int = 1,
- frontend_conf: dict = None,
- **kwargs,
- ):
- assert check_argument_types()
-
- # 1. Build ASR model
- scorers = {}
- asr_model, asr_train_args = ASRTask.build_model_from_file(
- asr_train_config, asr_model_file, device
- )
- if asr_model.frontend is None and frontend_conf is not None:
- frontend = WavFrontend(**frontend_conf)
- asr_model.frontend = frontend
- asr_model.to(dtype=getattr(torch, dtype)).eval()
-
- ctc = CTCPrefixScorer(ctc=asr_model.ctc, eos=asr_model.eos)
- token_list = asr_model.token_list
- scorers.update(
- ctc=ctc,
- length_bonus=LengthBonus(len(token_list)),
- )
-
- # 2. Build Language model
- if lm_train_config is not None:
- lm, lm_train_args = LMTask.build_model_from_file(
- lm_train_config, lm_file, device
- )
- scorers["lm"] = lm.lm
-
- # 3. Build ngram model
- # ngram is not supported now
- ngram = None
- scorers["ngram"] = ngram
-
- # 4. Build BeamSearch object
- # transducer is not supported now
- beam_search_transducer = None
-
- weights = dict(
- decoder=1.0 - ctc_weight,
- ctc=ctc_weight,
- lm=lm_weight,
- ngram=ngram_weight,
- length_bonus=penalty,
- )
- beam_search = BeamSearch(
- beam_size=beam_size,
- weights=weights,
- scorers=scorers,
- sos=asr_model.sos,
- eos=asr_model.eos,
- vocab_size=len(token_list),
- token_list=token_list,
- pre_beam_score_key=None if ctc_weight == 1.0 else "full",
- )
-
- beam_search.to(device=device, dtype=getattr(torch, dtype)).eval()
- for scorer in scorers.values():
- if isinstance(scorer, torch.nn.Module):
- scorer.to(device=device, dtype=getattr(torch, dtype)).eval()
- logging.info(f"Beam_search: {beam_search}")
- logging.info(f"Decoding device={device}, dtype={dtype}")
-
- # 5. [Optional] Build Text converter: e.g. bpe-sym -> Text
- if token_type is None:
- token_type = asr_train_args.token_type
- if bpemodel is None:
- bpemodel = asr_train_args.bpemodel
-
- if token_type is None:
- tokenizer = None
- elif token_type == "bpe":
- if bpemodel is not None:
- tokenizer = build_tokenizer(token_type=token_type, bpemodel=bpemodel)
- else:
- tokenizer = None
- else:
- tokenizer = build_tokenizer(token_type=token_type)
- converter = TokenIDConverter(token_list=token_list)
- logging.info(f"Text tokenizer: {tokenizer}")
-
- self.asr_model = asr_model
- self.asr_train_args = asr_train_args
- self.converter = converter
- self.tokenizer = tokenizer
- self.beam_search = beam_search
- self.beam_search_transducer = beam_search_transducer
- self.maxlenratio = maxlenratio
- self.minlenratio = minlenratio
- self.device = device
- self.dtype = dtype
- self.nbest = nbest
-
- @torch.no_grad()
- def __call__(
- self, speech: Union[torch.Tensor, np.ndarray]
- ):
- """Inference
-
- Args:
- speech: Input speech data
- Returns:
- text, token, token_int, hyp
-
- """
- assert check_argument_types()
-
- # Input as audio signal
- if isinstance(speech, np.ndarray):
- speech = torch.tensor(speech)
-
- # data: (Nsamples,) -> (1, Nsamples)
- speech = speech.unsqueeze(0).to(getattr(torch, self.dtype))
- lfr_factor = max(1, (speech.size()[-1] // 80) - 1)
- # lengths: (1,)
- lengths = speech.new_full([1], dtype=torch.long, fill_value=speech.size(1))
- batch = {"speech": speech, "speech_lengths": lengths}
-
- # a. To device
- batch = to_device(batch, device=self.device)
-
- # b. Forward Encoder
- enc, enc_len = self.asr_model.encode(**batch)
- if isinstance(enc, tuple):
- enc = enc[0]
- assert len(enc) == 1, len(enc)
-
- predictor_outs = self.asr_model.calc_predictor(enc, enc_len)
- pre_acoustic_embeds, pre_token_length = predictor_outs[0], predictor_outs[1]
- pre_token_length = torch.tensor([pre_acoustic_embeds.size(1)], device=pre_acoustic_embeds.device)
- decoder_outs = self.asr_model.cal_decoder_with_predictor(enc, enc_len, pre_acoustic_embeds, pre_token_length)
- decoder_out, ys_pad_lens = decoder_outs[0], decoder_outs[1]
-
- nbest_hyps = self.beam_search(
- x=enc[0], am_scores=decoder_out[0], maxlenratio=self.maxlenratio, minlenratio=self.minlenratio
- )
-
- nbest_hyps = nbest_hyps[: self.nbest]
- results = []
- for hyp in nbest_hyps:
- assert isinstance(hyp, (Hypothesis)), type(hyp)
-
- # remove sos/eos and get results
- last_pos = -1
- if isinstance(hyp.yseq, list):
- token_int = hyp.yseq[1:last_pos]
- else:
- token_int = hyp.yseq[1:last_pos].tolist()
-
- # remove blank symbol id, which is assumed to be 0
- token_int = list(filter(lambda x: x != 0, token_int))
-
- # Change integer-ids to tokens
- token = self.converter.ids2tokens(token_int)
-
- if self.tokenizer is not None:
- text = self.tokenizer.tokens2text(token)
- else:
- text = None
-
- results.append((text, token, token_int, hyp, speech.size(1), lfr_factor))
-
- # assert check_return_type(results)
- return results
-
-
-def inference(
- maxlenratio: float,
- minlenratio: float,
- batch_size: int,
- dtype: str,
- beam_size: int,
- ngpu: int,
- seed: int,
- ctc_weight: float,
- lm_weight: float,
- ngram_weight: float,
- penalty: float,
- nbest: int,
- num_workers: int,
- log_level: Union[int, str],
- data_path_and_name_and_type: list,
- audio_lists: Union[List[Any], bytes],
- key_file: Optional[str],
- asr_train_config: Optional[str],
- asr_model_file: Optional[str],
- lm_train_config: Optional[str],
- lm_file: Optional[str],
- word_lm_train_config: Optional[str],
- model_tag: Optional[str],
- token_type: Optional[str],
- bpemodel: Optional[str],
- output_dir: Optional[str],
- allow_variable_data_keys: bool,
- frontend_conf: dict = None,
- fs: Union[dict, int] = 16000,
- **kwargs,
-) -> List[Any]:
- assert check_argument_types()
- if batch_size > 1:
- raise NotImplementedError("batch decoding is not implemented")
- if word_lm_train_config is not None:
- raise NotImplementedError("Word LM is not implemented")
- if ngpu > 1:
- raise NotImplementedError("only single GPU decoding is supported")
-
- logging.basicConfig(
- level=log_level,
- format="%(asctime)s (%(module)s:%(lineno)d) %(levelname)s: %(message)s",
- )
-
- if ngpu >= 1:
- device = "cuda"
- else:
- device = "cpu"
- # data_path_and_name_and_type = data_path_and_name_and_type[0]
- features_type: str = data_path_and_name_and_type[1]
- hop_length: int = 160
- sr: int = 16000
- if isinstance(fs, int):
- sr = fs
- else:
- if 'model_fs' in fs and fs['model_fs'] is not None:
- sr = fs['model_fs']
- if features_type != 'sound':
- frontend_conf = None
- if frontend_conf is not None:
- if 'hop_length' in frontend_conf:
- hop_length = frontend_conf['hop_length']
-
- finish_count = 0
- file_count = 1
- if isinstance(audio_lists, bytes):
- file_count = 1
- else:
- file_count = len(audio_lists)
- if len(data_path_and_name_and_type) >= 3 and frontend_conf is not None:
- mvn_file = data_path_and_name_and_type[2]
- mvn_data = wav_utils.extract_CMVN_featrures(mvn_file)
- frontend_conf['mvn_data'] = mvn_data
-
- # 1. Set random-seed
- set_all_random_seed(seed)
-
- # 2. Build speech2text
- speech2text_kwargs = dict(
- asr_train_config=asr_train_config,
- asr_model_file=asr_model_file,
- lm_train_config=lm_train_config,
- lm_file=lm_file,
- token_type=token_type,
- bpemodel=bpemodel,
- device=device,
- maxlenratio=maxlenratio,
- minlenratio=minlenratio,
- dtype=dtype,
- beam_size=beam_size,
- ctc_weight=ctc_weight,
- lm_weight=lm_weight,
- ngram_weight=ngram_weight,
- penalty=penalty,
- nbest=nbest,
- frontend_conf=frontend_conf,
- )
- speech2text = Speech2Text(**speech2text_kwargs)
-
- data_path_and_name_and_type_new = [
- audio_lists, data_path_and_name_and_type[0], data_path_and_name_and_type[1]
- ]
-
- # 3. Build data-iterator
- loader = ASRTask.build_streaming_iterator_modelscope(
- data_path_and_name_and_type_new,
- dtype=dtype,
- batch_size=batch_size,
- key_file=key_file,
- num_workers=num_workers,
- preprocess_fn=ASRTask.build_preprocess_fn(speech2text.asr_train_args, False),
- collate_fn=ASRTask.build_collate_fn(speech2text.asr_train_args, False),
- allow_variable_data_keys=allow_variable_data_keys,
- inference=True,
- sample_rate=fs
- )
-
- forward_time_total = 0.0
- length_total = 0.0
- asr_result_list = []
- # 7 .Start for-loop
- # FIXME(kamo): The output format should be discussed about
- for keys, batch in loader:
- assert isinstance(batch, dict), type(batch)
- assert all(isinstance(s, str) for s in keys), keys
- _bs = len(next(iter(batch.values())))
- assert len(keys) == _bs, f"{len(keys)} != {_bs}"
- batch = {k: v[0] for k, v in batch.items() if not k.endswith("_lengths")}
-
- logging.info("decoding, utt_id: {}".format(keys))
- # N-best list of (text, token, token_int, hyp_object)
-
- try:
- time_beg = time.time()
- results = speech2text(**batch)
- time_end = time.time()
- forward_time = time_end - time_beg
- lfr_factor = results[0][-1]
- length = results[0][-2]
- results = [results[0][:-2]]
- forward_time_total += forward_time
- length_total += length
- logging.info(
- "decoding, feature length: {}, forward_time: {:.4f}, rtf: {:.4f}".
- format(length, forward_time, 100 * forward_time / (length * lfr_factor)))
- except TooShortUttError as e:
- logging.warning(f"Utterance {keys} {e}")
- hyp = Hypothesis(score=0.0, scores={}, states={}, yseq=[])
- results = [[" ", ["<space>"], [2], hyp]] * nbest
-
- # Only supporting batch_size==1
- key = keys[0]
- for n, (text, token, token_int, hyp) in zip(range(1, nbest + 1), results):
- if text is not None:
- text_postprocessed = postprocess_utils.sentence_postprocess(token)
- item = {'key': key, 'value': text_postprocessed}
- asr_result_list.append(item)
-
- logging.info("decoding, predictions: {}".format(text))
- finish_count += 1
- asr_utils.print_progress(finish_count / file_count)
-
- logging.info("decoding, feature length total: {}, forward_time total: {:.4f}, rtf avg: {:.4f}".
- format(length_total, forward_time_total, 100 * forward_time_total / (length_total * lfr_factor)))
- if features_type == 'sound':
- # data format is wav
- length_total_seconds = length_total / sr
- length_total_bytes = length_total * 2
- else:
- # data format is kaldi_ark
- length_total_seconds = length_total * hop_length / sr
- length_total_bytes = length_total * hop_length * 2
-
- logger.info(
- header_colors + # noqa: *
- 'decoding, feature length total: {}bytes, forward_time total: {:.4f}s, rtf avg: {:.4f}'
- .format(length_total_bytes, forward_time_total, forward_time_total /
- length_total_seconds) + end_colors)
-
- return asr_result_list
-
-
-def set_parameters(language: str = None,
- sample_rate: Union[int, Dict[Any, int]] = None):
- if language is not None:
- global global_asr_language
- global_asr_language = language
- if sample_rate is not None:
- global global_sample_rate
- global_sample_rate = sample_rate
-
-
-def asr_inference(maxlenratio: float,
- minlenratio: float,
- beam_size: int,
- ngpu: int,
- ctc_weight: float,
- lm_weight: float,
- penalty: float,
- name_and_type: list,
- audio_lists: Union[List[Any], bytes],
- asr_train_config: Optional[str],
- asr_model_file: Optional[str],
- nbest: int = 1,
- num_workers: int = 1,
- log_level: Union[int, str] = 'INFO',
- batch_size: int = 1,
- dtype: str = 'float32',
- seed: int = 0,
- key_file: Optional[str] = None,
- lm_train_config: Optional[str] = None,
- lm_file: Optional[str] = None,
- word_lm_train_config: Optional[str] = None,
- word_lm_file: Optional[str] = None,
- ngram_file: Optional[str] = None,
- ngram_weight: float = 0.9,
- model_tag: Optional[str] = None,
- token_type: Optional[str] = None,
- bpemodel: Optional[str] = None,
- allow_variable_data_keys: bool = False,
- transducer_conf: Optional[dict] = None,
- streaming: bool = False,
- frontend_conf: dict = None,
- fs: Union[dict, int] = None,
- lang: Optional[str] = None,
- outputdir: Optional[str] = None):
- if lang is not None:
- global global_asr_language
- global_asr_language = lang
- if fs is not None:
- global global_sample_rate
- global_sample_rate = fs
-
- # force use CPU if data type is bytes
- if isinstance(audio_lists, bytes):
- num_workers = 0
- ngpu = 0
-
- return inference(output_dir=outputdir,
- maxlenratio=maxlenratio,
- minlenratio=minlenratio,
- batch_size=batch_size,
- dtype=dtype,
- beam_size=beam_size,
- ngpu=ngpu,
- seed=seed,
- ctc_weight=ctc_weight,
- lm_weight=lm_weight,
- ngram_weight=ngram_weight,
- penalty=penalty,
- nbest=nbest,
- num_workers=num_workers,
- log_level=log_level,
- data_path_and_name_and_type=name_and_type,
- audio_lists=audio_lists,
- key_file=key_file,
- asr_train_config=asr_train_config,
- asr_model_file=asr_model_file,
- lm_train_config=lm_train_config,
- lm_file=lm_file,
- word_lm_train_config=word_lm_train_config,
- word_lm_file=word_lm_file,
- ngram_file=ngram_file,
- model_tag=model_tag,
- token_type=token_type,
- bpemodel=bpemodel,
- allow_variable_data_keys=allow_variable_data_keys,
- transducer_conf=transducer_conf,
- streaming=streaming,
- frontend_conf=frontend_conf)
-
-
-
-def get_parser():
- parser = config_argparse.ArgumentParser(
- description="ASR Decoding",
- formatter_class=argparse.ArgumentDefaultsHelpFormatter,
- )
-
- # Note(kamo): Use '_' instead of '-' as separator.
- # '-' is confusing if written in yaml.
- parser.add_argument(
- "--log_level",
- type=lambda x: x.upper(),
- default="INFO",
- choices=("CRITICAL", "ERROR", "WARNING", "INFO", "DEBUG", "NOTSET"),
- help="The verbose level of logging",
- )
-
- parser.add_argument("--output_dir", type=str, required=True)
- parser.add_argument(
- "--ngpu",
- type=int,
- default=0,
- help="The number of gpus. 0 indicates CPU mode",
- )
- parser.add_argument("--seed", type=int, default=0, help="Random seed")
- parser.add_argument(
- "--dtype",
- default="float32",
- choices=["float16", "float32", "float64"],
- help="Data type",
- )
- parser.add_argument(
- "--num_workers",
- type=int,
- default=1,
- help="The number of workers used for DataLoader",
- )
-
- group = parser.add_argument_group("Input data related")
- group.add_argument(
- "--data_path_and_name_and_type",
- type=str2triple_str,
- required=True,
- action="append",
- )
- group.add_argument("--audio_lists", type=list, default=[{'key':'EdevDEWdIYQ_0021','file':'/mnt/data/jiangyu.xzy/test_data/speech_io/SPEECHIO_ASR_ZH00007_zhibodaihuo/wav/EdevDEWdIYQ_0021.wav'}])
- group.add_argument("--key_file", type=str_or_none)
- group.add_argument("--allow_variable_data_keys", type=str2bool, default=False)
-
- group = parser.add_argument_group("The model configuration related")
- group.add_argument(
- "--asr_train_config",
- type=str,
- help="ASR training configuration",
- )
- group.add_argument(
- "--asr_model_file",
- type=str,
- help="ASR model parameter file",
- )
- group.add_argument(
- "--lm_train_config",
- type=str,
- help="LM training configuration",
- )
- group.add_argument(
- "--lm_file",
- type=str,
- help="LM parameter file",
- )
- group.add_argument(
- "--word_lm_train_config",
- type=str,
- help="Word LM training configuration",
- )
- group.add_argument(
- "--word_lm_file",
- type=str,
- help="Word LM parameter file",
- )
- group.add_argument(
- "--ngram_file",
- type=str,
- help="N-gram parameter file",
- )
- group.add_argument(
- "--model_tag",
- type=str,
- help="Pretrained model tag. If specify this option, *_train_config and "
- "*_file will be overwritten",
- )
-
- group = parser.add_argument_group("Beam-search related")
- group.add_argument(
- "--batch_size",
- type=int,
- default=1,
- help="The batch size for inference",
- )
- group.add_argument("--nbest", type=int, default=1, help="Output N-best hypotheses")
- group.add_argument("--beam_size", type=int, default=20, help="Beam size")
- group.add_argument("--penalty", type=float, default=0.0, help="Insertion penalty")
- group.add_argument(
- "--maxlenratio",
- type=float,
- default=0.0,
- help="Input length ratio to obtain max output length. "
- "If maxlenratio=0.0 (default), it uses a end-detect "
- "function "
- "to automatically find maximum hypothesis lengths."
- "If maxlenratio<0.0, its absolute value is interpreted"
- "as a constant max output length",
- )
- group.add_argument(
- "--minlenratio",
- type=float,
- default=0.0,
- help="Input length ratio to obtain min output length",
- )
- group.add_argument(
- "--ctc_weight",
- type=float,
- default=0.5,
- help="CTC weight in joint decoding",
- )
- group.add_argument("--lm_weight", type=float, default=1.0, help="RNNLM weight")
- group.add_argument("--ngram_weight", type=float, default=0.9, help="ngram weight")
- group.add_argument("--streaming", type=str2bool, default=False)
-
- group.add_argument(
- "--asr_model_config",
- default=None,
- help="",
- )
-
- group = parser.add_argument_group("Text converter related")
- group.add_argument(
- "--token_type",
- type=str_or_none,
- default=None,
- choices=["char", "bpe", None],
- help="The token type for ASR model. "
- "If not given, refers from the training args",
- )
- group.add_argument(
- "--bpemodel",
- type=str_or_none,
- default=None,
- help="The model path of sentencepiece. "
- "If not given, refers from the training args",
- )
-
- return parser
-
-
-def main(cmd=None):
- print(get_commandline_args(), file=sys.stderr)
- parser = get_parser()
- args = parser.parse_args(cmd)
- kwargs = vars(args)
- kwargs.pop("config", None)
- inference(**kwargs)
-
-
-if __name__ == "__main__":
- main()
diff --git a/funasr/bin/asr_inference_uniasr.py b/funasr/bin/asr_inference_uniasr.py
index 796c5b3..a1a23ba 100755
--- a/funasr/bin/asr_inference_uniasr.py
+++ b/funasr/bin/asr_inference_uniasr.py
@@ -8,6 +8,8 @@
from typing import Sequence
from typing import Tuple
from typing import Union
+from typing import Dict
+from typing import Any
import numpy as np
import torch
@@ -31,7 +33,21 @@
from funasr.utils.types import str2bool
from funasr.utils.types import str2triple_str
from funasr.utils.types import str_or_none
+from funasr.utils import asr_utils, wav_utils, postprocess_utils
+from funasr.models.frontend.wav_frontend import WavFrontend
+from modelscope.utils.logger import get_logger
+
+logger = get_logger()
+
+header_colors = '\033[95m'
+end_colors = '\033[0m'
+
+global_asr_language: str = 'zh-cn'
+global_sample_rate: Union[int, Dict[Any, int]] = {
+ 'audio_fs': 16000,
+ 'model_fs': 16000
+}
class Speech2Text:
"""Speech2Text class
@@ -66,6 +82,7 @@
token_num_relax: int = 1,
decoding_ind: int = 0,
decoding_mode: str = "model1",
+ frontend_conf: dict = None,
**kwargs,
):
assert check_argument_types()
@@ -75,6 +92,10 @@
asr_model, asr_train_args = ASRTask.build_model_from_file(
asr_train_config, asr_model_file, device
)
+ frontend = None
+ if asr_model.frontend is None and frontend_conf is not None:
+ frontend = WavFrontend(**frontend_conf)
+ # asr_model.frontend = frontend
asr_model.to(dtype=getattr(torch, dtype)).eval()
if decoding_mode == "model1":
decoder = asr_model.decoder
@@ -162,6 +183,7 @@
self.token_num_relax = token_num_relax
self.decoding_ind = decoding_ind
self.decoding_mode = decoding_mode
+ self.frontend = frontend
@torch.no_grad()
def __call__(
@@ -177,7 +199,7 @@
"""Inference
Args:
- data: Input speech data
+ speech: Input speech data
Returns:
text, token, token_int, hyp
@@ -190,14 +212,22 @@
# data: (Nsamples,) -> (1, Nsamples)
speech = speech.unsqueeze(0).to(getattr(torch, self.dtype))
+ lfr_factor = max(1, (speech.size()[-1] // 80) - 1)
# lengths: (1,)
lengths = speech.new_full([1], dtype=torch.long, fill_value=speech.size(1))
- batch = {"speech": speech, "speech_lengths": lengths}
+ speech_raw = speech.clone().to(self.device)
+ if self.frontend is not None:
+ feats, feats_len = self.frontend.forward(speech, lengths)
+ feats = to_device(feats, device=self.device)
+ feats_len = feats_len.int()
+ else:
+ feats = speech_raw
+ feats_len = lengths
+ batch = {"speech": feats, "speech_lengths": feats_len}
# a. To device
batch = to_device(batch, device=self.device)
# b. Forward Encoder
- speech_raw = speech.clone().to(self.device)
enc, enc_len = self.asr_model.encode(**batch, ind=self.decoding_ind)
if isinstance(enc, tuple):
enc = enc[0]
@@ -205,7 +235,7 @@
if self.decoding_mode == "model1":
predictor_outs = self.asr_model.calc_predictor_mask(enc, enc_len)
else:
- enc, enc_len = self.asr_model.encode2(enc, enc_len, speech_raw, lengths, ind=self.decoding_ind)
+ enc, enc_len = self.asr_model.encode2(enc, enc_len, feats, feats_len, ind=self.decoding_ind)
predictor_outs = self.asr_model.calc_predictor_mask2(enc, enc_len)
scama_mask = predictor_outs[4]
@@ -249,33 +279,37 @@
def inference(
- output_dir: str,
maxlenratio: float,
minlenratio: float,
batch_size: int,
- dtype: str,
beam_size: int,
ngpu: int,
- seed: int,
ctc_weight: float,
lm_weight: float,
- ngram_weight: float,
penalty: float,
- nbest: int,
- num_workers: int,
log_level: Union[int, str],
- data_path_and_name_and_type: Sequence[Tuple[str, str, str]],
- key_file: Optional[str],
+ data_path_and_name_and_type,
asr_train_config: Optional[str],
asr_model_file: Optional[str],
- lm_train_config: Optional[str],
- lm_file: Optional[str],
- word_lm_train_config: Optional[str],
- ngram_file: Optional[str],
- token_type: Optional[str],
- bpemodel: Optional[str],
- allow_variable_data_keys: bool,
- streaming: bool,
+ ngram_file: Optional[str] = None,
+ audio_lists: Union[List[Any], bytes] = None,
+ lm_train_config: Optional[str] = None,
+ lm_file: Optional[str] = None,
+ token_type: Optional[str] = None,
+ key_file: Optional[str] = None,
+ word_lm_train_config: Optional[str] = None,
+ bpemodel: Optional[str] = None,
+ allow_variable_data_keys: bool = False,
+ streaming: bool = False,
+ output_dir: Optional[str] = None,
+ dtype: str = "float32",
+ seed: int = 0,
+ ngram_weight: float = 0.9,
+ nbest: int = 1,
+ num_workers: int = 1,
+ frontend_conf: dict = None,
+ fs: Union[dict, int] = 16000,
+ lang: Optional[str] = None,
token_num_relax: int = 1,
decoding_ind: int = 0,
decoding_mode: str = "model1",
@@ -298,7 +332,46 @@
device = "cuda"
else:
device = "cpu"
+ hop_length: int = 160
+ sr: int = 16000
+ if isinstance(fs, int):
+ sr = fs
+ else:
+ if 'model_fs' in fs and fs['model_fs'] is not None:
+ sr = fs['model_fs']
+ # data_path_and_name_and_type for modelscope: (data from audio_lists)
+ # ['speech', 'sound', 'am.mvn']
+ # data_path_and_name_and_type for funasr:
+ # [('/mnt/data/jiangyu.xzy/exp/maas/mvn.1.scp', 'speech', 'kaldi_ark')]
+ if isinstance(data_path_and_name_and_type[0], Tuple):
+ features_type: str = data_path_and_name_and_type[0][1]
+ elif isinstance(data_path_and_name_and_type[0], str):
+ features_type: str = data_path_and_name_and_type[1]
+ else:
+ raise NotImplementedError("unknown features type:{0}".format(data_path_and_name_and_type))
+ if features_type != 'sound':
+ frontend_conf = None
+ flag_modelscope = False
+ else:
+ flag_modelscope = True
+ if frontend_conf is not None:
+ if 'hop_length' in frontend_conf:
+ hop_length = frontend_conf['hop_length']
+ finish_count = 0
+ file_count = 1
+ if flag_modelscope and not isinstance(data_path_and_name_and_type[0], Tuple):
+ data_path_and_name_and_type_new = [
+ audio_lists, data_path_and_name_and_type[0], data_path_and_name_and_type[1]
+ ]
+ if isinstance(audio_lists, bytes):
+ file_count = 1
+ else:
+ file_count = len(audio_lists)
+ if len(data_path_and_name_and_type) >= 3 and frontend_conf is not None:
+ mvn_file = data_path_and_name_and_type[2]
+ mvn_data = wav_utils.extract_CMVN_featrures(mvn_file)
+ frontend_conf['mvn_data'] = mvn_data
# 1. Set random-seed
set_all_random_seed(seed)
@@ -325,45 +398,66 @@
token_num_relax=token_num_relax,
decoding_ind=decoding_ind,
decoding_mode=decoding_mode,
+ frontend_conf=frontend_conf,
)
speech2text = Speech2Text(**speech2text_kwargs)
# 3. Build data-iterator
- loader = ASRTask.build_streaming_iterator(
- data_path_and_name_and_type,
- dtype=dtype,
- batch_size=batch_size,
- key_file=key_file,
- num_workers=num_workers,
- preprocess_fn=ASRTask.build_preprocess_fn(speech2text.asr_train_args, False),
- collate_fn=ASRTask.build_collate_fn(speech2text.asr_train_args, False),
- allow_variable_data_keys=allow_variable_data_keys,
- inference=True,
- )
+ if flag_modelscope:
+ loader = ASRTask.build_streaming_iterator_modelscope(
+ data_path_and_name_and_type_new,
+ dtype=dtype,
+ batch_size=batch_size,
+ key_file=key_file,
+ num_workers=num_workers,
+ preprocess_fn=ASRTask.build_preprocess_fn(speech2text.asr_train_args, False),
+ collate_fn=ASRTask.build_collate_fn(speech2text.asr_train_args, False),
+ allow_variable_data_keys=allow_variable_data_keys,
+ inference=True,
+ sample_rate=fs
+ )
+ else:
+ loader = ASRTask.build_streaming_iterator(
+ data_path_and_name_and_type,
+ dtype=dtype,
+ batch_size=batch_size,
+ key_file=key_file,
+ num_workers=num_workers,
+ preprocess_fn=ASRTask.build_preprocess_fn(speech2text.asr_train_args, False),
+ collate_fn=ASRTask.build_collate_fn(speech2text.asr_train_args, False),
+ allow_variable_data_keys=allow_variable_data_keys,
+ inference=True,
+ )
# 7 .Start for-loop
# FIXME(kamo): The output format should be discussed about
- with DatadirWriter(output_dir) as writer:
- for keys, batch in loader:
- assert isinstance(batch, dict), type(batch)
- assert all(isinstance(s, str) for s in keys), keys
- _bs = len(next(iter(batch.values())))
- assert len(keys) == _bs, f"{len(keys)} != {_bs}"
- batch = {k: v[0] for k, v in batch.items() if not k.endswith("_lengths")}
+ asr_result_list = []
+ if output_dir is not None:
+ writer = DatadirWriter(output_dir)
+ else:
+ writer = None
- # N-best list of (text, token, token_int, hyp_object)
- try:
- results = speech2text(**batch)
- except TooShortUttError as e:
- logging.warning(f"Utterance {keys} {e}")
- hyp = Hypothesis(score=0.0, scores={}, states={}, yseq=[])
- results = [[" ", ["<space>"], [2], hyp]] * nbest
+ for keys, batch in loader:
+ assert isinstance(batch, dict), type(batch)
+ assert all(isinstance(s, str) for s in keys), keys
+ _bs = len(next(iter(batch.values())))
+ assert len(keys) == _bs, f"{len(keys)} != {_bs}"
+ batch = {k: v[0] for k, v in batch.items() if not k.endswith("_lengths")}
- # Only supporting batch_size==1
- key = keys[0]
- logging.info(f"Utterance: {key}")
- for n, (text, token, token_int, hyp) in zip(range(1, nbest + 1), results):
- # Create a directory: outdir/{n}best_recog
+ # N-best list of (text, token, token_int, hyp_object)
+ try:
+ results = speech2text(**batch)
+ except TooShortUttError as e:
+ logging.warning(f"Utterance {keys} {e}")
+ hyp = Hypothesis(score=0.0, scores={}, states={}, yseq=[])
+ results = [[" ", ["<space>"], [2], hyp]] * nbest
+
+ # Only supporting batch_size==1
+ key = keys[0]
+ logging.info(f"Utterance: {key}")
+ for n, (text, token, token_int, hyp) in zip(range(1, nbest + 1), results):
+ # Create a directory: outdir/{n}best_recog
+ if writer is not None:
ibest_writer = writer[f"{n}best_recog"]
# Write the result to each file
@@ -371,8 +465,25 @@
ibest_writer["token_int"][key] = " ".join(map(str, token_int))
ibest_writer["score"][key] = str(hyp.score)
- if text is not None:
+ if text is not None:
+ text_postprocessed = postprocess_utils.sentence_postprocess(token)
+ item = {'key': key, 'value': text_postprocessed}
+ asr_result_list.append(item)
+ finish_count += 1
+ asr_utils.print_progress(finish_count / file_count)
+ if writer is not None:
ibest_writer["text"][key] = text
+ return asr_result_list
+
+
+def set_parameters(language: str = None,
+ sample_rate: Union[int, Dict[Any, int]] = None):
+ if language is not None:
+ global global_asr_language
+ global_asr_language = language
+ if sample_rate is not None:
+ global global_sample_rate
+ global_sample_rate = sample_rate
def get_parser():
@@ -419,6 +530,8 @@
required=True,
action="append",
)
+ group.add_argument("--audio_lists", type=list, default=None)
+ # example=[{'key':'EdevDEWdIYQ_0021','file':'/mnt/data/jiangyu.xzy/test_data/speech_io/SPEECHIO_ASR_ZH00007_zhibodaihuo/wav/EdevDEWdIYQ_0021.wav'}])
group.add_argument("--key_file", type=str_or_none)
group.add_argument("--allow_variable_data_keys", type=str2bool, default=False)
diff --git a/funasr/bin/modelscope_infer.py b/funasr/bin/modelscope_infer.py
index 74c2fb7..3be6d03 100755
--- a/funasr/bin/modelscope_infer.py
+++ b/funasr/bin/modelscope_infer.py
@@ -17,7 +17,7 @@
help="model name in modelscope")
parser.add_argument("--model_revision",
type=str,
- default="v1.0.3",
+ default="v1.0.4",
help="model revision in modelscope")
parser.add_argument("--local_model_path",
type=str,
diff --git a/funasr/models/e2e_asr_paraformer.py b/funasr/models/e2e_asr_paraformer.py
index 89f7cf0..3f8359d 100644
--- a/funasr/models/e2e_asr_paraformer.py
+++ b/funasr/models/e2e_asr_paraformer.py
@@ -493,7 +493,7 @@
def sampler(self, encoder_out, encoder_out_lens, ys_pad, ys_pad_lens, pre_acoustic_embeds):
tgt_mask = (~make_pad_mask(ys_pad_lens, maxlen=ys_pad_lens.max())[:, :, None]).to(ys_pad.device)
- ys_pad *= tgt_mask[:, :, 0]
+ ys_pad = ys_pad * tgt_mask[:, :, 0]
ys_pad_embed = self.decoder.embed(ys_pad)
with torch.no_grad():
decoder_outs = self.decoder(
diff --git a/funasr/models/predictor/cif.py b/funasr/models/predictor/cif.py
index ea41c6c..2eba4e2 100644
--- a/funasr/models/predictor/cif.py
+++ b/funasr/models/predictor/cif.py
@@ -2,6 +2,7 @@
from torch import nn
from funasr.modules.nets_utils import make_pad_mask
+from funasr.modules.streaming_utils.utils import sequence_mask
class CifPredictor(nn.Module):
def __init__(self, idim, l_order, r_order, threshold=1.0, dropout=0.1, smooth_factor=1.0, noise_threshold=0, tail_threshold=0.45):
@@ -14,6 +15,7 @@
self.threshold = threshold
self.smooth_factor = smooth_factor
self.noise_threshold = noise_threshold
+ self.tail_threshold = tail_threshold
def forward(self, hidden, target_label=None, mask=None, ignore_id=-1, mask_chunk_predictor=None,
target_label_length=None):
@@ -42,8 +44,39 @@
token_num = alphas.sum(-1)
if target_length is not None:
alphas *= (target_length / token_num)[:, None].repeat(1, alphas.size(1))
+ elif self.tail_threshold > 0.0:
+ hidden, alphas, token_num = self.tail_process_fn(hidden, alphas, token_num, mask=mask)
+
acoustic_embeds, cif_peak = cif(hidden, alphas, self.threshold)
+
+ if target_length is None and self.tail_threshold > 0.0:
+ token_num_int = torch.max(token_num).type(torch.int32).item()
+ acoustic_embeds = acoustic_embeds[:, :token_num_int, :]
+
return acoustic_embeds, token_num, alphas, cif_peak
+
+ def tail_process_fn(self, hidden, alphas, token_num=None, mask=None):
+ b, t, d = hidden.size()
+ tail_threshold = self.tail_threshold
+ if mask is not None:
+ zeros_t = torch.zeros((b, 1), dtype=torch.float32, device=alphas.device)
+ ones_t = torch.ones_like(zeros_t)
+ mask_1 = torch.cat([mask, zeros_t], dim=1)
+ mask_2 = torch.cat([ones_t, mask], dim=1)
+ mask = mask_2 - mask_1
+ tail_threshold = mask * tail_threshold
+ alphas = torch.cat([alphas, tail_threshold], dim=1)
+ else:
+ tail_threshold = torch.tensor([tail_threshold], dtype=alphas.dtype).to(alphas.device)
+ tail_threshold = torch.reshape(tail_threshold, (1, 1))
+ alphas = torch.cat([alphas, tail_threshold], dim=1)
+ zeros = torch.zeros((b, 1, d), dtype=hidden.dtype).to(hidden.device)
+ hidden = torch.cat([hidden, zeros], dim=1)
+ token_num = alphas.sum(dim=-1)
+ token_num_floor = torch.floor(token_num)
+
+ return hidden, alphas, token_num_floor
+
def gen_frame_alignments(self,
alphas: torch.Tensor = None,
@@ -120,10 +153,12 @@
alphas = torch.sigmoid(output)
alphas = torch.nn.functional.relu(alphas * self.smooth_factor - self.noise_threshold)
if mask is not None:
- alphas = alphas * mask.transpose(-1, -2).float()
+ mask = mask.transpose(-1, -2).float()
+ alphas = alphas * mask
if mask_chunk_predictor is not None:
alphas = alphas * mask_chunk_predictor
alphas = alphas.squeeze(-1)
+ mask = mask.squeeze(-1)
if target_label_length is not None:
target_length = target_label_length
elif target_label is not None:
@@ -134,7 +169,7 @@
if target_length is not None:
alphas *= (target_length / token_num)[:, None].repeat(1, alphas.size(1))
elif self.tail_threshold > 0.0:
- hidden, alphas, token_num = self.tail_process_fn(hidden, alphas, token_num)
+ hidden, alphas, token_num = self.tail_process_fn(hidden, alphas, token_num, mask=mask)
acoustic_embeds, cif_peak = cif(hidden, alphas, self.threshold)
if target_length is None and self.tail_threshold > 0.0:
@@ -143,12 +178,21 @@
return acoustic_embeds, token_num, alphas, cif_peak
- def tail_process_fn(self, hidden, alphas, token_num=None):
+ def tail_process_fn(self, hidden, alphas, token_num=None, mask=None):
b, t, d = hidden.size()
tail_threshold = self.tail_threshold
- tail_threshold = torch.tensor([tail_threshold], dtype=alphas.dtype).to(alphas.device)
- tail_threshold = tail_threshold.unsqueeze(0).repeat(b, 1)
- alphas = torch.cat([alphas, tail_threshold], dim=1)
+ if mask is not None:
+ zeros_t = torch.zeros((b, 1), dtype=torch.float32, device=alphas.device)
+ ones_t = torch.ones_like(zeros_t)
+ mask_1 = torch.cat([mask, zeros_t], dim=1)
+ mask_2 = torch.cat([ones_t, mask], dim=1)
+ mask = mask_2 - mask_1
+ tail_threshold = mask * tail_threshold
+ alphas = torch.cat([alphas, tail_threshold], dim=1)
+ else:
+ tail_threshold = torch.tensor([tail_threshold], dtype=alphas.dtype).to(alphas.device)
+ tail_threshold = torch.reshape(tail_threshold, (1, 1))
+ alphas = torch.cat([alphas, tail_threshold], dim=1)
zeros = torch.zeros((b, 1, d), dtype=hidden.dtype).to(hidden.device)
hidden = torch.cat([hidden, zeros], dim=1)
token_num = alphas.sum(dim=-1)
diff --git a/funasr/modules/streaming_utils/__init__.py b/funasr/modules/streaming_utils/__init__.py
new file mode 100644
index 0000000..e69de29
--- /dev/null
+++ b/funasr/modules/streaming_utils/__init__.py
diff --git a/funasr/version.txt b/funasr/version.txt
index b1e80bb..845639e 100644
--- a/funasr/version.txt
+++ b/funasr/version.txt
@@ -1 +1 @@
-0.1.3
+0.1.4
diff --git a/setup.py b/setup.py
index ac5960d..e17c6ae 100644
--- a/setup.py
+++ b/setup.py
@@ -45,7 +45,7 @@
"editdistance==0.5.2",
"wandb",
],
- # recipe: The modules actually are not invoked in the main module of espnet,
+ # recipe: The modules actually are not invoked in the main module of funasr,
# but are invoked for the python scripts in each recipe
"recipe": [
"espnet_model_zoo",
@@ -130,7 +130,7 @@
setup_requires=setup_requires,
tests_require=tests_require,
extras_require=extras_require,
- python_requires=">=3.6.0",
+ python_requires=">=3.7.0",
classifiers=[
"Programming Language :: Python",
"Programming Language :: Python :: 3",
--
Gitblit v1.9.1