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

Function tfds_dataset

axlearn/common/input_tf_data.py:285–388  ·  view source on GitHub ↗

Returns a BuildDatasetFn for the given TFDS dataset name and split. Args: dataset_name: The tensorflow dataset name. split: The dataset split. is_training: Whether the examples are used for training. If True, examples will be read in parallel and shuffled.

(
    dataset_name: str,
    *,
    split: str,
    is_training: bool,
    train_shuffle_buffer_size: Optional[int] = None,
    train_shuffle_files: Optional[bool] = None,
    data_dir: Optional[str] = None,
    download: bool = False,
    read_config: Optional[InstantiableConfig] = None,
    decoders: Optional[InstantiableConfig] = None,
)

Source from the content-addressed store, hash-verified

283
284
285def tfds_dataset(
286 dataset_name: str,
287 *,
288 split: str,
289 is_training: bool,
290 train_shuffle_buffer_size: Optional[int] = None,
291 train_shuffle_files: Optional[bool] = None,
292 data_dir: Optional[str] = None,
293 download: bool = False,
294 read_config: Optional[InstantiableConfig] = None,
295 decoders: Optional[InstantiableConfig] = None,
296) -> BuildDatasetFn:
297 """Returns a BuildDatasetFn for the given TFDS dataset name and split.
298
299 Args:
300 dataset_name: The tensorflow dataset name.
301 split: The dataset split.
302 is_training: Whether the examples are used for training.
303 If True, examples will be read in parallel and shuffled.
304 Otherwise examples will be read sequentially to ensure a deterministic order.
305 train_shuffle_buffer_size: The shuffle buffer size (required) when is_training=True.
306 If is_training=False or shuffle_buffer_size <= 0, no shuffling is done.
307 train_shuffle_files: Whether to shuffle files when is_training=True.
308 If is_training=False, no shuffling is done.
309 If is_training=True and train_shuffle_files is None, infer from shuffle_buffer_size > 0.
310 data_dir: Used for tfds.load. If None, use the value of the environment variable
311 DATA_DIR, TFDS_DATA_DIR, or "~/tensorflow_datasets" (in that order).
312 download: Whether to download the examples. If false, use the data under data_dir.
313 read_config: a Config that instantiates to a tfds.ReadConfig.
314 If None, constructs a default one with tfds_read_config().
315 decoders: An optional config instantiating a (nested) mapping of feature names to decoders.
316 See: https://www.tensorflow.org/datasets/api_docs/python/tfds/decode/Decoder
317
318 Returns:
319 A BuildDatasetFn, which returns a tf.data.Dataset for the specified TFDS.
320
321 Raises:
322 ValueError: if train_shuffle_buffer_size is None when is_training=True.
323 """
324 if is_training:
325 if train_shuffle_buffer_size is None:
326 raise ValueError("train_shuffle_buffer_size is required when is_training=True")
327 shuffle_files = (
328 train_shuffle_buffer_size > 0 if train_shuffle_files is None else train_shuffle_files
329 )
330 shuffle_buffer_size = train_shuffle_buffer_size
331 else:
332 # Disable shuffling.
333 #
334 # Note that it's OK for is_training=True and train_shuffle_buffer_size=0 so that the inputs
335 # are deterministic.
336 shuffle_buffer_size = 0
337 shuffle_files = False
338
339 if data_dir is None:
340 data_dir = get_data_dir()
341
342 if read_config is None:

Callers

nothing calls this directly

Calls 3

get_data_dirFunction · 0.90
config_for_functionFunction · 0.90
setMethod · 0.45

Tested by

no test coverage detected