使用PyTorch优化customIndexAdd函数,移除循环提升性能
优化 customIndexAdd 函数,移除循环提升效率
原函数的三重Python循环是核心性能瓶颈——Python循环在处理百万级维度时会产生巨大的开销,PyTorch的底层张量操作基于C++实现,完全可以通过向量化操作替代循环,把计算压到底层执行。
优化方案:利用 PyTorch 的 scatter_add_ 实现向量化累加
scatter_add_ 可以直接按指定索引对张量进行批量累加操作,完美匹配你要实现的 -2 维度的 index_add_ 逻辑:
import torch import numpy as np import time def customIndexAdd(x1, index, tensor): # 扩展index维度以匹配tensor形状,让广播机制生效 index_expanded = index.view(1, 1, -1, 1).expand_as(tensor) # 在-2维度执行批量累加 x1.scatter_add_(dim=-2, index=index_expanded, src=tensor) return x1 # 原测试代码保持不变 sequential_numbers = np.arange(1, 2*2*352798*2 + 1) tensor = sequential_numbers.reshape(2, 2, 352798, 2) t = torch.tensor(tensor).int() values = torch.arange(1, 352796 // 2 + 1) repeated_values = torch.repeat_interleave(values, repeats=2) final_values = torch.cat([torch.tensor([0]), repeated_values, torch.tensor([176399])]) index = final_values x = torch.ones(2, 2, 176400, 2).int() x.index_add_(-2, index, t) x1 = torch.ones(2, 2, 176400, 2).int() start = time.time() out1 = customIndexAdd(x1, index, t) end = time.time() print(f"优化后耗时: {end - start:.4f} 秒") print(torch.equal(x, out1))
优化原理
- 维度对齐:将
index扩展为和tensor相同的形状(通过view和expand_as),确保scatter_add_能在每个对应位置找到要累加的目标索引。 - 底层加速:
scatter_add_是PyTorch内置的底层操作,完全绕过Python循环,利用CPU/GPU的并行计算能力,对于大张量的提速效果极其明显。 - 内存效率:操作直接在输入张量
x1上原地修改(和原函数逻辑一致),避免额外内存开销。
额外提速建议
如果你的设备支持CUDA,把张量转移到GPU上执行会带来更显著的速度提升——只需在创建张量时加上.cuda():
t = torch.tensor(tensor).int().cuda() index = final_values.cuda() x = torch.ones(2, 2, 176400, 2).int().cuda() x1 = torch.ones(2, 2, 176400, 2).int().cuda()
内容的提问来源于stack exchange,提问作者Sai krishna
相关产品推荐
相关产品推荐

