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

使用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))

优化原理

  1. 维度对齐:将index扩展为和tensor相同的形状(通过view和expand_as),确保scatter_add_能在每个对应位置找到要累加的目标索引。
  2. 底层加速:scatter_add_是PyTorch内置的底层操作,完全绕过Python循环,利用CPU/GPU的并行计算能力,对于大张量的提速效果极其明显。
  3. 内存效率:操作直接在输入张量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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 08:45:22