如何使用含重复索引的索引列表正确更新Numpy数组?
解决Numpy批量更新重复索引的累加问题
Numpy的花式索引遇到重复位置时,直接用+=1只会对该位置执行一次加法,不会按索引重复次数累加。要实现预期的累加效果,需要先统计每个索引组合的出现次数,再按次数批量更新数组。
方案一:用collections.Counter统计索引次数(代码直观)
适合中小规模缓存数据,代码可读性强:
import numpy as np from collections import Counter def update(arr, index_counter): # 拆分统计后的索引与对应次数 indices, counts = zip(*index_counter.items()) rows, cols, chs = zip(*indices) # 按次数累加对应位置的值 arr[rows, cols, chs] += counts # 保持数值不超过10 arr[arr > 10] = 10 # 初始化数组(修正原代码的初始化错误) n_rows, n_cols, n_channels = 5, 5, 3 arr = np.zeros((n_rows, n_cols, n_channels)) cache_length = 3 cache = [] while True: # 示例:动态生成索引,替换为你的实际逻辑 row_idx = 0 col_idx = 3 ch_idx = 6 cache.append((row_idx, col_idx, ch_idx)) if len(cache) == cache_length: # 统计缓存中每个索引组合的出现次数 index_counter = Counter(cache) update(arr, index_counter) # 清空缓存 cache = []
方案二:用Numpy的unique处理(高效向量化)
适合大规模缓存数据,利用Numpy向量化操作提升效率:
import numpy as np def update(arr, cache): # 将缓存转为Numpy数组,形状为(缓存长度, 3) cache_np = np.array(cache) # 获取唯一索引组合及对应的出现次数 unique_indices, counts = np.unique(cache_np, axis=0, return_counts=True) # 拆分索引的三个维度 rows = unique_indices[:, 0] cols = unique_indices[:, 1] chs = unique_indices[:, 2] # 按次数累加 arr[rows, cols, chs] += counts # 裁剪数值上限 arr[arr > 10] = 10 # 初始化数组 n_rows, n_cols, n_channels = 5, 5, 3 arr = np.zeros((n_rows, n_cols, n_channels)) cache_length = 3 cache = [] while True: # 示例:动态生成索引 row_idx = 0 col_idx = 3 ch_idx = 6 cache.append((row_idx, col_idx, ch_idx)) if len(cache) == cache_length: update(arr, cache) cache = []
关键说明
- 原代码中
arr = np.array(n_rows, n_cols, n_channels)是错误的,需用np.zeros()或其他合理方式初始化数组。 - 两种方案都能解决重复索引的累加问题,按需选择:小规模数据用Counter更直观,大规模数据用Numpy unique效率更高。
内容的提问来源于stack exchange,提问作者firion
相关产品推荐
相关产品推荐

