如何让np.put_along_axis实现累加而非替换(类PyTorch scatter_add)
在NumPy中实现类似PyTorch scatter_add的累加功能
NumPy的np.put_along_axis默认是直接替换元素值,没有内置的累加模式,但可以通过构造索引直接赋值的方式实现和PyTorch scatter_add一样的累加效果。
针对你给出的示例实现:
import numpy as np frame = np.zeros((3, 2)) updates = np.array([[5,5], [10,10], [3,3]]) indices = np.array([[1,1], [1,1], [2,2]]) axis = 0 # 和PyTorch scatter_add的dim参数对应 # 构造完整的索引元组 idx = np.indices(frame.shape) idx[axis] = indices # 执行累加操作 frame[tuple(idx)] += updates print(frame) # 输出: # [[ 0. 0.] # [15. 15.] # [ 3. 3.]]
通用化实现(支持任意axis)
上面的方法可以适配任意维度和axis参数,核心是利用np.indices生成目标数组的网格索引,再将对应axis的位置替换为传入的indices,最后通过索引赋值完成累加。这种方式会自动处理重复索引的情况,多次指向同一位置时会持续累加更新值。
内容的提问来源于stack exchange,提问作者Craig
相关产品推荐
相关产品推荐

