如何通过1D张量元素修改2D张量对应索引的元素?
PyTorch张量按1D索引置0的实现方法
这里提供两种简洁的实现方式,直接解决你的需求:
方法一:索引赋值(最直观高效)
直接通过行索引和对应的目标列索引,定位到需要置0的位置并赋值:
import torch # 构造输入张量 x = torch.tensor([[1, 2], [3, 4], [5, 6], [7, 8]]) indices = torch.tensor([0, 1, 1, 0]) # 克隆原张量避免修改原数据 result = x.clone() # 生成行索引,和indices组合定位目标位置 rows = torch.arange(x.size(0)) result[rows, indices] = 0 print(result) # 输出: # tensor([[0, 2], # [3, 0], # [5, 0], # [0, 8]])
方法二:使用torch.where
通过构造掩码张量,结合torch.where完成替换:
import torch x = torch.tensor([[1, 2], [3, 4], [5, 6], [7, 8]]) indices = torch.tensor([0, 1, 1, 0]) # 生成列索引张量,通过广播和indices比较得到掩码 mask = torch.arange(x.size(1)) == indices.unsqueeze(1) # 掩码为True的位置置0,否则保留原张量值 result = torch.where(mask, torch.tensor(0), x) print(result) # 输出同上
说明
- 方法一的索引赋值方式更高效,直接定位目标位置操作,适合大多数场景;
- 方法二通过掩码实现,适合需要更复杂条件判断的扩展场景。
内容的提问来源于stack exchange,提问作者nameisnoname
相关产品推荐
相关产品推荐

