Get a DataLoader from a test dataset. This dataloader should be used when comparing WikiText2 perplexities with other papers, e.g. SparseGPT (arxiv.org/abs/2301.00774). Args: dataset: The dataset to create a dataloader from. tokenizer: The tokenizer to use. seqlen:
(
dataset: datasets.Dataset, tokenizer: PreTrainedTokenizerBase, seqlen: int = 2048, batch_size: int = 1
)
| 60 | |
| 61 | |
| 62 | def prepare_test_dataloader( |
| 63 | dataset: datasets.Dataset, tokenizer: PreTrainedTokenizerBase, seqlen: int = 2048, batch_size: int = 1 |
| 64 | ) -> DataLoader[dict[str, torch.Tensor]]: |
| 65 | """ |
| 66 | Get a DataLoader from a test dataset. This dataloader should be used when comparing WikiText2 perplexities with other papers, e.g. SparseGPT (arxiv.org/abs/2301.00774). |
| 67 | |
| 68 | Args: |
| 69 | dataset: The dataset to create a dataloader from. |
| 70 | tokenizer: The tokenizer to use. |
| 71 | seqlen: The sequence length of sequences in the dataset. |
| 72 | batch_size: The batch size. |
| 73 | |
| 74 | Returns: |
| 75 | A DataLoader. |
| 76 | """ |
| 77 | |
| 78 | logging.info(f"Preparing test dataloader") |
| 79 | |
| 80 | class TestDataset(Dataset): |
| 81 | def __init__(self, ds, tokenizer, seqlen=2048): |
| 82 | """Tokenize the entire dataset and reshape it into sequences of length seqlen.""" |
| 83 | |
| 84 | tokenized_ds = tokenizer("\n\n".join(ds['text']), return_tensors='pt') |
| 85 | nsamples = tokenized_ds.input_ids.numel() // seqlen |
| 86 | |
| 87 | input_ids = tokenized_ds.input_ids[0, : nsamples * seqlen] |
| 88 | input_ids = input_ids.reshape(nsamples, seqlen) |
| 89 | attn_mask = tokenized_ds.attention_mask[0, : nsamples * seqlen] |
| 90 | attn_mask = attn_mask.reshape(nsamples, seqlen) |
| 91 | |
| 92 | self.input_ids = input_ids |
| 93 | self.attn_mask = attn_mask |
| 94 | |
| 95 | def __getitem__(self, idx): |
| 96 | return {"input_ids": self.input_ids[idx], "attention_mask": self.attn_mask[idx]} |
| 97 | |
| 98 | def __len__(self): |
| 99 | return len(self.input_ids) |
| 100 | |
| 101 | test_ds = TestDataset(dataset, tokenizer, seqlen) |
| 102 | loader = DataLoader(test_ds, batch_size=batch_size) |
| 103 | logging.info(f"Preparing test dataloader done") |
| 104 | return loader |
| 105 | |
| 106 | |
| 107 | def prepare_dataloader( |