(self, input_ids, past_key_values, padding_mask=None)
| 586 | return |
| 587 | |
| 588 | def get_masks(self, input_ids, past_key_values, padding_mask=None): |
| 589 | batch_size, seq_length = input_ids.shape |
| 590 | full_attention_mask = torch.ones(batch_size, seq_length, seq_length, device=input_ids.device) |
| 591 | full_attention_mask.tril_() |
| 592 | past_length = 0 |
| 593 | if past_key_values: |
| 594 | past_length = past_key_values[0][0].shape[0] |
| 595 | if past_length: |
| 596 | full_attention_mask = torch.cat((torch.ones(batch_size, seq_length, past_length, |
| 597 | device=input_ids.device), full_attention_mask), dim=-1) |
| 598 | if padding_mask is not None: |
| 599 | full_attention_mask = full_attention_mask * padding_mask.unsqueeze(1) |
| 600 | if not past_length and padding_mask is not None: |
| 601 | full_attention_mask -= padding_mask.unsqueeze(-1) - 1 |
| 602 | full_attention_mask = (full_attention_mask < 0.5).bool() |
| 603 | full_attention_mask.unsqueeze_(1) |
| 604 | return full_attention_mask |
| 605 | |
| 606 | def get_position_ids(self, input_ids, device): |
| 607 | batch_size, seq_length = input_ids.shape |
nothing calls this directly
no outgoing calls
no test coverage detected