PyTorch连续网络中间数据处理与反向传播耗时问题求助
问题解答
反向传播是否会重跑function()?
不会直接重跑function(),但耗时根源在于function()的纯Python循环没有被PyTorch计算图高效处理。
PyTorch反向传播依赖前向过程构建的计算图,你用三重Python循环实现的邻域计算,每一步赋值操作都会被记录为计算图中的独立节点,导致计算图异常庞大。反向传播时需要遍历这海量节点逐一计算梯度,这才是耗时20分钟的原因——不是重跑function(),而是梯度计算的复杂度被循环拉到了极致。
优化方案
用PyTorch张量向量化操作替代三重循环(最核心优化)
你的循环本质是提取output1的5x5邻域,计算后赋值到input2的指定位置,完全可以用张量操作实现并行计算:import torch.nn.functional as F def function(self): batch_size, C, H, W = self.output1.shape kernel_size = 5 padding = 2 # 保证边缘位置能取到完整5x5邻域 # 提取所有位置的5x5邻域,形状变为[batch_size, C*5*5, H*W] unfolded = F.unfold(self.output1, kernel_size=kernel_size, padding=padding) # 生成idx对应的(i,j)索引对(i<=j) idx_pairs = torch.combinations(self.idx, r=2) # 把二维索引转为unfolded对应的线性索引 pos = idx_pairs[:, 0] * W + idx_pairs[:, 1] # 筛选出需要计算的邻域 selected_neighbors = unfolded[:, :, pos] # 批量应用some_criterion(确保该函数支持张量批量运算) criterion_results = some_criterion(selected_neighbors) # 初始化input2并赋值 input2 = torch.zeros(batch_size, H, W, device=self.output1.device) input2[:, idx_pairs[:, 0], idx_pairs[:, 1]] = criterion_results return input2向量化操作会利用CUDA并行计算,同时大幅减少计算图的节点数量,反向传播速度会呈数量级提升。
冻结net1的梯度(若无需再训练net1)
既然net1已预训练完成,如果训练net2时不需要更新net1的参数,直接冻结其参数的梯度计算:for param in net1.parameters(): param.requires_grad = False这样反向传播时不会计算
net1的梯度,减少不必要的计算开销。分离function()的计算图(若允许梯度不回传至net1)
如果some_criterion的计算不需要让梯度回传到net1,可以对output1做detach(),切断这部分的计算图连接:def function(self): output1_detached = self.output1.detach() # 后续用output1_detached进行计算,不再参与梯度回传此操作会让
function()的计算不加入主计算图,反向传播时不会处理这部分梯度,但仅适用于net1无需接收net2梯度的场景。
内容的提问来源于stack exchange,提问作者Yajin
相关产品推荐
相关产品推荐

