| 150 | """ |
| 151 | |
| 152 | def __init__(self, config): |
| 153 | super().__init__() |
| 154 | self.word_embeddings = nn.Embedding(config.vocab_size, config.hidden_size, padding_idx=config.pad_token_id) |
| 155 | self.position_embeddings = nn.Embedding(config.max_position_embeddings, config.hidden_size) |
| 156 | self.token_type_embeddings = nn.Embedding(config.type_vocab_size, config.hidden_size) |
| 157 | |
| 158 | # self.LayerNorm is not snake-cased to stick with TensorFlow model variable name and be able to load |
| 159 | # any TensorFlow checkpoint file |
| 160 | self.LayerNorm = BertLayerNorm(config.hidden_size, eps=config.layer_norm_eps) |
| 161 | self.dropout = nn.Dropout(config.hidden_dropout_prob) |
| 162 | |
| 163 | def forward(self, input_ids=None, token_type_ids=None, position_ids=None, inputs_embeds=None): |
| 164 | if input_ids is not None: |