PyTorch中如何用索引张量idx给二维张量arr指定位置赋值为1
问题根因
- 你定义的
idx默认是浮点类型张量,PyTorch不支持用浮点张量作为索引取值/赋值 - 直接使用
arr[idx]的索引逻辑不符合预期:你需要的是逐行对应索引位置赋值,这种写法触发的高级索引广播逻辑会把idx里的每个值当做行索引,完全不符合你的需求
正确实现方案
方案1:使用scatter_原地操作(最简洁)
scatter_是PyTorch专门用于按指定索引赋值的原地方法,完全匹配你的需求:
import torch arr = torch.zeros(size = (2,10)) # 注意要指定dtype为整数类型,torch.long是索引常用类型 idx = torch.tensor([ [0,2], [4,5] ], dtype=torch.long) # dim=1代表沿列方向匹配索引赋值 arr.scatter_(dim=1, index=idx, value=1) print(arr)
方案2:手动构造行+列索引配对
如果你需要更灵活的自定义索引逻辑,可以手动构造和列索引形状匹配的行索引再赋值:
import torch arr = torch.zeros(size = (2,10)) idx = torch.tensor([ [0,2], [4,5] ], dtype=torch.long) # 生成逐行匹配的行索引:[ [0,0], [1,1] ] row_idx = torch.arange(arr.shape[0]).unsqueeze(-1).expand_as(idx) arr[row_idx, idx] = 1 print(arr)
两种方案运行后都会输出你期望的结果:
tensor([[1., 0., 1., 0., 0., 0., 0., 0., 0., 0.], [0., 0., 0., 0., 1., 1., 0., 0., 0., 0.]])
内容的提问来源于stack exchange,提问作者Mr. Nobody
相关产品推荐
相关产品推荐

