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

如何高效对相同标签对应的张量元素求和(无需遍历标签)

无需遍历标签的快速实现方案

可以利用张量库(如PyTorch)的索引聚合操作来实现,这类操作是底层优化的向量化计算,完全不需要遍历标签值,效率远高于循环实现。以下是具体实现方式:

核心思路

将labels和source展平为一维标签和二维嵌入张量,然后通过index_add_或scatter_add这类聚合函数,直接按标签索引对嵌入向量进行求和。

代码示例(以PyTorch为例)

import torch

# 示例输入
labels = torch.tensor([[0, 1], [1, 2]])
source = torch.tensor([[[0, 1], [1, 2]], [[2, 3], [3, 4]]])

# 1. 展平张量:把(n,m)的labels转为(n*m,),(n,m,embdim)的source转为(n*m, embdim)
flatten_labels = labels.flatten()
flatten_source = source.flatten(0, 1)  # 仅展平前两个维度,保留嵌入维度

# 2. 确定标签范围,初始化输出张量
max_label = flatten_labels.max().item()
out = torch.zeros(max_label + 1, source.size(-1), dtype=source.dtype)

# 3. 按标签索引聚合求和
out.index_add_(0, flatten_labels, flatten_source)

print(out)
# 输出:tensor([[0, 1],
#               [3, 5],
#               [3, 4]])

替代实现(scatter_add)

如果偏好scatter_add,也可以这样写:

out = torch.zeros(max_label + 1, source.size(-1), dtype=source.dtype)
# 将标签扩展为与flatten_source同形状,用于scatter的索引定位
label_indices = flatten_labels.unsqueeze(1).expand_as(flatten_source)
out = out.scatter_add(0, label_indices, flatten_source)

关键说明

  • 这两种方法都是向量化操作,由张量库底层优化,处理大规模张量时性能远优于循环遍历。
  • 如果标签不是从0开始的连续值,输出张量中未出现的标签位置会保持0,可根据需求后续过滤或调整初始化方式。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.21 02:00:12