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

如何基于sample_weights预剪枝scikit-learn回归决策树?是否有内置方法?

带样本权重的回归树预剪枝实现方案

内置参数的加权支持

scikit-learn的DecisionTreeRegressor原生支持基于样本权重的预剪枝,部分参数在传入sample_weight时会自动切换为加权逻辑:

  • min_samples_leaf:当指定样本权重后,该参数代表节点中样本权重总和的最小值。若节点总权重小于设定值,会被强制设为叶节点。
  • min_samples_split:同理,该参数变为节点分裂所需的最小样本权重总和,总权重不足时不会触发分裂。

代码示例:

from sklearn.tree import DecisionTreeRegressor
import numpy as np

# 生成示例数据
X = np.random.rand(100, 2)
y = np.random.rand(100)
sample_weights = np.random.rand(100)

# 初始化带加权预剪枝的回归树
reg_tree = DecisionTreeRegressor(
    min_samples_leaf=5.0,  # 节点总权重<5.0则为叶节点
    min_samples_split=10.0,  # 节点总权重<10.0则不分裂
    random_state=42
)

# 拟合时传入样本权重
reg_tree.fit(X, y, sample_weight=sample_weights)

自定义预剪枝逻辑(进阶需求)

如果内置参数无法满足复杂加权条件,可以通过继承DecisionTreeRegressor并重写方法实现:

from sklearn.tree import DecisionTreeRegressor
from sklearn.tree._tree import Tree

class WeightedPrunedDecisionTreeRegressor(DecisionTreeRegressor):
    def __init__(self, min_weight_leaf=1.0, **kwargs):
        super().__init__(**kwargs)
        self.min_weight_leaf = min_weight_leaf

    def fit(self, X, y, sample_weight=None, check_input=True):
        super().fit(X, y, sample_weight=sample_weight, check_input=check_input)
        # 遍历树节点,将权重不足的内部节点转为叶节点
        tree = self.tree_
        stack = [0]
        while stack:
            node_id = stack.pop()
            if tree.children_left[node_id] != Tree.LEAF:
                node_weight = tree.weighted_n_node_samples[node_id]
                if node_weight < self.min_weight_leaf:
                    tree.children_left[node_id] = Tree.LEAF
                    tree.children_right[node_id] = Tree.LEAF
                else:
                    stack.append(tree.children_left[node_id])
                    stack.append(tree.children_right[node_id])
        return self

使用自定义类:

custom_reg_tree = WeightedPrunedDecisionTreeRegressor(
    min_weight_leaf=5.0,
    random_state=42
)
custom_reg_tree.fit(X, y, sample_weight=sample_weights)

关键说明

  • 优先使用内置参数方案,这是scikit-learn原生支持的稳定实现。
  • 自定义方法需依赖树的内部属性tree_.weighted_n_node_samples(存储节点总样本权重),修改树结构时需确保逻辑正确。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.14 16:25:00