-
Notifications
You must be signed in to change notification settings - Fork 2.1k
Commit
This commit does not belong to any branch on this repository, and may belong to a fork outside of the repository.
Add SLUE-VoxPopuli results for WavLM with mBART-50
- Loading branch information
Showing
7 changed files
with
155 additions
and
24 deletions.
There are no files selected for viewing
This file contains bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Original file line number | Diff line number | Diff line change |
---|---|---|
@@ -0,0 +1,3 @@ | ||
beam_size: 5 | ||
ctc_weight: 0.0 | ||
hugging_face_decoder: True |
79 changes: 79 additions & 0 deletions
79
egs2/slue-voxpopuli/asr1/conf/tuning/train_asr_branchformer_wavlm_mbart.yaml
This file contains bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Original file line number | Diff line number | Diff line change |
---|---|---|
@@ -0,0 +1,79 @@ | ||
encoder: branchformer | ||
encoder_conf: | ||
output_size: 1024 | ||
use_attn: true | ||
attention_heads: 8 | ||
attention_layer_type: rel_selfattn | ||
pos_enc_layer_type: rel_pos | ||
rel_pos_type: latest | ||
use_cgmlp: true | ||
cgmlp_linear_units: 4096 | ||
cgmlp_conv_kernel: 31 | ||
use_linear_after_conv: false | ||
gate_activation: identity | ||
merge_method: concat | ||
cgmlp_weight: 0.5 # used only if merge_method is "fixed_ave" | ||
attn_branch_drop_rate: 0.0 # used only if merge_method is "learned_ave" | ||
num_blocks: 18 | ||
dropout_rate: 0.1 | ||
positional_dropout_rate: 0.1 | ||
attention_dropout_rate: 0.1 | ||
input_layer: conv2d | ||
stochastic_depth_rate: 0.0 | ||
|
||
postencoder: hugging_face_transformers | ||
postencoder_conf: | ||
model_name_or_path: "akreal/mbart-large-50-finetuned-slue" | ||
length_adaptor_n_layers: 1 | ||
lang_token_id: 250004 | ||
|
||
decoder: hugging_face_transformers | ||
decoder_conf: | ||
model_name_or_path: "akreal/mbart-large-50-finetuned-slue" | ||
|
||
use_amp: true | ||
optim: adam | ||
batch_type: length | ||
batch_bins: 300000 | ||
accum_grad: 4 | ||
optim_conf: | ||
lr: 0.00005 | ||
weight_decay: 0.000001 | ||
scheduler: warmuplr # pytorch v1.1.0+ required | ||
scheduler_conf: | ||
warmup_steps: 40000 | ||
max_epoch: 100 | ||
|
||
freeze_param: [ | ||
"frontend.upstream" | ||
] | ||
|
||
frontend: s3prl | ||
frontend_conf: | ||
frontend_conf: | ||
upstream: wavlm_large # Note: If the upstream is changed, please change the input_size in the preencoder. | ||
download_dir: ./hub | ||
multilayer_feature: True | ||
|
||
preencoder: linear | ||
preencoder_conf: | ||
input_size: 1024 # Note: If the upstream is changed, please change this value accordingly. | ||
output_size: 80 | ||
|
||
model_conf: | ||
ctc_weight: 0.0 | ||
lsm_weight: 0.1 | ||
length_normalized_loss: false | ||
extract_feats_in_collect_stats: false # Note: "False" means during collect stats (stage 10), generating dummy stats files rather than extract_feats by forward frontend. | ||
# mBART dictionary customizations | ||
ignore_id: 1 | ||
sym_blank: "<pad>" | ||
sym_sos: "<s>" | ||
sym_eos: "</s>" | ||
lang_token_id: 250004 | ||
|
||
best_model_criterion: | ||
- - valid | ||
- acc | ||
- max | ||
keep_nbest_models: 10 |
This file contains bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Original file line number | Diff line number | Diff line change |
---|---|---|
@@ -0,0 +1,30 @@ | ||
#!/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 | ||
|
||
train_set="train" | ||
valid_set="devel" | ||
test_sets="test devel" | ||
|
||
asr_config=conf/tuning/train_asr_branchformer_wavlm_mbart.yaml | ||
inference_config=conf/decode_asr_hf.yaml | ||
|
||
./asr.sh \ | ||
--lang en \ | ||
--ngpu 1 \ | ||
--use_lm false \ | ||
--token_type hugging_face \ | ||
--hugging_face_model_name_or_path facebook/mbart-large-50-many-to-many-mmt \ | ||
--local_score_opts "--score_folder score_wer" \ | ||
--max_wav_duration 30 \ | ||
--speed_perturb_factors "0.9 1.0 1.1" \ | ||
--feats_normalize utterance_mvn \ | ||
--asr_config "${asr_config}" \ | ||
--inference_config "${inference_config}" \ | ||
--train_set "${train_set}" \ | ||
--valid_set "${valid_set}" \ | ||
--lm_train_text "data/${train_set}/text" \ | ||
--test_sets "${test_sets}" "$@" |
This file contains bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters