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

基于TensorFlow 2.x/Keras实现GradNorm遇梯度计算问题求助

GradNorm实现中梯度返回None的问题解决

问题背景

基于《GradNorm: Gradient Normalization for Adaptive Loss Balancing in Deep Multitask Networks》论文,用Keras的model.fit()训练双头部模型,尝试通过自定义Callback实现GradNorm梯度平衡时,计算得到的G1R是全None的数组。

问题原因

  1. 计算图依赖断裂:用tf.constant(loss_t1)包装日志中的损失值,这个常量和共享层权重没有任何计算图关联,梯度追踪无法找到依赖,自然返回None。
  2. Callback时机错误:on_batch_end是在当前批次的反向传播完成后触发的,此时当前批次的前向传播计算图已被销毁,无法基于该批次输入计算梯度。

解决方案:改用自定义训练循环

Callback并不适合GradNorm这种需要介入梯度计算的逻辑,更推荐用Keras自定义训练循环直接嵌入GradNorm的权重更新逻辑:

import tensorflow as tf

class GradNormTrainer:
    def __init__(self, model, optimizer, alpha=0.2):
        self.model = model
        self.optimizer = optimizer
        self.alpha = alpha
        # 根据你的模型结构,替换为实际的共享层权重
        self.shared_weights = [w for w in model.trainable_weights if 'shared' in w.name]
        # 初始化GradNorm参数
        self.w1 = tf.Variable(1.0, dtype=tf.float32)
        self.w2 = tf.Variable(1.0, dtype=tf.float32)
        self.l01 = None
        self.l02 = None
        self.task_loss_history = []

    def train_step(self, x, y1, y2):
        with tf.GradientTape(persistent=True) as tape:
            # 完整前向传播,保留计算图依赖
            pred1, pred2 = self.model(x)
            
            # 计算原始任务损失
            loss_t1 = tf.reduce_mean(tf.keras.losses.sparse_categorical_crossentropy(y1, pred1))
            loss_t2 = tf.reduce_mean(tf.keras.losses.sparse_categorical_crossentropy(y2, pred2))
            
            # 初始化第一批次的基准损失
            if self.l01 is None:
                self.l01 = loss_t1
                self.l02 = loss_t2
            
            # 计算归一化损失
            l_hat1 = loss_t1 / self.l01
            l_hat2 = loss_t2 / self.l02
            l_hat_avg = (l_hat1 + l_hat2) / 2
            
            # 计算共享层对每个任务损失的梯度范数
            grad1 = tape.gradient(loss_t1, self.shared_weights)
            grad2 = tape.gradient(loss_t2, self.shared_weights)
            G1 = tf.linalg.global_norm(grad1)
            G2 = tf.linalg.global_norm(grad2)
            
            # 计算平均梯度范数与梯度比率
            G_avg = (G1 + G2) / 2
            r1 = G1 / G_avg
            r2 = G2 / G_avg
            
            # 更新损失权重并归一化
            new_w1 = self.w1 * tf.math.pow(r1 / l_hat1, self.alpha)
            new_w2 = self.w2 * tf.math.pow(r2 / l_hat2, self.alpha)
            w_sum = new_w1 + new_w2
            self.w1.assign(new_w1 * 2 / w_sum)
            self.w2.assign(new_w2 * 2 / w_sum)
            
            # 计算加权总损失
            total_loss = self.w1 * loss_t1 + self.w2 * loss_t2
        
        # 更新模型权重
        model_grads = tape.gradient(total_loss, self.model.trainable_weights)
        self.optimizer.apply_gradients(zip(model_grads, self.model.trainable_weights))
        
        # 记录训练数据
        self.task_loss_history.append((loss_t1.numpy(), loss_t2.numpy()))
        return loss_t1, loss_t2, self.w1.numpy(), self.w2.numpy()

# 使用示例
# 假设train_dataset是包含(x, y1, y2)的tf.data.Dataset
trainer = GradNormTrainer(your_model, tf.keras.optimizers.Adam(), alpha=0.2)
for epoch in range(10):
    print(f"Epoch {epoch+1}/10")
    batch_count = 0
    for x, y1, y2 in train_dataset:
        loss1, loss2, w1, w2 = trainer.train_step(x, y1, y2)
        batch_count += 1
        if batch_count % 50 == 0:
            print(f"Batch {batch_count}: loss1={loss1:.4f}, loss2={loss2:.4f}, w1={w1:.4f}, w2={w2:.4f}")

关键修正点

  • 保留计算图依赖:直接基于模型前向传播计算损失,确保损失与共享层权重存在计算图关联,梯度能被正确追踪。
  • 正确时机计算梯度:在训练步骤中用tf.GradientTape包裹前向和损失计算,确保梯度计算在当前批次的计算图有效时执行。
  • 严格遵循论文公式:按GradNorm论文逻辑更新损失权重,并做归一化处理,保证权重总和与初始值一致。

内容的提问来源于stack exchange,提问作者samuraikmc

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.02 15:50:31