如何将PyTorch张量的切片求和操作向量化改写?
优化PyTorch列表推导式代码为向量化操作
原代码通过列表推导式结合Python原生sum实现张量计算,不仅代码冗余,还会中断梯度传播(Pythonsum无法追踪张量梯度),同时没利用PyTorch的向量化运算优势。可以通过张量重塑+内置聚合操作完全替代循环,优化后代码如下:
import torch n = 10 y = torch.rand(n ** 2, requires_grad=True) # 核心:将一维张量重塑为n×n的二维张量,简化行/列维度的聚合操作 y_2d = y.reshape(n, n) # 每行求和后减1(对应原one_node_per_position) one_node_per_position = y_2d.sum(dim=1) - 1 # 每列求和后减1(对应原one_node_per_point) one_node_per_point = y_2d.sum(dim=0) - 1 # 相邻行的和做差(对应原connectivity) connectivity = y_2d.sum(dim=1).diff()
关键优化点说明:
- 梯度保留:全程使用PyTorch张量的内置
sum()和diff()方法,完整保留y的梯度追踪能力,原代码用torch.FloatTensor包裹列表推导式会丢失梯度。 - 向量化效率:摆脱Python循环,利用PyTorch底层的C++优化实现批量计算,运行速度远高于列表推导式,尤其当
n较大时差异明显。 - 代码可读性:通过二维张量的维度语义(行/列)直接表达计算逻辑,比切片循环更直观。
内容的提问来源于stack exchange,提问作者Nourless
相关产品推荐
相关产品推荐

