| 43 | |
| 44 | |
| 45 | def get_c4(tokenizer, train_size, val_size, seed, seqlen, test_only): |
| 46 | print("get_c4") |
| 47 | try: |
| 48 | # set local path for faster loading |
| 49 | traindata = load_dataset("arrow", |
| 50 | data_files={ |
| 51 | "train": "/cpfs01/user/chenmengzhao/huggingface/datasets/allenai___json/allenai--c4-6fbe877195f42de5/0.0.0/0f7e3662623656454fcd2b650f34e886a7db4b9104504885bd462096cc7a9f51/json-train-00000-of-00002.arrow", |
| 52 | "validation": "/cpfs01/user/chenmengzhao/huggingface/datasets/allenai___json/allenai--c4-efc3d4f4606f44bd/0.0.0/fe5dd6ea2639a6df622901539cb550cf8797e5a6b2dd7af1cf934bed8e233e6e/json-validation.arrow", |
| 53 | },split='train' |
| 54 | ) |
| 55 | valdata = load_dataset("arrow", |
| 56 | data_files={ |
| 57 | "validation": "/cpfs01/user/chenmengzhao/huggingface/datasets/allenai___json/allenai--c4-efc3d4f4606f44bd/0.0.0/fe5dd6ea2639a6df622901539cb550cf8797e5a6b2dd7af1cf934bed8e233e6e/json-validation.arrow", |
| 58 | },split='validation' |
| 59 | ) |
| 60 | except: |
| 61 | traindata = load_dataset( |
| 62 | 'allenai/c4', 'allenai--c4', data_files={'train': 'en/c4-train.00000-of-01024.json.gz'}, split='train' |
| 63 | ) |
| 64 | valdata = load_dataset( |
| 65 | 'allenai/c4', 'allenai--c4', data_files={'validation': 'en/c4-validation.00000-of-00008.json.gz'}, split='validation' |
| 66 | ) |
| 67 | |
| 68 | random.seed(0) |
| 69 | valenc = [] |
| 70 | for _ in range(256): |
| 71 | while True: |
| 72 | i = random.randint(0, len(valdata) - 1) |
| 73 | tmp = tokenizer(valdata[i]['text'], return_tensors='pt') |
| 74 | if tmp.input_ids.shape[1] >= seqlen: |
| 75 | break |
| 76 | i = random.randint(0, tmp.input_ids.shape[1] - seqlen - 1) |
| 77 | j = i + seqlen |
| 78 | valenc.append(tmp.input_ids[:, i:j]) |
| 79 | valenc = torch.hstack(valenc) |
| 80 | if test_only: |
| 81 | return valenc |
| 82 | |
| 83 | random.seed(seed) |
| 84 | trainloader = [] |
| 85 | val_sample_ratio = 0.9 # sample train from [0:0.9] and val from [0.9:1.0] to avoid overlap |
| 86 | for _ in range(train_size): |
| 87 | while True: |
| 88 | i = random.randint(0, int(len(traindata)*val_sample_ratio) - 1) |
| 89 | trainenc = tokenizer(traindata[i]['text'], return_tensors='pt') |
| 90 | if trainenc.input_ids.shape[1] >= seqlen+1: |
| 91 | break |
| 92 | i = random.randint(0, trainenc.input_ids.shape[1] - seqlen - 1) |
| 93 | j = i + seqlen |
| 94 | inp = trainenc.input_ids[:, i:j] |
| 95 | tar = inp.clone() |
| 96 | tar[:, :-1] = -100 |
| 97 | trainloader.append((inp, tar)) |
| 98 | |
| 99 | valloader = [] |
| 100 | for _ in range(val_size): |
| 101 | while True: |
| 102 | i = random.randint(int(len(traindata)*val_sample_ratio),len(traindata)-1) |