Spaces:
Sleeping
Sleeping
| import torch | |
| from transformers import PreTrainedTokenizerBase | |
| def fix_pad_token(tokenizer: PreTrainedTokenizerBase) -> None: | |
| if tokenizer.pad_token_id is None: | |
| tokenizer.pad_token = tokenizer.eos_token | |
| def find_attention_start_and_end_positions( | |
| attention_mask: torch.Tensor, | |
| ) -> tuple[torch.Tensor, torch.Tensor]: | |
| # convoluted, and generated by ChatGPT, but seems to work | |
| indices = torch.arange(attention_mask.size(1)).to(attention_mask.device) | |
| # use broadcasting to expand indices to the shape of attention_mask | |
| indices = indices[None, :].expand_as(attention_mask) | |
| # set indices where attention_mask is 0 to -1 | |
| indices = torch.where(attention_mask == 1, indices, -1) | |
| max_indices = indices.max(dim=1).values | |
| # set indices where attention_mask is 0 to a large number | |
| # TODO: there's probably a more efficient way to get the min and the max in one go | |
| indices = torch.where(attention_mask == 1, indices, attention_mask.size(1) + 1) | |
| min_indices = indices.min(dim=1).values | |
| return min_indices, max_indices | |
| def adjust_read_indices_for_padding( | |
| read_indices: torch.Tensor, attention_mask: torch.Tensor | |
| ) -> torch.Tensor: | |
| """ | |
| Adjust read indices to account for padding in the input | |
| """ | |
| start_positions, end_positions = find_attention_start_and_end_positions( | |
| attention_mask | |
| ) | |
| tokenized_lengths = end_positions - start_positions + 1 | |
| # turn negative indices into positive indices | |
| fixed_read_indices = torch.where( | |
| read_indices < 0, read_indices + tokenized_lengths, read_indices | |
| ) | |
| return fixed_read_indices + start_positions | |