TensorFlow回归模型训练差、predict报Empty batch_outputs错误如何解决?
问题原因分析
两个问题的核心根源都是数据集构造和切分逻辑错误:
- 你调用
np.expand_dims(..., axis=0)是在第0维增加维度,最终生成的X_regression、y_regression形状都是(1, 200)(0到1000步长为5的序列共200个值)。 - 执行
X_regression[:150]切分数据集时,是在第0维度做切片,而该维度只有1个元素,最终X_reg_train形状仍为(1, 200)(相当于把200个样本当成1个样本的200个特征输入,训练逻辑完全错误,所以效果极差),X_reg_test则是空数组,调用predict时输入空数据就会触发Empty batch_outputs报错。
解决方法
仅需要修正数据集的维度扩展逻辑,把样本维度放在第0位,特征维度放在第1位,将axis=0改为axis=1即可:
# Set random seed tf.random.set_seed(42) # Create some regression data,修正expand_dims的轴参数 X_regression = np.expand_dims(np.arange(0, 1000, 5), axis=1) y_regression = np.expand_dims(np.arange(100, 1100, 5), axis=1) # Split it into training and test sets,现在第0维是200个样本,切分结果正常 X_reg_train = X_regression[:150] X_reg_test = X_regression[150:] y_reg_train = y_regression[:150] y_reg_test = y_regression[150:]
修正数据集后原有的模型编译、训练、预测代码不需要额外改动,训练效果会恢复正常,predict也不会再触发报错。不需要额外添加run_eagerly=True参数,该报错提示是TensorFlow的通用兜底提示,和你的场景无关。
内容的提问来源于stack exchange,提问作者user14653534
相关产品推荐
相关产品推荐

