MCPcopy Create free account
hub / github.com/THUDM/GLM / make_tfrecord_loaders

Function make_tfrecord_loaders

configure_data.py:211–245  ·  view source on GitHub ↗

Load train/val/test dataset from shuffled TFRecords

(args)

Source from the content-addressed store, hash-verified

209
210
211def make_tfrecord_loaders(args):
212 """Load train/val/test dataset from shuffled TFRecords"""
213
214 import data_utils.tf_dl
215 data_set_args = {'batch_size': args.batch_size,
216 'max_seq_len': args.seq_length,
217 'max_preds_per_seq': args.max_preds_per_seq,
218 'train': True,
219 'num_workers': max(args.num_workers, 1),
220 'seed': args.seed + args.rank + 1,
221 'threaded_dl': args.num_workers > 0
222 }
223 train = data_utils.tf_dl.TFRecordDataLoader(args.train_data,
224 **data_set_args)
225 data_set_args['train'] = False
226 if args.eval_seq_length is not None:
227 data_set_args['max_seq_len'] = args.eval_seq_length
228 if args.eval_max_preds_per_seq is not None:
229 data_set_args['max_preds_per_seq'] = args.eval_max_preds_per_seq
230 valid = None
231 if args.valid_data is not None:
232 valid = data_utils.tf_dl.TFRecordDataLoader(args.valid_data,
233 **data_set_args)
234 test = None
235 if args.test_data is not None:
236 test = data_utils.tf_dl.TFRecordDataLoader(args.test_data,
237 **data_set_args)
238 tokenizer = data_utils.make_tokenizer(args.tokenizer_type,
239 train,
240 args.tokenizer_path,
241 args.vocab_size,
242 args.tokenizer_model_type,
243 cache_dir=args.cache_dir)
244
245 return (train, valid, test), tokenizer
246
247
248def make_loaders(args, tokenizer):

Callers 1

make_loadersFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected