如何在PyTorch张量的指定索引位置添加标量?
解决二维张量指定索引位置加标量的问题
嘿,我完全懂你对scatter*和index*系列操作的困惑——其实你要的需求用非常简单的方法就能实现,甚至不需要用到那些复杂的scatter类操作!
方法一:直接索引赋值(最简洁高效)
你可以把二维索引拆成行和列两个一维张量,然后直接对这些位置执行加法操作,完全不用转换标量形状:
import torch # 初始化你的张量和索引 z = torch.tensor([[1., 1., 1., 1.], [1., 1., 1., 1.], [1., 1., 1., 1.]]) inds = torch.tensor([[0, 0], [1, 1], [1, 2]]) # 拆分索引为行和列维度 rows = inds[:, 0] cols = inds[:, 1] # 直接在指定位置加标量3 z[rows, cols] += 3 print(z) # 输出结果: # tensor([[4., 1., 1., 1.], # [1., 4., 4., 1.], # [1., 1., 1., 1.]])
这种方法的优势是直观、高效,PyTorch对索引操作做了深度优化,完全不用担心性能问题。如果你的索引里有重复位置(比如同一个坐标出现多次),这个方法也会自动累加数值,和scatter_add_的效果一致。
方法二:使用scatter_add_(如果你想尝试scatter系列操作)
如果你一定要用scatter*系列来实现,torch.scatter_add_也能完成这个需求,只是需要把标量转换成和索引匹配的形状(其实PyTorch的广播机制也能帮你简化这一步):
import torch z = torch.tensor([[1., 1., 1., 1.], [1., 1., 1., 1.], [1., 1., 1., 1.]]) inds = torch.tensor([[0, 0], [1, 1], [1, 2]]) # 准备要添加的数值,利用full_like生成和列索引同形状的全3张量 src = torch.full_like(inds[:, 1], 3) # 使用scatter_add_,dim=1表示按列维度进行映射 z.scatter_add_( dim=1, index=inds[:, 1].unsqueeze(1), # 需要把列索引变成二维张量匹配目标维度 src=src.unsqueeze(1) ) print(z) # 输出和上面完全一致
为什么不用纠结scatter系列?
scatter*和index*操作更多是为了处理跨张量的元素映射/聚合场景(比如把一个张量的元素根据索引放到另一个张量里),而你的需求只是对现有张量的指定位置做简单加法,直接索引操作显然更贴合需求,代码也更易读。
内容的提问来源于stack exchange,提问作者Colin
相关产品推荐
相关产品推荐

