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

如何用NumPy加速自定义DecisionTreeClassifier的fit方法以提升效率?

优化决策树运行速度的NumPy向量化方案

你的核心问题在于双重循环遍历特征和阈值——Python层面的循环是性能瓶颈,我们可以用NumPy的向量化操作完全替代这些循环,同时优化计算逻辑,把速度拉上来。下面是具体的优化方案:

一、先优化__info函数的计算效率

原函数里的x参数根本没用到,而且权重计算可以简化,同时用字典映射替代冗余的分支判断,速度会更快:

def __info(self, y, criterion):
    """ 简化不确定性度量计算,仅传入y即可 """
    y_len = y.shape[0]
    if y_len == 0:
        return 0.0
    # 直接用bincount结果除以总数,比手动传weights高效
    prob_distr = np.bincount(y, minlength=self.num_class) / y_len
    # 用字典映射替代if-elif-else,分支判断更高效
    criterion_calcs = {
        'error': lambda p: 1 - p.max(),
        'gini': lambda p: 1 - np.sum(p ** 2),
        'entropy': lambda p: -np.sum(p * np.log2(p + 1e-10))
    }
    return criterion_calcs.get(criterion, lambda _: 0.1)(prob_distr)

二、核心优化:用向量化替代__find_threshold的双重循环

这是速度提升的关键!我们把特征遍历、阈值遍历里的Python循环,全部换成NumPy的向量化操作(底层是C实现,速度远超Python循环):

def __find_threshold(self, x, y):
    max_info_gain = -1.0
    n_samples, n_features = x.shape
    if n_samples == 0:
        raise RuntimeError("Received an empty sample!")
    
    # 预计算当前数据集的基础不纯度,避免重复计算
    base_info = self.__info(y, self.criterion)
    
    # 1. 批量生成所有特征的候选阈值
    thresholds_list = []
    for feat_idx in range(n_features):
        vals = np.unique(x[:, feat_idx])
        if len(vals) < 2:
            thresholds_list.append(vals)
        else:
            # 用滑动窗口计算相邻均值,比手动切片更简洁高效
            thresholds = (vals[:-1] + vals[1:]) / 2
            thresholds_list.append(thresholds)
    
    # 2. 遍历每个特征,向量化计算所有阈值的信息增益
    for feat_idx in range(n_features):
        thresholds = thresholds_list[feat_idx]
        if len(thresholds) == 0:
            continue
        
        # 向量化生成分割mask:shape (n_thresholds, n_samples)
        masks = x[:, feat_idx] <= thresholds[:, np.newaxis]
        # 计算每个阈值对应的左右样本数量
        left_counts = np.sum(masks, axis=1)
        right_counts = n_samples - left_counts
        
        # 跳过分割后某一侧为空的无效阈值
        valid_mask = (left_counts > 0) & (right_counts > 0)
        if not np.any(valid_mask):
            continue
        valid_thresholds = thresholds[valid_mask]
        valid_left_counts = left_counts[valid_mask]
        valid_right_counts = right_counts[valid_mask]
        
        # 向量化统计左右样本的类别概率分布
        # 用apply_along_axis批量计算bincount(NumPy 1.23+支持更高效的axis参数)
        left_y_list = [y[m] for m in masks[valid_mask]]
        left_probs = np.array([
            np.bincount(ly, minlength=self.num_class) / lc 
            for ly, lc in zip(left_y_list, valid_left_counts)
        ])
        right_y_list = [y[~m] for m in masks[valid_mask]]
        right_probs = np.array([
            np.bincount(ry, minlength=self.num_class) / rc 
            for ry, rc in zip(right_y_list, valid_right_counts)
        ])
        
        # 批量计算每个阈值的信息增益
        left_info = np.array([self.__info_from_probs(p, self.criterion) for p in left_probs])
        right_info = np.array([self.__info_from_probs(p, self.criterion) for p in right_probs])
        info_gains = base_info - (valid_left_counts / n_samples) * left_info - (valid_right_counts / n_samples) * right_info
        
        # 更新最优特征和阈值
        best_feat_threshold_idx = np.argmax(info_gains)
        current_max_gain = info_gains[best_feat_threshold_idx]
        if current_max_gain > max_info_gain:
            max_info_gain = current_max_gain
            best_feature_id = feat_idx
            best_threshold = valid_thresholds[best_feat_threshold_idx]
    
    # 执行最终分割并返回结果
    x_left, x_right, y_left, y_right = self.__div_samples(x, y, best_feature_id, best_threshold)
    return best_feature_id, best_threshold, x_left, x_right, y_left, y_right

# 新增辅助函数:直接从概率分布计算不纯度,避免重复统计类别
def __info_from_probs(self, prob_distr, criterion):
    criterion_calcs = {
        'error': lambda p: 1 - p.max(),
        'gini': lambda p: 1 - np.sum(p ** 2),
        'entropy': lambda p: -np.sum(p * np.log2(p + 1e-10))
    }
    return criterion_calcs.get(criterion, lambda _: 0.1)(prob_distr)

三、额外小优化点

  • 提前初始化准则函数:把criterion_calcs字典初始化到类的__init__方法里,不用每次调用都重新创建,减少开销。
  • 减少数据复制:确保__div_samples函数用索引而非复制数组(比如x[x[:, feat_idx] <= threshold]),避免不必要的内存操作。
  • 过滤无效特征:提前检查特征是否所有值相同,直接跳过这类特征的阈值计算,节省时间。

这些优化把最耗时的Python循环替换成了NumPy的底层C实现操作,应该能让你的代码速度大幅接近sklearn的水平。

内容的提问来源于stack exchange,提问作者DavidFromMoscowStateUniv

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.07 12:12:57