ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

scikit-learn 决策树原生支持类别特征:DecisionTreeClassifier 与 DecisionTreeRegressor 的 categorical_features 参数详解

2026/9/18 23:33:51 拓冰建站 浏览量
scikit-learn 决策树原生支持类别特征:DecisionTreeClassifier 与 DecisionTreeRegressor 的 categorical_features 参数详解 scikit-learn 决策树原生支持类别特征DecisionTreeClassifier 与 DecisionTreeRegressor 的 categorical_features 参数详解【免费下载链接】scikit-learnscikit-learn: machine learning in Python项目地址: https://gitcode.com/gh_mirrors/sc/scikit-learnscikit-learn 的树模块正在迎来一项重要能力DecisionTreeClassifier与DecisionTreeRegressor原生支持类别categorical特征不再强制要求先做 one-hot/ordinal 编码。本文基于当前仓库中的变更说明与源码实现详解categorical_features参数的四种指定方式、类别数量上限256、criterionabsolute_error的例外限制以及底层的比特位集bitset编码原理帮助你在混合类型数据上直接训练决策树并理解其约束边界。功能概述变更说明了什么该功能记录在待发布变更 33354.major-feature.rst 中原文要点如下tree.DecisionTreeRegressor与tree.DecisionTreeClassifier原生支持类别特征适用于**二分类binary classification与单输出回归single-output regression**场景支持所有可用的分裂准则criteria唯一例外是absolute_error类别特征通过categorical_features参数指定每个特征最多支持 256 个类别。需要说明的适用前提该条目位于doc/whats_new/upcoming_changes/目录属于尚未发布的新功能下文所有行为描述均以当前仓库源码为准。categorical_features参数四种指定方式categorical_features参数同时出现在DecisionTreeClassifier与DecisionTreeRegressor的构造函数中见 sklearn/tree/_classes.py 中分类器的参数定义与 sklearn/tree/_classes.py 中回归器的参数定义文档字符串统一描述为categorical_features : array-like of {bool, int, str} of shape (n_features,) or (n_categorical_features,), or from_dtype, defaultNone即支持None、布尔掩码、整数下标、字符串列名、from_dtype五种取值。这套解析逻辑集中在sklearn/utils/validation.py的_check_categorical_features函数中见 sklearn/utils/validation.py其行为规则可以逐条对应取值形式含义源码中的校验规则None没有特征被视为类别特征直接返回None走纯数值分裂路径布尔数组shape(n_features,)布尔掩码逐列标记是否为类别特征形状必须等于n_features否则抛ValueError整数数组类别特征的下标必须落在[0, n_features - 1]区间内字符串数组类别特征的列名要求训练数据带列名DataFrame未知列名会抛错并列出观察到的全部列名from_dtype自动识别 DataFrame 中 dtype 为Categorical/Enum的列要求输入是 narwhals 支持的 DataFrame如 pandas、polars对纯 ndarray 该取值等价于None一个典型的端到端用法import pandas as pd from sklearn.tree import DecisionTreeClassifier # 混合类型数据数值特征 类别特征pandas Categorical dtype df pd.DataFrame({ age: [25, 45, 35, 50, 28], # 数值 city: pd.Series([A, B, A, C, A], dtypecategory), # 类别 tier: pd.Series([low, high, mid, high, low], dtypecategory), # 类别 }) y [0, 1, 0, 1, 0] # 二分类目标 # 方式一按列名指定 clf DecisionTreeClassifier(categorical_features[city, tier]).fit(df, y) # 方式二布尔掩码第 1、2 列为类别特征 clf DecisionTreeClassifier(categorical_features[False, True, True]).fit(df, y) # 方式三整数下标 clf DecisionTreeClassifier(categorical_features[1, 2]).fit(df, y) # 方式四从 DataFrame dtype 自动推断列 dtype 为 Categorical/Enum 时生效 clf DecisionTreeClassifier(categorical_featuresfrom_dtype).fit(df, y)从_check_categorical_features的源码可以看到几个细节值得注意字符串列名方式中如果传入的数据是没有列名的数组会抛出明确的ValueErrorcategorical_features should be passed as an array of integers or as a boolean mask...见 sklearn/utils/validation.pyfrom_dtype的实现通过 narwhalsnw.from_native读取 schema 中的 dtype只对Categorical与Enum两种类型置位见 sklearn/utils/validation.py如果最终没有任何一列被判为类别特征函数返回None即树退化为常规数值树。源码中的校验时机fit 阶段的is_categorical_属性在BaseDecisionTree.fit中类别特征的解析发生在正式训练之前结果缓存为估计器的is_categorical_属性见 sklearn/tree/_classes.pyself.is_categorical_ _check_categorical_features(X, self.categorical_features)这一设计意味着类别特征的识别使用的是原始 dtype此时validate_data尚未把 DataFrame 强转成统一的数值数组因此from_dtype才能在数值化之前捕获到类别信息。训练完成后可通过该属性检查解析结果None表示无类别特征。256 个类别上限的来源bitset 编码与两种 splitter 的容量差异每特征最多 256 个类别这一限制来自底层 C 实现中的比特位集容量。从源码结构看sklearn/tree/_tree.pyx 中定义了MAX_NUM_CATEGORIES_PY N_BITSETS即比特位组的数量对应 256 个类别位槽sklearn/tree/_classes.py 将其导入为 Python 层可见的MAX_NUM_CATEGORIES用于 fit 时检查训练数据中各类别特征实际出现的取值数量是否超限sklearn/tree/_partitioner.pyx 中节点分裂时使用的counts、weighted_counts、means辅助数组均以MAX_NUM_CATEGORIES为长度分配这正是一个类别占一个位/槽的编码方式。同时sklearn/tree/_classes.py 定义了另一个常量MAX_NUM_CATEGORIES_RANDOM 2**24配合注释sklearn/tree/_classes.py可以看出使用splitterbest的经典分裂器受 bitset 容量限制256而splitterrandom的随机分裂器可以处理多得多的类别2^24。也就是说256 的上限主要约束的是默认splitterbest路径。为什么absolute_error不支持类别特征变更说明中明确指出所有准则均支持除absolute_error外。在源码中这一限制以运行时错误形式强制执行见 sklearn/tree/_classes.pyif has_categorical and self.criterion absolute_error: raise ValueError(...) # 提示 criterionabsolute_error 与类别特征不兼容回归器参数文档中也写明了同样的约束Categorical features are not supported withcriterion\absolute_error\见 sklearn/tree/_classes.py。从准则定义看sklearn/tree/_classes.py 中absolute_error: _criterion.MAEMAE 的最优叶预测值是样本中位数而非均值而类别分裂路径下的辅助统计量partitioner 中按类别累计的means等是为均值型目标预计算的因此该组合被整体排除。回归器的可选准则为{squared_error, absolute_error, poisson}见 sklearn/tree/_classes.py其中squared_error与poisson均可与类别特征配合使用分类器侧的二分类场景则可用其支持的划分准则如 gini、entropy、log_loss。适用范围与使用建议结合变更说明与源码可以总结出当前实现的边界场景范围面向二分类与单输出回归超出该范围例如多输出的组合请以运行时的实际校验为准类别数量splitterbest默认下单特征不超过 256 个类别若数据类别更多可考虑splitterrandom或预先合并低频类别与criterion的组合回归任务避免criterionabsolute_error与高基数特征的关系类别特征会被编码进分裂结构相比先 one-hot 再建树它避免了把单一类别特征展开成大量 0/1 列后每棵树只使用其中一列的问题也避免了 ordinal 编码引入的错误序关系数据输入使用from_dtype时建议提供带Categorical/Enumdtype 的 pandas 或 polars DataFrame这是自动识别路径的唯一信息来源。小结该变更让 scikit-learn 的DecisionTreeClassifier/DecisionTreeRegressor具备了原生类别分裂能力通过categorical_features参数布尔掩码、整数下标、列名或from_dtype四种方式声明类别列底层按比特位集编码、默认上限 256 类别除absolute_error外全部准则可用。关键实现分布在 sklearn/tree/_classes.py参数与约束、sklearn/utils/validation.py参数解析、sklearn/tree/_tree.pyx 与 sklearn/tree/_partitioner.pyx编码容量与分裂统计几个文件中可作为后续深入阅读与验证行为时的索引入口。【免费下载链接】scikit-learnscikit-learn: machine learning in Python项目地址: https://gitcode.com/gh_mirrors/sc/scikit-learn创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考