如何向量化这段使用二维间接索引的PyTorch代码片段?
批量成对交互计算的高效向量化实现
我有一段用三重循环实现的代码,无法直接改写成向量化形式:
for b in range(B): for i in range(M - 1): for j in range(i + 1, M): interaction = (vectors[b, fields[b, i], fields[b, j], :] * vectors[b, fields[b, j], fields[b, i], :]).sum() pairwise[b] += interaction
其中:
fields是形状为 B×M 的数组vectors是形状为 B×M×M×D 的数组- 输出
pairwise是形状为 B 的数组
我原本考虑用gather/scatter方法,但它们仅支持一维索引。有没有高效的实现方式?哪怕无法完全向量化,仅在B维度实现向量化也能满足需求。
更新:我尝试过一种向量化实现,但运行速度极慢:
comb_indices = torch.tril_indices(M, M, -1) components = torch.arange(D, device=vectors.device).view(1, 1, -1) batches = torch.arange(B, device=vectors.device).view(-1, 1, 1) this_indices = fields[:, comb_indices[0, :]] that_indices = fields[:, comb_indices[1, :]] linearized_this_fields = (this_indices + M * that_indices).view(B, -1, 1) linearized_that_fields = (that_indices + M * this_indices).view(B, -1, 1) linearized_fields = vectors \ .permute([0, 3, 1, 2]) \ .reshape(batch_size, self.field_dim, M * M) \ .permute([0, 2, 1]) this = linearized_fields[batches, linearized_this_fields, components] that = linearized_fields[batches, linearized_that_fields, components] pairwise = (this * that).sum(dim=[-1, -2])
高效实现方案
可以利用PyTorch的广播和优化索引特性,避免冗余的维度变换,充分利用硬件加速。核心思路是直接批量提取所有需要的交互对张量,再通过点积求和得到结果。
实现代码
import torch B, M, D = vectors.shape[0], vectors.shape[1], vectors.shape[3] # 生成所有i<j的索引对,形状为(2, K),其中K = M*(M-1)/2 i_idx, j_idx = torch.tril_indices(M, M, -1) # 批量获取每个样本的fields[i]和fields[j],形状均为(B, K) f_i = fields[:, i_idx] f_j = fields[:, j_idx] # 批量提取对应交互的张量: # vec1: (B, K, D) 对应 vectors[b, f_i[b,k], f_j[b,k], :] vec1 = vectors[torch.arange(B)[:, None], f_i, f_j] # vec2: (B, K, D) 对应 vectors[b, f_j[b,k], f_i[b,k], :] vec2 = vectors[torch.arange(B)[:, None], f_j, f_i] # 计算所有交互项的点积,再在K(对数量)和D(特征维度)上求和 pairwise = (vec1 * vec2).sum(dim=[1, 2])
方案优势
- 减少中间操作:避免原方案中多次permute、reshape带来的内存开销和延迟,直接通过高级索引提取目标张量。
- 利用PyTorch索引优化:
torch.arange(B)[:, None]配合批量索引的方式是框架优化过的,比手动线性化索引效率更高。 - 降低内存占用:无需创建大尺寸的中间张量(如原方案的
linearized_fields),仅生成必要的交互对张量,节省内存带宽。
正确性验证
可以用小批量测试数据对比原循环和向量化代码的输出:
# 生成测试数据 B_test, M_test, D_test = 2, 3, 4 fields_test = torch.randint(0, M_test, (B_test, M_test)) vectors_test = torch.randn(B_test, M_test, M_test, D_test) # 原循环计算 pairwise_loop = torch.zeros(B_test) for b in range(B_test): for i in range(M_test - 1): for j in range(i + 1, M_test): interaction = (vectors_test[b, fields_test[b,i], fields_test[b,j], :] * vectors_test[b, fields_test[b,j], fields_test[b,i], :]).sum() pairwise_loop[b] += interaction # 向量化计算 i_idx, j_idx = torch.tril_indices(M_test, M_test, -1) f_i = fields_test[:, i_idx] f_j = fields_test[:, j_idx] vec1 = vectors_test[torch.arange(B_test)[:, None], f_i, f_j] vec2 = vectors_test[torch.arange(B_test)[:, None], f_j, f_i] pairwise_vec = (vec1 * vec2).sum(dim=[1,2]) # 验证结果一致性 print(torch.allclose(pairwise_loop, pairwise_vec)) # 输出应为True
内容的提问来源于stack exchange,提问作者Alex Shtoff
相关产品推荐
相关产品推荐

