| 806 | "The bare BART Model outputting raw hidden-states without any specific head on top.", BART_START_DOCSTRING, |
| 807 | ) |
| 808 | class BartModel(PretrainedBartModel): |
| 809 | def __init__(self, config: BartConfig): |
| 810 | super().__init__(config) |
| 811 | |
| 812 | padding_idx, vocab_size = config.pad_token_id, config.vocab_size |
| 813 | self.shared = nn.Embedding(vocab_size, config.d_model, padding_idx) |
| 814 | |
| 815 | self.encoder = BartEncoder(config, self.shared) |
| 816 | self.decoder = BartDecoder(config, self.shared) |
| 817 | |
| 818 | self.init_weights() |
| 819 | |
| 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) |
no outgoing calls