scikit-learn 决策树原生支持类别特征:DecisionTreeClassifier 与 DecisionTreeRegressor 的 categorical_features 参数详解
【免费下载链接】scikit-learnscikit-learn: machine learning in Python项目地址: https://gitcode.com/gh_mirrors/sc/scikit-learn
scikit-learn 的树模块正在迎来一项重要能力:DecisionTreeClassifier与DecisionTreeRegressor原生支持类别(categorical)特征,不再强制要求先做 one-hot/ordinal 编码。本文基于当前仓库中的变更说明与源码实现,详解categorical_features参数的四种指定方式、类别数量上限(256)、criterion="absolute_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", default=None即支持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"], dtype="category"), # 类别 "tier": pd.Series(["low", "high", "mid", "high", "low"], dtype="category"), # 类别 }) 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_features="from_dtype").fit(df, y)从_check_categorical_features的源码可以看到几个细节值得注意:
- 字符串列名方式中,如果传入的数据是没有列名的数组,会抛出明确的
ValueError("categorical_features should be passed as an array of integers or as a boolean mask..."),见 sklearn/utils/validation.py; "from_dtype"的实现通过 narwhals(nw.from_native)读取 schema 中的 dtype,只对Categorical与Enum两种类型置位(见 sklearn/utils/validation.py);- 如果最终没有任何一列被判为类别特征,函数返回
None,即树退化为常规数值树。
源码中的校验时机:fit 阶段的is_categorical_属性
在BaseDecisionTree.fit中,类别特征的解析发生在正式训练之前,结果缓存为估计器的is_categorical_属性(见 sklearn/tree/_classes.py):
self.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)可以看出:使用splitter="best"的经典分裂器受 bitset 容量限制(256),而splitter="random"的随机分裂器可以处理多得多的类别(2^24)。也就是说,256 的上限主要约束的是默认splitter="best"路径。
为什么absolute_error不支持类别特征
变更说明中明确指出"所有准则均支持,除'absolute_error'外"。在源码中,这一限制以运行时错误形式强制执行(见 sklearn/tree/_classes.py):
if has_categorical and self.criterion == "absolute_error": raise ValueError(...) # 提示 criterion='absolute_error' 与类别特征不兼容回归器参数文档中也写明了同样的约束:"Categorical features are not supported withcriterion=\"absolute_error\""(见 sklearn/tree/_classes.py)。
从准则定义看(sklearn/tree/_classes.py 中"absolute_error": _criterion.MAE),MAE 的最优叶预测值是样本中位数,而非均值;而类别分裂路径下的辅助统计量(partitioner 中按类别累计的means等)是为均值型目标预计算的,因此该组合被整体排除。回归器的可选准则为{"squared_error", "absolute_error", "poisson"}(见 sklearn/tree/_classes.py),其中squared_error与poisson均可与类别特征配合使用;分类器侧的二分类场景则可用其支持的划分准则(如 gini、entropy、log_loss)。
适用范围与使用建议
结合变更说明与源码,可以总结出当前实现的边界:
- 场景范围:面向二分类与单输出回归;超出该范围(例如多输出)的组合请以运行时的实际校验为准;
- 类别数量:
splitter="best"(默认)下单特征不超过 256 个类别;若数据类别更多,可考虑splitter="random"或预先合并低频类别; - 与
criterion的组合:回归任务避免criterion="absolute_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),仅供参考