【发布时间】:2020-04-01 00:27:26
【问题描述】:
我正在使用 keras 训练神经网络,由于我的数据集非常大,我使用 fit_generator 将数据提供给网络。
作为fit_generator 的第一个参数,我必须提供一个生成器来为我的模型生成数据块。
我使用tf.data.Dataset 来制作数据集并使用make_one_shot_iterator 和调用get_next 方法为网络提供数据。
这是代码
def generator():
dataset_iterator = DatasetGenerator(...) # defined class to returns a tf iterator
with tf.Session() as sess:
next_batch = dataset_iterator.get_next()
while True:
img, label = sess.run(next_batch)
# some process on label
yield img, label
# down in the code for training:
model.fit_generator(generator=generator(), ...)
这工作得很好。
当我尝试将dataset_iterator 作为generator 方法的参数发送时,问题就开始了,如下所示:
def generator(dataset_iterator):
with tf.Session() as sess:
next_batch = dataset_iterator.get_next()
while True:
img, label = sess.run(next_batch)
# some process on label
yield img, label
# down in the code for training:
dataset_iterator = DatasetGenerator(...)
model.fit_generator(generator=generator(dataset_iterator), ...)
现在,我收到以下错误:
RuntimeError: The Session graph is empty. Add operations to the graph before calling run().
【问题讨论】:
-
在创建
tf.Session()之前添加next_batch = dataset_iterator.get_next()行,使其包含在图表中并且图表不会为空。 -
@ShubhamPanchal 感谢您的回复。但它没有帮助。同样的错误。
标签: python tensorflow