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

PyTorch:如何基于索引张量高效设置张量元素值?

高效实现PyTorch张量按索引批量赋值

原始张量定义

import torch

tensor_to_change = torch.tensor([[-36.9127, -45.6596, -47.1595],
        [-36.9409, -45.7024, -47.2050],
        [-36.9865, -45.7665, -47.2711],
        [-36.3202, -36.9561, -47.2066],
        [-36.2929, -36.9333, -47.1702]])
tensor_of_indices = torch.tensor([[0],
        [0],
        [0],
        [1],
        [1]])
tensor_of_values = torch.tensor([[-37.9409],
        [-38.4865],
        [-36.9561],
        [-34.9561],
        [-38.7562]])

现有低效实现

目前通过Python for循环完成按索引张量给目标张量赋值,但运行速度极慢:

for i, a in enumerate(tensor_of_indices):
    tensor_to_change[i][a] = tensor_of_values[i]

高效替代方案

可以直接使用PyTorch内置的索引机制或专用函数完成批量赋值,完全规避Python循环的性能瓶颈,以下两种方法都可行:

方法一:高级索引赋值

通过生成行索引,结合压缩后的列索引直接批量赋值:

# 压缩索引和值张量的多余维度(从(5,1)转为(5))
tensor_of_indices = tensor_of_indices.squeeze(-1)
tensor_of_values = tensor_of_values.squeeze(-1)
# 生成每行的索引(0到4)
row_indices = torch.arange(tensor_to_change.size(0), device=tensor_to_change.device)
# 批量赋值
tensor_to_change[row_indices, tensor_of_indices] = tensor_of_values

方法二:使用scatter_函数

scatter_是PyTorch专门用于按索引分散赋值的内置函数,无需手动处理维度:

# dim=1表示按列维度进行分散赋值,index为目标位置索引,src为待赋值的张量
tensor_to_change.scatter_(dim=1, index=tensor_of_indices, src=tensor_of_values)

两种方法都能实现和原循环完全一致的赋值效果,且基于PyTorch底层优化,在张量规模越大时,性能提升越显著。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.25 08:54:19