如何高效对相同标签对应的张量元素求和(无需遍历标签)
无需遍历标签的快速实现方案
可以利用张量库(如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
相关产品推荐
相关产品推荐

