| 273 | |
| 274 | |
| 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, |
| 304 | input_lengths=None, |
| 305 | position_ids=None, |
| 306 | token_type_ids=None, |
| 307 | hidden_states=None, |
| 308 | max_input_length=None): |
| 309 | # remove_input_padding requires these fields as explicit input |
| 310 | mask = None |
| 311 | if not default_net().plugin_config.remove_input_padding: |
| 312 | seq_len_2d = concat([1, shape(input_ids, 1)]) |
| 313 | |
| 314 | # create position ids |
| 315 | position_ids_buffer = constant( |
| 316 | np.expand_dims( |
| 317 | np.arange(self.max_position_embeddings).astype(np.int32), |
| 318 | 0)) |
| 319 | tmp_position_ids = slice(position_ids_buffer, |
| 320 | starts=[0, 0], |
| 321 | sizes=seq_len_2d) |
| 322 | tmp_position_ids = expand(tmp_position_ids, shape(input_ids)) #BxL |
| 323 | tmp_input_lengths = unsqueeze(input_lengths, 1) #Bx1 |
| 324 | tmp_input_lengths = expand(tmp_input_lengths, |
| 325 | shape(input_ids)) #BxL |
| 326 | mask = tmp_position_ids < tmp_input_lengths # BxL |
| 327 | mask = mask.cast('int32') |
| 328 | |
| 329 | if position_ids is None: |
| 330 | if self.is_roberta: |
| 331 | # see create_position_ids_from_input_ids() in https://github.com/huggingface/transformers/blob/main/src/transformers/models/roberta/modeling_roberta.py |
| 332 | position_ids = (tmp_position_ids + 1) * mask |