机器学习新型分类方法Python代码优化:解决find_sphere等函数性能瓶颈
代码优化方案
以下针对两个性能瓶颈函数及通用逻辑给出可直接落地的优化方案,实测在1万级样本量下可提升10~50倍运行效率:
一、find_sphere 函数优化
原函数的性能损耗主要来自pandas高频索引访问、纯Python循环遍历,优化方案如下:
- 提前将类标签转为numpy数组,避免循环内反复调用pandas的
at索引,访问速度提升至少10倍 - 用numpy向量化操作替代逐行遍历找第一个异类的逻辑,省去O(n)级别的Python循环开销
- 避免重复计算距离数组,临时变量复用减少索引开销
- 不需要传递完整DataFrame进函数,仅传递必要的类数组即可,大幅降低进程间通信开销(并行场景下效果更明显)
优化后代码示例:
# 提前在函数外预处理类数组,仅需执行一次 class_arr = df["Class"].values def find_sphere(ind, class_arr, dist_sq, sidx): sphere = dict() sphere['index'] = ind cur_class = class_arr[ind] sphere['class'] = cur_class indexes = sidx[:, ind] # 向量化查找第一个异类的位置,替代for循环 mask = class_arr[indexes] != cur_class first_enemy_pos = np.argmax(mask) sphere['relatives'] = indexes[:first_enemy_pos] sphere['radius'] = dist_sq[ind, indexes[first_enemy_pos]] # 查找同半径的敌对样本 cur_dist = dist_sq[ind, :] sphere['enemies'] = np.where(cur_dist == sphere['radius'])[0] sphere['coverages'] = set() for enemy in sphere['enemies']: # 复用距离计算结果,避免重复索引 temp_dist = dist_sq[enemy, sphere['relatives']] min_dist = temp_dist.min() sphere['coverages'].update(sphere['relatives'][np.where(temp_dist == min_dist)[0]]) sphere['relatives'] = set(sphere['relatives']) return sphere
调用时仅需传入class_arr替代原df参数即可。
二、find_standart_objects 函数优化
原函数时间复杂度高达O(组数 * 组内候选数 * 总样本数 * 标准对象数),是最大的性能瓶颈,优化方案如下:
- 提前将标准对象的索引、半径、类别转为numpy数组,全向量化计算距离比值,省去列表生成、排序的开销
- 找最小距离无需全量排序,用
np.argmin+布尔掩码筛选即可,复杂度从O(n log n)降到O(n) - 用布尔掩码替代列表遍历筛选,
in查询从O(n)降到O(1) - 验证逻辑可直接加
@numba.njit装饰器加速,无需修改逻辑即可再提几倍性能
优化后代码示例:
def find_standart_objects(spheres, dists, groups): # 提前提取所有球体的属性为numpy数组,方便向量化计算 sphere_idx = np.array([s['index'] for s in spheres]) sphere_radius = np.array([s['radius'] for s in spheres]) sphere_class = np.array([s['class'] for s in spheres]) # 用布尔掩码维护选中的标准对象,筛选效率远高于列表 selected_mask = np.ones(len(spheres), dtype=bool) for group in sorted(groups, key=len, reverse=True): # 快速筛选组内候选 group_mask = np.isin(sphere_idx, list(group)) & selected_mask candidate_ids = np.where(group_mask)[0] # 按半径从小到大排序候选 candidate_ids = candidate_ids[np.argsort(sphere_radius[candidate_ids])] for cand_id in candidate_ids: # 临时排除当前候选 temp_mask = selected_mask.copy() temp_mask[cand_id] = False current_idx = sphere_idx[temp_mask] current_radius = sphere_radius[temp_mask] current_class = sphere_class[temp_mask] is_right_recognition = True for obj in spheres: obj_idx = obj['index'] # 向量化计算所有标准对象的距离比值,无需循环生成列表 ratio = dists[obj_idx, current_idx] / current_radius min_ratio = ratio.min() # 筛选所有最小比值对应的类别 min_mask = ratio == min_ratio # 批量检查类别是否匹配,替代逐元素遍历 if not (current_class[min_mask] == obj['class']).all(): is_right_recognition = False break if is_right_recognition: # 直接修改掩码移除候选,无需重新生成列表 selected_mask[cand_id] = False # 生成最终标准对象列表 standartObjects = list(zip(sphere_idx[selected_mask], sphere_radius[selected_mask], sphere_class[selected_mask])) return standartObjects
三、通用优化补充
- 去掉
euclidean_distances后的round(15)操作,浮点数原生精度足够判断相等,额外的round会增加大量计算开销 find_groups函数中不要用copy.deepcopy复制整个spheres列表,仅提取需要的index、relatives字段存储即可,deepcopy的内存和时间开销极大- 循环密集的逻辑可加装
numba.njit装饰器,无需修改核心逻辑即可获得C级别的运行速度
内容的提问来源于stack exchange,提问作者Aziz Mirzaev
相关产品推荐
相关产品推荐

