()
| 345 | read_config = read_config.set(is_training=is_training) |
| 346 | |
| 347 | def fn() -> tf.data.Dataset: |
| 348 | local_read_config = read_config.clone() |
| 349 | builder = tfds.builder(dataset_name, data_dir=data_dir) |
| 350 | if download: |
| 351 | logging.info("Downloading %s", dataset_name) |
| 352 | builder.download_and_prepare() |
| 353 | logging.info("Downloading %s done", dataset_name) |
| 354 | |
| 355 | # If we can infer that the split doesn't have enough shards, fallback to using subsplit API. |
| 356 | # Note: Without the alias, pytype will complain if we modify the nonlocal `split` below. |
| 357 | per_process_split = split |
| 358 | available_shards = _infer_num_shards(builder, split) |
| 359 | if available_shards is not None: |
| 360 | required_shards = ( |
| 361 | read_config.num_shards |
| 362 | if read_config.num_shards is not None |
| 363 | else jax.process_count() |
| 364 | ) |
| 365 | if available_shards < required_shards: |
| 366 | per_process_split = _maybe_shard_examples( |
| 367 | builder=builder, |
| 368 | read_config=read_config, |
| 369 | split=split, |
| 370 | required_shards=required_shards, |
| 371 | is_training=is_training, |
| 372 | dataset_name=dataset_name, |
| 373 | ) |
| 374 | local_read_config.set(num_shards=1, shard_index=0) |
| 375 | |
| 376 | ds: tf.data.Dataset = builder.as_dataset( |
| 377 | split=per_process_split, |
| 378 | shuffle_files=shuffle_files, |
| 379 | read_config=local_read_config.instantiate(), |
| 380 | decoders=maybe_instantiate(decoders), |
| 381 | ) |
| 382 | if shuffle_buffer_size > 0: |
| 383 | # Subsequent processing may merge/split examples (e.g. for T5), so shuffle examples |
| 384 | # during training before any processing. |
| 385 | ds = ds.shuffle(shuffle_buffer_size, reshuffle_each_iteration=True) |
| 386 | return ds |
| 387 | |
| 388 | return fn |
| 389 |
no test coverage detected