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

如何修改SkLearn RandomForestClassifier以使用序列自助采样方法

自定义RandomForestClassifier实现序列自助采样方案

核心实现思路

你定位的_generate_sample_indices是sklearn私有模块的工具函数,直接修改全局函数会污染环境,最优方案是在自定义CustomRF类中重写fit方法的树构建逻辑,直接注入你提前生成的序列自助采样索引,完全兼容原有RandomForestClassifier的所有参数和后续调用方法。

完整代码实现

from sklearn.ensemble import RandomForestClassifier
from sklearn.utils import check_random_state
from sklearn.ensemble._forest import _get_n_samples_bootstrap
import numpy as np

class CustomRF(RandomForestClassifier):
    def __init__(self, sequential_bootstrap_indices=None, **kwargs):
        super().__init__(**kwargs)
        # 入参要求:sequential_bootstrap_indices为列表,每个元素是对应单棵树的采样索引数组,长度等于n_estimators
        self.sequential_bootstrap_indices = sequential_bootstrap_indices

    def fit(self, X, y, sample_weight=None):
        # 前置参数校验
        if self.sequential_bootstrap_indices is not None:
            if len(self.sequential_bootstrap_indices) != self.n_estimators:
                raise ValueError(f"采样索引数量{len(self.sequential_bootstrap_indices)}与树的数量{self.n_estimators}不匹配")
        
        # 继承父类的输入数据校验逻辑
        X, y = self._validate_data(
            X, y, multi_output=True, accept_sparse="csc", dtype=None
        )
        if sample_weight is not None:
            sample_weight = np.asarray(sample_weight)
        y = np.atleast_1d(y)
        if y.ndim == 2 and y.shape[1] == 1:
            y = y.ravel()

        self.n_outputs_ = 1 if y.ndim == 1 else y.shape[1]
        self.classes_ = np.unique(y)
        self.n_classes_ = len(self.classes_)
        n_samples = X.shape[0]
        self.n_samples_bootstrap = _get_n_samples_bootstrap(n_samples, self.max_samples)

        # 初始化基学习器
        self._validate_estimator()
        random_state = check_random_state(self.random_state)
        trees = [
            self._make_estimator(append=False, random_state=random_state)
            for _ in range(self.n_estimators)
        ]

        # 核心修改:替换原随机采样逻辑,注入预生成的序列自助采样索引
        for tree_idx, tree in enumerate(trees):
            # 优先使用自定义采样索引,无自定义索引时兼容原随机采样逻辑
            if self.sequential_bootstrap_indices is not None:
                sample_indices = self.sequential_bootstrap_indices[tree_idx]
            else:
                sample_indices = random_state.randint(0, n_samples, self.n_samples_bootstrap)
            
            # 生成样本权重对齐原训练逻辑
            if sample_weight is None:
                curr_sample_weight = np.ones(n_samples, dtype=np.float64)
            else:
                curr_sample_weight = sample_weight.copy()
            sample_counts = np.bincount(sample_indices, minlength=n_samples)
            curr_sample_weight *= sample_counts

            # 训练单棵决策树
            tree.fit(X, y, sample_weight=curr_sample_weight, check_input=False)
        
        self.estimators_ = trees
        return self

使用说明

  • 你预生成的序列自助采样索引每个子数组的长度,要和CustomRF初始化时的max_samples参数配置的采样数量一致,避免长度不匹配报错。
  • 若需要使用OOB得分功能,需要基于自定义的采样索引额外实现OOB样本的计算逻辑,原有自带的OOB逻辑依赖原生采样结果,无法直接复用。
  • 训练完成后原有predict、score等方法都可以直接正常调用,不需要额外修改。

内容的提问来源于stack exchange,提问作者Javier C Salaverri

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.28 21:15:06