PyTorch批量重排2D张量:寻求替代CPU循环的向量化方案
向量化实现批量2D张量的索引重排
需求描述
现有尺寸为(batch_size, N, N)的initial_tensor张量,以及尺寸为(batch_size, N)的indexes张量,其中indexes的每一行指定对应批量中2D张量的元素新顺序。需要根据indexes重排批量内各2D张量的元素,替代以下CPU上的嵌套循环实现:
for batch in range(batch_size): old_ids = indexes[batch] for i in range(N): for j in range(N): target[batch][i][j] = initial_tensor[batch][old_ids[i]][old_ids[j]]
向量化解决方案(以PyTorch为例)
利用张量的高级索引特性可以直接实现等价的向量化操作,彻底摆脱循环,同时支持GPU加速:
方法1:使用gather方法分步索引
import torch # 获取张量维度 batch_size, N = initial_tensor.shape[0], initial_tensor.shape[1] # 扩展索引维度,分别适配行和列的索引需求 row_idx = indexes.unsqueeze(2) # 形状变为 (batch_size, N, 1) col_idx = indexes.unsqueeze(1) # 形状变为 (batch_size, 1, N) # 先按行索引,再按列索引完成重排 target = initial_tensor.gather(1, row_idx).gather(2, col_idx)
方法2:直接使用广播式高级索引
import torch # 生成批量维度的索引,保持与其他维度的广播兼容 batch_idx = torch.arange(batch_size)[:, None, None] # 扩展索引维度实现广播,得到(N,N)形状的索引矩阵 row_idx = indexes[:, :, None] col_idx = indexes[:, None, :] # 直接索引得到目标张量 target = initial_tensor[batch_idx, row_idx, col_idx]
原理说明
两种方法本质都是将indexes的每一行扩展为(N,N)的索引矩阵:
- 行索引矩阵中,第
i行的所有元素都是old_ids[i] - 列索引矩阵中,第
j列的所有元素都是old_ids[j]
这样每个位置(i,j)就对应原循环中initial_tensor[batch][old_ids[i]][old_ids[j]]的取值,完全等价于嵌套循环的逻辑,但通过向量化操作可以利用GPU并行计算大幅提升效率。
内容的提问来源于stack exchange,提问作者Denis Sapegin
相关产品推荐
相关产品推荐

