Buckets:
| import torch | |
| from torch.utils.data import Dataset | |
| from torch.nn.utils.rnn import pad_sequence | |
| from typing import List, Tuple | |
| from data.indirect_idx.tokenizer import CharacterTokenizer | |
| class IndirectIdxDataset(Dataset): | |
| def __init__(self, data: List[str], tokenizer: CharacterTokenizer): | |
| self.data = data | |
| self.tokenizer = tokenizer | |
| self.tokenized_data = [] | |
| # tokenize all data once | |
| for sample in data: | |
| parts = sample.split(', ') | |
| source_str, source_char, shift_str, target_char = parts | |
| delimiter = ", " | |
| input_str = delimiter.join([source_str, source_char, shift_str]) | |
| input_tokens = self.tokenizer.encode_input(input_str) | |
| target_token = self.tokenizer.encode_target(target_char.replace('\n', '')) | |
| self.tokenized_data.append((input_tokens, target_token)) | |
| def __len__(self) -> int: | |
| return len(self.tokenized_data) | |
| def __getitem__(self, idx: int) -> Tuple[List[int], int]: | |
| return self.tokenized_data[idx] | |
| def collate_fn(batch: List[Tuple[List[int], int]], pad_idx: int = 0): | |
| inputs, targets = zip(*batch) | |
| input_tensors = [torch.tensor(seq, dtype=torch.long) for seq in inputs] | |
| padded_inputs = pad_sequence(input_tensors, batch_first=True, padding_value=pad_idx) | |
| target_tensor = torch.full_like(padded_inputs, fill_value=-1, dtype=torch.long) | |
| target_masks = [] | |
| for i, (padded_input, input_tensor) in enumerate(zip(padded_inputs, input_tensors)): | |
| mask = torch.zeros_like(padded_input, dtype=torch.bool) | |
| indices = torch.where(padded_input == input_tensor[-1])[0] | |
| last_idx = indices[-1] | |
| mask[last_idx] = True | |
| target_tensor[i, last_idx] = targets[i] | |
| target_masks.append(mask) | |
| # attention_mask = (padded_inputs != pad_idx).long() | |
| target_masks = torch.stack(target_masks, dim=0) | |
| return padded_inputs, target_tensor, target_masks | |
Xet Storage Details
- Size:
- 1.97 kB
- Xet hash:
- 5739c39b1d73bcf8fe30dce8e7e22e806b25fdbabec5f62b3ae35611c8ae0706
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.