TensorFlow TPU运行MNIST DCGAN报Dimensions must be equal错误如何解决
问题原因
这个维度不匹配报错来自两个核心问题:
- 输入数据集未丢弃最后一个不完整批次:MNIST训练集共60000个样本,你设置的单批次大小为256,60000除以256后余数为96,因此最后一个批次仅包含96个样本。而你在训练步中写死了生成噪声的批次大小为固定256,导致判别器输出的真实样本损失维度为[96]、生成样本损失维度为[256],二者相加时维度不匹配触发报错。
- TPU训练要求所有输入批次的形状完全固定,不支持动态批次大小,不满的批次会直接导致分布计算异常。
解决方案
按以下步骤修改代码即可正常运行:
- 第一步:修改数据集构造逻辑,添加
drop_remainder=True参数丢弃最后一个不完整批次,保证所有批次大小均为256:
# 修改前 train_dataset = tf.data.Dataset.from_tensor_slices(train_images).shuffle(BUFFER_SIZE).batch(BATCH_SIZE) # 修改后 train_dataset = tf.data.Dataset.from_tensor_slices(train_images).shuffle(BUFFER_SIZE).batch(BATCH_SIZE, drop_remainder=True)
- 第二步:修改训练步中噪声生成的逻辑,动态读取输入图片的批次维度,不要写死固定值,提升代码鲁棒性:
# 修改前 noise = tf.random.normal([BATCH_SIZE, noise_dim]) # 修改后 noise = tf.random.normal([tf.shape(images)[0], noise_dim])
- 第三步(可选优化,适配分布式训练的损失计算):你当前使用了无聚合的损失计算模式,需手动对损失做跨样本、跨TPU副本的平均,避免梯度计算出现偏差:
def generator_loss(fake_output): per_sample_loss = cross_entropy(tf.ones_like(fake_output), fake_output) return tf.nn.compute_average_loss(per_sample_loss, global_batch_size=BATCH_SIZE) def discriminator_loss(real_output, fake_output): real_loss = cross_entropy(tf.ones_like(real_output), real_output) fake_loss = cross_entropy(tf.zeros_like(fake_output), fake_output) per_sample_loss = real_loss + fake_loss return tf.nn.compute_average_loss(per_sample_loss, global_batch_size=BATCH_SIZE)
内容的提问来源于stack exchange,提问作者Tahmid Faisal
相关产品推荐
相关产品推荐

