PyTorch传入前层激活最大值归一化ReLU触发二次反向传播报错
错误原因
报错的核心是你存在self.thresh_list里的层最大激活值是带autograd计算图历史的张量:
- 每次迭代执行
loss.backward()后,PyTorch会默认释放当前迭代前向传播生成的所有计算图中间缓存,避免显存占用过高 - 你没有对存到
thresh_list的最大值做任何切断计算图的处理,下一次迭代前向传播时,旧迭代生成的、已经被释放的计算图节点会被接入新的计算流 - 反向传播遍历新计算图时,会试图访问已经被释放的旧节点,直接触发二次反向传播的报错
当你把prev_layer_max换成固定值1时,常量不属于任何历史计算图,自然不会触发该问题。
修复方法
不要使用给loss.backward()加retain_graph=True的方案,该方案会禁止PyTorch释放计算图缓存,导致显存随迭代持续上涨,最终触发显存溢出,仅适合特殊的多损失反向场景,不适用常规训练。
正确修复只需要保证存入thresh_list的值是脱离计算图的纯数值即可:
- 模型初始化时,将
thresh_list初始化为存普通浮点数的列表,不要存张量 - 每次记录层最大激活值时,调用
.item()将单元素张量转为普通Python数值,彻底切断和历史计算图的关联
修正后的前向传播阈值记录代码如下:
def forward(self, input, epoch): x = self.conv1(input) x = self.relu(x, 1) # 取纯数值存入阈值列表,不保留计算图关联 self.thresh_list[0] = max(self.thresh_list[0], x.max().item()) x = self.conv_dropout(x) x = self.conv2(x) x = self.relu(x, self.thresh_list[0]) self.thresh_list[1] = max(self.thresh_list[1], x.max().item()) x = self.pool1(x) x = self.conv_dropout(x) x = self.conv3(x) x = self.relu(x, self.thresh_list[1]) self.thresh_list[2] = max(self.thresh_list[2], x.max().item()) # 后续原有逻辑保持不变
对应模型__init__方法中,thresh_list初始化代码修改为:
# 根据你的层数量对应初始化长度,初始值设为0.0即可 self.thresh_list = [0.0, 0.0, 0.0]
另外你的训练循环存在一处隐性bug:labels.to(device)不是原地操作,这行代码不会把标签迁移到GPU上,需要修改为labels = labels.to(device),避免后续计算损失时出现设备不匹配错误。
内容的提问来源于stack exchange,提问作者Musheer Abdullah
相关产品推荐
相关产品推荐

