Why is reading from Tensorflow Record files so slow for big tensors?

Viewed 876

I have a few tensorflow (v2.1.1) record files, with about a 200K examples in each. Each tensor is 300x1500 in dimension. Each tfrecord file would be about 60GB. When trying to read a batch of size 2048, the read latency encountered is 70-80 seconds. Not quite sure what is wrong. The input is totally preprocessed so except deserialisation, there are no other transformations.

Even while reading a single example, I see a 100-150 ms latency which is again too much.

import tensorflow as tf

def load_dataset():
    filenames = ["A.tfrecord", "B.tfrecord", "C.tfrecord", "D.tfrecord"]
    print("Training on record names: ", filenames)
    raw_dataset = tf.data.TFRecordDataset(filenames, buffer_size=100, num_parallel_reads=tf.data.experimental.AUTOTUNE)
    parsed_dataset = raw_dataset.map(parse_example,
                                     num_parallel_calls=tf.data.experimental.AUTOTUNE)
    return parsed_dataset


def parse_example(serialized_example):
    parse_dict = {
        'X': tf.io.FixedLenFeature([], tf.string),
        'X_lengths': tf.io.FixedLenFeature([], tf.string),
        'Y': tf.io.FixedLenFeature([], tf.string),
        'Y_lengths': tf.io.FixedLenFeature([], tf.string),
    }

    example = tf.io.parse_single_example(serialized_example, parse_dict)

    X = tf.io.parse_tensor(example['X'], out_type=tf.float32)
    X.set_shape([300, 1500])

    X_lengths = tf.io.parse_tensor(example['X_lengths'], out_type=tf.int32)
    X_lengths.set_shape([])

    Y = tf.io.parse_tensor(example['Y'], out_type=tf.int32)
    Y.set_shape([40])

    Y_lengths = tf.io.parse_tensor(example['Y_lengths'], out_type=tf.int32)
    Y_lengths.set_shape([])

    return X, X_lengths, Y, Y_lengths


def get_dataset(strategy=None):
    dataset = load_dataset()

    dataset = dataset.batch(2048, drop_remainder=True)

    dataset = dataset.prefetch(tf.data.experimental.AUTOTUNE)

    dataset = strategy.experimental_distribute_dataset(dataset)

    return dataset

My dataset object is provided by the get_dataset method. Which I am using in my training loop:

while True:
    train_dataset = get_dataset(strategy=mirrored_strategy)
    before_next_batch_time = time.time()
    for batch in train_dataset:
        print("Time taken for next batch: {} ".format(time.time() - before_next_batch_time))
        x, x_lengths, y, y_lengths = batch

        before_train_step = time.time()

        mean_train_loss = train_step(x, x_lengths, y, y_lengths)

        print("Time taken for train step: {}".format(time.time() - before_train_step))
        before_next_batch_time = time.time()
        print("****")

Actual run logs:

****
Time taken for next batch: 76.75251317024231
Time taken for train step: 2.1996893882751465
****
Time taken for next batch: 76.99043083190918
Time taken for train step: 2.192229747772217
****
Time taken for next batch: 76.46133637428284
Time taken for train step: 2.2198166847229004
****
Time taken for next batch: 76.34514284133911
Time taken for train step: 2.1689696311950684
****
Time taken for next batch: 77.19221472740173
Time taken for train step: 2.2315127849578857
****

Edit:

The TF record file is kept on a mounted disk on the same vm. There are no network calls, as the distribution is happening within the same vm, this vm has 4 GPUs.

1 Answers

Try increasing buffer size to some bigger number. From 100 to 10000, for example. It helped me decrease training time by an order of magnitude on dgx-2!

Related