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

TensorFlow中两个拆分CNN模型如何手动传递梯度联合训练

拆分双模型手动反向传播实现方案

问题根因

你当前写法的核心问题是两个GradientTape的计算完全独立,model_0的梯度计算没有和最终分类损失挂钩,相当于没有拿到model_1回传的上游梯度,自然无法实现端到端的训练效果。

实现方案

核心逻辑是先通过tape_1计算损失对model_0输出的上游梯度,再将该梯度传入tape_0计算model_0参数的实际更新梯度,完全符合链式法则要求,且保留两个独立的tape:

for epoch in range(epochs):
    for step, (x_batch, y_batch) in enumerate(train_generator):
        # tape_0 追踪输入到model_0输出的计算路径
        with tf.GradientTape() as tape_0:
            y_pred_0 = model_0(x_batch, training=True)

        # tape_1 追踪model_0输出到最终损失的计算路径
        with tf.GradientTape() as tape_1:
            # 手动监控model_0的输出,用于提取上游梯度
            tape_1.watch(y_pred_0)
            y_pred_1 = model_1(y_pred_0, training=True)
            loss_value = loss_fn(y_batch, y_pred_1)
        
        # 同时计算model_1参数梯度 + 损失对model_0输出的上游梯度
        grads_1, grads_upstream = tape_1.gradient(
            loss_value,
            [model_1.trainable_weights, y_pred_0]
        )
        # 将上游梯度传入tape_0,计算model_0参数的实际梯度
        grads_0 = tape_0.gradient(
            y_pred_0,
            model_0.trainable_weights,
            output_gradients=grads_upstream
        )

        # 统一更新参数,和端到端训练效果完全一致
        optimizer.apply_gradients(zip(grads_1, model_1.trainable_weights))
        optimizer.apply_gradients(zip(grads_0, model_0.trainable_weights))

原理解释

  1. tape_1.watch(y_pred_0)是因为y_pred_0不是在tape_1上下文内生成的张量,需要手动声明监控才能计算损失对它的梯度
  2. output_gradients参数是GradientTape提供的原生接口,作用是传入反向传播的上游梯度值,自动和当前路径的梯度做乘积得到最终参数梯度,等价于完整计算图的链式求导结果
  3. 该写法和直接拼接两个模型训练的更新效果完全一致,同时满足你保留两个独立tape、手动实现反向传播的要求。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 02:54:04