numpy add.at函数索引越界问题:为何与循环结果不一致?
问题解答
错误原因
你在np.add.at中用错了索引格式:当用**列表[idx1, idx2]传递索引时,numpy会把两个数组的所有元素都当成对第一维度(axis0)**的索引,而非配对的(行, 列)坐标。
你的idx2数组里存在元素9,但arr的axis0维度大小是idx1.max()+1 = 8+1 = 9,索引范围仅为0~8,用9去索引axis0自然触发越界异常。
修正写法
要实现和for循环一致的配对索引逻辑,需要把索引改成元组(idx1, idx2)——这样numpy会将idx1作为axis0的索引数组,idx2作为axis1的索引数组,逐个配对执行累加操作:
np.add.at(arr, (idx1, idx2), 1)
运行修正后的代码,arr和arr_via_for的结果会完全一致。
内容的提问来源于stack exchange,提问作者I. Dakhtin
相关产品推荐
相关产品推荐

