我认为第二个建议是最简单的方法。为了避免最后一批的拆分问题,您可以使用
drop_remainder
选择权
dataset.batch
dataset = dataset.batch(batch_size * multiple_gpus)
iterator = dataset.make_one_shot_iterator()
batches = iterator.get_next()
split_dims = [0] * multiple_gpus
drawn_batch_size = tf.shape(batches)[0]
以贪婪的方式,也就是说,适合
batch_size
每个装置上的张量,直到用完为止
#### Solution 1 [Greedy]:
for i in range(multiple_gpus):
split_dims[i] = tf.maximum(0, tf.minimum(batch_size, drawn_batch_size))
drawn_batch_size -= batch_size
或者以更分散的方式,确保每个设备至少获得一个样本(假设
multiple_gpus
drawn_batch_size
)
### Solution 2 [Spread]
drawn_batch_size -= - multiple_gpus
for i in range(multiple_gpus):
split_dims[i] = tf.maximum(0, tf.minimum(batch_size - 1, drawn_batch_size)) + 1
drawn_batch_size -= batch_size
## Split batches
batches = tf.split(batches, split_dims)