I'm using tensorflow 2.8.0 to build a neural network. My dataset consists of 16000 3D images of 61x61x61 resolution as inputs. Is use a generator and dataset.map()to create the dataset. To speed things up I would like to use the FILE autosharding policy, which works only if I use one GPU. If I use more than one GPU (using tf.distribute.MirroredStrategy()) I get the following message:
AUTO sharding policy will apply DATA sharding policy as it failed to apply FILE sharding policy because of the following reason: Did not find a shardable source, walked to a node which is not a dataset: name: "FlatMapDataset/_2"
and the DATA policy is applied. Does anyone know how I can use the FILE policy on multiple GPUs?
Here is the code I use to create the dataset:
def get_dataset(self):
output_signature = tf.TensorSpec(shape=(), dtype=tf.string), \
tf.TensorSpec(shape=(), dtype=tf.float32),\
tf.TensorSpec(shape=(), dtype=tf.float32)
AUTOTUNE = tf.data.experimental.AUTOTUNE
dataset = tf.data.Dataset.from_generator(self.generator, output_signature=output_signature)
if self.shuffle is True:
dataset = dataset.shuffle(self.num_IDs)
dataset = dataset.map(self.map_generator, num_parallel_calls=self.num_threads)
if self.cache is True:
dataset = dataset.cache(self.cache_path)
dataset = dataset.batch(self.batch_size, drop_remainder=self.drop_remainder)
if self.prefetch is True:
dataset = dataset.prefetch(buffer_size=AUTOTUNE)
# AutoShard
options = tf.data.Options()
options.experimental_distribute.auto_shard_policy = tf.data.experimental.AutoShardPolicy.AUTO
dataset = dataset.with_options(options)
return dataset
def generator(self):
for sample in self.list_IDs:
query_ = int( sample[sample.find('query-') + len('query-') : len(sample)]) # get query from string
if self.queries_rescaled is not None:
query = self.queries_rescaled[query_]
else:
query = query_
yield sample, query, self.labels[sample]
def map_generator(self, x_elem, query_elem, label_elem):
x_input = tf.numpy_function(func=self.get_input, inp=[x_elem], Tout=self.Tout_dtype)
return {"encoder_input": x_input, "query": query_elem}, label_elem