MCPcopy Create free account
hub / github.com/AnswerDotAI/ModernBERT / MaskedLMOutput

Class MaskedLMOutput

src/bert_layers/model.py:786–817  ·  view source on GitHub ↗

Base class for masked language models outputs. Args: loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided): Masked language modeling (MLM) loss. logits (`torch.FloatTensor` of shape `(batch_size, sequence_length, config.vocab

Source from the content-addressed store, hash-verified

784
785@dataclass
786class MaskedLMOutput(ModelOutput):
787 """
788 Base class for masked language models outputs.
789
790 Args:
791 loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided):
792 Masked language modeling (MLM) loss.
793 logits (`torch.FloatTensor` of shape `(batch_size, sequence_length, config.vocab_size)`):
794 Prediction scores of the language modeling head (scores for each vocabulary token before SoftMax).
795 hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
796 Tuple of `torch.FloatTensor` (one for the output of the embeddings, if the model has an embedding layer, +
797 one for the output of each layer) of shape `(batch_size, sequence_length, hidden_size)`.
798
799 Hidden-states of the model at the output of each layer plus the optional initial embedding outputs.
800 attentions (`tuple(torch.FloatTensor)`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):
801 Tuple of `torch.FloatTensor` (one for each layer) of shape `(batch_size, num_heads, sequence_length,
802 sequence_length)`.
803
804 Attentions weights after the attention softmax, used to compute the weighted average in the self-attention
805 heads.
806 """
807
808 loss: Optional[torch.FloatTensor] = None
809 logits: torch.FloatTensor = None
810 hidden_states: Optional[Tuple[torch.FloatTensor, ...]] = None
811 attentions: Optional[Tuple[torch.FloatTensor, ...]] = None
812 indices: Optional[torch.LongTensor] = None
813 cu_seqlens: Optional[torch.LongTensor] = None
814 max_seqlen: Optional[int] = None
815 batch_size: Optional[int] = None
816 seq_len: Optional[int] = None
817 labels: Optional[torch.LongTensor] = None
818
819
820@dataclass

Callers 2

forwardMethod · 0.85
forwardMethod · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected