多任务对抗训练中,如何在TensorFlow实现PyTorch式高阶梯度反向传播
我的目标是在多任务模型中,借助对抗学习思路,将不同任务的梯度输入判别器,使这些梯度在统计上无法区分,以此约束梯度对齐。但在TensorFlow中实现时遇到了问题。
模型结构概述
share_embedding = tf.concat(embeddings) concat_embedding_task1 = tf.concat([share_embedding, diff_embedding1], -1) concat_embedding_task2 = tf.concat([share_embedding, diff_embedding2], -1) output_task1 = dense_tower(name='tower1', concat_embedding_task1, nn_units_task1) output_task2 = dense_tower(name='tower2', concat_embedding_task2, nn_units_task2) loss_task1 = BCELoss(pred=output_task1, label=label1) loss_task2 = BCELoss(pred=output_task2, label=label2) total_loss = loss_task1 + loss_task2
当前实现困境
我尝试用tf.GradientTape获取每个任务回流至share_embedding的梯度,代码如下:
with tf.GradientTape() as tape: # 模型前向计算逻辑 grad_1 = tape.gradient(loss_task1, share_embedding) grad_2 = tape.gradient(loss_task2, share_embedding) grad_share = tf.where(flag_task, grad_1, grad_2) disc_logits = dense_tower(name='discriminator_layer', grad_share, nn_units_disc) disc_loss = BCELoss(pred=disc_logits, label=flag_task)
得到disc_loss后,不清楚如何让其梯度回流至share_embedding。判别器输入grad_share已处于反向传播流程中,需要让disc_loss的梯度沿着原路径再次反向传播。
调研后认为高阶导数可解决:令share_embedding为x,常规反向传播梯度为dLoss/dx,判别器梯度需计算dLoss_disc/dx,但不确定能否用tape.gradient(disc_loss, share_embedding)实现,以及如何与原梯度一同作用于share_embedding。
参考原论文的PyTorch实现,关键代码通过create_graph=True构建导数计算图,允许高阶导数计算,让梯度从判别器回流至共享特征:
grads[task] = grad(curr_loss, features[task], create_graph=True)[0] grads_norm = grads[task].norm(p=2, dim=1).unsqueeze(1) + 1e-10 input_dscr = grads[task] / grads_norm outputs_dscr[task] = self.discriminator(input_dscr)
核心问题:如何在TensorFlow中记录梯度自身的二阶梯度,实现与PyTorch中grad(loss, feature, create_graph=True)相同的效果?
解决方案
在TensorFlow中,要实现类似PyTorchcreate_graph=True的高阶导数追踪,需要使用嵌套的tf.GradientTape,并正确设置persistent和watch参数,具体步骤如下:
1. 外层Tape追踪共享特征的梯度计算图
用持久化的外层GradientTape包裹整个前向和梯度计算过程,确保能多次调用gradient方法,并且追踪梯度本身的计算图(用于后续二阶导计算)。
2. 内层Tape计算任务损失对共享特征的梯度
外层Tape开启后,在内层计算任务损失对share_embedding的梯度,此时外层Tape会记录梯度的计算过程,为后续计算判别器损失对share_embedding的二阶导做准备。
3. 计算判别器损失并反向传播二阶导
得到梯度后喂给判别器计算disc_loss,再通过外层Tape计算disc_loss对share_embedding的梯度(即二阶导),最后将原任务梯度与判别器的二阶导合并,更新share_embedding。
完整示例代码
# 外层设置persistent=True,允许多次调用gradient;手动watch share_embedding确保被追踪 with tf.GradientTape(persistent=True, watch_accessed_variables=False) as outer_tape: outer_tape.watch(share_embedding) # 第一步:计算任务前向和损失 with tf.GradientTape() as inner_tape: # 模型前向计算 concat_embedding_task1 = tf.concat([share_embedding, diff_embedding1], -1) concat_embedding_task2 = tf.concat([share_embedding, diff_embedding2], -1) output_task1 = dense_tower(name='tower1', inputs=concat_embedding_task1, units=nn_units_task1) output_task2 = dense_tower(name='tower2', inputs=concat_embedding_task2, units=nn_units_task2) loss_task1 = BCELoss(pred=output_task1, label=label1) loss_task2 = BCELoss(pred=output_task2, label=label2) total_loss = loss_task1 + loss_task2 # 第二步:计算任务损失对share_embedding的梯度(一阶导) grad_1 = inner_tape.gradient(loss_task1, share_embedding) grad_2 = inner_tape.gradient(loss_task2, share_embedding) grad_share = tf.where(flag_task, grad_1, grad_2) # 第三步:判别器前向计算,得到disc_loss disc_logits = dense_tower(name='discriminator_layer', inputs=grad_share, units=nn_units_disc) disc_loss = BCELoss(pred=disc_logits, label=flag_task) # 第四步:计算disc_loss对share_embedding的梯度(二阶导,对应PyTorch create_graph=True的效果) disc_grad_to_share = outer_tape.gradient(disc_loss, share_embedding) # 第五步:合并原任务梯度与判别器梯度,更新share_embedding total_task_grad = outer_tape.gradient(total_loss, share_embedding) # 可根据需求为判别器梯度添加权重系数,比如lambda * disc_grad_to_share final_grad = total_task_grad + disc_grad_to_share # 应用梯度更新(假设optimizer是你定义的优化器) optimizer.apply_gradients([(final_grad, share_embedding)]) # 释放持久化的tape,避免内存泄漏 del outer_tape
关键细节说明
persistent=True:外层Tape必须设置该参数,因为需要多次调用gradient方法(计算任务总梯度和判别器的二阶导)。watch_accessed_variables=False+ 手动watch:默认Tape会自动追踪所有可训练变量,手动指定watch(share_embedding)能精准控制追踪对象,减少不必要的计算开销。- 嵌套Tape:内层Tape计算一阶导,外层Tape记录一阶导的计算过程,从而能计算
disc_loss对share_embedding的二阶导,实现梯度从判别器回流到共享特征的效果,完全对应PyTorch中create_graph=True的作用。
内容的提问来源于stack exchange,提问作者Irene Ma

