Load train/val/test dataset from shuffled TFRecords
(args)
| 209 | |
| 210 | |
| 211 | def 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 | |
| 248 | def make_loaders(args, tokenizer): |