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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.14 07:43:40