What is the proper use of Tensorflow dataset prefetch and cache options?

Viewed 8580

I have read TF pages and some posts and about the use of prefetch() and cache() to speed up model input pipeline and tried to implement it on my data. Cache() worked for me as expected, i.e. reading data from the dist in the first epoch, and in the all subsequent epochs it is just reading data from memory. But I have many difficulties using prefetch() and I really don't understand when and how to use it. Can anybody help me with it? I really need some help. My application is like this: I have a set of large TFRecord files, each includes some raw records to be processed before feeding my net. They are going to be mixed (different sample streams) so what I do is:

def read_datasets(pattern, numFiles, numEpochs=125, batchSize=1024, take=dataLength):

    files = tf.data.Dataset.list_files(pattern)

    def _parse(x):
        x = tf.data.TFRecordDataset(x, compression_type='GZIP')
    return x

    np = 4 # half of the number of CPU cores
    dataset = files.interleave(_parse, cycle_length=numFiles, block_length=1, num_parallel_calls=np)\
    .map(lambda x: parse_tfrecord(x), num_parallel_calls=np)
    dataset = dataset.take(take)
    dataset = dataset.batch(batchSize)
    dataset = dataset.cache()
    dataset = dataset.prefetch(buffer_size=10)
    dataset = dataset.repeat(numEpochs)
    return dataset

parse_tfrecord(x) function in interleave function is the required preprocessing of data before it applies to the model, and my guess is that the preprocesing time is comparable to the batch processing time by network. My whole dataset (including all input files) contains about 500 batch of 1024 samples. My questions are:

1- If I do caching, do I really need prefetching?

2- Is it the right sequence to do mapping, batching, caching, prefetching, and repeating?

3- Tensorflow documentation says that the buffer size of prefetch refers to the dataset elements and if it is batched, to the number of batches. So in this case I will read 10 batches of 1024 examples, right? My problem is that I don't see any difference in run time by changing prefetching buffer size and the memory consumption is not changed much even by setting buffer size to 1000 or bigger.

2 Answers

I found this great explanation for Andrew Nu from Stanford. https://cs230.stanford.edu/blog/datapipeline/#best-practices

"When the GPU is working on forward / backward propagation on the current batch, we want the CPU to process the next batch of data so that it is immediately ready. As the most expensive part of the computer, we want the GPU to be fully used all the time during training. We call this consumer / producer overlap, where the consumer is the GPU and the producer is the CPU.

With tf.data, you can do this with a simple call to dataset.prefetch(1) at the end of the pipeline (after batching). This will always prefetch one batch of data and make sure that there is always one ready.

In some cases, it can be useful to prefetch more than one batch. For instance if the duration of the preprocessing varies a lot, prefetching 10 batches would average out the processing time over 10 batches, instead of sometimes waiting for longer batches.

To give a concrete example, suppose than 10% of the batches take 10s to compute, and 90% take 1s. If the GPU takes 2s to train on one batch, by prefetching multiple batches you make sure that we never wait for these rare longer batches."

I'm not quite sure how to determine processing time of each batch but that's the next step. If your batches are roughly taking the same amount of time to process then I believe prefetch(batch_size=1) should suffice as your GPU wouldn't be waiting for the cPU to finish processing a computationally expensive batch.

Can you have a look into this Stackoverflow Answer to get a quick idea about TensorFlow Dataset's functions cache() and prefetch().

Also, I found this Tensorflow Documentation very helpful to optimize the performance of the tf.Data Api. They have specified the benchmark and the execution time for various ways of execution. You can also find information about serialized and parallelized load and transformation of data and their execution time respectively.

Related