I have a tf.dataset which is just a list of integer ranges from which to select data in a static array. For performance, it's key the static array lives on the device and not in host memory, since copying data between host and device otherwise takes up 50% of training time. The static array fits comfortably in GPU memory. (For background on why I'm doing this see Unable to get tf.data to match keras Sequence performance).
host_array = np.concatenate([self.cache[filename] for filename in self.filenames])
with tf.device('/gpu:0'):
device_array = tf.constant(host_array)
frag_length = tf.constant([self.fragment_length])
ds = tf.data.Dataset.from_tensor_slices(index_list)
@tf.function
def get_segment_tuple(index):
sample = tf.slice(device_array, [index], frag_length)
return (sample, sample)
if self.shuffle:
ds = ds.shuffle(len(index_list))
x_ds = ds.map(get_segment_tuple, num_parallel_calls=tf.data.AUTOTUNE)
return x_ds.batch(self.batch_size)
I'd like this to work on a multi-gpu machine. tf.keras.utils.experimental.DatasetCreator seems like the right tool for the job since I can feed that to Keras' model.fit. Then I could do something like,
def make_local_ds(host_array):
device_array = tf.constant(host_array)
frag_length = tf.constant([self.fragment_length])
ds = tf.data.Dataset.from_tensor_slices(index_list)
@tf.function
def get_segment_tuple(index):
sample = tf.slice(device_array, [index], frag_length)
return (sample, sample)
if self.shuffle:
ds = ds.shuffle(len(index_list)) # FIXME: does this shuffle the same on each worker?
return ds.map(get_segment_tuple, num_parallel_calls=tf.data.AUTOTUNE)
def dataset_fn(input_context):
with tf.device(??):
dataset = make_local_ds(host_array)
batch_size = input_context.get_per_replica_batch_size(global_batch_size)
dataset = dataset.shard(input_context.num_input_pipelines, input_context.input_pipeline_id)
return dataset.batch(batch_size)
But what do I put into tf.device? There's no device in InputContext. I need to know which GPU to stick the backing store data into for this generator, so that get_segment_tuple slices local device memory rather than causing cross device copying, or worse, HtoD copying.