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和输出层的参数,无法优化前端的定位逻辑。
- 输入到每个ROI模块的 landmarks 都做了
- 原SLPT代码保留了整个前向过程的计算图连续性,梯度可以从最后一个stage一直回传到初始特征提取模块和所有ROI模块,让所有组件都能通过损失迭代优化,最终收敛到高精度。而你的修改破坏了这种连续性,模型失去了关键组件的优化能力,自然无法达到预期精度。
内容的提问来源于stack exchange,提问作者FurxDev
相关产品推荐
相关产品推荐

