关于torch.scatter()实现多维权重矩阵赋值的技术疑问
解决torch.scatter_()赋值张量而非固定值的问题
要实现将张量b的元素对应填充到weights的指定位置,核心是正确使用torch.scatter_()的src参数而非固定值,具体步骤如下:
初始化权重矩阵
先创建形状为[B, N, V]的初始权重张量,通常用全零初始化:weights = torch.zeros(B, N, V, device=b.device, dtype=b.dtype)调用scatter_完成赋值
直接将b作为src参数传入,指定dim=2(对应特征维度),index设为a:weights.scatter_(dim=2, index=a, src=b)此时,
b中每个位置b[i,j,m]的值会被赋值到weights[i,j,a[i,j,m]]的位置。
示例验证
假设参数为:
B, N, V, k = 1, 2, 3, 2 a = torch.tensor([[[0, 2], [1, 0]]]) # 特征索引 b = torch.tensor([[[0.5, 0.8], [0.3, 0.6]]]) # 对应权重
执行上述代码后,weights的结果为:
tensor([[[0.5, 0.0, 0.8], [0.6, 0.3, 0.0]]])
weights[0,0,0]来自b[0,0,0](对应a[0,0,0]=0)weights[0,0,2]来自b[0,0,1](对应a[0,0,1]=2)weights[0,1,1]来自b[0,1,0](对应a[0,1,0]=1)weights[0,1,0]来自b[0,1,1](对应a[0,1,1]=0)
处理重复索引
如果a中存在重复的索引(即同一个特征位置被多次赋值),默认行为是覆盖最后一次赋值的值。若需要累加权重,可添加reduce='add'参数:
weights.scatter_(dim=2, index=a, src=b, reduce='add')
内容的提问来源于stack exchange,提问作者YuxuanSnow
相关产品推荐
相关产品推荐

