You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

机器学习新型分类方法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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.28 11:06:02