如何修改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
相关产品推荐
相关产品推荐

