How do I distribute a static tensor to each GPU and use the local copy for my tf.data?

Viewed 68

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.

0 Answers
Related