多输出生成模型的损失融合与反向传播问题求助
问题与解决方案:TensorFlow多输出生成模型的损失融合与反向传播
核心问题
- 生成模型输出多个结果,需融合对应损失,但当前融合后的损失无法影响权重更新
- 希望实现类似PyTorch中
loss.backward()+optimizer.step()的流程,对融合后的损失g_loss = g_loss + (LAMBDA * l_loss)进行反向传播
用户原代码
def train(self, epochs, batch_size=1, sample_interval=50): LAMBDA = 0.1 start_time = datetime.datetime.now() for epoch in range(epochs): if epoch % 500 == 0: optimizer = Adam(0.0001, 0.5) self.combined.compile(loss=['mae'], loss_weights=[1, 100], optimizer=optimizer) for batch_i, (imgs_A, imgs_B) in enumerate(self.data_loader.load_batch(batch_size)): # os.environ['CUDA_VISIBLE_DEVICES'] = '-1' fake_A1, fake_A2, fake_A3 = self.generator.predict(imgs_B) g1_loss = self.combined.train_on_batch([imgs_A, imgs_B], [fake_A1]) fake_A2 = tf.image.resize(fake_A2, [256, 256]) fake_A3 = tf.image.resize(fake_A3, [256, 256]) g2_loss = self.combined.train_on_batch([imgs_A, imgs_B], [fake_A2]) g3_loss = self.combined.train_on_batch([imgs_A, imgs_B], [fake_A3]) g1_loss = tf.cast(g1_loss, tf.double) g2_loss = tf.cast(g2_loss, tf.double) g3_loss = tf.cast(g3_loss, tf.double) g_loss = g1_loss + g2_loss + g3_loss imgs_A = np.fft.fft2(imgs_A) fake_A1 = np.fft.fft2(fake_A1) fake_A2 = np.fft.fft2(fake_A2) fake_A3 = np.fft.fft2(fake_A3) mse = tf.keras.losses.MeanSquaredError() l1_loss = mse(imgs_A, fake_A1).numpy() l2_loss = mse(imgs_A, fake_A2).numpy() l3_loss = mse(imgs_A, fake_A3).numpy() l1_loss = tf.cast(l1_loss, tf.double) l2_loss = tf.cast(l2_loss, tf.double) l3_loss = tf.cast(l3_loss, tf.double) l_loss = l1_loss + l2_loss + l3_loss #Normalized loss g_loss = g_loss + (LAMBDA * l_loss) elapsed_time = datetime.datetime.now() - start_time print("[Epoch %d/%d] [Batch %d/%d] [G loss: %f] time: %s" % (epoch, epochs, batch_i, self.data_loader.n_batches, g_loss, elapsed_time))
原代码核心问题分析
梯度链路断裂
- 使用
generator.predict()获取输出属于推理模式,不会记录梯度信息 - 傅里叶域损失计算中调用
.numpy(),将张量转为NumPy数组,彻底断开计算图,无法反向传播 - 三次
train_on_batch是独立的反向传播,仅数值上合并损失,未将总损失与模型权重绑定
- 使用
训练状态不稳定
- 每个epoch重新
compile模型,会重置训练状态 - 每500epoch重建优化器,丢失动量等优化状态,影响收敛
- 每个epoch重新
解决方案:手动梯度追踪与反向传播
使用TensorFlow的tf.GradientTape上下文管理器,手动记录梯度并更新权重,完全对应PyTorch的反向传播流程。
修正后完整代码
import tensorflow as tf from tensorflow.keras.optimizers import Adam import datetime import numpy as np def train(self, epochs, batch_size=1, sample_interval=50): LAMBDA = 0.1 start_time = datetime.datetime.now() # 仅初始化一次优化器,避免重置状态 optimizer = Adam(0.0001, 0.5) mae_loss = tf.keras.losses.MeanAbsoluteError() mse_loss = tf.keras.losses.MeanSquaredError() for epoch in range(epochs): # 可选:每500epoch调整学习率,而非重建优化器 if epoch % 500 == 0 and epoch != 0: current_lr = optimizer.learning_rate.numpy() optimizer.learning_rate.assign(current_lr * 0.5) for batch_i, (imgs_A, imgs_B) in enumerate(self.data_loader.load_batch(batch_size)): # 将输入转为TensorFlow张量,确保在计算图内 imgs_A = tf.convert_to_tensor(imgs_A, dtype=tf.float32) imgs_B = tf.convert_to_tensor(imgs_B, dtype=tf.float32) with tf.GradientTape() as tape: # 直接调用generator获取输出(训练模式,记录梯度) fake_A1, fake_A2, fake_A3 = self.generator(imgs_B, training=True) # 调整输出尺寸 fake_A2 = tf.image.resize(fake_A2, [256, 256]) fake_A3 = tf.image.resize(fake_A3, [256, 256]) # 计算MAE损失 g1_loss = mae_loss(imgs_A, fake_A1) g2_loss = mae_loss(imgs_A, fake_A2) g3_loss = mae_loss(imgs_A, fake_A3) g_loss_total = g1_loss + g2_loss + g3_loss # 傅里叶域MSE损失(全程在计算图内) imgs_A_fft = tf.signal.fft2d(tf.cast(imgs_A, tf.complex64)) fake_A1_fft = tf.signal.fft2d(tf.cast(fake_A1, tf.complex64)) fake_A2_fft = tf.signal.fft2d(tf.cast(fake_A2, tf.complex64)) fake_A3_fft = tf.signal.fft2d(tf.cast(fake_A3, tf.complex64)) # 复数损失取绝对值后计算MSE l1_loss = mse_loss(tf.abs(imgs_A_fft), tf.abs(fake_A1_fft)) l2_loss = mse_loss(tf.abs(imgs_A_fft), tf.abs(fake_A2_fft)) l3_loss = mse_loss(tf.abs(imgs_A_fft), tf.abs(fake_A3_fft)) l_loss_total = l1_loss + l2_loss + l3_loss # 合并总损失 total_loss = g_loss_total + (LAMBDA * l_loss_total) # 计算梯度(对应PyTorch的loss.backward()) gradients = tape.gradient(total_loss, self.generator.trainable_variables) # 更新权重(对应PyTorch的optimizer.step()) optimizer.apply_gradients(zip(gradients, self.generator.trainable_variables)) elapsed_time = datetime.datetime.now() - start_time print("[Epoch %d/%d] [Batch %d/%d] [Total loss: %f] time: %s" % (epoch, epochs, batch_i, self.data_loader.n_batches, total_loss.numpy(), elapsed_time))
关键修改说明
tf.GradientTape追踪梯度:上下文内所有运算都会被记录,实现自动微分- 直接调用生成器:
self.generator(imgs_B, training=True)启用训练模式并记录梯度,替代推理模式的predict - 计算图内完成傅里叶变换:用
tf.signal.fft2d替代np.fft.fft2d,避免转NumPy断开梯度链路 - 统一梯度更新:
tape.gradient()计算总损失对可训练变量的梯度,optimizer.apply_gradients()更新权重,完全复刻PyTorch的反向传播流程 - 稳定训练状态:仅初始化一次优化器,如需调整学习率直接修改
optimizer.learning_rate
内容的提问来源于stack exchange,提问作者Sanjay Gupta
相关产品推荐
相关产品推荐

