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

PyTorch实现2D张量按索引聚合:重复索引取均值其余保留

PyTorch按重复索引对2D张量分组聚合求均值实现

核心逻辑分三步:

  • 将二维索引idx展平为一维分组标签,长度和输入张量x的行数完全对齐,每个标签对应该行所属的聚合分组
  • 统计每个分组下所有x行的元素和、分组包含的行数
  • 分组和除以分组行数得到均值,按索引升序排列即为最终结果。不需要单独写分支判断索引是否重复:无重复索引的分组行数为1,计算结果就是原值;重复索引的分组自动得到平均值,完全匹配需求规则。

通用兼容实现代码

兼容所有PyTorch 1.0以上版本,支持CPU/GPU张量、自动微分:

import torch

# 题面定义的输入张量
idx = torch.tensor([[0, 1, 2], [1, 2, 3], [3, 4, 5]])
x = torch.tensor([[10, 10, 10], [11, 11, 11], [12, 12, 12],
                  [13, 13, 13], [14, 14, 14], [15, 15, 15],
                  [16, 16, 16], [17, 17, 17], [18, 18, 18]])

# 展平二维索引为一维分组标签
flat_idx = idx.flatten()
# 按升序获取所有唯一索引,确定输出顺序
unique_idx = torch.unique(flat_idx)
# 初始化分组求和、计数缓冲
sum_buffer = torch.zeros((len(unique_idx), x.shape[1]), dtype=torch.float32, device=x.device)
count_buffer = torch.zeros(len(unique_idx), dtype=torch.float32, device=x.device)
# 按索引累加对应行的值、统计每组样本数
sum_buffer.scatter_add_(0, flat_idx.unsqueeze(1).expand_as(x), x.to(torch.float32))
count_buffer.scatter_add_(0, flat_idx, torch.ones_like(flat_idx, dtype=torch.float32))
# 计算分组均值
result = sum_buffer / count_buffer.unsqueeze(1)

运行后输出的result与题面期望结果完全一致:

tensor([[10.0000, 10.0000, 10.0000],
        [12.0000, 12.0000, 12.0000],
        [13.0000, 13.0000, 13.0000],
        [15.5000, 15.5000, 15.5000],
        [17.0000, 17.0000, 17.0000],
        [18.0000, 18.0000, 18.0000]])

简化版本(PyTorch 1.12+)

高版本PyTorch提供了自带均值聚合的index_reduce_接口,可以省去手动计数步骤:

flat_idx = idx.flatten()
group_num = flat_idx.max().item() + 1
result = torch.zeros((group_num, x.shape[1]), dtype=torch.float32, device=x.device).index_reduce_(
    dim=0,
    index=flat_idx.unsqueeze(1).expand_as(x),
    source=x.to(torch.float32),
    reduce="mean",
    include_self=False
)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.29 06:18:22