MCPcopy Create free account
hub / github.com/apple/axlearn / fn

Function fn

axlearn/common/input_tf_data.py:347–386  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

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

Callers 1

process_dataset_fnFunction · 0.70

Calls 15

maybe_instantiateFunction · 0.90
set_recursivelyFunction · 0.90
get_recursivelyFunction · 0.90
_infer_num_shardsFunction · 0.85
_maybe_shard_examplesFunction · 0.85
_pad_for_evaluationFunction · 0.85
_pad_logical_to_physicalFunction · 0.85
_infer_cardinalityFunction · 0.85
_init_stateFunction · 0.85
pad_to_batchFunction · 0.85
define_shapeFunction · 0.85
cloneMethod · 0.80

Tested by

no test coverage detected