【发布时间】:2020-07-20 04:56:41
【问题描述】:
我正在尝试使用tf.data.Dataset.from_generator() 来生成训练和验证数据。
我有自己的数据生成器,可以即时进行功能准备:
def data_iterator(self, input_file_list, ...):
for f in input_file_list:
X, y = get_feature(f)
yield X, y
最初我将其直接提供给 tensorflow keras 模型,但在第一批之后我遇到数据超出范围错误。然后我决定将它封装在 tensorflow 数据生成器中:
train_gen = lambda: data_iterator(train_files, ...)
valid_gen = lambda: data_iterator(valid_files, ...)
output_types = (tf.float32, tf.float32)
output_shapes = (tf.TensorShape([499, 13]), tf.TensorShape([2]))
train_dat = tf.data.Dataset.from_generator(train_gen,
output_types=output_types,
output_shapes=output_shapes)
valid_dat = tf.data.Dataset.from_generator(valid_gen,
output_types=output_types,
output_shapes=output_shapes)
train_dat = train_dat.repeat().batch(batch_size=128)
valid_dat = valid_dat.repeat().batch(batch_size=128)
然后拟合:
model.fit(x=train_dat,
validation_data=valid_dat,
steps_per_epoch=train_steps,
validation_steps=valid_steps,
epochs=100,
callbacks=callbacks)
但是,尽管生成器中有.repeat(),但我仍然收到错误消息:
BaseCollectiveExecutor::StartAbort 超出范围:序列结束
我的问题是:
- 为什么
.repeat()在这里不起作用? - 我应该在自己的迭代器中添加
while True以避免这种情况吗?我觉得这可以解决它,但它看起来不像是正确的做法。
【问题讨论】:
-
你能提供完整的堆栈跟踪吗?
-
@thushv89 不幸的是,我不再拥有我的了,因为我开始进行另一次测试运行并覆盖了日志。但这与此处相同:github.com/tensorflow/tensorflow/issues/31509 或可能报告此问题的任何帖子。
标签: python tensorflow keras generator tensorflow-datasets