Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
43 commits
Select commit Hold shift + click to select a range
c2ec7ad
add mbart
mhn226 Oct 9, 2023
9f94d8f
Add tristage scheduler
mhn226 Oct 17, 2023
4fc90ca
Add mbart beam search
mhn226 Oct 17, 2023
4268806
Add IWLST recipes
mhn226 Oct 17, 2023
63eb6eb
Add new models' inteference interface
mhn226 Oct 19, 2023
41c1d63
Add info of new models
mhn226 Oct 19, 2023
ad11b9f
Add nllb scores
mhn226 Oct 20, 2023
b2132ab
Add new models' info
mhn226 Oct 20, 2023
71e7ed1
Add test info IWSLT recipe
mhn226 Oct 23, 2023
e4f2f0a
Merge branch 'text_based_HF' of https://github.com/mhn226/speechbrain…
mhn226 Oct 23, 2023
462340b
Add test info IWSLT recipe
mhn226 Oct 23, 2023
1adbc73
add docstrings for S2STransformerBeamSearcher
mhn226 Oct 23, 2023
acfe5cc
Update IWSLT recipes
mhn226 Oct 23, 2023
9d31ecf
Update IWSLT recipes
mhn226 Oct 23, 2023
5e149fe
fix doctest
mhn226 Oct 23, 2023
a9fa0d9
add requirements
mhn226 Oct 23, 2023
da3fba3
add protobuf
mhn226 Oct 23, 2023
886f9e4
fix doctest
mhn226 Oct 23, 2023
3ebdd76
small fixes
mravanelli Oct 26, 2023
76d44a9
Add protobuf install
mhn226 Oct 30, 2023
1d14263
Minor reform
mhn226 Oct 30, 2023
eb0489d
Remove protobuf
mhn226 Oct 30, 2023
0119545
Fix docstings
mhn226 Oct 30, 2023
93ffdf8
Fix docstrings
mhn226 Oct 30, 2023
b3012a3
minor fix
mhn226 Oct 30, 2023
dc74e29
minor reform
mhn226 Oct 30, 2023
48a6316
remove labse
mhn226 Oct 30, 2023
2bba87e
Add attention pooling
mhn226 Oct 30, 2023
3c9660a
Add labse
mhn226 Oct 30, 2023
7a4d68f
Add info about SAMU
mhn226 Oct 31, 2023
45e531c
add iwslt recipes with samu
mhn226 Nov 1, 2023
c390f92
fix recipe test
mhn226 Nov 2, 2023
c104afe
fix comments
mhn226 Nov 2, 2023
62e6df3
fix recipe test
mhn226 Nov 2, 2023
46e43e8
Merge branch 'unstable-v0.6' into samu
mhn226 Nov 3, 2023
1d1aeaf
change recipe structure
mhn226 Nov 3, 2023
b770551
fix test recipe
mhn226 Nov 3, 2023
aaf6ca3
Add new recipes
mhn226 Nov 6, 2023
7c31da0
Merge branch 'unstable-v0.6' into samu
mhn226 Nov 6, 2023
5babf01
minor doctest change
mhn226 Nov 6, 2023
84562bd
minor doctest change
mhn226 Nov 6, 2023
739a6d8
small changes
mravanelli Nov 6, 2023
2385a9f
add dropbox links
mravanelli Nov 6, 2023
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 24 additions & 0 deletions recipes/IWSLT22_lowresource/AST/transformer/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -48,13 +48,37 @@ python train_with_w2v_mbart.py hparams/train_w2v2_mbart_st.yaml --root_data_fold

One should change hparams/train_w2v2_mbart_st.yaml to hparams/train_w2v2_nllb_st.yaml in the above training command for using NLLB model instead.

## Pre-training Semantically-Aligned Multimodal Utterance-level (SAMU) wav2vec

Inspired by [SAMU-XLSR](https://arxiv.org/abs/2205.08180), a model that unifies speech and text modality for making the pre-trained speech foundation model more semantically aware, we introduce here a recipe for fine-tuning a pre-trained wav2vec 2.0 model in the same manner. Training data can be paired speech/text data of the kind used by ASR or AST. In this recipe, we use directly the IWSLT2022_Tamasheq_data AST data.

For launching SAMU training:
```
python train_samu.py hparams/train_samu.yaml --root_data_folder=your/data/path # e.g., /workspace/speechbrain/recipes/IWSLT22_lowresource/IWSLT2022_Tamasheq_data/taq_fra_clean
```

After the SAMU model is pre-trained, one can use it in the same manner as wav2vec 2.0 model. We found that using SAMU model as speech encoder coupled with a decoder from mBART or NLLB helps further improve BLEU scores on this challenging dataset.

For launching AST training:
```
train_with_samu_mbart.py hparams/train_samu_mbart_st.yaml --root_data_folder=your/data/path --pre_trained_samu=your/samu/ckpt
```

Examples of the two parameters:
--root_data_folder=/workspace/speechbrain/recipes/IWSLT22_lowresource/IWSLT2022_Tamasheq_data/taq_fra_clean
--pre_trained_samu=/workspace/speechbrain/recipes/IWSLT22_lowresource/results/samu_pretraining/7777/save/CKPT+checkpoint_epoch100/wav2vec2.ckpt

One should change hparams/train_samu_mbart_st.yaml to hparams/train_samu_nllb_st.yaml in the above training command for using NLLB model instead.

# Results

| No. | hyperparams file | dev BLEU | test BLEU | Model Link |
| --- |:----------------:|:---------:|:--------:|:--------:|
| 1 | train_w2v2_st.yaml | 7.63 | 5.38 | Not avail. | Not avail. |
| 2 | train_w2v2_mbart_st.yaml | 9.62 | 7.73 | [DropBox](https://www.dropbox.com/sh/xjo0ou739oksnus/AAAgyrCwywmDRRuUiDnUva2za?dl=0) |
| 3 | train_w2v2_nllb_st.yaml | 11.09 | 8.70 | [DropBox](https://www.dropbox.com/sh/spp2ijgfdbzuz26/AABkJ97e72D7aKzNLTm1qmWEa?dl=0) |
| 4 | train_samu_mbart_st.yaml | 13.41 | 10.28 | [DropBox](https://www.dropbox.com/sh/98s1xyc3chreaw6/AABom3FnwY5SsIvg4en9tWC2a?dl=0) |
| 5 | train_samu_nllb_st.yaml | 13.89 | 11.32 | [DropBox](https://www.dropbox.com/sh/ekkpl9c3kxsgllj/AABa0q2LrJe_o7JF-TTbfxZ-a?dl=0) |

## Citation
```
Expand Down
121 changes: 121 additions & 0 deletions recipes/IWSLT22_lowresource/AST/transformer/hparams/train_samu.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,121 @@
# ############################################################################
# Model: SAMU model
# losses: cosine similarity
# Training: Tamasheq-French corpus
# Author: Ha Nguyen, 2023
# ############################################################################

# Seed needs to be set at top of yaml, before objects with parameters are made
seed: 7777
__set_seed: !!python/object/apply:torch.manual_seed [!ref <seed>]
debug: False
output_folder: !ref results/samu_pretraining/<seed>
save_folder: !ref <output_folder>/save
train_log: !ref <output_folder>/train_log.txt
wer_file: !ref <output_folder>/wer.txt

# root data folder points to 17h version inside the github folder (IWSLT2022_Tamasheq_data/taq_fra_clean/)
root_data_folder: !PLACEHOLDER # e.g., /users/hnguyen/IWSLT2022_Tamasheq_data/taq_fra_clean
# data folder is the place where the json files will be stored prior to training
data_folder: !ref <root_data_folder>/json_version/
# Data files
train_set: !ref <data_folder>/train.json
valid_set: !ref <data_folder>/valid.json
test_set: !ref <data_folder>/test.json
skip_prep: False

# URL for the HuggingFace model we want to load (BASE here)
wav2vec2_hub: LIA-AvignonUniversity/IWSLT2022-tamasheq-only

# wav2vec 2.0 specific parameters
wav2vec2_frozen: False

# Training parameters
number_of_epochs: 100
lr: 0.001
lr_wav2vec: 0.00001
lr_labse: 0.00001
sorting: ascending
batch_size: 2
test_batch_size: 1
ckpt_interval_minutes: 15 # save checkpoint every N min

epoch_counter: !new:speechbrain.utils.epoch_loop.EpochCounter
limit: !ref <number_of_epochs>

dataloader_options:
batch_size: !ref <batch_size>
num_workers: 4

test_dataloader_options:
batch_size: !ref <test_batch_size>
num_workers: 4

# Transformer
d_model: 768
loss_scale: 50

wav2vec2: !new:speechbrain.lobes.models.huggingface_transformers.wav2vec2.Wav2Vec2
source: !ref <wav2vec2_hub>
output_norm: False
freeze: !ref <wav2vec2_frozen>
save_path: !ref <save_folder>/wav2vec2_checkpoint

attn_pooling: !new:speechbrain.nnet.pooling.AttentionPooling
input_dim: !ref <d_model>

#LaBSE
labse_path: setu4993/LaBSE
labse_frozen: True
LaBSE: !new:speechbrain.lobes.models.huggingface_transformers.labse.LaBSE
source: !ref <labse_path>
freeze: !ref <labse_frozen>
output_norm: True
save_path: !ref <save_folder>/labse_checkpoint

modules:
wav2vec2: !ref <wav2vec2>
attn_pooling: !ref <attn_pooling>
LaBSE: !ref <LaBSE>

model: !new:torch.nn.ModuleList
- [!ref <attn_pooling>, !ref <attn_pooling>]

adam_opt_class: !name:torch.optim.Adam
lr: !ref <lr>

wav2vec_opt_class: !name:torch.optim.Adam
lr: !ref <lr_wav2vec>

labse_opt_class: !name:torch.optim.Adam
lr: !ref <lr_labse>

lr_annealing_adam: !new:speechbrain.nnet.schedulers.NewBobScheduler
initial_value: !ref <lr>
improvement_threshold: 0.0025
annealing_factor: 0.5
patient: 2

lr_annealing_wav2vec: !new:speechbrain.nnet.schedulers.NewBobScheduler
initial_value: !ref <lr_wav2vec>
improvement_threshold: 0.0025
annealing_factor: 0.9

lr_annealing_labse: !new:speechbrain.nnet.schedulers.NewBobScheduler
initial_value: !ref <lr_labse>
improvement_threshold: 0.0025
annealing_factor: 0.9

checkpointer: !new:speechbrain.utils.checkpoints.Checkpointer
checkpoints_dir: !ref <save_folder>
recoverables:
model: !ref <model>
wav2vec2: !ref <wav2vec2>
LaBSE: !ref <LaBSE>
lr_annealing_adam: !ref <lr_annealing_adam>
lr_annealing_wav2vec: !ref <lr_annealing_wav2vec>
lr_annealing_labse: !ref <lr_annealing_labse>
counter: !ref <epoch_counter>

train_logger: !new:speechbrain.utils.train_logger.FileTrainLogger
save_file: !ref <train_log>
Original file line number Diff line number Diff line change
@@ -0,0 +1,198 @@
# ############################################################################
# Model: E2E ST with SAMU encoder and mBART decoder
# Encoder: SAMU
# Decoder: mBART decoder
# losses: NLL
# Training: Tamasheq-French corpus
# Author: Ha Nguyen, 2023
# ############################################################################

# Seed needs to be set at top of yaml, before objects with parameters are made
seed: 1337 #7777
__set_seed: !!python/object/apply:torch.manual_seed [!ref <seed>]
debug: False
output_folder: !ref results/samu_mbart/<seed>
save_folder: !ref <output_folder>/save
train_log: !ref <output_folder>/train_log.txt
wer_file: !ref <output_folder>/wer.txt
bleu_file: !ref <output_folder>/bleu.txt

# root data folder points to 17h version inside the github folder (IWSLT2022_Tamasheq_data/taq_fra_clean/)
root_data_folder: !PLACEHOLDER # e.g., /users/hnguyen/IWSLT2022_Tamasheq_data/taq_fra_clean
# data folder is the place where the json files will be stored prior to training
data_folder: !ref <root_data_folder>/json_version/
lang: "fr" #for the BLEU score detokenization
target_lang: "fr_XX" # for mbart initialization

annotation_train: !ref <data_folder>/train.json
annotation_valid: !ref <data_folder>/valid.json
annotation_test: !ref <data_folder>/test.json
skip_prep: False

# URL for the HuggingFace model we want to load (BASE here)
wav2vec2_hub: LIA-AvignonUniversity/IWSLT2022-tamasheq-only
wav2vec2_folder: !ref <save_folder>/wav2vec2_checkpoint

# wav2vec 2.0 specific parameters
wav2vec2_frozen: False

# Training parameters
number_of_epochs: 500
lr: 0.001
lr_wav2vec: 0.0001
lr_mbart: 0.0001
batch_size: 2
test_batch_size: 1
gradient_accumulation: 6
valid_search_interval: 4
loss_reduction: batchmean
ckpt_interval_minutes: 15 # save checkpoint every N min

# Data sorting parameters: sorting_debug_duration replaces sorting_min_duration in debug mode
sorting: ascending

epoch_counter: !new:speechbrain.utils.epoch_loop.EpochCounter
limit: !ref <number_of_epochs>

dataloader_options:
batch_size: !ref <batch_size>
num_workers: 4

test_dataloader_options:
batch_size: !ref <test_batch_size>
num_workers: 4

# Feature parameters (W2V2 etc)
features_dim: 768 # base wav2vec output dimension, for large replace by 1024

#projection for w2v
enc_dnn_layers: 1
enc_dnn_neurons: 1024 #256

# Transformer
activation: !name:torch.nn.GELU

# Outputs
label_smoothing: 0.1
pad_index: 1 # pad_index defined by mbart model
bos_index: 250008 # fr_XX bos_index defined by mbart model
eos_index: 2

# Decoding parameters
# Be sure that the bos and eos index match with the BPEs ones
min_decode_ratio: 0.0
max_decode_ratio: 0.25
valid_beam_size: 5

############################## models ################################
#wav2vec model
wav2vec2: !new:speechbrain.lobes.models.huggingface_transformers.wav2vec2.Wav2Vec2
source: !ref <wav2vec2_hub>
output_norm: True
freeze: !ref <wav2vec2_frozen>
save_path: !ref <wav2vec2_folder>

#linear projection
enc: !new:speechbrain.lobes.models.VanillaNN.VanillaNN
input_shape: [null, null, !ref <features_dim>]
activation: !ref <activation>
dnn_blocks: !ref <enc_dnn_layers>
dnn_neurons: !ref <enc_dnn_neurons>

#mBART
mbart_path: facebook/mbart-large-50-many-to-many-mmt
mbart_frozen: False
vocab_size: 250054
mBART: !new:speechbrain.lobes.models.huggingface_transformers.mbart.mBART
source: !ref <mbart_path>
freeze: !ref <mbart_frozen>
save_path: !ref <save_folder>/mbart_checkpoint
target_lang: !ref <target_lang>

log_softmax: !new:speechbrain.nnet.activations.Softmax
apply_log: True

modules:
wav2vec2: !ref <wav2vec2>
enc: !ref <enc>
mBART: !ref <mBART>

model: !new:torch.nn.ModuleList
- [!ref <enc>]

adam_opt_class: !name:torch.optim.Adam
lr: !ref <lr>

wav2vec_opt_class: !name:torch.optim.Adam
lr: !ref <lr_wav2vec>

mbart_opt_class: !name:torch.optim.Adam
lr: !ref <lr_mbart>

seq_cost: !name:speechbrain.nnet.losses.nll_loss
label_smoothing: !ref <label_smoothing>
reduction: !ref <loss_reduction>

lr_annealing_adam: !new:speechbrain.nnet.schedulers.NewBobScheduler
initial_value: !ref <lr>
improvement_threshold: 0.0025
annealing_factor: 0.5
patient: 2

warmup: 8000
hold: 32000
cooldown: 40000
optimizer_step_limit: 80000

lr_annealing_wav2vec: !new:speechbrain.nnet.schedulers.TriStageLRSchedule
lr: !ref <lr_wav2vec>
warmup_steps: !ref <warmup>
hold_steps: !ref <hold>
decay_steps: !ref <cooldown>
total_steps: !ref <optimizer_step_limit>

lr_annealing_mbart: !new:speechbrain.nnet.schedulers.TriStageLRSchedule
lr: !ref <lr_mbart>
warmup_steps: !ref <warmup>
hold_steps: !ref <hold>
decay_steps: !ref <cooldown>
total_steps: !ref <optimizer_step_limit>

checkpointer: !new:speechbrain.utils.checkpoints.Checkpointer
checkpoints_dir: !ref <save_folder>
recoverables:
model: !ref <model>
wav2vec2: !ref <wav2vec2>
mBART: !ref <mBART>
lr_annealing_wav2vec: !ref <lr_annealing_wav2vec>
lr_annealing_mbart: !ref <lr_annealing_mbart>
counter: !ref <epoch_counter>

valid_search: !new:speechbrain.decoders.S2SHFTextBasedBeamSearcher
modules: [!ref <mBART>, null, null]
vocab_size: !ref <vocab_size>
bos_index: !ref <bos_index>
eos_index: !ref <eos_index>
min_decode_ratio: !ref <min_decode_ratio>
max_decode_ratio: !ref <max_decode_ratio>
beam_size: !ref <valid_beam_size>
using_eos_threshold: True
length_normalization: True

train_logger: !new:speechbrain.utils.train_logger.FileTrainLogger
save_file: !ref <train_log>

bleu_computer: !name:speechbrain.utils.bleu.BLEUStats
merge_words: False
lang: !ref <lang>

acc_computer: !name:speechbrain.utils.Accuracy.AccuracyStats

# Path to the samu checkpoint
pre_trained_samu: !PLACEHOLDER # e.g., /users/hnguyen/output_samu_pretraining/7777/save/CKPT+checkpoint_epoch100/wav2vec2.ckpt
pretrainer: !new:speechbrain.utils.parameter_transfer.Pretrainer
collect_in: !ref <save_folder>
loadables:
wav2vec: !ref <wav2vec2>
paths:
wav2vec: !ref <pre_trained_samu>
Loading