You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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,提问作者박병주

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.28 02:13:29