| 11 | import os |
| 12 | |
| 13 | def 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 | |
| 45 | def get_c4(tokenizer, train_size, val_size, seed, seqlen, test_only): |