MCPcopy Create free account
hub / github.com/microsoft/Cream / get_wds_dataset

Function get_wds_dataset

TinyCLIP/src/training/data.py:346–467  ·  view source on GitHub ↗
(args, preprocess_img, is_train, epoch=0, floor=False, tokenizer=None)

Source from the content-addressed store, hash-verified

344
345
346def 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:

Callers

nothing calls this directly

Calls 8

get_dataset_sizeFunction · 0.85
SharedEpochClass · 0.85
ResampledShards2Class · 0.85
detshuffle2Class · 0.85
expand_urlsFunction · 0.85
DataInfoClass · 0.85
decodeMethod · 0.80
to_tupleMethod · 0.80

Tested by

no test coverage detected