TensorFlow使用add_loss后首个epoch结束无法拼接批次报错咨询
你遇到的形状不匹配错误是因为tf.data.Dataset.batch()默认会保留样本数不足批次大小的最后一个批次,你设置的批次大小为128,训练样本总数除以128的余数是117,所以最后一批只有117个样本。Keras内部对add_loss传入的损失值做全局聚合时,默认假设所有批次的损失张量形状一致,碰到形状不同的最后一批就触发了广播错误。
解决方案
- 方案1:丢弃不完整的最后一批(最简便)
在调用batch方法时添加drop_remainder=True参数,直接过滤掉样本数不足的最后一批,训练时损失少量样本对整体模型效果几乎无影响,修改代码如下:
autoencoder.fit(train_dataset_x_x.batch(AUTOENCODER_BATCH_SIZE, drop_remainder=True), epochs=AUTOENCODER_NUM_EPOCHS, shuffle=True)
- 方案2:保留所有样本,手动聚合损失为标量
如果你不想丢弃最后一批样本,可以手动对单批次内的样本损失做平均,确保每次add_loss传入的是标量值,避免Keras内部聚合时出现形状不匹配,修改call方法中的损失计算逻辑:
# 先算每个样本的损失,再对单批次内的所有样本求平均,得到标量损失 r_loss = tf.math.reduce_mean(tf.math.reduce_sum(tf.math.square(x - decoded), axis=[1, 2, 3])) self.add_loss(r_loss)
- 可选优化:移除不必要的标签返回
你的损失完全在模型内部基于输入和输出计算,不需要外部传入标签,可以修改数据读取逻辑,避免Keras额外处理标签带来的潜在问题:
# 调整数据解析逻辑,仅返回输入图像 def read_tfrecord(example): example = tf.io.parse_single_example(example, CELEB_A_FORMAT) image = decode_image(example['image']) return image
调整后fit调用不需要修改,可正常运行。
内容的提问来源于stack exchange,提问作者under_the_sea_salad
相关产品推荐
相关产品推荐

