PyTorch可微分基于索引的条件求和实现方法问询
用PyTorch实现可微分的按索引分组求和
当然有!PyTorch原生就提供了完全可微分的分组求和功能,正好匹配你的需求——torch.scatter_add_(或其非原地版本torch.scatter_add)就是你要找的类似scatter的sum操作,完美替代不可微分的掩码求和方法。
具体实现代码
先把你的输入转换成PyTorch张量,然后通过scatter_add_完成分组求和:
import torch # 你的输入数据 idx = torch.tensor([0, 1, 0, 2, 3, 1]) data = torch.tensor([[0, 1, 2], [3, 4, 5], [6, 7, 8], [9, 10, 11], [12, 13, 14], [15, 16, 17]]) # 初始化输出张量:形状为(max_idx+1, 特征维度),初始值为0 max_idx = idx.max().item() output = torch.zeros(max_idx + 1, data.shape[1], dtype=data.dtype, device=data.device) # 执行可微分的分组求和 # 关键:将idx扩展为与data相同的维度,确保每个元素都能映射到output的对应行 output.scatter_add_( dim=0, # 沿第0维度(行)进行聚合 index=idx.unsqueeze(1).expand_as(data), # 匹配data维度的索引矩阵 src=data # 待聚合的原始数据 ) print(output)
输出结果
运行后会得到你期望的4×3数组:
tensor([[ 6, 8, 10], [18, 20, 22], [ 9, 10, 11], [12, 13, 14]])
关键参数解释
dim=0:指定沿着行维度(第0维)进行聚合,把data中对应idx的行累加到output的对应行。index:必须和src(即data)维度一致,所以我们先把一维的idx转为列向量,再扩展成和data一样的6×3形状,这样每个元素都能准确找到要累加的目标位置。src:就是需要分组求和的原始特征数据。
额外说明
这个操作是完全可微分的,如果你后续需要对output进行反向传播,梯度会正确传递到data张量上,完美解决你之前掩码方法不可微分的痛点。
另外,在PyTorch 1.12及以上版本,你也可以使用torch.scatter并指定reduce='sum'参数来实现相同效果:
output = torch.scatter( torch.zeros(max_idx + 1, data.shape[1]), dim=0, index=idx.unsqueeze(1).expand_as(data), src=data, reduce='sum' )
内容的提问来源于stack exchange,提问作者Mehdi Saman Booy
相关产品推荐
相关产品推荐

