PyTorch中如何高效按给定行索引将张量对应位置值设为0
高效实现方案
核心思路是通过构造与B形状匹配的行索引张量,使用PyTorch原生高级索引直接定位目标位置并赋值,全程无Python层循环,无额外大张量构造开销,性能接近你误写的全局索引赋值的速度。
具体实现代码
# 构造和B形状对齐的行索引,自动适配A所在的设备(CPU/CUDA) row_indices = torch.arange(A.shape[0], device=A.device)[:, None] # 直接对目标位置批量赋值为0 A[row_indices, B] = 0
代码说明
torch.arange(A.shape[0], device=A.device)生成0到M-1的行号序列- 末尾的
[:, None]是将一维行号序列升维为形状为(M, 1)的张量,触发广播机制后和形状为(M, P)的B维度对齐,最终每个索引对(row_indices[i][j], B[i][j])刚好对应A中第i行需要置0的第j个目标位置 - 整个操作完全在PyTorch后端实现,没有Python和设备的交互开销,也不需要像scatter方案那样构造和A同尺寸的全零张量,内存和计算效率都更高
示例验证
用你给出的测试用例运行后,得到的A和预期结果完全一致:
import torch A = torch.tensor([list(range(1,11)), list(range(1,11)), list(range(1,11))]) B = torch.tensor([[1,2], [2,3], [3,5]]) row_indices = torch.arange(A.shape[0], device=A.device)[:, None] A[row_indices, B] = 0 print(A) # 输出: # tensor([[ 1, 0, 0, 4, 5, 6, 7, 8, 9, 10], # [ 1, 2, 0, 0, 5, 6, 7, 8, 9, 10], # [ 1, 2, 3, 0, 5, 0, 7, 8, 9, 10]])
内容的提问来源于stack exchange,提问作者BobHU
相关产品推荐
相关产品推荐

