(self, config: BERTConfig)
| 275 | class BertModel(BertBase): |
| 276 | |
| 277 | def __init__(self, config: BERTConfig): |
| 278 | super().__init__(config) |
| 279 | |
| 280 | self.config = config |
| 281 | self.max_position_embeddings = config.max_position_embeddings |
| 282 | self.padding_idx = config.pad_token_id |
| 283 | self.is_roberta = config.is_roberta |
| 284 | self.embedding = BertEmbedding( |
| 285 | vocab_size=config.vocab_size, |
| 286 | hidden_size=config.hidden_size, |
| 287 | max_position_embeddings=config.max_position_embeddings, |
| 288 | type_vocab_size=config.type_vocab_size, |
| 289 | dtype=config.dtype) |
| 290 | |
| 291 | self.layers = ModuleList([ |
| 292 | BertEncoderLayer( |
| 293 | hidden_size=config.hidden_size, |
| 294 | num_attention_heads=config.num_attention_heads, |
| 295 | max_position_embeddings=config.max_position_embeddings, |
| 296 | hidden_act=config.hidden_act, |
| 297 | tp_group=config.mapping.tp_group, |
| 298 | tp_size=config.mapping.tp_size, |
| 299 | dtype=config.dtype) for _ in range(config.num_hidden_layers) |
| 300 | ]) |
| 301 | |
| 302 | def forward(self, |
| 303 | input_ids=None, |
nothing calls this directly
no test coverage detected