PyTorch如何实现批量index_fill 按批次索引给张量对应位置赋值
解决方案
你之前写法报错有两个核心原因:
- 生成的
index是浮点型张量,PyTorch要求索引必须是整数类型,需要先转成long类型 - 你用的
range(2)是一维结构,和形状为(2,3)的index无法正确广播对齐维度
方法1:高级索引(无额外大内存开销)
直接构造行索引的广播结构批量赋值,全程只需要创建一个长度等于value行数的一维索引张量,内存开销可忽略:
# 先把index转成整数类型 index = index.long() # 行索引升维到(行数, 1),和index的(行数, 索引数)自动广播匹配 value[torch.arange(value.shape[0])[:, None], index] = 1
方法2:torch.scatter_(原地操作,无需创建全1张量)
你之前对scatter的使用存在误解,scatter_的src参数直接支持传入标量,完全不需要额外创建和value同尺寸的全1张量,直接原地修改,开销极低:
index = index.long() value.scatter_(dim=-1, index=index, src=1)
两种方法都可以得到预期输出,性能差异极小,都属于PyTorch内置的高效实现。
内容的提问来源于stack exchange,提问作者namespace-Pt
相关产品推荐
相关产品推荐

