MCPcopy Create free account
hub / github.com/OpenGVLab/EfficientQAT / get_wikitext2

Function get_wikitext2

datautils_block.py:13–42  ·  view source on GitHub ↗
(tokenizer, train_size, val_size, seed, seqlen, test_only)

Source from the content-addressed store, hash-verified

11import os
12
13def get_wikitext2(tokenizer, train_size, val_size, seed, seqlen, test_only):
14 print("get_wikitext2")
15 traindata = load_dataset('wikitext', 'wikitext-2-raw-v1', split='train')
16 testdata = load_dataset('wikitext', 'wikitext-2-raw-v1', split='test')
17
18 testenc = tokenizer("\n\n".join(testdata['text']), return_tensors='pt')
19 if test_only:
20 return testenc
21 trainenc = tokenizer("\n\n".join(traindata['text']), return_tensors='pt')
22
23
24 random.seed(seed)
25 trainloader = []
26 val_sample_ratio = 0.9 # sample train from [0:0.9] and val from [0.9:1.0] to avoid overlap
27 for _ in range(train_size):
28 i = random.randint(0, int(trainenc.input_ids.shape[1]*val_sample_ratio) - seqlen - 1)
29 j = i + seqlen
30 inp = trainenc.input_ids[:, i:j]
31 tar = inp.clone()
32 tar[:, :-1] = -100
33 trainloader.append((inp, tar))
34 valloader = []
35 for _ in range(val_size):
36 i = random.randint(int(trainenc.input_ids.shape[1]*val_sample_ratio) - seqlen - 1, trainenc.input_ids.shape[1] - seqlen - 1)
37 j = i + seqlen
38 inp = trainenc.input_ids[:, i:j]
39 tar = inp.clone()
40 tar[:, :-1] = -100
41 valloader.append((inp, tar))
42 return trainloader, valloader
43
44
45def get_c4(tokenizer, train_size, val_size, seed, seqlen, test_only):

Callers 1

get_loadersFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected