基于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的数组。
问题原因
- 计算图依赖断裂:用
tf.constant(loss_t1)包装日志中的损失值,这个常量和共享层权重没有任何计算图关联,梯度追踪无法找到依赖,自然返回None。 - 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
相关产品推荐
相关产品推荐

