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

PyTorch如何实现批量index_fill 按批次索引给张量对应位置赋值

解决方案

你之前写法报错有两个核心原因:

  1. 生成的index是浮点型张量,PyTorch要求索引必须是整数类型,需要先转成long类型
  2. 你用的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.06 16:57:03