如何高效生成逐行循环移位首行的2D张量或numpy数组?
循环移位堆叠3D张量的高效实现方案
原实现中for循环+逐次拼接的写法开销很高,可通过向量化索引完全消除循环,实现数倍到数十倍的性能提升,具体实现如下:
PyTorch 版本实现
import torch import numpy as np a = np.random.rand(33, 11) miss_size = 64 lp_order = a.shape[1] - 1 inv_a = -np.flip(a, axis=1) mtx_size = miss_size + lp_order # 初始行向量生成逻辑和原代码完全一致 mtx_row = torch.cat((torch.from_numpy(inv_a), torch.zeros((a.shape[0], miss_size - 1 + a.shape[1]))), dim=1) # 以下替换原for循环逻辑 seq_length = mtx_row.shape[1] # 生成移位索引矩阵,广播适配维度 shift_idx = (torch.arange(seq_length) - torch.arange(mtx_size + 1).unsqueeze(1)) % seq_length # 直接索引得到完整3D张量,和原输出维度、数值完全一致 mtx_full = mtx_row[:, shift_idx]
实现原理
- 核心利用模运算直接计算每个移位步对应的原始元素索引,一次性完成所有移位结果的提取,完全避免了循环移位、逐次张量拼接的冗余开销
- 向量化操作底层调用框架优化后的算子,CPU下性能远超Python级循环,GPU下可完全并行计算,
mtx_size越大性能提升越明显 - 输出结果与原代码完全等价,可直接替换使用,不需要调整后续逻辑
NumPy 版本实现(可选)
import numpy as np a = np.random.rand(33, 11) miss_size = 64 lp_order = a.shape[1] - 1 inv_a = -np.flip(a, axis=1) mtx_size = miss_size + lp_order mtx_row = np.concatenate([inv_a, np.zeros((a.shape[0], miss_size - 1 + a.shape[1]))], axis=1) seq_length = mtx_row.shape[1] shift_idx = (np.arange(seq_length) - np.arange(mtx_size + 1)[:, None]) % seq_length mtx_full = mtx_row[:, shift_idx]
内容的提问来源于stack exchange,提问作者jBloodless
相关产品推荐
相关产品推荐

