(args, preprocess_img, is_train, epoch=0, floor=False, tokenizer=None)
| 344 | |
| 345 | |
| 346 | def get_wds_dataset(args, preprocess_img, is_train, epoch=0, floor=False, tokenizer=None): |
| 347 | input_shards = args.train_data if is_train else args.val_data |
| 348 | assert input_shards is not None |
| 349 | resampled = getattr(args, 'dataset_resampled', False) and is_train |
| 350 | |
| 351 | num_shards = None |
| 352 | if is_train: |
| 353 | if args.train_num_samples is not None: |
| 354 | num_samples = args.train_num_samples |
| 355 | else: |
| 356 | num_samples, num_shards = get_dataset_size(input_shards) |
| 357 | if not num_samples: |
| 358 | raise RuntimeError( |
| 359 | 'Currently, the number of dataset samples must be specified for the training dataset. ' |
| 360 | 'Please specify it via `--train-num-samples` if no dataset length info is present.') |
| 361 | else: |
| 362 | # Eval will just exhaust the iterator if the size is not specified. |
| 363 | num_samples = args.val_num_samples or 0 |
| 364 | |
| 365 | # create a shared epoch store to sync epoch to dataloader worker proc |
| 366 | shared_epoch = SharedEpoch(epoch=epoch) |
| 367 | |
| 368 | if is_train and args.train_data_upsampling_factors is not None: |
| 369 | assert resampled, "--train_data_upsampling_factors is only supported when sampling with replacement (with --dataset-resampled)." |
| 370 | |
| 371 | if resampled: |
| 372 | pipeline = [ResampledShards2( |
| 373 | input_shards, |
| 374 | weights=args.train_data_upsampling_factors, |
| 375 | deterministic=True, |
| 376 | epoch=shared_epoch, |
| 377 | )] |
| 378 | else: |
| 379 | pipeline = [wds.SimpleShardList(input_shards)] |
| 380 | |
| 381 | # at this point we have an iterator over all the shards |
| 382 | if is_train: |
| 383 | if not resampled: |
| 384 | pipeline.extend([ |
| 385 | detshuffle2( |
| 386 | bufsize=_SHARD_SHUFFLE_SIZE, |
| 387 | initial=_SHARD_SHUFFLE_INITIAL, |
| 388 | seed=args.seed, |
| 389 | epoch=shared_epoch, |
| 390 | ), |
| 391 | wds.split_by_node, |
| 392 | wds.split_by_worker, |
| 393 | ]) |
| 394 | pipeline.extend([ |
| 395 | # at this point, we have an iterator over the shards assigned to each worker at each node |
| 396 | # wds.tarfile_to_samples(handler=log_and_continue), |
| 397 | tarfile_to_samples_nothrow, |
| 398 | wds.shuffle( |
| 399 | bufsize=_SAMPLE_SHUFFLE_SIZE, |
| 400 | initial=_SAMPLE_SHUFFLE_INITIAL, |
| 401 | ), |
| 402 | ]) |
| 403 | else: |
nothing calls this directly
no test coverage detected