(token_ids)
| 49 | return extended_attention_mask |
| 50 | |
| 51 | def bert_position_ids(token_ids): |
| 52 | # Create position ids |
| 53 | seq_length = token_ids.size(1) |
| 54 | position_ids = torch.arange(seq_length, dtype=torch.long, |
| 55 | device=token_ids.device) |
| 56 | position_ids = position_ids.unsqueeze(0).expand_as(token_ids) |
| 57 | |
| 58 | return position_ids |
| 59 | |
| 60 | |
| 61 | class BertLMHead(MegatronModule): |