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

PyTorch前向传播重分配变量为何导致模型训练不收敛?

SLPT模型修改前向传播后无法收敛的原因分析

问题描述

我是PyTorch新手,在使用开源项目SLPT训练模型时遇到问题。修改前向传播的变量重分配逻辑后,训练很快稳定在NME≈30%,无法收敛到预期的4.1%;但原代码训练能达到预期效果,想询问两种情况训练结果不同的原因。

修改后的前向传播代码

# stage_1
ROI_anchor_1, bbox_size, start_anchor = self.ROI_1(initial_landmarks.detach())
ROI_anchor_1 = ROI_anchor_1.view(bs, self.num_point * self.Sample_num * self.Sample_num, 2)
ROI_feature_1 = self.interpolation(feature_map, ROI_anchor_1.detach()).view(bs, self.num_point, self.Sample_num,
                                                                            self.Sample_num, self.d_model)
ROI_feature_1 = ROI_feature_1.view(bs * self.num_point, self.Sample_num, self.Sample_num,
                                 self.d_model).permute(0, 3, 2, 1)

transformer_feature_1 = self.feature_extractor(ROI_feature_1).view(bs, self.num_point, self.d_model)

offset_1 = self.Transformer(transformer_feature_1)
offset_1 = self.out_layer(offset_1)

landmarks_1 = start_anchor.unsqueeze(1) + bbox_size.unsqueeze(1) * offset_1
output_list.append(landmarks_1)

# stage_2
ROI_anchor_2, bbox_size, start_anchor = self.ROI_2(landmarks_1[:, -1, :, :].detach())
ROI_anchor_2 = ROI_anchor_2.view(bs, self.num_point * self.Sample_num * self.Sample_num, 2)
ROI_feature_2 = self.interpolation(feature_map, ROI_anchor_2.detach()).view(bs, self.num_point, self.Sample_num,
                                                                             self.Sample_num, self.d_model)
ROI_feature_2 = self.feature_extractor(ROI_feature_2).view(bs, self.num_point, self.d_model)

offset_2 = self.Transformer(transformer_feature_2)
offset_2 = self.out_layer(offset_2)

landmarks_2 = start_anchor.unsqueeze(1) + bbox_size.unsqueeze(1) * offset_2
output_list.append(landmarks_2)

# stage_3
ROI_anchor_3, bbox_size, start_anchor = self.ROI_3(landmarks_2[:, -1, :, :].detach())
ROI_anchor_3 = ROI_anchor_3.view(bs, self.num_point * self.Sample_num * self.Sample_num, 2)
ROI_feature_3= self.interpolation(feature_map, ROI_anchor_3.detach()).view(bs, self.num_point, self.Sample_num,
                                                                               self.Sample_num, self.d_model)
ROI_feature_3 = self.feature_extractor(ROI_feature_3).view(bs, self.num_point, self.d_model)

offset_3 = self.Transformer(transformer_feature_3)
offset_3 = self.out_layer(offset_3)

landmarks_3 = start_anchor.unsqueeze(1) + bbox_size.unsqueeze(1) * offset_3
output_list.append(landmarks_3)

原因分析

  • 多处滥用detach()切断了梯度反向传播路径:
    • 输入到每个ROI模块的 landmarks 都做了detach(),比如self.ROI_1(initial_landmarks.detach())、ROI_2(landmarks_1[:, -1, :, :].detach()),导致ROI模块的参数无法通过后续损失进行更新。ROI模块是SLPT中负责精准定位特征区域的核心组件,参数无法优化直接影响后续特征提取质量。
    • 对ROI_anchor_1/2/3做detach()后再传入插值模块,使得插值得到的ROI特征完全脱离计算图,后续Transformer和输出层的梯度无法回传到插值模块和ROI模块,模型只能更新Transformer和输出层的参数,无法优化前端的定位逻辑。
  • 原SLPT代码保留了整个前向过程的计算图连续性,梯度可以从最后一个stage一直回传到初始特征提取模块和所有ROI模块,让所有组件都能通过损失迭代优化,最终收敛到高精度。而你的修改破坏了这种连续性,模型失去了关键组件的优化能力,自然无法达到预期精度。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 08:20:32