PyTorch中torch.argsort()反向传播问题:重排feature时卷积权重不更新
我想利用torch.argsort()的输出生成新张量,示例代码如下:
def reodering (self, x) : feature = x.clone().detach().requires_grad_(True) weight = self.feature_conv(feature) indices = weight.argsort(dim = 1) sorted_feature = torch.gather(input=feature, dim=1, index=indices) return sorted_feature
其中x是尺寸为(b, c, h, w)的张量,self.feature_conv()是nn.Conv2d()层,损失函数采用MSELoss(sorted_feature, label)。但运行代码时,self.feature_conv()的权重无法更新。
训练代码示例如下:
if __name__ == '__main__': ... model = reodering() criterion = nn.MSELoss() optimizer = optim.Adam(model.parameters(), lr=args.lr) ... #-------------------- # training #-------------------- for epoch in range(args.num_epochs): model.train() for i, data in enumerate(inputs_set, start = 1) : inputs = np.load(os.path.join(inputs_path, data)) labels = np.load(os.path.join(labels_path, data)) inputs = torch.from_numpy(inputs).to(device) labels = torch.from_numpy(labels).to(device) preds = model(inputs) loss = criterion(preds, labels) optimizer.zero_grad() loss.backward() optimizer.step() ...
输入和标签是分割为17x17的图像张量。我知道PyTorch中torch.long类型不支持反向传播,尝试把torch.gather()的输入改成weight时权重能正常更新,但我需要对feature进行重排,请问有其他可行方法吗?
问题核心是你手动对feature执行了detach()操作,切断了它与输入x的梯度链路;同时argsort返回的indices是整数类型,本身无法参与反向传播,导致feature_conv的梯度无法回传。以下是两个可行的解决思路:
方法1:用可微分软排序替代硬排序
硬排序(argsort)不可微分,我们可以用软排序近似,它能保留梯度信息。如果不想引入第三方库,可基于Gumbel-Softmax实现近似硬排序:
def reodering (self, x) : weight = self.feature_conv(x) b, c, h, w = weight.shape # 将h、w维度合并,方便排序处理 weight_flat = weight.view(b, c, h*w) x_flat = x.view(b, c, h*w) # 先得到硬排序的索引 sorted_indices = torch.argsort(weight_flat, dim=1) # 生成one-hot矩阵并通过Gumbel-Softmax软化,保留梯度 one_hot = torch.zeros_like(weight_flat).scatter(1, sorted_indices, 1.0) soft_one_hot = torch.nn.functional.gumbel_softmax(one_hot.log(), tau=0.1, hard=False, dim=1) # 用软化后的权重对x进行重排,再恢复原维度 sorted_feature_flat = torch.einsum('bci,bci->bci', soft_one_hot, x_flat) sorted_feature = sorted_feature_flat.view(b, c, h, w) return sorted_feature
其中tau是温度系数,值越小越接近硬排序,同时梯度稳定性会稍差,可根据训练情况调整。
方法2:手动注册梯度钩子修复回传路径
如果必须使用硬排序,可以通过梯度钩子手动将sorted_feature的梯度映射回weight,从而让梯度正常回传到feature_conv:
def reodering (self, x) : weight = self.feature_conv(x) indices = weight.argsort(dim = 1) sorted_feature = torch.gather(input=x, dim=1, index=indices) # 注册梯度钩子,将sorted_feature的梯度映射回weight的原始位置 def grad_hook(grad): # indices.argsort(dim=1)能得到排序前的位置索引 grad_weight = torch.gather(grad, dim=1, index=indices.argsort(dim=1)) return grad_weight weight.register_hook(grad_hook) return sorted_feature
同时注意:你原代码中feature = x.clone().detach().requires_grad_(True)完全多余,直接使用x即可,detach()是导致梯度链路断裂的原因之一。
内容的提问来源于stack exchange,提问作者박병주

