(tokenizer, num_samples=8192)
| 53 | } |
| 54 | |
| 55 | def create_variable_long_dataset(tokenizer, num_samples=8192): |
| 56 | torch.manual_seed(42) |
| 57 | torch.cuda.manual_seed_all(42) |
| 58 | np.random.seed(42) |
| 59 | random.seed(42) |
| 60 | lengths = torch.normal(mean=4096, std=1024, size=(num_samples,)).int().clamp(16, 8192) |
| 61 | tokens_list = [] |
| 62 | masks_list = [] |
| 63 | for length in lengths: |
| 64 | tokens = torch.randint(100, 16000, (length.item(),)) |
| 65 | mask = torch.ones(length.item()) |
| 66 | padded_tokens = torch.full((8192,), tokenizer.pad_token_id, dtype=torch.long) |
| 67 | padded_mask = torch.zeros(8192) |
| 68 | padded_tokens[:length] = tokens |
| 69 | padded_mask[:length] = mask |
| 70 | tokens_list.append(padded_tokens) |
| 71 | masks_list.append(padded_mask) |
| 72 | |
| 73 | return { |
| 74 | 'input_ids': torch.stack(tokens_list), |
| 75 | 'attention_mask': torch.stack(masks_list) |
| 76 | } |
| 77 | |
| 78 | def create_all_datasets(tokenizer, num_samples=8192): |
| 79 | return { |
no outgoing calls
no test coverage detected