MCPcopy Create free account
hub / github.com/huggingface/transformers / _prepare_bart_decoder_inputs

Function _prepare_bart_decoder_inputs

src/transformers/modeling_bart.py:115–133  ·  view source on GitHub ↗

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
)

Source from the content-addressed store, hash-verified

113
114
115def _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
136class PretrainedBartModel(PreTrainedModel):

Callers 3

test_advanced_inputsMethod · 0.90
forwardMethod · 0.85

Calls 5

shift_tokens_rightFunction · 0.85
make_padding_maskFunction · 0.85
invert_maskFunction · 0.85
fill_with_neg_infFunction · 0.85
toMethod · 0.80

Tested by 2

test_advanced_inputsMethod · 0.72