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

如何对PyTorch中torch.gather选取的元素赋值并修改原始张量

解决方案

你可以直接用PyTorch提供的scatter_原地方法实现需求,它是gather的逆操作,完全不需要写Python层循环,是官方推荐的对应实现:

import torch
data = torch.tensor([[0,1,2], [3,4,5], [6,7,8]])
indices = torch.tensor([1,2,1], dtype=torch.long).unsqueeze(-1)

# 直接原地赋值,dim、index参数和你调用gather时的参数完全一致
data.scatter_(
    dim=1,
    index=indices,
    src=torch.full_like(indices, fill_value=42)
)
print(data)

运行后输出就是你需要的结果:

tensor([[ 0, 42,  2],
        [ 3,  4, 42],
        [ 6, 42,  8]])

方法说明

PyTorch中所有后缀带_的方法都是原地操作,会直接修改张量本身的存储内容。scatter_的参数要求和gather完全对齐,你之前用gather时传的dim和index参数可以直接复用,只需要通过src传入和index形状相同的待赋值张量即可。

如果你觉得scatter_的逻辑不够直观,也可以手动构造行、列索引直接对原张量赋值,写法更易懂:

data = torch.tensor([[0,1,2], [3,4,5], [6,7,8]])
indices = torch.tensor([1,2,1], dtype=torch.long)

# 生成对应行的索引
row_idx = torch.arange(data.shape[0])
# 直接对原张量的对应位置赋值
data[row_idx, indices] = 42
print(data)

这种写法适合indices是1维的场景,不需要做unsqueeze操作,可读性更高。


内容的提问来源于stack exchange,提问作者Luchspeter

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.04 22:54:04