实现带值阈值的Torch Softmax遇RuntimeError求解决方案
解决带阈值Softmax的梯度计算inplace错误
错误根源
你遇到的RuntimeError是因为inplace操作修改了梯度计算图中的张量:Clip模块里的x[i] = new_x_i直接修改了Softmax输出的张量(该张量属于反向传播依赖的计算图节点),导致反向传播时无法找到原始版本的张量进行梯度计算。
修改方案
核心是避免inplace修改输入张量,通过创建新张量来实现阈值裁剪逻辑,同时用向量化操作替代循环,提升效率并确保计算图完整。
修改后的代码
import torch import torch.nn as nn class Clip(nn.Module): def __init__(self, threshold): super().__init__() self.threshold = threshold def forward(self, x): # 复制输入张量,完全避免修改原计算图中的张量 x_clipped = x.clone() # 生成每个元素是否大于等于阈值的掩码 mask_above_threshold = x >= self.threshold # 计算每个样本中大于阈值部分的总和 sum_above = torch.sum(x * mask_above_threshold, dim=-1, keepdim=True) # 筛选出需要调整的样本(存在大于阈值的元素) need_adjust = sum_above > 1e-8 # 用小epsilon避免除0 if need_adjust.any(): # 对大于阈值的部分进行归一化 x_clipped[mask_above_threshold] = x[mask_above_threshold] / sum_above.expand_as(x)[mask_above_threshold] # 小于阈值的部分先设为阈值 x_clipped[~mask_above_threshold] = self.threshold # 重新归一化确保总和为1 current_sum = torch.sum(x_clipped, dim=-1, keepdim=True) scale = 1.0 / current_sum x_clipped = x_clipped * scale return x_clipped class GradientPolicy(nn.Module): def __init__(self, n_cols=5, time_window=50, threshold = 0.2): """DDPG policy network initializer.""" super().__init__() self.threshold = threshold self.sequential = nn.Sequential( nn.Conv2d(in_channels=n_cols, out_channels=2, kernel_size=(1, 3)), nn.ReLU(), nn.Conv2d(in_channels=2, out_channels=20, kernel_size=(1, time_window-2)), nn.ReLU() ) self.final_convolution = nn.Conv2d(in_channels=21, out_channels=1, kernel_size=(1, 1)) self.softmax = nn.Softmax(dim=-1) # 不需要嵌套Sequential self.clip = Clip(threshold) def mu(self, x, x2, x3): output = self.sequential(x) output = torch.cat([output, x2], dim=1) output = self.final_convolution(output) output = torch.cat([output, x3], dim=2) output = torch.squeeze(output, 3) output = torch.squeeze(output, 1) # shape [N, M + 1] output = self.softmax(output) output = self.clip(output) return output def forward(self, observation, last_action): mu = self.mu(observation, last_action, last_action) action = mu.cpu().detach().numpy().squeeze() return action
关键修改点说明
- 移除inplace操作:用
x.clone()创建输入张量的副本,所有修改都在副本上进行,原张量(计算图依赖的节点)保持不变。 - 向量化替代循环:用张量掩码和广播操作实现逐样本的阈值处理,比循环更高效,同时避免了逐元素修改带来的inplace风险。
- 确保总和为1:在裁剪后重新计算总和并做缩放,严格保证输出满足概率分布的总和为1的要求。
- 简化Softmax定义:直接使用
nn.Softmax(dim=-1)替代嵌套的nn.Sequential,减少不必要的模块嵌套。
内容的提问来源于stack exchange,提问作者leeway00
相关产品推荐
相关产品推荐

