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
相关产品推荐
相关产品推荐

