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

PyTorch中基于类别张量高效实现数据按类分组求和(低内存)

内存高效实现按类别张量求和(无需显式One-Hot编码)

问题描述

现有两个形状均为(B, 1, N)的张量:

  • x:数据张量
  • y:类别标签张量,取值范围为0到C-1(对应示例代码中的设置)

需要生成形状为(B, C)的张量z,满足:

  1. 将x中的数据按y的类别划分到对应通道
  2. 沿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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.26 01:30:09