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

