(self, input_ids, position_ids, attention_mask, *,
output_hidden_states=False, **kw_args)
| 569 | self.final_layernorm = layernorm(hidden_size, eps=layernorm_epsilon) |
| 570 | |
| 571 | def forward(self, input_ids, position_ids, attention_mask, *, |
| 572 | output_hidden_states=False, **kw_args): |
| 573 | # sanity check |
| 574 | assert len(input_ids.shape) >= 2 |
| 575 | batch_size, query_length = input_ids.shape[:2] |
| 576 | |
| 577 | if attention_mask is None: |
| 578 | # Definition: None means full attention |
| 579 | attention_mask = torch.ones(1, 1, device=input_ids.device) |
| 580 | elif isinstance(attention_mask, int) and (attention_mask < 0): |
| 581 | # Definition: -1 means lower triangular attention mask |
| 582 | attention_mask = torch.ones(query_length, query_length, |
| 583 | device=input_ids.device).tril() |
| 584 | |
| 585 | attention_mask = attention_mask.type_as( |
| 586 | next(self.parameters()) |
| 587 | ) |
| 588 | assert len(attention_mask.shape) == 2 or \ |
| 589 | len(attention_mask.shape) == 4 and attention_mask.shape[1] == 1 |
| 590 | |
| 591 | # initial output_cross_layer might be generated by word/position_embedding_forward |
| 592 | output_cross_layer = {} |
| 593 | |
| 594 | # embedding part |
| 595 | if 'word_embedding_forward' in self.hooks: |
| 596 | hidden_states = self.hooks['word_embedding_forward'](input_ids, output_cross_layer=output_cross_layer, **kw_args) |
| 597 | else: # default |
| 598 | hidden_states = HOOKS_DEFAULT['word_embedding_forward'](self, input_ids, output_cross_layer=output_cross_layer,**kw_args) |
| 599 | |
| 600 | # handle position embedding |
| 601 | if 'position_embedding_forward' in self.hooks: |
| 602 | position_embeddings = self.hooks['position_embedding_forward'](position_ids, output_cross_layer=output_cross_layer, **kw_args) |
| 603 | else: |
| 604 | assert len(position_ids.shape) <= 2 |
| 605 | assert position_ids.shape[-1] == hidden_states.shape[1], (position_ids.shape, hidden_states.shape) |
| 606 | position_embeddings = HOOKS_DEFAULT['position_embedding_forward'](self, position_ids, output_cross_layer=output_cross_layer, **kw_args) |
| 607 | if position_embeddings is not None: |
| 608 | # import pdb; pdb.set_trace() # !! DEBUG |
| 609 | # hs: torch.Size([2, 88034, 1920]) |
| 610 | # pos: torch.Size([1, 23522, 1920]) |
| 611 | hidden_states = hidden_states + position_embeddings |
| 612 | |
| 613 | hidden_states = self.embedding_dropout(hidden_states) |
| 614 | |
| 615 | output_per_layers = [] |
| 616 | if self.checkpoint_activations: |
| 617 | # define custom_forward for checkpointing |
| 618 | def custom(start, end, kw_args_index, cross_layer_index): |
| 619 | def custom_forward(*inputs): |
| 620 | layers_ = self.layers[start:end] |
| 621 | x_, mask = inputs[0], inputs[1] |
| 622 | |
| 623 | # recover kw_args and output_cross_layer |
| 624 | flat_inputs = inputs[2:] |
| 625 | kw_args, output_cross_layer = {}, {} |
| 626 | for k, idx in kw_args_index.items(): |
| 627 | kw_args[k] = flat_inputs[idx] |
| 628 | for k, idx in cross_layer_index.items(): |
nothing calls this directly
no test coverage detected