You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何让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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.23 08:47:18