You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

多输出生成模型的损失融合与反向传播问题求助

问题与解决方案: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))

原代码核心问题分析

  1. 梯度链路断裂

    • 使用generator.predict()获取输出属于推理模式,不会记录梯度信息
    • 傅里叶域损失计算中调用.numpy(),将张量转为NumPy数组,彻底断开计算图,无法反向传播
    • 三次train_on_batch是独立的反向传播,仅数值上合并损失,未将总损失与模型权重绑定
  2. 训练状态不稳定

    • 每个epoch重新compile模型,会重置训练状态
    • 每500epoch重建优化器,丢失动量等优化状态,影响收敛

解决方案:手动梯度追踪与反向传播

使用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.18 16:45:58