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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.25 10:45:01