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

Class BartModel

src/transformers/modeling_bart.py:808–894  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

806 "The bare BART Model outputting raw hidden-states without any specific head on top.", BART_START_DOCSTRING,
807)
808class 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)

Callers 6

convert_bart_checkpointFunction · 0.90
test_advanced_inputsMethod · 0.90
__init__Method · 0.85
__init__Method · 0.85
__init__Method · 0.85

Calls

no outgoing calls

Tested by 2

test_advanced_inputsMethod · 0.72