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

Function get_redpajama

datautils_block.py:117–154  ·  view source on GitHub ↗
(tokenizer, train_size, val_size, seed, seqlen)

Source from the content-addressed store, hash-verified

115 return trainloader, valloader
116
117def get_redpajama(tokenizer, train_size, val_size, seed, seqlen):
118 print("get_redpajama")
119 try:
120 loacal_dataset = "/cpfs01/user/chenmengzhao/huggingface/datasets/togethercomputer___red_pajama-data-1_t-sample"
121 traindata = load_dataset(loacal_dataset,split='train')
122 except:
123 traindata = load_dataset("togethercomputer/RedPajama-Data-1T-Sample",split='train')
124 random.seed(seed)
125 traindata = traindata.shuffle(seed=seed)
126 trainloader = []
127 val_sample_ratio = 0.9
128 for _ in range(train_size):
129 while True:
130 i = random.randint(0, int(len(traindata)*val_sample_ratio) - 1)
131 trainenc = tokenizer(traindata[i]['text'], return_tensors='pt')
132 if trainenc.input_ids.shape[1] >= seqlen+1:
133 break
134 i = random.randint(0, trainenc.input_ids.shape[1] - seqlen - 1)
135 j = i + seqlen
136 inp = trainenc.input_ids[:, i:j]
137 tar = inp.clone()
138 tar[:, :-1] = -100
139 trainloader.append((inp, tar))
140
141 valloader = []
142 for _ in range(val_size):
143 while True:
144 i = random.randint(int(len(traindata)*val_sample_ratio),len(traindata)-1)
145 trainenc = tokenizer(traindata[i]['text'], return_tensors='pt')
146 if trainenc.input_ids.shape[1] >= seqlen+1:
147 break
148 i = random.randint(0, trainenc.input_ids.shape[1] - seqlen - 1)
149 j = i + seqlen
150 inp = trainenc.input_ids[:, i:j]
151 tar = inp.clone()
152 tar[:, :-1] = -100
153 valloader.append((inp, tar))
154 return trainloader, valloader
155
156
157

Callers 1

get_loadersFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected