(tokenizer, num_samples=8192)
| 30 | } |
| 31 | |
| 32 | def create_variable_short_dataset(tokenizer, num_samples=8192): |
| 33 | torch.manual_seed(42) |
| 34 | torch.cuda.manual_seed_all(42) |
| 35 | np.random.seed(42) |
| 36 | random.seed(42) |
| 37 | lengths = torch.normal(mean=256, std=64, size=(num_samples,)).int().clamp(16, 512) |
| 38 | tokens_list = [] |
| 39 | masks_list = [] |
| 40 | for length in lengths: |
| 41 | tokens = torch.randint(100, 16000, (length.item(),)) |
| 42 | mask = torch.ones(length.item()) |
| 43 | padded_tokens = torch.full((512,), tokenizer.pad_token_id, dtype=torch.long) |
| 44 | padded_mask = torch.zeros(512) |
| 45 | padded_tokens[:length] = tokens |
| 46 | padded_mask[:length] = mask |
| 47 | tokens_list.append(padded_tokens) |
| 48 | masks_list.append(padded_mask) |
| 49 | |
| 50 | return { |
| 51 | 'input_ids': torch.stack(tokens_list), |
| 52 | 'attention_mask': torch.stack(masks_list) |
| 53 | } |
| 54 | |
| 55 | def create_variable_long_dataset(tokenizer, num_samples=8192): |
| 56 | torch.manual_seed(42) |
no outgoing calls
no test coverage detected