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

如何使用torch.unique去重计算高价函数并实现可导的张量梯度求解?

解决方案:支持微分的唯一元素优化计算

问题根源在于torch.unique及其返回的inverse_indices不支持自动微分,导致梯度无法通过索引映射传递。以下是两种可行的实现方式:

方案一:自定义Autograd Function(推荐)

通过自定义torch.autograd.Function,手动实现前向计算逻辑和反向梯度传递逻辑,既能保留"仅计算唯一元素昂贵函数"的性能优势,又能支持完整的微分。

import torch

class UniqueExp(torch.autograd.Function):
    @staticmethod
    def forward(ctx, inputs):
        # 前向:提取唯一元素,计算exp后映射回原张量形状
        unique_inputs, inverse_indices = torch.unique(inputs, return_inverse=True)
        unique_exp = torch.exp(unique_inputs)
        full_exp = unique_exp[inverse_indices]
        # 保存反向所需的中间张量
        ctx.save_for_backward(unique_inputs, inverse_indices)
        return full_exp
    
    @staticmethod
    def backward(ctx, grad_output):
        # 反向:将梯度按唯一元素分组求和,再映射回原输入
        unique_inputs, inverse_indices = ctx.saved_tensors
        # 计算每个唯一元素对应的梯度总和,再乘以exp的导数(即自身)
        grad_unique = grad_output[inverse_indices].bincount(minlength=unique_inputs.numel()) * torch.exp(unique_inputs)
        # 将唯一元素的梯度映射回原输入的每个位置
        grad_input = grad_unique[inverse_indices]
        return grad_input

# 测试代码
inputs = torch.rand(100)
inputs = torch.round(inputs, decimals=2)
inputs.requires_grad_(True)

full_exp = UniqueExp.apply(inputs)
# 计算梯度
grad = torch.autograd.grad(full_exp[0], inputs)
print(grad[0][0])  # 输出应为exp(inputs[0]),与直接计算的梯度一致

方案二:手动梯度映射(适合简单场景)

如果不想自定义Function,可在保留原前向逻辑的基础上,手动计算梯度传递路径:

import torch

inputs = torch.rand(100)
inputs = torch.round(inputs, decimals=2)
inputs.requires_grad_(True)

# 前向逻辑(与原代码一致)
unique_inputs, inverse_indices = torch.unique(inputs, return_inverse=True)
unique_exp = torch.exp(unique_inputs)
full_exp = unique_exp[inverse_indices]

# 手动计算梯度(以full_exp[0]的梯度为例)
grad_output = torch.zeros_like(full_exp)
grad_output[0] = 1.0  # 目标位置的梯度设为1

# 计算唯一元素的梯度:每个唯一元素的梯度等于对应所有原元素的梯度之和乘以exp(unique_inputs)
grad_unique = torch.zeros_like(unique_inputs)
for idx in range(unique_inputs.numel()):
    mask = inverse_indices == idx
    grad_unique[idx] = grad_output[mask].sum() * torch.exp(unique_inputs[idx])

# 将梯度映射回原输入张量
grad_input = grad_unique[inverse_indices]
print(grad_input[0])  # 输出与直接计算的梯度一致

说明

  • 方案一的自定义Function效率更高,适合大规模张量场景,且能无缝集成到PyTorch的自动微分流程中。
  • 方案二的手动梯度映射无需自定义类,但循环操作在数据量大时性能较差,仅适合小规模场景。

内容的提问来源于stack exchange,提问作者Peter234

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.15 19:53:27