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
相关产品推荐
相关产品推荐

