TensorFlow/Keras升级后model.fit报错:tf.data.Dataset迭代问题求助
解决TensorFlow/Keras升级后model.fit()的迭代错误
错误原因
升级后出现的tf.data.Dataset only supports Python-style iteration in eager mode or within tf.function错误,核心原因是混用numpy数组输入与手动指定的validation_steps参数。当使用numpy数组作为model.fit()输入时,Keras会自动将其转换为tf.data.Dataset对象,但手动指定validation_steps会触发Dataset的迭代检查,在非eager模式下引发冲突。
修复方案
1. 移除validation_steps参数
由于你使用了validation_split=0.2,Keras会自动划分验证集并处理批次迭代,无需手动指定validation_steps。原代码中该参数基于训练集长度计算,反而会干扰Keras的自动逻辑。
修改后的训练代码:
# fit network history = model.fit(train_x, train_y, validation_split=0.2, epochs=epochs, batch_size=batch_size, verbose=verbose, callbacks=keras_callbacks)
2. 显式转换数据类型(推荐)
升级后的TensorFlow对数据类型要求更严格,建议显式将输入数据转换为float32,避免隐性类型不匹配问题:
# 在to_supervised函数返回后添加类型转换 train_x, train_y = to_supervised(train, n_input, n_output) train_x = np.asarray(train_x).astype('float32') train_y = np.asarray(train_y).astype('float32')
3. 修正train_y的维度匹配
原代码中对train_y的两次reshape存在逻辑矛盾:先转为(样本数, 时间步, 1),又直接展平为一维数组。需根据模型输出形状调整标签维度:
- 如果是多步时序预测,模型最后一层为
Dense(n_output),则标签应保持(样本数, n_output):# 替换原有的两次reshape train_y = train_y.reshape((train_y.shape[0], train_y.shape[1])) - 如果模型输出是
(样本数, n_output, 1),则保留train_y = train_y.reshape((train_y.shape[0], train_y.shape[1], 1))即可,去掉reshape(-1)。
验证修改效果
完成上述调整后,重新运行训练代码,即可避免Dataset迭代相关错误。核心思路是让Keras自动处理数据集的批次划分与迭代,减少手动参数干预带来的版本兼容问题。
内容的提问来源于stack exchange,提问作者Celia Chen
相关产品推荐
相关产品推荐

