如何使用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
相关产品推荐
相关产品推荐

