-
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.
Merge branch 'espnet:master' into master
- Loading branch information
Showing
25 changed files
with
872 additions
and
27 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
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 |
82 changes: 82 additions & 0 deletions
82
egs2/slurp_entity/asr1/conf/tuning/train_asr_branchformer_xlsr_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,82 @@ | ||
# network architecture | ||
# encoder related | ||
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-slurp" | ||
length_adaptor_n_layers: 1 | ||
lang_token_id: 250004 | ||
|
||
decoder: hugging_face_transformers | ||
decoder_conf: | ||
model_name_or_path: "akreal/mbart-large-50-finetuned-slurp" | ||
|
||
use_amp: true | ||
num_workers: 2 | ||
optim: adam | ||
batch_type: length | ||
batch_bins: 170000 | ||
accum_grad: 4 | ||
optim_conf: | ||
lr: 0.00005 | ||
weight_decay: 0.000001 | ||
scheduler: warmuplr # pytorch v1.1.0+ required | ||
scheduler_conf: | ||
warmup_steps: 25000 | ||
max_epoch: 50 | ||
|
||
freeze_param: [ | ||
"frontend.upstream" | ||
] | ||
|
||
frontend: s3prl | ||
frontend_conf: | ||
frontend_conf: | ||
upstream: xls_r_300m # 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
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,33 @@ | ||
#!/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_char" | ||
valid_set="devel_char" | ||
test_sets="test_char devel_char" | ||
|
||
asr_config=conf/tuning/train_asr_branchformer_xlsr_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_data_opts "--token_type_bpe false" \ | ||
--local_score_opts "--token_type_bpe false" \ | ||
--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}" \ | ||
--inference_nj 1 \ | ||
--gpu_inference true \ | ||
--train_set "${train_set}" \ | ||
--valid_set "${valid_set}" \ | ||
--lm_train_text "data/${train_set}/text" \ | ||
--test_sets "${test_sets}" "$@" |
Oops, something went wrong.