You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.10.06 18:48:03