Skip to content

fix(transformer): enforce max_len in make_transformer_src_tgt_masks - #3086

Open
FinalSunFlower wants to merge 1 commit into
speechbrain:developfrom
FinalSunFlower:fix/transformer-mask-size
Open

FinalSunFlower wants to merge 1 commit into
speechbrain:developfrom
FinalSunFlower:fix/transformer-mask-size

Conversation

@FinalSunFlower

Copy link
Copy Markdown

Resolves #2344.

Problem
In make_transformer_src_tgt_masks, length_to_mask(abs_len) infers its sequence width from abs_len.max(). When a batch is padded to a fixed sequence length or aligned to a hardware boundary where wav_len.max() < 1.0, the resulting key-padding mask width is strictly smaller than src.shape[1], causing shape mismatch errors in downstream multi-head attention layers.

Solution
Explicitly pass max_len=src.shape[1] to length_to_mask when constructing src_key_padding_mask.

Added unit regression test verifying that mask dimensions remain aligned to the full padded sequence width when relative lengths are strictly less than 1.0. The test also executes MultiheadAttention with the generated mask.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

Incorrect transformer mask size

1 participant