【发布时间】:2020-11-18 09:59:10
【问题描述】:
这就是我的推理设置的样子
autotune = tf.data.experimental.AUTOTUNE
with strategy.scope():
model = LoadModel()
raw_dataset = tf.data.TFRecordDataset(tfRecordAddress)
train_dataset = raw_dataset.map(_parse_example, num_parallel_calls=autotune)
train_dataset = train_dataset.padded_batch(batch_size, padding_values=(1, 1, b'-'), padded_shapes=(512, 512, 1))
# train_dataset = train_dataset.repeat()
train_dataset = train_dataset.prefetch(autotune)
train_dataset = strategy.experimental_distribute_dataset(train_dataset)
def per_core_inference_fn(inputIds,attnIds ):
return model.inference((inputIds, attnIds))
@tf.function
def inference_fn(inputIds, attnIds):
return strategy.run(per_core_inference_fn, args=(inputIds,attnIds))
results = []
for x in train_dataset:
t0 = time.time()
results.append(inference_fn(x[0], x[1]))
t1 = time.time()
print('time is :', t1-t0)
有了巨大的 batch_size,推理速度非常快,大约 0.0003 秒。但是,下一批的提取需要很长时间,for x in train_dataset:,大概需要 60-80 秒。
据我所知,我的推理是正确的,但不知何故,TPU 的 CPU 在批量检索时遇到了巨大的瓶颈。
我在训练期间没有看到这个瓶颈。所以看起来model.fit 正在做一些我没有做的事情。
【问题讨论】:
标签: tensorflow google-compute-engine tpu google-cloud-tpu