MCPcopy Create free account
hub / github.com/MuLabPKU/TransArch / prepare_test_dataloader

Function prepare_test_dataloader

CLOVER_ICML_2025/src/data.py:62–104  ·  view source on GitHub ↗

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
)

Source from the content-addressed store, hash-verified

60
61
62def 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
107def prepare_dataloader(

Callers 2

mainFunction · 0.90
mainFunction · 0.90

Calls 1

TestDatasetClass · 0.70

Tested by 1

mainFunction · 0.72