How to use TF datasets with images for multi-GPU training in Keras?

Viewed 271

I am creating my TensorFlow Dataset like this:

@tf.function
def load_images(image_path, label):
    image = tf.io.read_file(image_path)
    image = tf.image.decode_jpeg(image, channels=3)
    image = vgg16.preprocess_input(image)  # VGG16 is my base model
    image = tf.image.resize(image, (IMG_SIZE, IMG_SIZE))
    return (image, label)

training_dataset = tf.data.Dataset.from_tensor_slices((train_paths, train_labels_le))
training_dataset = (
    training_dataset.shuffle(1024)
    .map(load_images, num_parallel_calls=tf.data.AUTOTUNE)
    .batch(BATCH_SIZE)
    .prefetch(tf.data.AUTOTUNE)
)

Then I build and fit my Keras model with a mirrored strategy for multi-GPU training:

strategy = tf.distribute.MirroredStrategy()
with strategy.scope():
    keras_model = build_model()  # based on tf.keras.applications.VGG16
    keras_model.compile(loss="categorical_crossentropy", optimizer="adam", metrics=["accuracy"])
keras_model.fit(training_dataset, batch_size=BATCH_SIZE, epochs=EPOCHS)

However, calling .fit() leads to the following warning:

AUTO sharding policy will apply DATA sharding policy as it failed to apply FILE sharding policy because of the following reason: Found an unsharable source dataset: name: "TensorSliceDataset/_2"

So I try to set these options, as suggested by the warning:

options = tf.data.Options()
options.experimental_distribute.auto_shard_policy = tf.data.experimental.AutoShardPolicy.DATA
training_dataset.with_options(options)

But this does not solve it, the warning is still being printed.
So how is one supposed to correctly use a TF dataset for Keras multi-GPU training?

EDIT: Link to minimal working example source code.

0 Answers
Related