TFRecord reading pipeline slows down after PREFETCH samples

Viewed 482

I have divided my training data into multiple tf-record files, and read them using this piece of code:

SHUFFLE_BUFFER = 64
PREFETCH = 256
dataset = tf.data.TFRecordDataset(filenames)
dataset = dataset.shuffle(SHUFFLE_BUFFER) 
dataset = dataset.map(_parse_image_function, num_parallel_calls=tf.data.experimental.AUTOTUNE)
dataset = dataset.batch(BATCH_SIZE)
dataset = dataset.prefetch(PREFETCH)
dataset = dataset.repeat()

This dataset is fed directly to model.fit(dataset).

The first PREFETCH samples are loaded quickly, and GPU utilization is constantly above 80%. However, after that the fast reading seems to stop, GPU utilization drops, and training time slows down massively. Anyone know what might be going wrong?

1 Answers

This is kind of hard to diagnose without knowing more details (storage backend, record size, number of records per file, number of files, any io operations in _parse_image_function?, ....)

My first suspicion is on tf.data.TFRecordDataset(filenames) - opening one file after the next may introduce latency spikes that could temporarily starve the dataset cpu pipeline. (Multiple smaller files may also have lower benefits from automatic read aheads)

I would try to add an additional prefetch right after tf.data.TFRecordDataset(filenames) to decouple IO (and maybe interleave records from different files (num_parallel_reads argument)).

If the prefetch does not help I would try to hard code num_parallel_calls (mostly because I have not read the autotune code yet - and maybe use a private thread pool if your pipeline needs more then the default parallelism).

Depending on your storage backend - repeated training restarts once training slows down (to test/optimize the dataset) may just pull data from various caches and may slow down once the used dataset exceeds the caches.

Related