PyTorch中基于类别张量高效实现数据按类分组求和(低内存)
内存高效实现按类别张量求和(无需显式One-Hot编码)
问题描述
现有两个形状均为(B, 1, N)的张量:
x:数据张量y:类别标签张量,取值范围为0到C-1(对应示例代码中的设置)
需要生成形状为(B, C)的张量z,满足:
- 将
x中的数据按y的类别划分到对应通道 - 沿
N维度对同一类别的数据求和
常规的One-Hot编码实现会创建(B, C, N)的中间张量,当C或N较大时会快速耗尽GPU内存,因此需要更内存高效的实现方式。
常规实现(内存开销大)
import torch B, C, N = 2, 10, 1000 x = torch.randn(B, 1, N) y = torch.randint(low=0, high=C, size=(B, 1, N)) one_hot = torch.nn.functional.one_hot(y, C) # 形状: (B, 1, N, C) one_hot = one_hot.squeeze().permute(0, -1, 1) # 形状: (B, C, N) z = x * one_hot # 形状: (B, C, N) z = z.sum(-1) # 形状: (B, C)
高效解决方案
可以使用PyTorch的torch.scatter_add_方法,直接通过索引完成按类别求和,无需创建庞大的One-Hot中间张量。
实现代码
import torch B, C, N = 2, 10, 1000 x = torch.randn(B, 1, N) y = torch.randint(low=0, high=C, size=(B, 1, N)) # 初始化结果张量为全0,形状(B, C) z = torch.zeros(B, C, device=x.device) # 调整张量形状后,通过scatter_add_直接按类别累加 z.scatter_add_( dim=1, index=y.squeeze(1), # 调整为(B, N),每个位置对应类别索引 src=x.squeeze(1) # 调整为(B, N),每个位置对应待累加的数据 )
原理说明
scatter_add_的核心是直接通过索引将源张量的值累加到目标张量的对应位置,完全避免了生成稀疏的One-Hot中间张量。- 原One-Hot方法需要占用
B*C*N的内存,而高效方法仅需要B*C + B*N的内存(目标张量z加上源张量x、y),当C较大时,内存开销的差距会非常显著。
正确性验证
可以对比两种方法的结果,确认输出一致:
# 常规方法生成的结果 one_hot = torch.nn.functional.one_hot(y, C).squeeze().permute(0, -1, 1) z_onehot = (x * one_hot).sum(-1) # 高效方法生成的结果 z_scatter = torch.zeros(B, C, device=x.device) z_scatter.scatter_add_(1, y.squeeze(1), x.squeeze(1)) # 验证结果是否一致 print(torch.allclose(z_onehot, z_scatter)) # 输出: True
内容的提问来源于stack exchange,提问作者Vivek
相关产品推荐
相关产品推荐

