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

多任务对抗训练中,如何在TensorFlow实现PyTorch式高阶梯度反向传播

问题:TensorFlow中实现多任务对抗梯度对齐的高阶导数问题

我的目标是在多任务模型中,借助对抗学习思路,将不同任务的梯度输入判别器,使这些梯度在统计上无法区分,以此约束梯度对齐。但在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.19 05:39:53