news 2026/9/18 23:33:47

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
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-learn

scikit-learn 的树模块正在迎来一项重要能力:DecisionTreeClassifierDecisionTreeRegressor原生支持类别(categorical)特征,不再强制要求先做 one-hot/ordinal 编码。本文基于当前仓库中的变更说明与源码实现,详解categorical_features参数的四种指定方式、类别数量上限(256)、criterion="absolute_error"的例外限制,以及底层的比特位集(bitset)编码原理,帮助你在混合类型数据上直接训练决策树并理解其约束边界。

功能概述:变更说明了什么

该功能记录在待发布变更 33354.major-feature.rst 中,原文要点如下:

  • tree.DecisionTreeRegressortree.DecisionTreeClassifier原生支持类别特征,适用于**二分类(binary classification)与单输出回归(single-output regression)**场景;
  • 支持所有可用的分裂准则(criteria),唯一例外是'absolute_error'
  • 类别特征通过categorical_features参数指定;
  • 每个特征最多支持 256 个类别

需要说明的适用前提:该条目位于doc/whats_new/upcoming_changes/目录,属于尚未发布的新功能;下文所有行为描述均以当前仓库源码为准。

categorical_features参数:四种指定方式

categorical_features参数同时出现在DecisionTreeClassifierDecisionTreeRegressor的构造函数中(见 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,只对CategoricalEnum两种类型置位(见 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 中,节点分裂时使用的countsweighted_countsmeans辅助数组均以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_errorpoisson均可与类别特征配合使用;分类器侧的二分类场景则可用其支持的划分准则(如 gini、entropy、log_loss)。

适用范围与使用建议

结合变更说明与源码,可以总结出当前实现的边界:

  1. 场景范围:面向二分类与单输出回归;超出该范围(例如多输出)的组合请以运行时的实际校验为准;
  2. 类别数量splitter="best"(默认)下单特征不超过 256 个类别;若数据类别更多,可考虑splitter="random"或预先合并低频类别;
  3. criterion的组合:回归任务避免criterion="absolute_error"
  4. 与高基数特征的关系:类别特征会被编码进分裂结构,相比"先 one-hot 再建树",它避免了把单一类别特征展开成大量 0/1 列后每棵树只使用其中一列的问题,也避免了 ordinal 编码引入的错误序关系;
  5. 数据输入:使用"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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/18 23:31:25

RevokeMsgPatcher 微信防撤回补丁完整安装指南

RevokeMsgPatcher 微信防撤回补丁完整安装指南 【免费下载链接】RevokeMsgPatcher :trollface: A hex editor for WeChat/QQ/TIM - PC版微信/QQ/TIM防撤回补丁(我已经看到了,撤回也没用了) 项目地址: https://gitcode.com/GitHub_Trending/…

作者头像 李华
网站建设 2026/9/18 23:31:23

智慧矿山解决方案PPT怎么写?从数据链路到架构设计全解析

简介:一份38页的基于工业互联网的智慧矿山解决方案PPT,面向矿业企业管理者、智慧矿山项目规划人员及信息化从业者,系统阐述矿山数字化转型的整体路径。内容围绕政策背景、行业痛点与智慧矿山建设目标展开,重点讲解云计算、物联网、…

作者头像 李华
网站建设 2026/9/18 23:30:43

用 kohya_ss 训练你的 AI 画师:零基础 LoRA 训练完整教程

用 kohya_ss 训练你的 AI 画师:零基础 LoRA 训练完整教程 【免费下载链接】kohya_ss 项目地址: https://gitcode.com/GitHub_Trending/ko/kohya_ss 你手里有 20 张喜欢的图,想让 AI 学会其中的"味道",再生成同风格的新图。…

作者头像 李华
网站建设 2026/9/18 23:30:20

IDEA+HBuilderX全栈开发:Spring Boot与uni-app打包

这套组合拳我打了不下几十个项目了——IntelliJ IDEA 负责后端(Java/Spring Boot 那一摊),HBuilderX 负责前端(uni-app 这一摊),两边各自启动、各自打包,最后在服务器或者手机上拼成一个完整产品…

作者头像 李华