numpy.add.at能否使用二维索引?二维索引累加结果异常咨询
解决Numpy中
np.add.at二维索引累加的问题 你碰到的问题是因为np.add.at的索引参数格式不对——当你直接传入二维的indices数组时,Numpy会把它当成第一个维度(行)的索引数组,然后广播到所有列,导致整行都被累加,而不是你期望的单个坐标累加。
正确的实现代码
要精准定位每个(行,列)坐标进行累加,你需要把indices拆分成对应两个维度的索引元组,最简单的方式是转置后转成元组:
import numpy as np image = np.zeros((5,5), dtype=np.int32) indices = np.array([[1,1], [1,1], [3,3]]) # 将二维索引转换为(行索引数组, 列索引数组)的元组 np.add.at(image, tuple(indices.T), 1) print(image)
运行后会得到你期望的输出:
[[0 0 0 0 0] [0 2 0 0 0] [0 0 0 0 0] [0 0 0 1 0] [0 0 0 0 0]]
另一种等价写法
你也可以直接提取行和列的索引数组,组成元组传入:
np.add.at(image, (indices[:, 0], indices[:, 1]), 1)
为什么原来的代码会出错?
当你直接传入indices二维数组时,Numpy会将其视为行维度的索引数组(形状为(3,2)),数组里的每个元素都是行索引值(1或3)。然后Numpy会对这些索引对应的行,给所有列都累加1:
- 索引数组里共有4个
1(来自[[1,1],[1,1]]),所以第1行的每个元素都加了4 - 索引数组里共有2个
3(来自[3,3]),所以第3行的每个元素都加了2
这就是你看到错误输出的原因。
内容的提问来源于stack exchange,提问作者user1411900
相关产品推荐
相关产品推荐

