Tensorflow报错TypeError: 'RepeatDataset' object is not callable的原因
错误原因
报错的核心触发点是代码中两处多余的括号调用:
tf.data.Dataset经过链式调用后返回的RepeatDataset是可迭代的数据集实例,并非可执行函数,原代码中return ds()尝试将数据集对象作为函数调用,直接触发TypeError: 'RepeatDataset' object is not callable报错。- TensorFlow输入函数生成器的设计要求是返回可调用的输入函数本身,原代码末尾
return input_function()会在调用make_input_fn时直接执行内部逻辑返回数据集,不符合API规范要求。
修复方案
删除两处多余的括号即可,修复后的完整代码如下:
def make_input_fn(data_df, label_df, num_epochs=10, shuffle=True, batch_size=32): def input_function(): ds = tf.data.Dataset.from_tensor_slices((dict(data_df), label_df)) if shuffle: ds = ds.shuffle(1000) ds = ds.batch(batch_size).repeat(num_epochs) return ds return input_function
原有调用逻辑train_input_fn = make_input_fn(dftrain, y_train)无需修改,调用后将得到符合要求的可调用输入函数,传入TensorFlow模型的训练接口即可正常运行。
内容的提问来源于stack exchange,提问作者mwckres0
相关产品推荐
相关产品推荐

