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

Function get_c4

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

Source from the content-addressed store, hash-verified

43
44
45def 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)

Callers 1

get_loadersFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected