Prepare masks that ignore padding tokens in the decoder and a causal mask for the decoder if none are provided. This mimics the default behavior in fairseq. To override it pass in masks. Note: this is not called during generation
(
config, input_ids, decoder_input_ids=None, decoder_padding_mask=None, causal_mask_dtype=torch.float32
)
| 113 | |
| 114 | |
| 115 | def _prepare_bart_decoder_inputs( |
| 116 | config, input_ids, decoder_input_ids=None, decoder_padding_mask=None, causal_mask_dtype=torch.float32 |
| 117 | ): |
| 118 | """Prepare masks that ignore padding tokens in the decoder and a causal mask for the decoder if |
| 119 | none are provided. This mimics the default behavior in fairseq. To override it pass in masks. |
| 120 | Note: this is not called during generation |
| 121 | """ |
| 122 | pad_token_id = config.pad_token_id |
| 123 | if decoder_input_ids is None: |
| 124 | decoder_input_ids = shift_tokens_right(input_ids, pad_token_id) |
| 125 | bsz, tgt_len = decoder_input_ids.size() |
| 126 | if decoder_padding_mask is None: |
| 127 | decoder_padding_mask = make_padding_mask(decoder_input_ids, pad_token_id) |
| 128 | else: |
| 129 | decoder_padding_mask = invert_mask(decoder_padding_mask) |
| 130 | causal_mask = torch.triu(fill_with_neg_inf(torch.zeros(tgt_len, tgt_len)), 1).to( |
| 131 | dtype=causal_mask_dtype, device=decoder_input_ids.device |
| 132 | ) |
| 133 | return decoder_input_ids, decoder_padding_mask, causal_mask |
| 134 | |
| 135 | |
| 136 | class PretrainedBartModel(PreTrainedModel): |