(
self,
input_ids,
attention_mask=None,
decoder_input_ids=None,
encoder_outputs: Optional[Tuple] = None,
decoder_attention_mask=None,
decoder_cached_states=None,
use_cache=None,
output_attentions=None,
output_hidden_states=None,
)
| 820 | @add_start_docstrings_to_callable(BART_INPUTS_DOCSTRING) |
| 821 | @add_code_sample_docstrings(tokenizer_class=_TOKENIZER_FOR_DOC, checkpoint="facebook/bart-large") |
| 822 | def forward( |
| 823 | self, |
| 824 | input_ids, |
| 825 | attention_mask=None, |
| 826 | decoder_input_ids=None, |
| 827 | encoder_outputs: Optional[Tuple] = None, |
| 828 | decoder_attention_mask=None, |
| 829 | decoder_cached_states=None, |
| 830 | use_cache=None, |
| 831 | output_attentions=None, |
| 832 | output_hidden_states=None, |
| 833 | ): |
| 834 | |
| 835 | if decoder_input_ids is None: |
| 836 | use_cache = False |
| 837 | |
| 838 | output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions |
| 839 | output_hidden_states = ( |
| 840 | output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states |
| 841 | ) |
| 842 | use_cache = use_cache if use_cache is not None else self.config.use_cache |
| 843 | |
| 844 | # make masks if user doesn't supply |
| 845 | if not use_cache: |
| 846 | decoder_input_ids, decoder_padding_mask, causal_mask = _prepare_bart_decoder_inputs( |
| 847 | self.config, |
| 848 | input_ids, |
| 849 | decoder_input_ids=decoder_input_ids, |
| 850 | decoder_padding_mask=decoder_attention_mask, |
| 851 | causal_mask_dtype=self.shared.weight.dtype, |
| 852 | ) |
| 853 | else: |
| 854 | decoder_padding_mask, causal_mask = None, None |
| 855 | |
| 856 | assert decoder_input_ids is not None |
| 857 | |
| 858 | if encoder_outputs is None: |
| 859 | encoder_outputs = self.encoder( |
| 860 | input_ids=input_ids, |
| 861 | attention_mask=attention_mask, |
| 862 | output_attentions=output_attentions, |
| 863 | output_hidden_states=output_hidden_states, |
| 864 | ) |
| 865 | assert isinstance(encoder_outputs, tuple) |
| 866 | # decoder outputs consists of (dec_features, layer_state, dec_hidden, dec_attn) |
| 867 | decoder_outputs = self.decoder( |
| 868 | decoder_input_ids, |
| 869 | encoder_outputs[0], |
| 870 | attention_mask, |
| 871 | decoder_padding_mask, |
| 872 | decoder_causal_mask=causal_mask, |
| 873 | decoder_cached_states=decoder_cached_states, |
| 874 | output_attentions=output_attentions, |
| 875 | output_hidden_states=output_hidden_states, |
| 876 | use_cache=use_cache, |
| 877 | ) |
| 878 | |
| 879 | # Attention and hidden_states will be [] or None if they aren't needed |
no test coverage detected