Numpy代码优化:简化SOM中BMU邻域查找的冗余代码行
简化SOM BMU邻域处理的Python代码
嘿,我明白你现在的困扰——用np.argwhere循环处理BMU索引确实会让代码又长又啰嗦,而且数据量大的时候效率还不高。这里有几个更简洁高效的方法来重构你的代码,让逻辑更清晰:
方法1:字典映射(直观易读)
把每个BMU坐标转成可哈希的元组作为键,直接映射对应的原始数据索引列表,代码简洁且容易理解:
import numpy as np # 先把som.bmus转成numpy数组(如果还不是的话) bmus = np.array(som.bmus) # 构建坐标到索引的映射字典 bmu_index_map = {} for idx, coord in enumerate(bmus): coord_key = tuple(coord) bmu_index_map.setdefault(coord_key, []).append(idx) # 获取唯一BMU坐标,并快速拿到对应的索引列表 unique_bmus = np.unique(bmus, axis=0) bmu_idxs = [bmu_index_map[tuple(coord)] for coord in unique_bmus]
方法2:纯Numpy操作(高性能)
如果你的数据集很大,用Numpy的排序和分割操作可以避免Python循环,效率更高:
import numpy as np bmus = np.array(som.bmus) # 按BMU坐标排序,同时记录原始索引 sorted_idx = np.lexsort(bmus.T) sorted_bmus = bmus[sorted_idx] # 找到唯一坐标的分割位置 _, split_pos = np.unique(sorted_bmus, axis=0, return_index=True) split_pos = np.append(split_pos, len(sorted_bmus)) # 提取每个唯一坐标对应的原始索引 bmu_idxs = [sorted_idx[split_pos[i]:split_pos[i+1]] for i in range(len(split_pos)-1)] unique_bmus = sorted_bmus[split_pos[:-1]]
方法3:用Pandas一行搞定(极简风格)
如果项目里已经在用Pandas,用分组操作可以把代码压缩到几行,非常清爽:
import pandas as pd import numpy as np bmus = np.array(som.bmus) df = pd.DataFrame(bmus, columns=["x", "y"]) df["original_idx"] = df.index # 直接分组提取索引,同时得到唯一BMU坐标 bmu_groups = df.groupby(["x", "y"])["original_idx"].apply(list) unique_bmus = bmu_groups.index.to_numpy() bmu_idxs = bmu_groups.tolist()
这三种方法都能替代繁琐的np.argwhere循环,你可以根据自己的代码环境和数据规模选择最合适的方案~
内容的提问来源于stack exchange,提问作者Matti Wens
相关产品推荐
相关产品推荐

