Sklearn-源码解析-书-v1-0-二十一-
Sklearn 源码解析(书)v1.0(二十一)
实现细节
-
首行携带元信息(样本数、特征数、标签名称),通过手动解析获取这些元数据。
-
随后逐行读取数值部分,填充到预先分配好的
np.empty矩阵中,保持了对大文件的内存友好。 -
若提供
descr_file_name,会调用load_descr读取对应 RST 说明文档并一并返回。
之所以不直接使用
np.loadtxt,是因为np.loadtxt只能一次性读取数值矩阵,无法捕获首行的元信息。手动解析后,后续循环只处理数值,兼容了“约定优于配置”的资源组织方式。
46.6.2 load_gzip_compressed_csv_data()
def load_gzip_compressed_csv_data(
data_file_name,
*,
data_module=DATA_MODULE,
descr_file_name=None,
descr_module=DESCR_MODULE,
encoding="utf-8",
**kwargs,
):
"""Loads gzip-compressed CSV with `importlib.resources`."""
data_path = resources.files(data_module) / data_file_name
with data_path.open("rb") as compressed_file:
compressed_file = gzip.open(compressed_file, mode="rt", encoding=encoding)
data = np.loadtxt(compressed_file, **kwargs)
if descr_file_name is None:
return data
else:
descr = load_descr(descr_module=descr_module,
descr_file_name=descr_file_name)
return data, descr
-
**kwargs直接透传给np.loadtxt,调用者可以自定义分隔符、跳过行数等参数,避免为每种 CSV 格式额外编写包装函数。 -
读取过程先以二进制打开资源文件,再通过
gzip.open解压为文本流,最后交给np.loadtxt完成一次性加载。
46.6.3 load_descr()
def load_descr(descr_file_name, *, descr_module=DESCR_MODULE, encoding="utf-8"):
"""Load `descr_file_name` from `descr_module` with `importlib.resources`."""
path = resources.files(descr_module) / descr_file_name
return path.read_text(encoding=encoding)
函数直接读取 RST 文本并返回,供 Bunch 的 DESCR 字段使用。
46.7 小结
三个装载器统一遵循 资源包 (data/, descr/, images/) + importlib.resources 的约定,使得代码在打包、分发以及跨平台使用时都保持一致。
46.8 经典玩具数据集加载器 —— 即拿即用的标准样品箱
下面分别解析六个玩具数据集的实现细节,展示它们如何复用上述装载器并完成统一的 Bunch 封装。
46.8.1 load_wine()
data, target, target_names, fdescr = load_csv_data(
data_file_name="wine_data.csv", descr_file_name="wine_data.rst"
)
feature_names = [
"alcohol", "malic_acid", "ash", "alcalinity_of_ash", "magnesium",
"total_phenols", "flavanoids", "nonflavanoid_phenols",
"proanthocyanins", "color_intensity", "hue",
"od280/od315_of_diluted_wines", "proline",
]
frame = None
target_columns = ["target"]
if as_frame:
frame, data, target = _convert_data_dataframe(
"load_wine", data, target, feature_names, target_columns
)
-
通过
load_csv_data获得data(178×13)、target(178)和target_names。 -
feature_names硬编码对应 CSV 列顺序。 -
当
as_frame=True时,调用_convert_data_dataframe将 NumPy 数组零拷贝转换为 Pandas DataFrame/Series,返回完整frame。
46.8.2 load_iris()
data_file_name = "iris.csv"
data, target, target_names, fdescr = load_csv_data(
data_file_name=data_file_name, descr_file_name="iris.rst"
)
feature_names = [
"sepal length (cm)", "sepal width (cm)",
"petal length (cm)", "petal width (cm)",
]
frame = None
target_columns = ["target"]
if as_frame:
frame, data, target = _convert_data_dataframe(
"load_iris", data, target, feature_names, target_columns
)
- 与
load_wine类似,只是额外在返回的Bunch中加入了filename与data_module字段,帮助用户追溯原始 CSV 所在的包路径。
46.8.3 load_breast_cancer()
实现与前两者相同,只是 feature_names 使用 np.array([...]) 以保持与原始 RST 文档中列顺序的一致性。
46.8.4 load_digits()
data, fdescr = load_gzip_compressed_csv_data(
data_file_name="digits.csv.gz", descr_file_name="digits.rst", delimiter=","
)
target = data[:, -1].astype(int, copy=False)
flat_data = data[:, :-1]
images = flat_data.view()
images.shape = (-1, 8, 8)
feature_names = [
f"pixel_{row_idx}_{col_idx}"
for row_idx in range(8)
for col_idx in range(8)
]
-
图像重塑:
flat_data.view()创建对原始flat_data的视图(不复制内存),随后images.shape = (-1, 8, 8)把 64 维特征重新解释为 8×8 图像矩阵。 -
特征名生成:列表推导遍历行、列顺序,生成
pixel_0_0 … pixel_7_7,确保特征名与图像空间布局一一对应。
46.8.5 load_diabetes()
data = load_gzip_compressed_csv_data("diabetes_data_raw.csv.gz")
target = load_gzip_compressed_csv_data("diabetes_target.csv.gz")
if scaled:
data = scale(data, copy=False)
data /= data.shape[0] ** 0.5
scale把每列均值置零、方差归一;随后再除以sqrt(N)(N为样本数),实现 均值中心化 + 样本规模归一化,这在统计学上等价于把每列除以sqrt(N),以便后续模型对不同规模的数据保持数值稳定。
46.8.6 load_linnerud()
data_filename = "linnerud_exercise.csv"
target_filename = "linnerud_physiological.csv"
data_path = resources.files(DATA_MODULE) / data_filename
with data_path.open("r", encoding="utf-8") as f:
header_exercise = f.readline().split()
f.seek(0)
data_exercise = np.loadtxt(f, skiprows=1)
target_path = resources.files(DATA_MODULE) / target_filename
with target_path.open("r", encoding="utf-8") as f:
header_physiological = f.readline().split()
f.seek(0)
data_physiological = np.loadtxt(f, skiprows=1)
- 因为特征文件和目标文件结构略有差异(各自拥有独立的表头),这里没有直接使用
load_csv_data,而是手动读取头部获得feature_names与target_names,随后分别加载数值矩阵。
46.8.7 统一返回
所有玩具数据集最终返回 Bunch,包含 data、target、feature_names、target_names(若有)以及 DESCR。当 as_frame=True 时,还会返回 frame,实现 NumPy 与 Pandas 双向兼容。
46.9 文件夹数据集加载器 —— 批量分拣与格式转换
46.9.1 load_files()
def load_files(
container_path,
*,
description=None,
categories=None,
load_content=True,
shuffle=True,
encoding=None,
decode_error="strict",
random_state=0,
allowed_extensions=None,
):
"""Load text files with categories as subfolder names."""
target = []
target_names = []
filenames = []
folders = [
f for f in sorted(listdir(container_path)) if isdir(join(container_path, f))
]
if categories is not None:
folders = [f for f in folders if f in categories]
if allowed_extensions is not None:
allowed_extensions = frozenset(allowed_extensions)
for label, folder in enumerate(folders):
target_names.append(folder)
folder_path = join(container_path, folder)
files = sorted(listdir(folder_path))
if allowed_extensions is not None:
documents = [
join(folder_path, file)
for file in files
if os.path.splitext(file)[1] in allowed_extensions
]
else:
documents = [join(folder_path, file) for file in files]
target.extend(len(documents) * [label])
filenames.extend(documents)
filenames = np.array(filenames)
target = np.array(target)
if shuffle:
random_state = check_random_state(random_state)
indices = np.arange(filenames.shape[0])
random_state.shuffle(indices)
filenames = filenames[indices]
target = target[indices]
if load_content:
data = []
for filename in filenames:
data.append(Path(filename).read_bytes())
if encoding is not None:
data = [d.decode(encoding, decode_error) for d in data]
return Bunch(
data=data,
filenames=filenames,
target_names=target_names,
target=target,
DESCR=description,
)
return Bunch(
filenames=filenames,
target_names=target_names,
target=target,
DESCR=description,
)
-
目录结构:
container_path下每个子目录被视为 一个类别(target_names),文件路径收集到filenames。 -
过滤:
categories可限制读取的子目录;allowed_extensions限制文件扩展名(如仅.txt)。 -
打乱:当
shuffle=True且提供random_state时,使用check_random_state生成确定性的随机数生成器,确保跨调用的一致顺序。 -
内容读取:
load_content=False只返回路径;True时读取字节流并根据encoding解码为 Unicode 文本。
46.9.2 流程图
46.10 实用工具函数
46.10.1 _convert_data_dataframe()
def _convert_data_dataframe(
caller_name, data, target, feature_names, target_names, sparse_data=False
):
pd = check_pandas_support(f"{caller_name} with as_frame=True")
if not sparse_data:
data_df = pd.DataFrame(data, columns=feature_names, copy=False)
else:
data_df = pd.DataFrame.sparse.from_spmatrix(data, columns=feature_names)
target_df = pd.DataFrame(target, columns=target_names)
combined_df = pd.concat([data_df, target_df], axis=1)
X = combined_df[feature_names]
y = combined_df[target_names]
if y.shape[1] == 1:
y = y.iloc[:, 0]
return combined_df, X, y
-
将 NumPy(稠密或稀疏)数组转为 Pandas DataFrame,并且在创建时使用
copy=False实现 零拷贝。 -
target先包装为 DataFrame,再与data合并,返回完整combined_df、特征子集X与目标子集y。
46.10.2 _pkl_filepath()
def _pkl_filepath(*args, **kwargs):
"""Return filename for Python 3 pickles."""
py3_suffix = kwargs.get("py3_suffix", "_py3")
basename, ext = splitext(args[-1])
basename += py3_suffix
new_args = args[:-1] + (basename + ext,)
return join(*new_args)
- 为 Python 3 的 pickle 文件添加后缀
_py3,确保向后兼容旧版的.pkl文件。
46.10.3 模块级常量
DATA_MODULE = "sklearn.datasets.data"
DESCR_MODULE = "sklearn.datasets.descr"
IMAGES_MODULE = "sklearn.datasets.images"
RemoteFileMetadata = namedtuple(
"RemoteFileMetadata", ["filename", "url", "checksum"]
)
这些常量统一了资源定位的命名空间,所有加载函数均基于它们查找数据、说明文档或示例图片。
46.11 设计中的取舍 —— 一问一答
为什么不直接使用 np.loadtxt 读取 CSV?
load_csv_data 需要在首行获取 样本数、特征数、类别名称,这些元信息是 np.loadtxt 无法捕获的。手动解析首行后,后续循环只处理数值部分,保持了对元数据的完整支持。这种设计遵循了 “约定优于配置” 的原则,使得数据文件只需放在约定好的 data/ 包路径下即可被统一读取。
这种设计的 trade‑off 是什么?
-
优势:明确的元数据约定、一次性读取全部数据、对稀疏/稠密矩阵统一处理、兼容历史数据文件。
-
劣势:在极大 CSV(>10⁶ 行)时手动循环略慢于一次性
np.loadtxt;代码维护成本稍高。总体而言,这种取舍符合 scikit‑learn “易用性 > 极端性能” 的哲学。
46.12 动手练习
-
追踪远程文件下载的完整生命周期
-
阅读
src/sklearn/datasets/_base.py中_fetch_remote()(约第 680‑760 行) -
关键点:
-
临时文件命名模式
prefix=remote.filename + '.part_'的并发安全作用。 -
shutil.move(temp_file_path, file_path)实现原子写入的原理。 -
SHA256 校验在下载前(已有文件)和下载后(新文件)的两次检查。
-
n_retries与delay参数如何实现指数退避重试。
-
-
思考:
-
为什么使用
NamedTemporaryFile(delete=False)而不是直接写目标文件? -
若下载中断留下
.part_临时文件,下次调用会如何处理? -
URLError和TimeoutError之外的异常(如KeyboardInterrupt)如何被捕获并清理?
-
-
-
对比三种本地数据装载器的设计差异
-
阅读
load_csv_data()、load_gzip_compressed_csv_data()、load_descr()。 -
思考:
-
load_csv_data为什么不直接用np.loadtxt而要手动解析首行? -
load_gzip_compressed_csv_data中的**kwargs透传给np.loadtxt的设计意图是什么? -
三个函数如何体现 “约定优于配置” 的资源组织规范(
data/、descr/、images/子包)?
-
-
-
剖析 Digits 数据集的图像重塑与特征名生成
-
阅读
load_digits()实现。 -
思考:
-
flat_data.view()后直接修改images.shape为什么不拷贝内存?这依赖什么条件? -
如果去掉
copy=False,target的内存行为会有什么不同? -
特征名列表推导式中
row_idx与col_idx的嵌套顺序决定了什么?如何验证它与images的空间布局一致?
-
-
-
探究玩具数据集加载器的共性与差异
-
对比阅读
load_wine、load_iris、load_breast_cancer、load_diabetes、load_linnerud。 -
思考:
-
load_linnerud为何不复用load_csv_data而是手动读取两个文件? -
load_diabetes中scaled=True时的data /= data.shape[0] ** 0.5有何统计学含义? -
load_iris与load_breast_cancer返回的 Bunch 多出filename和data_module字段,用途是什么?
-
-
-
实战
load_files目录树加载器- 假设目录结构如下:
container/
├── class_a/
│ ├── doc1.txt
│ └── doc2.txt
└── class_b/
└── doc3.txt
-
思考:
-
categories=['class_a']参数如何过滤目标类别? -
allowed_extensions=['.txt']与encoding='utf-8'如何协作完成文件筛选与解码? -
shuffle=True且random_state=42时,返回的filenames与target的对应关系如何保证? -
load_content=False时返回的 Bunch 缺少哪些字段?适用于什么场景?
-
46.13 本章小结
本章系统地剖析了 scikit‑learn 数据集模块的内部实现,从 数据主目录管理、远程文件安全下载、本地 CSV 与压缩数据的装载、经典玩具数据集的封装,到 目录树文本加载器 与 实用工具函数,形成了完整的“数据物流配送中心”。首先学习了缓存目录的创建与清理机制,其次了解了原子下载与校验的细节,接着掌握了 CSV 与 GZIP 数据的读取方式,随后深入了六大玩具数据集的加载逻辑与 Bunch 封装,随后探索了 load_files 的目录树加载策略,最后了解了 DataFrame 转换、兼容 pickle 路径以及模块级常量的设计。
| 概念 | 解释 |
|------|------|
| get_data_home() | 返回或创建数据缓存根目录,支持环境变量配置 |
| clear_data_home() | 递归删除缓存目录,释放磁盘空间 |
| _fetch_remote() | 原子下载、SHA256 校验、指数退避重试的核心实现 |
| fetch_file() | 高层 API,负责 URL 解析、目录创建、调用 _fetch_remote |
| load_csv_data() | 手动解析首行元信息后逐行读取 CSV |
| load_gzip_compressed_csv_data() | 直接使用 np.loadtxt 读取 gzip 解压后的 CSV |
| load_descr() | 读取 RST 说明文档 |
| load_wine/iris/... | 统一返回 Bunch,支持 as_frame 与 return_X_y |
| load_digits() | 零拷贝视图重塑为 8×8 图像,自动特征名生成 |
| load_files() | 目录树文本加载器,支持过滤、编码、打乱 |
| _convert_data_dataframe() | 将 NumPy 数据转为 Pandas DataFrame/Series(as_frame) |
| _pkl_filepath() | 兼容 Python 3 pickle 路径的后缀处理 |
| 常量 DATA_MODULE / DESCR_MODULE / IMAGES_MODULE | 统一资源定位的命名空间 |
下一章将学习 远程真实数据集的取件之旅,解析
fetch_*系列函数如何从网络获取大型数据集,并了解缓存、校验与并行下载的高级技巧。
46.14 模块地图/架构图
sklearn/datasets/_base.py
├── 缓存目录管理
│ ├── get_data_home() # 获取/创建数据主目录(支持环境变量SCIKIT_LEARN_DATA)
│ └── clear_data_home() # 递归删除缓存目录
├── 远程文件获取与校验
│ ├── _fetch_remote() # 核心下载器:原子写入+SHA256校验+指数退避重试
│ ├── _sha256() # 计算文件SHA256哈希值
│ ├── _derive_folder_and_filename_from_url() # 从URL推导安全的本地目录名与文件名
│ ├── _filter_filename() # 字符串安全化为合法文件名
│ └── fetch_file() # 公共API:按URL下载文件至缓存目录
├── 本地数据装载器
│ ├── load_csv_data() # 解析包内CSV:首行含(n_samples, n_features, target_names...)
│ ├── load_gzip_compressed_csv_data() # 解析包内GZIP压缩CSV(np.loadtxt)
│ ├── load_descr() # 读取包内RST说明文档
│ ├── load_sample_images() # 加载随包附带的JPG示例图片(依赖PIL)
│ └── load_sample_image() # 按名称加载单张示例图片
├── 经典数据集加载器(玩具数据集)
│ ├── load_wine() # 葡萄酒分类数据集(178×13)
│ ├── load_iris() # 鸢尾花分类数据集(150×4)
│ ├── load_breast_cancer() # 乳腺癌分类数据集(569×30)
│ ├── load_digits() # 手写数字图像数据集(1797×64,含8×8 images)
│ ├── load_diabetes() # 糖尿病回归数据集(442×10,可选标准化)
│ └── load_linnerud() # 体能训练多输出回归数据集(20×3→3)
├── 文件夹数据集加载器
│ └── load_files() # 从目录树加载文本文件(子目录=类别)
├── 实用工具函数
│ ├── _convert_data_dataframe() # numpy数组转pandas DataFrame/Series(as_frame=True时用)
│ ├── _pkl_filepath() # 生成带_py3后缀的pickle路径(兼容旧版本)
│ └── __main__ # 模块级全局常量:DATA_MODULE/DESCR_MODULE/IMAGES_MODULE/RemoteFileMetadata
46.15 设计中的取舍 —— 一问一答
为什么采用当前方案而不是更复杂的替代方案? 本章实现优先保证与既有 API 的一致性、可维护性与运行效率。这意味着在少数极端场景下,调用者需要自行在灵活性、内存与速度之间做取舍,换取默认路径的清晰与稳定。
第 47 章 —— California Housing 数据集深度解析
47.1 学习目标
-
掌握scikit-learn远程真实数据集的下载、校验、缓存与加载全流程
-
理解结构化表格数据(California Housing、Covertype)的特征工程与归一化处理
-
掌握异常检测数据集(KDDCUP99)的子集切片、标签二值化与缓存自检机制
-
理解非结构化图像数据(LFW、Olivetti)的解码管线、内存映射缓存与任务模式组织差异
-
掌握异质地理数据(物种分布)的双ZIP归档解析、栅格重建与空间坐标网格构建
-
了解Bunch对象统一封装接口、参数校验装饰器与as_frame/return_X_y灵活返回机制
-
难度:★★★☆☆(3/5)
-
预备知识:Python 基础、面向对象编程与代码阅读基础
47.2 数据集概述
California Housing 数据集是机器学习中经典的 回归 基准题目,记录了加州 20,640 条住宅信息以及对应的房价(中位数)。
它最早出自 StatLib,现在由 scikit‑learn 统一维护并提供一键下载接口。
| 项目 | 说明 |
|-----|------|
| 样本总数 | 20,640 |
| 特征维度 | 8(输入) + 1(目标) |
| 特征类型 | 实数(float) |
| 目标变量 | MedHouseVal(单位:10 万美元),范围约 0.15–5.0 |
| 适用任务 | 回归、特征重要性分析、空间可视化等 |
生活类比
想象你是一位房产经纪人,需要在一张巨大的表格中快速判断某个街区的房价。表格的每一行对应一个街区(样本),列则是影响房价的因素(特征),而目标列是该街区的中位房价。我们通过 fetch_california_housing 把这张表格搬到本地,随后进行“数据清洗”(如把总房间数除以户数得到每户平均房间数),最后交给机器学习模型帮助我们完成定价预测。
就像挑选二手房时会同时看「户型」「房龄」「学区」这些维度,这个数据集把这些维度都量化了,让模型能从数字中学出「哪些因素真的影响房价」。
47.3 源码逐行解析(含关键实现细节)
下面是 sklearn/datasets/_california_housing.py 中核心函数 fetch_california_housing 的完整源码及其逐行中文注释。我们不仅解释「做了什么」,更深入探讨「为什么这样设计」。
import logging
import tarfile
from numbers import Integral, Real
from os import PathLike, remove
from os.path import exists
import joblib
import numpy as np
# 第 47 章 —— ------------------- 1️⃣ 通用工具导入 -------------------
from sklearn.datasets import get_data_home # 获取/创建默认缓存目录
from sklearn.datasets._base import ( # 私有 API,供内部使用
RemoteFileMetadata, # 描述远程文件的元信息
_convert_data_dataframe, # 将 numpy 数据转为 pandas DataFrame
_fetch_remote, # 下载并校验远程文件
_pkl_filepath, # 生成本地 .pkz 缓存文件路径
load_descr, # 读取数据集描述文件
)
from sklearn.utils import Bunch # 类似 dict 的轻量容器
from sklearn.utils._param_validation import ( # 参数校验装饰器
Interval, validate_params,
)
# 第 47 章 —— ------------------- 2️⃣ 远程文件元信息 -------------------
ARCHIVE = RemoteFileMetadata(
filename="cal_housing.tgz",
url="https://ndownloader.figshare.com/files/5976036",
checksum="aaa5c9a6afe2225cc2aed2723682ae403280c4a3695a2ddda4ffb5d8215ea681",
)
logger = logging.getLogger(__name__)
# 第 47 章 —— ==================== 参数校验装饰器 ====================
@validate_params(
{
"data_home": [str, PathLike, None],
"download_if_missing": ["boolean"],
"return_X_y": ["boolean"],
"as_frame": ["boolean"],
"n_retries": [Interval(Integral, 1, None, closed="left")],
"delay": [Interval(Real, 0.0, None, closed="neither")],
},
prefer_skip_nested_validation=True,
)
def fetch_california_housing(
*,
data_home=None,
download_if_missing=True,
return_X_y=False,
as_frame=False,
n_retries=3,
delay=1.0,
):
"""
主函数:下载/读取 California Housing 数据集并返回 Bunch 或 (X, y)。
"""
# ------------------- 3️⃣ 确定缓存根目录 -------------------
data_home = get_data_home(data_home=data_home)
# ------------------- 4️⃣ 本地缓存文件路径 -------------------
filepath = _pkl_filepath(data_home, "cal_housing.pkz")
# ------------------- 5️⃣ 若缓存不存在则下载 -------------------
if not exists(filepath):
if not download_if_missing:
# 用户选择不自动下载时抛明确异常
raise OSError("Data not found and `download_if_missing` is False")
logger.info(
f"Downloading Cal. housing from {ARCHIVE.url} to {data_home}"
)
# 5.1 下载压缩包(支持重试、延时)
archive_path = _fetch_remote(
ARCHIVE,
dirname=data_home,
n_retries=n_retries,
delay=delay,
)
# 5.2 解压并读取原始 CSV(使用 tarfile + np.loadtxt)
with tarfile.open(mode="r:gz", name=archive_path) as f:
# 文件路径在压缩包内部
cal_housing = np.loadtxt(
f.extractfile("CaliforniaHousing/cal_housing.data"),
delimiter=",",
)
# ---------- 列重排 ----------
# 原始列顺序不符合文档约定,需要硬编码映射
columns_index = [8, 7, 2, 3, 4, 5, 6, 1, 0]
# 将目标列(房价)搬到第 0 位,后面的特征依次排列
cal_housing = cal_housing[:, columns_index]
# 将处理好的 ndarray 持久化到本地缓存(压缩级别 6)
joblib.dump(cal_housing, filepath, compress=6)
# 下载完成后删除临时压缩包,节省磁盘空间
remove(archive_path)
else:
# ------------------- 6️⃣ 直接读取本地缓存 -------------------
cal_housing = joblib.load(filepath)
# ------------------- 7️⃣ 定义特征名称(保持顺序) -------------------
feature_names = [
"MedInc", "HouseAge", "AveRooms", "AveBedrms",
"Population", "AveOccup", "Latitude", "Longitude",
]
# ------------------- 8️⃣ 拆分目标与特征 -------------------
target, data = cal_housing[:, 0], cal_housing[:, 1:]
# ------------------- 9️⃣ 特征工程(平均值化) -------------------
# total_rooms / households → avg rooms per household
data[:, 2] /= data[:, 5] # AveRooms
# total_bedrooms / households → avg bedrooms per household
data[:, 3] /= data[:, 5] # AveBedrms
# population / households → avg occupancy per household
data[:, 5] = data[:, 4] / data[:, 5] # AveOccup (覆盖原始 households 列)
# ------------------- 10️⃣ 目标尺度缩放 -------------------
target = target / 100000.0 # 从美元 → 10 万美元单位
# ------------------- 11️⃣ 加载数据集描述 -------------------
descr = load_descr("california_housing.rst")
X, y = data, target
frame = None
target_names = ["MedHouseVal"]
# ------------------- 12️⃣ 可选的 DataFrame 包装 -------------------
if as_frame:
# 将 ndarray → pandas DataFrame/Series,并返回合并的 frame
frame, X, y = _convert_data_dataframe(
"fetch_california_housing", data, target, feature_names, target_names
)
# ------------------- 13️⃣ 返回值决策 -------------------
if return_X_y:
return X, y
return Bunch(
data=X,
target=y,
frame=frame,
target_names=target_names,
feature_names=feature_names,
DESCR=descr,
)
47.3.1 代码要点详解
| 代码片段 | 作用 | 设计思路与取舍 |
|----------|------|----------------|
| RemoteFileMetadata 定义(第 23~28 行) | 描述远程文件名、URL、SHA256 校验和 | 硬编码 URL 与校验和 能一键下载、防范中间人攻击;若数据源更换,仅需改此处。未采用动态检测是为了简化逻辑,官方数据极少变更。 |
| @validate_params 装饰器(第 31~42 行) | 自动检查参数类型与取值范围 | 使用 sklearn 统一参数校验框架,避免重复写 if not isinstance(...);prefer_skip_nested_validation=True 提升嵌套对象校验效率。 |
| get_data_home(第 46 行) | 返回默认缓存目录(如 ~/scikit_learn_data) | 集中管理所有数据集缓存路径,用户可通过 data_home 参数覆盖;比硬编码路径更灵活。 |
| _pkl_filepath(第 49 行) | 生成类似 ~/scikit_learn_data/cal_housing.pkz 的完整路径 | 封装路径拼接逻辑,确保跨平台(Windows/Linux/macOS)一致;.pkz 是 joblib 的压缩 pickle 后缀。 |
| not exists(filepath) 分支(第 52~83 行) | 缓存不存在时触发下载流程 | “缓存优先”策略:避免重复下载,提升二次加载速度;仅在真正需要时联网,符合“离线优先”设计。 |
| _fetch_remote(第 64~68 行) | 下载远程文件并校验 checksum、支持重试 | 内建重试机制(默认 3 次)和指数退避延时,应对暂时性网络波动;校验 checksum 确保数据完整性,防止文件损坏。 |
| tarfile.open + np.loadtxt(第 71~74 行) | 解压 .tgz 并直接读取 CSV 为 numpy array | 为何不用 pandas.read_csv? 本数据集纯数值、无缺失值,np.loadtxt 更快、占用更少内存;配合 joblib 压缩缓存后,I/O 开销极低。 |
| columns_index = [8, 7, 2, 3, 4, 5, 6, 1, 0](第 77 行) | 硬编码列顺序映射 | 原始数据来自 StatLib,列顺序与文档不符;硬编码映射能一次性统一格式,避免运行时解析列名的开销。取舍:若上游数据格式变更需同步修改此列表,但官方数据极稳定。 |
| joblib.dump(..., compress=6)(第 80 行) | 将处理好的数据写入本地缓存 | 为何选 compress=6? joblib 压缩级别 0~9,6 在压缩率与 CPU 开销之间取得平衡:文件变小约 60%,解压速度仍足够快;若追求极速读取可设 0,若磁盘紧张可调至 9。 |
| remove(archive_path)(第 83 行) | 下载后删除临时 .tgz 压缩包 | 节省磁盘空间;压缩包仅用于一次性解压,缓存后即可删除,避免重复存储。 |
| joblib.load(filepath)(第 86 行) | 直接读取本地缓存的 .pkz 文件 | 与 dump 互为逆过程;joblib 支持内存映射(mmap),大数据可部分加载,但本数据集仅 20k 行,直接读取即可。 |
| feature_names 列表(第 89~96 行) | 保持特征顺序与名称一致 | 按照 MedInc, HouseAge, ..., Longitude 的顺序列出,特征工程后仍对应;此顺序必须与 columns_index 重排后的数据列保持一致。 |
| target, data = cal_housing[:, 0], cal_housing[:, 1:](第 99 行) | 拆分标签(目标)与特征矩阵 | 第 0 列为房价目标,其余 8 列为原始特征;这种约定在 sklearn 数据集中非常常见(target 在第 0 列)。 |
| 特征工程三行(第 102~107 行) | 将总量特征转为“每户平均”特征 | - AveRooms = total_rooms / households
- AveBedrms = total_bedrooms / households
- AveOccup = population / households
为什么要这么做? 原始数据提供的是街区总和(如 total_rooms),但模型更关心「每户」情况;除以户数可消除街区规模影响,使特征具有可比性。例如:一个大社区总房间数多,但不一定每户房间多;此步骤赋予特征真正的「密度」含义。 |
| target = target / 100000.0(第 110 行) | 目标从美元转为 10 万美元单位 | 原始数据单位是美元,房价范围约 15000~500000;除以 10 万后落在 0.15~5.0,数值更易于建模(避免大数值导致梯度爆炸);同时保持与文档描述一致。 |
| _convert_data_dataframe(第 123~126 行) | 可选地将数据转为 pandas 对象 | 当 as_frame=True 时,调用 sklearn 内部统一函数生成 DataFrame/Series,保留列名和 dtype;这让探索性分析(如 df.describe()、df.plot())更便利,而纯模型训练可保持 ndarray 以获得更佳性能。 |
| Bunch 返回(第 129~134 行) | 封装所有输出为类 dict 对象 | Bunch 兼容 sklearn API,支持属性访问(如 housing.data)和字典式访问(如 housing['data']);内含 data, target, feature_names, DESCR 等标准字段,frame 仅在 as_frame=True 时存在。 |
💡 设计中的取舍(一问一答形式)
问:为什么使用硬编码的
columns_index而非动态检测列顺序?
答:硬编码映射能一次性确保特征顺序统一,运行时开销几乎为零;动态检测(如读取 CSV 首行)虽然更鲁棒,但需要额外 IO 和解析逻辑,而此数据集多年未变更格式,故牺牲极少的鲁棒性换取更高的加载速度是值得的。
问:为什么特征工程要除以
data[:, 5](即原始 households 列)?
答:因为原始数据中,
total_rooms,total_bedrooms,population都是街区级总和;只有除以户数(households)后,才能得到「每户平均房间数」、「每户平均卧室数」、「每户平均居住人数」,这些才是衡量街区居住密度与舒适度的有效特征,否则大社区会被误判为「房间更多」而非「每户更宽敞」。
问:为何选用 joblib 而非原生 pickle 进行缓存?
答:joblib 对大型 numpy 数组有优化的压缩和内存映射机制,读写速度比 pickle 快数倍,尤其在反复加载同一数据集时优势明显;同时其压缩格式仍保持与 numpy 兼容,解压后直接可用。
问:目标除以 10 万这一步是否必须?
答:从数学上说不是必须的(模型能学习到尺度),但从工程角度看能提升数值稳定性:梯度下降在特征尺度相同时收敛更快;同时保持与文档、教材一致,减少使用者混淆。
47.4 数据流与架构图(按小节配置图表)
47.5 生活类比贯穿全文:从表格到模型的隐喻
正如房产经纪人不会只看总价,而是综合「户型」「房龄」「每户人数」等维度判断性价比,我们的数据处理流程也是从「原始总和」→「平均化特征」→「尺度归一化」,一步步让特征更贴近真实决策维度。
这一思想贯穿于特征工程(第 9️⃣ 步)和目标缩放(第 🔟 步),它们共同把难以直接比较的街区总量转化为可用于建模的「每户」指标。
47.5.1 源码地图:函数调用与数据转换流程(对应「源码逐行解析」小节)
47.5.2 特征工程逻辑图:从原始总和到每户平均值(对应「特征工程」段落)
47.5.3 目标尺度变换图:美元 → 10 万美元(对应「目标尺度缩放」段落)
47.6 使用示例(含生活类比延伸)
>>> from sklearn.datasets import fetch_california_housing
>>> # 场景:你是房产经纪人,先拿到数据看看街区特征分布
>>> housing = fetch_california_housing()
>>> print("样本数:", housing.data.shape[0]) # 20640 条街区记录
>>> print("特征均值(每户平均值后):")
>>> for name, val in zip(housing.feature_names, housing.data.mean(axis=0)):
... print(f" {name}: {val:.2f}")
...
样本数: 20640
特征均值(每户平均值后):
MedInc: 3.87
HouseAge: 28.64
AveRooms: 5.43
AveBedrms: 1.09
Population: 1425.59
AveOccup: 2.50
Latitude: 35.63
Longitude: -119.57
>>> # 再看看目标房价分布:是否符合你的预期?
>>> import numpy as np
>>> print("房价中位数:", np.median(housing.target)) # ~1.80(即 18 万美元)
>>> print("房价均值:", housing.target.mean()) # ~1.90(19 万美元)
>>> # 如果你想用 pandas 做探索性分析(如画直方图):
>>> housing_df = fetch_california_housing(as_frame=True)
>>> housing_df.data.hist(bins=30, figsize=(12, 8))
>>> # 标题可设为:「加州各街区核心特征分布(经纪人视角)」
生活类比延伸
假设你看到某街区
AveRooms=6.0(每户平均 6 间房)但MedInc=2.0(中位收入仅 2 万美元),你会怎么判断?
可能是「豪华大院但本地人收入低」(如退休社区),或者「数据异常」。
这个数据集正是通过把「总房间数」转化为「每户平均房间数」,让你的判断不再被街区规模所误导——这正是特征工程的价值所在。
47.7 小结
-
通过 统一的下载‑缓存‑加载 流程,用户只需一次函数调用即可获得完整、预处理好的回归数据集。
-
关键实现围绕 列重排、特征工程(平均值化)、尺度缩放 三个步骤展开,保证特征语义与文档描述一致,并通过生活类比强化「每户平均」这一核心思想。
-
采用 joblib 高效缓存、validate_params 参数校验、以及 Bunch 统一返回结构,使得 API 与 scikit‑learn 其他数据集保持一致,易于上手和二次开发。
-
设计中多处取舍均衡了 性能、鲁棒性 和 易用性:如硬编码列顺序提升速度、joblib 压缩级别 6 平衡空间与时间、特征除以户数消除规模偏差等。
实战建议:在初次实验时使用
as_frame=True快速浏览特征分布(如收入 vs 房价散点图);正式模型训练时切换回默认ndarray以获得更佳性能。记住:好特征不是凭空捏造的,而是从原始数据中通过合理的数学变换(如除以户数)挖掘出来的——就像经纪人不会只看挂牌价,而是会算「每平方英尺单价」。
47.8 设计中的取舍
为什么采用当前方案,而不是更复杂的替代方案? 本章源码优先选择清晰、可维护且与既有 API 兼容的实现;这降低了使用和调试成本,但也意味着部分极端场景需要调用者自行权衡性能、灵活性与实现复杂度。
47.9 动手练习
47.9.1 对比California Housing与Covertype的缓存策略差异
阅读 _california_housing.py 与 _covtype.py 的缓存相关代码:
-
California Housing 使用单文件
cal_housing.pkz存储合并数组,而 Covertype 分离samples与targets两个文件,这种设计差异的工程考量是什么? -
Covertype 使用
TemporaryDirectory(dir=covtype_dir) + os.rename实现原子写入,California Housing 直接joblib.dump到目标路径,前者解决了什么并发风险? -
两者
compress参数分别为 6 和 9,压缩级别选择的权衡因素是什么?
47.9.2 分析KDDCUP99的子集构建逻辑与工程妥协
阅读 _kddcup99.py 中 fetch_kddcup99 的 subset 参数处理逻辑:
-
subset='SA'时,为何保留所有正常样本而仅随机抽取 3377 条异常样本?这种极度不平衡构造的基准意义是什么? -
subset='SF'/'http'/'smtp'共享logged_in==1筛选与log(x+0.1)变换,但列索引硬编码(如data[:,11]),这种硬编码在工程上为何可接受? -
_fetch_brute_kddcup99中定义的结构化dtype dt如何支撑混合类型(int/float/bytes)的精确解析?
47.9.3 探究LFW双任务模式的数据组织差异与缓存复用
对比 _lfw.py 中 _fetch_lfw_people 与 _fetch_lfw_pairs 的实现:
-
两者均使用
joblib.Memory(location=lfw_home, compress=6).cache装饰加载器,这种缓存粒度(函数级)与California Housing的文件级缓存有何异同? -
fetch_lfw_people返回images(n, H, W) 与data(n, HW),而fetch_lfw_pairs返回pairs(n, 2, H, W) 与data(n, 2H*W),这种形状设计如何服务于分类与验证两种任务的典型模型输入? -
slice_与resize参数如何正交控制图像预处理管线?默认slice_(70:195, 78:172) + resize=0.5的几何含义是什么?
47.9.4 解析物种分布数据集的地理栅格重建机制
阅读 _species_distributions.py 的核心解析流程:
-
SAMPLES.zip与COVERAGES.zip均通过np.load作为 npz 读取,再用BytesIO包装内部文件,这种“嵌套归档”解析模式的优势是什么? -
_load_coverage中NODATA_value替换为-9999的处理,如何配合dtype=np.int16节省内存?若原始NODATA_value已为-9999会怎样? -
construct_grids根据x_left_lower_corner/Nx/grid_size重建坐标网格,这种“元数据驱动几何重建”模式在GIS数据工程中的通用性体现在哪里?
47.10 源码地图:函数调用与数据转换流程(对应「源码逐行解析」小节)
sklearn/datasets/_california_housing.py
├── ARCHIVE # RemoteFileMetadata: 远程文件元数据
├── fetch_california_housing() # 公共入口函数
│ ├── @validate_params # 参数校验装饰器
│ ├── get_data_home() # 获取数据根目录
│ ├── _pkl_filepath() # 生成缓存路径
│ ├── _fetch_remote() # 下载远程文件(含重试/校验)
│ ├── tarfile.open() # 解压.tgz归档
│ ├── np.loadtxt() # 解析CSV数据
│ ├── columns_index重排 # 列顺序调整 [8,7,2,3,4,5,6,1,0]
│ ├── joblib.dump(compress=6) # 序列化缓存
│ ├── 领域特征工程 # 计算人均指标、目标缩放
│ ├── _convert_data_dataframe() # 可选DataFrame转换
│ └── Bunch封装 # 统一返回结构
sklearn/datasets/_covtype.py
├── ARCHIVE # RemoteFileMetadata
├── FEATURE_NAMES/TARGET_NAMES # 硬编码特征/目标名
├── fetch_covtype() # 公共入口函数
│ ├── @validate_params
│ ├── get_data_home()
│ ├── _pkl_filepath() # 分离samples/targets路径
│ ├── TemporaryDirectory(dir=covtype_dir) # 原子写入保证
│ ├── _fetch_remote()
│ ├── GzipFile + np.genfromtxt() # 流式解压解析
│ ├── joblib.dump(compress=9) # 高压缩缓存
│ ├── os.rename() # 原子移动防半写
│ ├── check_random_state + shuffle # 可选打乱
│ ├── _convert_data_dataframe()
│ └── Bunch封装
sklearn/datasets/_kddcup99.py
├── ARCHIVE / ARCHIVE_10_PERCENT # 双数据源元数据
├── fetch_kddcup99() # 公共入口(含subset逻辑)
│ ├── @validate_params(StrOptions subset)
│ ├── _fetch_brute_kddcup99() # 核心加载器
│ ├── subset='SA' 逻辑 # 正常全保留+异常抽样3377
│ ├── subset='SF'/'http'/'smtp' 逻辑 # logged_in筛选+log变换+协议切片
│ ├── shuffle_method() # 可选打乱
│ └── Bunch封装
├── _fetch_brute_kddcup99() # 底层加载实现
│ ├── 目录版本隔离 kddcup99-py3 / _10-py3
│ ├── 结构化dtype定义 dt (42列混合类型)
│ ├── joblib.load 缓存读取 + 异常自检自愈
│ ├── GzipFile 逐行解码 + np.asarray(object)
│ ├── 逐列astype(DT[j]) 精确类型转换
│ └── joblib.dump(compress=3) 缓存写入
├── _mkdirp() # 兼容mkdir -p
sklearn/datasets/_lfw.py
├── ARCHIVE / FUNNELED_ARCHIVE / TARGETS # 三档数据源元数据
├── _check_fetch_lfw() # 统一下载调度器
│ ├── 并行下载3标注文件 + 按funneled选图像包
│ ├── tarfile_extractall 解压 + 删除压缩包
├── _load_imgs() # 图像解码管线
│ ├── PIL.Image 裁剪/缩放/归一化/灰度化
│ ├── 预分配float32连续内存
│ └── 批量解码填充
├── _fetch_lfw_people() # 人脸识别任务加载器
│ ├── 扫描子目录名为人名 + min_faces过滤
│ ├── np.searchsorted 标签编码
│ ├── RandomState(42) 固定种子打破IID
│ └── joblib.Memory(location=lfw_home, compress=6) 缓存装饰
├── fetch_lfw_people() # 公共入口
├── _fetch_lfw_pairs() # 人脸验证任务加载器
│ ├── 解析pairsDevTrain.txt等索引文件
│ ├── 3列=同人/4列=异人 构造pairs张量
│ └── reshape至 (n_pairs, 2, H, W)
├── fetch_lfw_pairs() # 公共入口
│ ├── subset选择索引文件
│ └── Bunch封装(data/pairs/target/target_names)
sklearn/datasets/_olivetti_faces.py
├── FACES # RemoteFileMetadata .mat文件
├── fetch_olivetti_faces() # 公共入口
│ ├── @validate_params
│ ├── _pkl_filepath(olivetti.pkz)
│ ├── scipy.io.loadmat 读取faces键
│ ├── .T.copy() 转置拷贝 (400, 4096)
│ ├── 归一化: float32 -> [0,1]
│ ├── reshape(400,64,64).transpose(0,2,1) # HWC顺序
│ ├── target = i//10 隐含标签
│ ├── check_random_state shuffle
│ └── Bunch(data/images/target/DESCR)
sklearn/datasets/_species_distributions.py
├── SAMPLES / COVERAGES # 双ZIP归档元数据
├── DATA_ARCHIVE_NAME = 'species_coverage.pkz'
├── __main__ # 全局代码:extra_params/grid元数据定义
├── _load_coverage() # 栅格文件解析
│ ├── 读取6行头提取NODATA_value
│ ├── np.loadtxt 读取数据矩阵
│ └── 异常值替换为 -9999
├── _load_csv() # 训练/测试点解析
│ ├── np.loadtxt dtype='S22,f4,f4'
├── construct_grids() # 空间坐标网格重建
│ ├── 根据元数据计算xgrid/ygrid
├── fetch_species_distributions() # 公共入口
│ ├── extra_params 固定地理元数据
│ ├── np.load(zip) 作为npz读取
│ ├── BytesIO 包装文件句柄
│ ├── 14个栅格堆叠 (14, 1592, 1212)
│ ├── joblib.dump(compress=9) 最高压缩
│ └── 返回Bunch(coverages/train/test/地理元数据)
第 48 章 —— OpenML 数据生态 —— 接入“机器学习数据集商店”
48.1 学习目标
-
难度:★★★☆☆(3/5)
-
预备知识:Python 基础、面向对象编程与 Markdown/代码阅读基础
-
完整获取链路:从 OpenML API 查询元数据 → 下载 ARFF(gzip)文件 → MD5 完整性校验 → 双解析器(LIAC‑ARFF 与 Pandas)自动调度 → 生成统一的
Bunch对象。 -
双解析器架构:了解 LIAC‑ARFF(支持稀疏 ARFF、纯 Python 实现)与 Pandas(基于
read_csv的高速 C 加速实现)之间的权衡、自动选择策略以及何时强制指定解析器。 -
网络层韧性:掌握分层缓存、原子写入、失败缓存自动清理、指数退避重试等机制,确保在并发、网络抖动或磁盘异常下仍能可靠获取数据。
-
数据完整性保障:学习目标列同质性检查、缺失值检测、特征列筛选(排除
is_ignore、is_row_identifier),以及分类标签的逆映射与categorydtype 统一。 -
离线测试夹具:了解
tests/data/openml/目录下的离线数据如何配合 monkey‑patch 实现零网络、确定性单元测试。
48.2 生活类比 —— “跨国精密仪器的进口报关与组装”
把 OpenML 数据获取想象成一次跨国精密仪器的 进口报关与组装 流程:
| 步骤 | 类比对象 | 作用 |
|------|----------|------|
| 元数据 API (_get_data_info_by_name、_get_data_description_by_id) | 电子提单,列出仪器型号、版本、配件清单、校验码(MD5) | 为后续下载指明 哪款仪器、哪一批次、哪些组件 必须采购。 |
| 保税仓库 + 原子签收 (_open_openml_url + TemporaryDirectory + shutil.move) | 在目的港的保税仓库,采用临时箱子装载,装满后一次性搬入正式仓库 | 防止并发写入导致的 半成品(损坏缓存)。 |
| 双装配线 (LIAC‑ARFF 与 Pandas) | 两条装配线:手工工艺(通用,能处理稀疏部件) vs 自动流水线(高速,仅稠密) | 根据 元数据中的 format 决定走哪条线;parser='auto' 自动分流。 |
| 质检站 (_load_arff_response – MD5 校验、ParserError 重试) | 检测包装是否完整、部件是否缺失,若不合格则退回重新装箱 | MD5 不匹配 → 抛异常 → _retry_with_clean_cache 自动删除并重下。 |
| 入库标准化 (_verify_target_data_type、_valid_data_column_names) | 检查仪器关键部件(目标列)是否齐全、材质统一,剔除标记为 不参与测量 的配件 | 确保交付给车间的 标准化组件箱(Bunch)符合质量要求。 |
| 离线样本库 (tests/data/openml/) | 实验室的标准样本库,用于离线验证装配流程 | 单元测试直接读取本地 ARFF/JSON,避免真实海关与运输的波动。 |
在后续章节,类比会持续出现:缓存 如同 保税仓库,重试 如同 重新报关,解析器 如同 装配线,帮助读者对抽象代码保持直观感知。
48.3 OpenML 源码地图
sklearn/datasets/_arff_parser.py
├─ _split_sparse_columns() # 稀疏列筛选 + 索引重映射
├─ _sparse_data_to_array() # 稀疏三元组 → 稠密 NumPy
├─ _post_process_frame() # DataFrame → X / y 切分
├─ _liac_arff_parser() # 纯 Python 解析器(稀疏/稠密)
│ ├─ _io_to_generator()
│ ├─ _arff.load(...)
│ ├─ dense: itertools.chain + np.fromiter
│ ├─ sparse: _split_sparse_columns → coo_matrix
│ ├─ 分类目标逆映射
│ └─ pandas 输出:块读取 + concat + dtype 转换
├─ _pandas_arff_parser() # Pandas.read_csv 实现(仅稠密)
│ ├─ 跳过 @data 行
│ ├─ 构造 dtype(Int64 / category)
│ ├─ read_csv 参数(na_values、quotechar、…)
│ ├─ 单引号剥离
│ └─ _post_process_frame()
└─ load_arff_from_gzip_file() # 统一入口,调度两种解析器
sklearn/datasets/_openml.py
├─ _get_local_path()
├─ _retry_with_clean_cache()
├─ _retry_on_network_error()
├─ _open_openml_url() # 网络下载 + 原子缓存
├─ _get_* 系列 API # data、features、qualities、qualities → 样本数
├─ _load_arff_response() # MD5 校验 → 解析器调度 → ParserError 重试
├─ _download_data_to_bunch() # 目标列检查、特征列过滤、Bunch 组装
├─ _verify_target_data_type()
├─ _valid_data_column_names()
└─ fetch_openml() # 公共入口:参数校验 → 元数据获取 → 下载 → 返回
48.4 ARFF 解析核心 —— 双引擎驱动的数据翻译器
48.4.1 为什么需要双解析器架构?
-
LIAC‑ARFF:纯 Python 实现,能够解析 稀疏 ARFF(COO 三元组),对
nominal列保留原始字符串。代价是 CPU 与内存消耗大,但通用性强。 -
Pandas 解析器:利用
pandas.read_csv的 C 加速路径,速度快、自动 dtype 推断、支持Int64扩展整数与category。但只能处理 稠密 ARFF,不支持 sparse COO 语法。 -
自动选择 (
parser='auto'):当元数据中的format为sparse_arff时,强制使用liac‑arff;否则默认走pandas,在性能与通用性之间取得平衡。
48.4.2 逐行注释 — load_arff_from_gzip_file
def load_arff_from_gzip_file(
gzip_file,
parser,
output_type,
openml_columns_info,
feature_names_to_select,
target_names_to_select,
shape=None,
read_csv_kwargs=None,
):
"""
Load a compressed ARFF file using a given parser.
参数解释
----------
gzip_file : GzipFile 实例
已经解压缩的二进制流,来源可以是本地缓存也可以是网络请求。
parser : {"pandas", "liac-arff"}
选择的解析器;"pandas" 推荐用于稠密数据,"liac-arff" 用于稀疏或兼容性需求。
output_type : {"numpy", "sparse", "pandas"}
决定返回的 X / y 结构:NumPy 数组、稀疏 CSR 矩阵或 Pandas DataFrame/Series。
openml_columns_info : dict
OpenML 提供的列级元信息(数据类型、是否目标、索引等),后续用于 dtype 映射与过滤。
feature_names_to_select / target_names_to_select : list[str]
需要保留的特征列与目标列名称,过滤掉 `is_ignore`、`is_row_identifier` 等。
shape : tuple[int, int] | None
当解析器返回 generator 时必须显式提供数据形状;稠密/稀疏 ARFF 均可能需要此信息。
read_csv_kwargs : dict | None
传给 `pandas.read_csv` 的可选参数,允许用户覆盖默认的 CSV 选项。
"""
# 分流调度:根据 parser 参数调用对应的底层实现
if parser == "liac-arff":
# 调用纯 Python 实现,负责稀疏 & 稠密两种情况
return _liac_arff_parser(
gzip_file,
output_type,
openml_columns_info,
feature_names_to_select,
target_names_to_select,
shape,
)
elif parser == "pandas":
# 调用基于 pandas.read_csv 的实现,仅限稠密 ARFF
return _pandas_arff_parser(
gzip_file,
output_type,
openml_columns_info,
feature_names_to_select,
target_names_to_select,
read_csv_kwargs,
)
else:
# 防御性检查:若用户提供了未知的 parser,抛出明确异常
raise ValueError(
f"Unknown parser: '{parser}'. Should be 'liac-arff' or 'pandas'."
)
代码作用概述:本函数是 统一入口,屏蔽不同解析器的实现细节,确保
fetch_openml只需要传递统一的参数即可完成 ARFF 解析。它首先检查parser,随后把所有上下文(gzip 流、列信息、返回类型等)原原本本转交给对应的实现函数。
48.4.3 逐行注释 — _liac_arff_parser
def _liac_arff_parser(
gzip_file,
output_arrays_type,
openml_columns_info,
feature_names_to_select,
target_names_to_select,
shape=None,
):
"""
ARFF parser using the LIAC-ARFF library coded purely in Python.
负责稀疏与稠密 ARFF 的统一解析,返回 NumPy / SciPy 稀疏 / Pandas 对象。
"""
# --------------------------------------------------------------
# 1️⃣ 将 gzip 流转换为 UTF‑8 文本生成器,按行解码,避免一次性解压到内存
def _io_to_generator(gzip_file):
for line in gzip_file:
# 每行都是 bytes,需要转成 str
yield line.decode("utf-8")
stream = _io_to_generator(gzip_file)
# --------------------------------------------------------------
# 2️⃣ 决定 ARFF 的内部表示方式:
# * 如果 output_arrays_type=="sparse" → 返回稀疏 COO 三元组 (COO)
# * 否则返回密集生成器 (DENSE_GEN)
return_type = _arff.COO if output_arrays_type == "sparse" else _arff.DENSE_GEN
# --------------------------------------------------------------
# 3️⃣ 禁用 LIAC‑ARFF 默认的 nominal 编码,以保留类别原始字符串
encode_nominal = not (output_arrays_type == "pandas")
# --------------------------------------------------------------
# 4️⃣ 调用外部库解析 ARFF
arff_container = _arff.load(
stream,
return_type=return_type,
encode_nominal=encode_nominal,
)
# --------------------------------------------------------------
# 5️⃣ 合并需要返回的列(特征 + 目标),并收集 nominal (categorical) 列的映射
columns_to_select = feature_names_to_select + target_names_to_select
categories = {
name: cat
for name, cat in arff_container["attributes"]
if isinstance(cat, list) and name in columns_to_select
}
# --------------------------------------------------------------
# 6️⃣ Pandas 返回路径(output_arrays_type == "pandas")
if output_arrays_type == "pandas":
pd = check_pandas_support("fetch_openml with as_frame=True")
# a. 把属性信息转成有序字典,获取列名顺序
columns_info = OrderedDict(arff_container["attributes"])
columns_names = list(columns_info.keys())
# b. 读取第一行数据,用于估算每行所占内存,进而计算块大小(chunksize)
first_row = next(arff_container["data"])
first_df = pd.DataFrame([first_row], columns=columns_names, copy=False)
row_bytes = first_df.memory_usage(deep=True).sum()
chunksize = get_chunk_n_rows(row_bytes)
# c. 逐块读取,保留需要的列
columns_to_keep = [col for col in columns_names if col in columns_to_select]
dfs = [first_df[columns_to_keep]]
for data in chunk_generator(arff_container["data"], chunksize):
dfs.append(
pd.DataFrame(data, columns=columns_names, copy=False)[columns_to_keep]
)
# 若块数>1,使用第二块推断 dtype,修正第一块的 dtype
if len(dfs) >= 2:
dfs[0] = dfs[0].astype(dfs[1].dtypes)
# d. LIAC‑ARFF 用 None 表示缺失值,这里统一转成 np.nan
frame = pd.concat(dfs, ignore_index=True)
frame = pd_fillna(pd, frame)
# e. 根据 OpenML 元信息为整数列使用 Pandas 扩展类型 Int64,
# 为 nominal 列使用 category,以保留缺失值信息
dtypes = {}
for name in frame.columns:
column_dtype = openml_columns_info[name]["data_type"]
if column_dtype.lower() == "integer":
dtypes[name] = "Int64" # 支持缺失值的整数 dtype
elif column_dtype.lower() == "nominal":
dtypes[name] = "category"
else:
dtypes[name] = frame.dtypes[name]
frame = frame.astype(dtypes)
# f. 使用通用的切分函数得到 X / y
X, y = _post_process_frame(
frame, feature_names_to_select, target_names_to_select
)
# --------------------------------------------------------------
# 7️⃣ 非 Pandas 路径:返回 NumPy / 稀疏矩阵
else:
arff_data = arff_container["data"]
# a. 把特征 / 目标列名称映射到在 ARFF 文件中的整数索引
feature_indices_to_select = [
int(openml_columns_info[col_name]["index"])
for col_name in feature_names_to_select
]
target_indices_to_select = [
int(openml_columns_info[col_name]["index"])
for col_name in target_names_to_select
]
# b. 当 ARFF 数据是生成器(稠密)时,需要提前知道 shape,以便一次性读取
if isinstance(arff_data, Generator):
if shape is None:
raise ValueError(
"shape must be provided when arr['data'] is a Generator"
)
# 计算总元素数目;shape[0]==-1 表示未知样本数
count = -1 if shape[0] == -1 else shape[0] * shape[1]
data = np.fromiter(
itertools.chain.from_iterable(arff_data),
dtype="float64",
count=count,
)
data = data.reshape(*shape)
X = data[:, feature_indices_to_select]
y = data[:, target_indices_to_select]
# c. 当 ARFF 数据是稀疏三元组时,需要先筛选列、重映射索引并构造 COO → CSR
elif isinstance(arff_data, tuple):
# 只保留需要的特征列
arff_data_X = _split_sparse_columns(arff_data, feature_indices_to_select)
num_obs = max(arff_data[1]) + 1
X_shape = (num_obs, len(feature_indices_to_select))
X = sp.sparse.coo_matrix(
(arff_data_X[0], (arff_data_X[1], arff_data_X[2])),
shape=X_shape,
dtype=np.float64,
).tocsr()
# 目标列转换为稠密数组
y = _sparse_data_to_array(arff_data, target_indices_to_select)
else:
# 理论上永远不应该进入此分支
raise ValueError(
f"Unexpected type for data obtained from arff: {type(arff_data)}"
)
# d. 分类目标逆映射:把整数编码恢复为原始标签字符串
is_classification = {
col_name in categories for col_name in target_names_to_select
}
if not is_classification:
pass # 没有目标列
elif all(is_classification):
# 对每个目标列做逆映射
y = np.hstack(
[
np.take(
np.asarray(categories.pop(col_name), dtype="O"),
y[:, i : i + 1].astype(int, copy=False),
)
for i, col_name in enumerate(target_names_to_select)
]
)
elif any(is_classification):
# 混合 nominal 与 numeric 目标目前不被支持
raise ValueError(
"Mix of nominal and non-nominal targets is not currently supported"
)
# e. 若只有单目标列,压平成 1‑D;若无目标列,则返回 None
if y.shape[1] == 1:
y = y.reshape((-1,))
elif y.shape[1] == 0:
y = None
# --------------------------------------------------------------
# 8️⃣ 根据请求的返回类型统一返回四元组
if output_arrays_type == "pandas":
# 当返回 Pandas 时,categories 已经包含在 DataFrame 的 dtype 中
return X, y, frame, None
# 否则返回 X / y / None / categories(非 Pandas 场景需要 categories)
return X, y, None, categories
代码作用概述:本函数是 双解析器的核心,通过 流式生成器、块读取、稀疏三元组 → COO → CSR 的转换,既保证了对大规模稀疏数据的内存友好,又能够在
output_arrays_type="pandas"时直接返回高效的DataFrame。分类目标逆映射确保即使 LIAC‑ARFF 把 nominal 变成整数,也能在返回前恢复原始标签。
48.4.4 逐行注释 — _pandas_arff_parser
def _pandas_arff_parser(
gzip_file,
output_arrays_type,
openml_columns_info,
feature_names_to_select,
target_names_to_select,
read_csv_kwargs=None,
):
"""
ARFF parser using `pandas.read_csv`.
只适用于稠密 ARFF;在读取前手动跳过 ARFF 头部(@attribute、@data)。
"""
import pandas as pd
# 1️⃣ 跳过 ARFF 元数据,定位到真正的 CSV 数据行(@data 之后)
for line in gzip_file:
if line.decode("utf-8").lower().startswith("@data"):
break
# 2️⃣ 根据 OpenML 元信息准备 dtype 映射
dtypes = {}
for name in openml_columns_info:
column_dtype = openml_columns_info[name]["data_type"]
if column_dtype.lower() == "integer":
dtypes[name] = "Int64" # 支持缺失值的整数扩展类型
elif column_dtype.lower() == "nominal":
dtypes[name] = "category"
# 3️⃣ pandas.read_csv 要求以列索引的方式传递 dtypes;
# 把列名 → 整数索引的映射准备好
dtypes_positional = {
col_idx: dtypes[name]
for col_idx, name in enumerate(openml_columns_info)
if name in dtypes
}
# 4️⃣ 默认的 read_csv 参数(符合 ARFF 规范)
default_read_csv_kwargs = {
"header": None,
"index_col": False,
"na_values": ["?"],
"keep_default_na": False,
"comment": "%",
"quotechar": '"',
"skipinitialspace": True,
"escapechar": "\\",
"dtype": dtypes_positional,
}
# 用户提供的参数会覆盖默认值
read_csv_kwargs = {**default_read_csv_kwargs, **(read_csv_kwargs or {})}
# 5️⃣ 读取 CSV 数据
frame = pd.read_csv(gzip_file, **read_csv_kwargs)
# 6️⃣ 为了保持与 LIAC‑ARFF 行为一致,需要把单引号包裹的字符串去除
single_quote_pattern = re.compile(r"^'(?P<contents>.*)'$")
def strip_single_quotes(input_string):
match = re.search(single_quote_pattern, input_string)
return input_string if match is None else match.group("contents")
# 只对 categorical 列执行去引号操作
categorical_columns = [
name
for name, dtype in frame.dtypes.items()
if isinstance(dtype, pd.CategoricalDtype)
]
for col in categorical_columns:
frame[col] = frame[col].cat.rename_categories(strip_single_quotes)
# 7️⃣ 切分特征 / 目标
X, y = _post_process_frame(
frame, feature_names_to_select, target_names_to_select
)
# 8️⃣ 根据要求返回 Pandas 对象或转成 NumPy / 稀疏
if output_arrays_type == "pandas":
return X, y, frame, None
else:
# 转成 NumPy(稠密)后返回,categories 用于非 Pandas 场景
X, y = X.to_numpy(), y.to_numpy()
categories = {
name: dtype.categories.tolist()
for name, dtype in frame.dtypes.items()
if isinstance(dtype, pd.CategoricalDtype)
}
return X, y, None, categories
代码作用概述:本函数把 ARFF 当作 CSV 读取,利用 Pandas 的 C 加速路径实现 数倍的速度提升。它在读取后手动剥离单引号,确保 类别标签 与 LIAC‑ARFF 保持一致。若用户请求
output_arrays_type != "pandas",函数会把 DataFrame 转回 NumPy,并返回categories供后续逆映射。
48.4.5 架构视图——解析器调度图
说明:
fetch_openml根据元信息决定parser(auto时会检查format),随后_load_arff_response调用统一入口load_arff_from_gzip_file,最终的返回结构取决于output_arrays_type(pandas、numpy、sparse)。
48.5 OpenML API 交互与缓存策略 —— 网络层的“重试与熔断”设计
48.5.1 分层缓存机制
-
本地目录结构:
~/scikit_learn_data/openml.org/<url_path>.gz,保持与服务器路径一一对应,便于直接定位下载文件。 -
原子写入:在
TemporaryDirectory中下载至临时文件,成功后使用shutil.move原子搬入缓存目录,防止并发写入产生的 半成品。 -
缓存失效装饰器
_retry_with_clean_cache:捕获除URLError之外的所有异常,删除本地缓存后仅一次重新执行函数,解决“毒药缓存”。
48.5.1.1 关键函数示例
def _open_openml_url(
url: str, data_home: Optional[str], n_retries: int = 3, delay: float = 1.0
):
"""
负责下载资源并写入本地缓存。若 data_home 为 None 则直接返回网络流。
- 采用 Accept‑encoding: gzip 请求压缩数据。
- 若缓存不存在,则在目标目录的子目录 tmpdir 中以原子方式写入。
- 下载异常时会清理残留文件并向上抛出,配合 @_retry_with_clean_cache 重试。
"""
# …实现代码略…
作用概述:该函数实现了 “先缓存后网络” 的双层策略,确保在多进程/多线程环境下缓存文件的完整性与一致性,同时对网络波动提供自动重试。
48.5.2 网络韧性:指数退避与错误过滤
def _retry_on_network_error(
n_retries: int = 3, delay: float = 1.0, url: str = ""
):
"""
对网络错误进行指数退避重试(URLError、TimeoutError),
但对 OpenML 专有的 412 错误直接抛出。
"""
# …实现代码略…
作用概述:在下载阶段若出现暂时的网络故障(如 DNS 超时、临时 5xx)会自动 重试 三次,每次间隔
delay秒;而对 412 错误(业务层错误)则不重试,直接让调用者处理。
48.5.3 流程图——缓存与重试
48.6 数据下载、校验与装箱 —— 从字节流到 Bunch 的完整管线
48.6.1 完整性校验 (_load_arff_response)
-
流式 MD5:读取 gzip 流时每次读取 4096 B,实时更新
hashlib.md5,避免一次性将完整文件读入内存。 -
校验失败处理:若校验不匹配抛出
ValueError,包装信息提示“清理缓存并重试”。随后装饰器_retry_with_clean_cache捕获异常、删除缓存、再次执行下载。
def _load_arff_response(...):
gzip_file = _open_openml_url(url, data_home, n_retries=n_retries, delay=delay)
with closing(gzip_file):
md5 = hashlib.md5()
for chunk in iter(lambda: gzip_file.read(4096), b""):
md5.update(chunk)
actual_md5_checksum = md5.hexdigest()
if actual_md5_checksum != md5_checksum:
raise ValueError(...)
# 之后调用 load_arff_from_gzip_file 进行实际解析
代码作用概述:它把 网络下载 → 完整性校验 → 解析器调度 合二为一,确保任何文件损坏都能在同一次 API 调用中被发现并自动恢复。
48.6.2 目标列合法性检查 (_verify_target_data_type)
-
检查目标列是否 全部为 numeric 或 全部为 nominal(同质性),不支持混合。
-
若出现缺失值,立即抛出
ValueError。 -
对被标记为
is_ignore或is_row_identifier的目标列,仅 warn 而不阻止加载。
def _verify_target_data_type(features_dict, target_columns):
if not isinstance(target_columns, list):
raise ValueError(...)
found_types = set()
for target_column in target_columns:
if target_column not in features_dict:
raise KeyError(...)
if features_dict[target_column]["data_type"] == "numeric":
found_types.add(np.float64)
else:
found_types.add(object)
if features_dict[target_column]["is_ignore"] == "true":
warn(...)
if features_dict[target_column]["is_row_identifier"] == "true":
warn(...)
if len(found_types) > 1:
raise ValueError("Can only handle homogeneous multi-target datasets...")
代码作用概述:确保模型训练时目标列的 类型统一,防止出现混合 float 与 object 导致的 fit 错误。
48.6.3 特征列筛选 (_valid_data_column_names)
def _valid_data_column_names(features_list, target_columns):
valid_data_column_names = []
for feature in features_list:
if (
feature["name"] not in target_columns
and feature["is_ignore"] != "true"
and feature["is_row_identifier"] != "true"
):
valid_data_column_names.append(feature["name"])
return valid_data_column_names
-
排除 目标列、
is_ignore="true"、is_row_identifier="true"。 -
当
as_frame=False时,还会检查data_type=="string",如果出现则抛出错误提醒使用as_frame=True。
48.6.4 Bunch 装箱 (_download_data_to_bunch)
-
output_type 决策:
sparse→"sparse";as_frame=True→"pandas";否则"numpy"。 -
调用
_load_arff_response(已被_retry_with_clean_cache包装)获取(X, y, frame, categories)。 -
最后把 数据、目标、完整 DataFrame、类别映射以及 元信息(
feature_names、target_names、DESCR、details、url)封装进Bunch。
return Bunch(
data=X,
target=y,
frame=frame,
categories=categories,
feature_names=data_columns,
target_names=target_columns,
)
代码作用概述:它是从 底层网络流 到 用户可直接使用的 Bunch 的桥梁,负责 校验、返回类型选择、异常统一处理 与 最终包装。
48.6.5 流程图——完整管线
48.7 OpenML 测试数据夹具 —— 离线测试的“标本库”
-
位置:
sklearn/datasets/tests/data/openml/id_<id>/,每个子目录包含 预下载的.arff.gz与 对应的 JSON 元信息。 -
离线模式:
_monkey_patch_webbased_functions用urlopen替换为读取本地夹具的函数,所有网络请求都被拦截。 -
覆盖场景:
-
稠密数据 (
id_1,id_2,id_3) -
稀疏 ARFF (
id_61,id_62) → 验证 LIAC‑ARFF 分支 -
多目标 (
id_292,id_561) -
边界情况:缺失值、字符串特征、不同
quotechar/escapechar(id_1119,id_1590) -
新版兼容性 (
id_40589、id_40675、id_40945、id_40966、id_42074、id_42585)
-
48.7.1 夹具在测试中的角色
说明:通过这些本地资源,单元测试在 不依赖网络 的情况下,实现 确定性 与 高速,保证所有边界场景均得到完整覆盖。
48.8 设计取舍 —— “性能 vs 通用性” 的一问一答
Q:为什么不直接把 Pandas 解析器也扩展到稀疏 ARFF?
A:Pandas 的 read_csv 只能处理普通 CSV,缺少对 ARFF 中 sparse_arff 语法({index value, ...})的解析能力。实现完整的稀疏解析会导致内部重写 CSV 解析器,违背 Pandas 轻量、维护成本低的设计原则。因此保留 LIAC‑ARFF 作为专门的稀疏处理模块。
Q:双解析器会导致维护成本与代码重复吗?
A:是的,维护两套代码会增加测试工作量。但这是一种 向后兼容 的折中:OpenML 数据集仍然会出现稀疏 ARFF,完全去除 LIAC‑ARFF 将导致老数据无法加载。通过 严密的单元测试(见 tests/test_openml.py)我们能够保持两套实现的同步可靠。
Q:为何在 as_frame=False 时仍返回 categories?
A:当返回的是纯 NumPy/稀疏矩阵时,类别信息会被编码为整数。categories 用于把这些整数映射回原始标签,保持 API 与 as_frame=True 下的行为一致。Pandas 直接保留 category dtype,故返回 None。
Q:如果用户强制使用 parser='pandas' 加载稀疏 ARFF,会发生什么?
A:fetch_openml 在解析前会检测 format=="sparse_arff" 并抛出 ValueError,提示用户改用 liac‑arff。这样可以提前拦截不兼容的配置,避免隐藏的运行时错误。
48.9 动手练习
-
追踪双解析器的分流与汇合逻辑
-
查看
fetch_openml(约 578‑800 行)与load_arff_from_gzip_file(约 363‑430 行)。 -
关键判断:
parser='auto'时,parser_ = "liac-arff"当return_sparse=True,否则为"pandas"。 -
当
parser='pandas'但数据为稀疏时,fetch_openml会提前抛出ValueError("Sparse ARFF datasets cannot be loaded with parser='pandas'")。
-
-
剖析缓存原子写入与 MD5 校验的协同机制
-
_open_openml_url使用TemporaryDirectory+shutil.move实现 原子写入。 -
MD5 在
_load_arff_response中 下载完成后 校验,若不匹配触发_retry_with_clean_cache删除缓存并重下。
-
-
扩展新输出格式(Polars)
-
在
fetch_openml的StrOptions中加入'polars'。 -
在
_download_data_to_bunch中,当as_frame=True且用户指定output_type='polars'时,将output_type设为'polars'并传递给解析器。 -
解析器返回的 Pandas
DataFrame/Series在_download_data_to_bunch中使用polars.from_pandas完成转换;categories在 Polars 场景下返回None,因为 Polars 已保留Categoricaldtype。
-
48.10 本章小结
本章系统地解构了 OpenML 数据获取的全链路:从 元数据 API 抽取数据集 ID 与下载 URL、通过 分层缓存与原子写入 确保并发安全、利用 MD5 流式校验 防止文件损坏、在 双解析器(LIAC‑ARFF 与 Pandas)之间实现 自动调度,并在 目标列合法性 与 特征列筛选 层面保证数据完整性。随后我们把解析得到的矩阵、标签及丰富的元信息装箱进 Bunch,为后续模型训练提供统一的输入接口。最后,通过 离线测试夹具 实现了 确定性、零网络依赖 的单元测试,覆盖了稀疏、稠密、多目标、缺失值、特殊字符等所有关键边界场景。
48.10.1 概念速查表
| 概念 | 解释 |
|------|------|
| fetch_openml | OpenML 数据获取入口,统一处理缓存、解析器选择、格式转换与元信息注入。 |
| _liac_arff_parser | 纯 Python ARFF 解析器,支持稀疏和稠密,返回 NumPy / SciPy / Pandas。 |
| _pandas_arff_parser | 基于 pandas.read_csv 的高速解析器,仅限稠密 ARFF。 |
| parser='auto' | 自动调度器:稀疏 ARFF → LIAC‑ARFF,其他 → Pandas。 |
| _open_openml_url | 网络下载入口,使用 TemporaryDirectory + shutil.move 实现原子写入,确保缓存并发安全。 |
| MD5 校验 | 下载后流式计算哈希,确保文件未被篡改或下载不完整。 |
| _retry_with_clean_cache | 捕获异常 → 删除本地缓存 → 再次尝试,防止“毒药缓存”。 |
| _verify_target_data_type | 检查目标列同质性与缺失值,必要时发出警告或抛异常。 |
| _valid_data_column_names | 从完整特征列表中剔除目标列、is_ignore、is_row_identifier 标记的列。 |
| Bunch | 包含 data、target、frame、categories、DESCR、details、url 的字典式容器。 |
| 离线测试夹具 | tests/data/openml/ 目录下的预下载 ARFF / JSON,用于脱网单元测试,确保所有边界场景均得到可靠覆盖。 |
下一章,我们将继续深入 概率校准(第 49 章),学习
sklearn.calibration.CalibratedClassifierCV如何通过交叉验证和 Sigmoid / Temperature 缩放把分类器的原始输出转化为可靠的概率估计。
48.11 OpenML 源码地图
sklearn/datasets/_arff_parser.py
├── _split_sparse_columns() # 稀疏矩阵列筛选与索引重映射
├── _sparse_data_to_array() # 稀疏三元组转稠密目标数组
├── _post_process_frame() # DataFrame 列切分逻辑 (X/y 分离)
├── _liac_arff_parser() # LIAC-ARFF 纯 Python 解析器 (支持稀疏/稠密)
│ ├── 文本流生成器 _io_to_generator
│ ├── 调用 sklearn.externals._arff.load (COO/DENSE_GEN 模式)
│ ├── 稠密模式: itertools.chain + np.fromiter 高效构建数组
│ ├── 稀疏模式: _split_sparse_columns + scipy.sparse.coo_matrix
│ ├── 分类目标后处理: categories 映射还原原始标签
│ └── Pandas 输出模式: 分块读取 + pd.concat + dtype 转换
├── _pandas_arff_parser() # Pandas 基于 read_csv 的高效解析器 (仅稠密)
│ ├── 跳过 ARFF 头部定位 @data 行
│ ├── 基于 OpenML 元信息预定义 dtype (Int64/category)
│ ├── pd.read_csv 配置 na_values/comment/quotechar 等 ARFF 规范参数
│ ├── 单引号剥离后处理 (分类列 rename_categories)
│ └── _post_process_frame 切分 X/y
└── load_arff_from_gzip_file() # 统一入口,根据 parser 参数分发
sklearn/datasets/_openml.py
├── _get_local_path() # 缓存路径构建 (保留服务端目录结构)
├── _retry_with_clean_cache() # 缓存失效装饰器: 捕获异常 -> 删除损坏缓存 -> 重试
├── _retry_on_network_error() # 网络错误指数退避重试装饰器 (排除 HTTP 412)
├── _open_openml_url() # 统一网络入口: gzip 解压、原子写入 (TemporaryDirectory + shutil.move)
├── _get_json_content_from_openml_api() # JSON API 封装: 统一错误码 412 转 OpenMLError
├── _get_data_info_by_name() # 按名称/版本搜索数据集 ID (支持 version='active')
├── _get_data_description_by_id() # 获取数据集描述、URL、MD5、格式等核心元信息
├── _get_data_features() # 获取列级元数据 (名称、类型、是否目标、缺失值、忽略标记)
├── _get_data_qualities() # 获取质量指标 (NumberOfInstances 等,用于预分配形状)
├── _get_num_samples() # 从 qualities 提取样本数
├── _load_arff_response() # 下载 ARFF + MD5 校验 + 调用解析器 + Pandas ParserError 兜底重试
├── _verify_target_data_type() # 目标列合法性: 无缺失值、同质性 (全数值/全分类)、警告 ignore/row_identifier
├── _valid_data_column_names() # 特征列筛选: 排除目标/忽略/行标识列,非 DataFrame 模式拦截 string 类型
├── _download_data_to_bunch() # 核心管线: 决定 output_type -> 校验目标 -> 重试下载解析 -> 构建 Bunch
├── fetch_openml() # 公共 API: 参数校验 -> 元数据获取 -> 确定 parser/output_type -> 调用下载管线 -> 注入 DESCR/details/url
└── __main__ # 全局代码块 (用于直接运行脚本调试)
sklearn/datasets/tests/data/openml/
├── id_1/__init__.py # 基础稠密分类/回归测试夹具
├── id_2/__init__.py # 基础稠密测试夹具
├── id_3/__init__.py # 基础稠密测试夹具
├── id_61/__init__.py # 稀疏 ARFF 格式测试夹具 (验证 LIAC-ARFF 路径)
├── id_62/__init__.py # 稀疏 ARFF 格式测试夹具
├── id_292/__init__.py # 多目标数据集测试夹具
├── id_561/__init__.py # 多目标数据集测试夹具
├── id_1119/__init__.py # 边界情况: 缺失值、字符串特征
├── id_1590/__init__.py # 边界情况测试夹具
├── id_40589/__init__.py # 新版本数据集 API 兼容性验证
├── id_40675/__init__.py # 新版本测试夹具
├── id_40945/__init__.py # 新版本测试夹具
├── id_40966/__init__.py # 新版本测试夹具
├── id_42074/__init__.py # 新版本测试夹具
└── id_42585/__init__.py # 新版本测试夹具
sklearn/datasets/tests/test_openml.py
└── 测试用例 * # 覆盖 parser/as_frame/sparse 等参数组合,通过 data_home 注入本地夹具实现离线测试
48.12 设计取舍 —— “性能 vs 通用性” 的一问一答
为什么采用当前方案而不是更复杂的替代方案? 本章实现优先保证与既有 API 的一致性、可维护性与运行效率。这意味着在少数极端场景下,调用者需要自行在灵活性、内存与速度之间做取舍,换取默认路径的清晰与稳定。
第 49 章 —— 文本与多标签数据 —— 处理“非结构化数据的翻译机”
49.1 学习目标
-
难度:★★★☆☆(3/5)
-
预备知识:Python 基础、面向对象编程与 Markdown/代码阅读基础
-
理解 20 Newsgroups 数据集的下载、缓存、文本清洗与向量化全流程实现。
-
掌握 RCV1 多标签数据集的分片下载、稀疏矩阵拼装、样本 ID 对齐置换与时间序列分割机制。
-
了解 joblib 压缩缓存、原子性文件操作、增量下载重试等工程化数据加载模式。
-
能阅读并复现大规模文本数据集的预处理与特征工程管道。
49.2 生活类比(完整段落)
想象数据加载器是一家跨国物流与加工厂:远程仓库存放来自 Figshare 服务器的原始压缩包,清关与解压由 _download_20newsgroups / fetch_rcv1 负责(支持断点续传与校验);原料预处理车间通过 strip_newsgroup_header、strip_newsgroup_footer、strip_newsgroup_quoting 去除邮件头、签名档和引用行这些“包装箱标签、说明书、填充物”;分拣流水线依据 categories 筛选感兴趣的主题、subset 分割训练/测试集并通过 shuffle 打乱顺序;精加工车间则使用 CountVectorizer 将清洗后的文本转为特征向量,并做 L2 归一化;所有中间产物最终存入 joblib 压缩缓存(如 .pkz、.pkl 文件),实现“零拷贝、按需加载、可复现”的高效周转——就像现代工厂追求零库存与柔性制造,数据加载器也通过缓存与流式解析避免重复劳动,只在真正需要时才进行昂贵的下载与解压操作。
扩展:在随后的小节中,每一个关键步骤都会对应一个“工厂环节”的比喻,帮助读者把抽象的代码流程映射到真实的物流场景中,从原材料进库到成品出库,一气呵成。
49.3 源码地图
sklearn/datasets/_twenty_newsgroups.py
├── _download_20newsgroups() # 下载、解压、序列化缓存、清理原始文件
├ ├── strip_newsgroup_header() # 基于 '\\n\\n' 分割剥离邮件头
├── │ ├── strip_newsgroup_footer() # 启发式识别并移除签名档
├── │ │ ├── strip_newsgroup_quoting() # 正则匹配移除引用行
├── │ │ │ ├── fetch_20newsgroups() # 主入口:缓存加载、文本清洗、类别筛选、打乱、返回 Bunch
├── │ │ │ │ └── fetch_20newsgroups_vectorized() # 向量化入口:CountVectorizer 拟合/转换、稀疏矩阵缓存、归一化、DataFrame 支持
sklearn/datasets/_rcv1.py
├── fetch_rcv1() # 主入口:分片下载 X/y、样本 ID 对齐、稀疏矩阵构建、时间/随机分割
├── _find_permutation() # 计算从数组 a 到 b 的排序置换索引
├── _inverse_permutation() # 计算置换索引的逆置换
└── __main__ # 模块级全局常量定义(ARCHIVE、XY_METADATA 等)
49.4 20 Newsgroups 文本语料 —— 从原始新闻到特征向量的“翻译流水线”
49.4.1 为什么需要压缩缓存 (CACHE_NAME)?
原始解压后的 20 Newsgroups 数据集约 86 MB。若每次都重新下载、解压并加载,启动开销极大。scikit‑learn 使用 joblib + zlib 将缓存的字典序列化为压缩 pickle(.pkz),体积仅为原始的约 1/6。二次加载时直接反序列化此缓存,跳过网络传输与解压阶段,显著提升响应速度。缓存键中会嵌入 remove 参数(如 headers、quotes),确保不同预处理选项产生的缓存互不干扰。
49.4.2 fetch_20newsgroups 如何实现“延迟下载”
def fetch_20newsgroups(...):
# ① 获取数据根目录并构建缓存路径
data_home = get_data_home(data_home=data_home)
cache_path = _pkl_filepath(data_home, CACHE_NAME)
# ② 若本地缓存已存在则直接读取
if os.path.exists(cache_path):
with open(cache_path, "rb") as f:
compressed_content = f.read()
uncompressed = codecs.decode(compressed_content, "zlib_codec")
cache = pickle.loads(uncompressed) # ← 读取缓存
# ③ 缓存缺失或失效则触发下载
if cache is None:
if download_if_missing:
cache = _download_20newsgroups(
target_dir=twenty_home,
cache_path=cache_path,
n_retries=n_retries,
delay=delay,
)
else:
raise OSError("20Newsgroups dataset not found")
# ……后续按 subset、categories、shuffle 等继续处理
逐行解释
-
第 1‑2 行:获取统一的
data_home,保证所有 scikit‑learn 数据都落在同一根目录。 -
第 3‑4 行:使用
_pkl_filepath将CACHE_NAME拼接成完整路径(如~/scikit_learn_data/20news-bydate.pkz)。 -
第 5‑10 行:如果缓存文件已存在,先读取二进制内容并用
zlib_codec解压,再用pickle反序列化为 Python 对象。此过程耗时极短。 -
第 11‑19 行:若缓存失效或不存在且
download_if_missing=True,调用_download_20newsgroups完成一次性下载、解压、缓存并清理临时目录的原子操作。否则抛出错误,提醒使用者手动准备数据。
49.4.3 文本清洗三件套 (remove 参数) 的设计哲学
def strip_newsgroup_header(text):
_before, _blankline, after = text.partition("\n\n")
return after
_QUOTE_RE = re.compile(r"(writes in|writes:|wrote:|says:|said:|^In article|^Quoted from|^\||^>)")
def strip_newsgroup_quoting(text):
good_lines = [line for line in text.split("\n") if not _QUOTE_RE.search(line)]
return "\n".join(good_lines)
def strip_newsgroup_footer(text):
lines = text.strip().split("\n")
for line_num in range(len(lines) - 1, -1, -1):
line = lines[line_num]
if line.strip().strip("-") == "":
break
return "\n".join(lines[:line_num]) if line_num > 0 else text
解释
-
strip_newsgroup_header使用str.partition("\n\n")将文本在首个双换行符处分割,返回后半段即去掉所有邮件头部元信息。 -
strip_newsgroup_quoting通过预编译正则_QUOTE_RE匹配以>、|、writes:等开头的行,过滤掉这些可能是引用的内容。 -
strip_newsgroup_footer从文本末尾向前扫描,寻找仅包含空白或连字符的行(常见的签名分隔--),并截断其之后的内容。
设计取舍:采用轻量级字符串操作而非 BeautifulSoup 等 HTML 解析库,牺牲对极端异常格式的鲁棒性,换取零外部依赖、极高执行效率以及行为可预期性。
49.4.4 fetch_20newsgroups_vectorized 如何实现“开箱即用的向量化”
def fetch_20newsgroups_vectorized(...):
# ① 构建向量化缓存文件名,包含 remove 状态保证唯一性
filebase = "20newsgroup_vectorized"
if remove:
filebase += "remove-" + "-".join(remove)
target_file = _pkl_filepath(data_home, filebase + ".pkl")
# ② 复用 fetch_20newsgroups 获得原始文本(固定 random_state=12 保证缓存一致)
data_train = fetch_20newsgroups(..., subset="train", remove=remove, random_state=12)
data_test = fetch_20newsgroups(..., subset="test", remove=remove, random_state=12)
# ③ 若缓存存在则直接加载
if os.path.exists(target_file):
X_train, X_test, feature_names = joblib.load(target_file)
else: # 否则现场拟合 CountVectorizer 并缓存
vectorizer = CountVectorizer(dtype=np.int16)
X_train = vectorizer.fit_transform(data_train.data).tocsr()
X_test = vectorizer.transform(data_test.data).tocsr()
feature_names = vectorizer.get_feature_names_out()
joblib.dump((X_train, X_test, feature_names), target_file, compress=9)
# ④ 可选 L2 归一化
if normalize:
X_train = X_train.astype(np.float64)
X_test = X_test.astype(np.float64)
preprocessing.normalize(X_train, copy=False)
preprocessing.normalize(X_test, copy=False)
# ⑤ 根据 subset 返回对应矩阵或拼接全体
if subset == "train":
data, target = X_train, data_train.target
elif subset == "test":
data, target = X_test, data_test.target
else: # all
data = sp.vstack((X_train, X_test)).tocsr()
target = np.concatenate((data_train.target, data_test.target))
# ⑥ 可选 DataFrame 包装
if as_frame:
frame, data, target = _convert_data_dataframe(
"fetch_20newsgroups_vectorized",
data,
target,
feature_names,
target_names=["category_class"],
sparse_data=True,
)
# …返回 Bunch 或 (data, target)…
逐行注释
-
缓存键:
filebase加上remove参数,确保不同清洗配置对应不同缓存文件,避免混用。 -
复用:先调用
fetch_20newsgroups获得已经清洗好的文本,保持下载与清洗只进行一次。 -
缓存检查:若向量化结果已存在,直接
joblib.load,省去重新拟合的成本。 -
现场拟合:
CountVectorizer(dtype=np.int16)采用 16 位整数压缩计数矩阵,随后保存。 -
归一化:
preprocessing.normalize原地修改矩阵,避免额外拷贝。 -
子集返回:依据
subset选择训练、测试或全部数据,sp.vstack用于稀疏矩阵高效拼接。 -
DataFrame:如果用户需要 pandas 接口,调用内部
_convert_data_dataframe完成包装。
49.4.5 架构图:20 Newsgroups 加载 Pipeline
49.5 RCV1 多标签分类数据 —— 百万级稀疏矩阵的“分布式拼装术”
49.5.1 分片下载与流式解析
if download_if_missing and (not exists(samples_path) or not exists(sample_id_path)):
files = []
for each in XY_METADATA: # 5 个特征分片
logger.info("Downloading %s" % each.url)
file_path = _fetch_remote(each, dirname=rcv1_dir,
n_retries=n_retries, delay=delay)
files.append(GzipFile(filename=file_path))
# 使用 load_svmlight_files 逐文件流式读取,避免一次性载入整个 gzip
Xy = load_svmlight_files(files, n_features=N_FEATURES)
# 按官方顺序垂直堆叠训练/测试分片并拼接 sample_id
X = sp.vstack([Xy[8], Xy[0], Xy[2], Xy[4], Xy[6]]).tocsr()
sample_id = np.hstack((Xy[9], Xy[1], Xy[3], Xy[5], Xy[7])).astype(np.uint32)
joblib.dump(X, samples_path, compress=9)
joblib.dump(sample_id, sample_id_path, compress=9)
for f in files: # 关闭句柄并删除临时 gzip
f.close()
remove(f.name)
-
分片:5 个
.dat.gz文件分别对应训练集与测试集的不同时间段。 -
流式读取:
load_svmlight_files接受文件对象列表,以生成器方式逐行解码,显著降低内存峰值。 -
拼装:
sp.vstack按官方文档指定的顺序(test_pt0, train, test_pt1, …)堆叠,确保时间顺序保持不变。
49.5.2 标签文件的稠密‑稀疏转换
y = np.zeros((N_SAMPLES, N_CATEGORIES), dtype=np.uint8) # 稠密矩阵占用约 80 MB
sample_id_bis = np.zeros(N_SAMPLES, dtype=np.int32)
category_names = {}
with GzipFile(filename=topics_archive_path, mode="rb") as f:
for line in f:
cat, doc, _ = line.decode("ascii").split(" ")
if cat not in category_names:
n_cat += 1
category_names[cat] = n_cat
doc = int(doc)
if doc != doc_previous:
doc_previous = doc
n_doc += 1
sample_id_bis[n_doc] = doc
y[n_doc, category_names[cat]] = 1
y = sp.csr_matrix(y[:, order]) # 转为 CSR,压缩稀疏结构
-
构建稠密矩阵:在解析期间直接对
y[n_doc, col] = 1赋值,逻辑简洁。 -
转换 CSR:完成后一次性
csr_matrix转换,显著压缩存储(稀疏率约 3.15%),随后joblib.dump(..., compress=9)进行二次压缩。
49.5.3 _find_permutation 与 _inverse_permutation 的协作
def _inverse_permutation(p):
n = p.size
s = np.zeros(n, dtype=np.int32)
i = np.arange(n, dtype=np.int32)
np.put(s, p, i) # s[p] = i
return s
def _find_permutation(a, b):
t = np.argsort(a) # a 排序后的索引
u = np.argsort(b) # b 排序后的索引
u_ = _inverse_permutation(u) # b 的逆置换
return t[u_] # 组合得到从 a 到 b 的置换
- 作用:
sample_id_bis(标签文件中的文档 ID 顺序)与sample_id(特征文件顺序)不一致。先对两者分别argsort,得到对应的排序索引t与u。_inverse_permutation(u)将u逆转,使其能够映射回原始位置,最后t[u_]即为把y按sample_id_bis重新排列,以匹配特征矩阵的顺序。
49.5.4 架构图:RCV1 加载 Pipeline
49.6 设计中的取舍
在 20 Newsgroups 的文本清洗实现上,作者选择了 纯标准库字符串操作(partition、split、正则)而非像 BeautifulSoup 那样的通用 HTML 解析器。此举带来的 trade‑off 如下:
-
优势:零外部依赖、执行速度极快、行为透明、便于审计,尤其在缓存后只执行一次。
-
劣势:对极端、畸形邮件格式的鲁棒性有限;若数据源结构发生变化,需手动调整正则或分割逻辑。
结论:对于已知且相对固定的文本数据集(如 20 Newsgroups),“恰好够用”的轻量实现更符合工程效率与可维护性的目标;若面向更通用的邮件或 HTML 内容,则应考虑更健壮的解析库。
49.7 动手练习
49.7.1 练习 1:fetch_20newsgroups 细节
-
缓存键为何要包含
remove参数?因为不同的文本清洗组合会产生不同的实际内容。如果缓存键不区分
remove,则在一次清洗后缓存的结果会在后续请求中被错误复用,导致模型看到与预期不一致的文本。将remove作为键的一部分保证每种清洗配置都有独立缓存,避免交叉污染。 -
categories筛选后为何使用np.searchsorted重新编码target?原始
target是全局的类别索引(0‑19),当只保留子集categories时,这些索引会出现间隙。np.searchsorted(labels, data.target)将剩余标签重新映射为连续的整数序列(0, 1, …),确保后续模型训练时标签空间紧凑且不出现空洞。
49.7.2 练习 2:fetch_rcv1 细节
-
分片下载与
load_svmlight_files拼装 X 的流程-
循环遍历
XY_METADATA中的 5 条RemoteFileMetadata,调用_fetch_remote下载每个 gzip 包。 -
将每个文件包装为
GzipFile对象,形成文件对象列表files。 -
load_svmlight_files(files, n_features=N_FEATURES)逐文件流式读取 SVMLight 格式,返回一个列表,其中奇数索引为特征矩阵、偶数索引为对应的sample_id。 -
使用
sp.vstack按官方顺序垂直堆叠特征子矩阵,得到完整的稀疏矩阵X;使用np.hstack合并所有sample_id。
-
-
标签稠密‑稀疏构建过程
-
初始化全零稠密矩阵
y(np.uint8)以及sample_id_bis用于记录文档 ID 的顺序。 -
逐行解析
.qrels.gz,为每个文档‑主题对在稠密矩阵中设置1。 -
完成后通过
sp.csr_matrix将稠密矩�压缩为 CSR 稀疏格式,再使用joblib.dump(..., compress=9)缓存。
-
-
sample_id_bis与sample_id顺序不一致的现实场景-
特征文件 按 下载顺序(即时间戳)排列,保证训练样本在前。
-
标签文件 则是按 文档 ID(内部编号)排序,可能与特征的时间顺序不匹配。此不一致来源于原始数据集的历史生成方式,需要通过置换对齐。
-
-
subset='train'与shuffle=True同时指定时的先后顺序-
首先依据
subset对X、y、sample_id进行 时间序列切分(训练集为前 23 149 条)。 -
若
shuffle=True,随后在切分后的子集上执行随机打乱(shuffle_),这一步会破坏时间顺序,但仅影响已划定的训练或测试集合。
-
49.7.3 练习 3:向量化与缓存策略比较
-
两种策略的优缺点
-
20 Newsgroups(现场拟合 + 缓存)
-
优点:灵活,可随
remove、stop_words、ngram_range等参数自由组合;首次运行后缓存即可重复使用。 -
缺点:首次向量化需要遍历全部文本,计算量随文本规模线性增长;不适合极大数据集(数十 GB)会导致内存压力。
-
-
RCV1(预计算稀疏向量)
-
优点:特征矩阵已在服务器端完成向量化,下载后直接使用,极大降低本地计算成本,适合 百万级样本。
-
缺点:缺乏自定义空间;如果需要不同的特征抽取方式(如 TF‑IDF、子词),只能重新下载或自行重算。
-
-
-
若在
fetch_20newsgroups_vectorized中支持自定义TfidfVectorizer,缓存键的设计建议将关键的向量化参数(如
use_idf、norm、sublinear_tf、ngram_range、max_df、min_df)拼接进缓存文件名,例如:params_key = "-".join([ f"idf{use_idf}", f"norm{norm}", f"sub{int(sublinear_tf)}", f"ngram{ngram_range[0]}x{ngram_range[1]}", f"maxdf{max_df}", f"mindf{min_df}" ]) filebase = f"20newsgroup_vectorized_{params_key}" if remove: filebase += "_remove-" + "-".join(remove) target_file = _pkl_filepath(data_home, filebase + ".pkl")这样每一种向量化配置都有唯一的缓存文件,避免因参数变更而误读旧缓存,同时保持键的可读性与可追溯性。
49.8 本章小结
下面表格对本章的关键概念进行归纳。
| 概念 | 解释 |
|------|------|
| CACHE_NAME (.pkz) | 压缩 pickle 缓存,二次加载避免重复下载解压。 |
| strip_newsgroup_header/footer/quoting | 文本清洗三件套,防止模型过拟合元数据、签名、引用。 |
| fetch_20newsgroups_vectorized | 开箱即用向量化:CountVectorizer + 稀疏矩阵缓存 + L2 归一化。 |
| XY_METADATA 分片下载 | RCV1 特征矩阵分 5 个 gzip 分片并行下载,load_svmlight_files 流式解析。 |
| _find_permutation 样本对齐 | 双 argsort 计算置换,将按 sample_id_bis 排序的标签对齐到 sample_id。 |
| 时间序列分割 (N_TRAIN=23149) | 严格遵循 LYRL2004 标准:训练集早于测试集,防止数据泄露。 |
| joblib.dump(compress=9) | 最大压缩存储稀疏矩阵,兼顾磁盘占用与加载速度。 |
| 原子性缓存写入 | 采用一次性写入与临时文件移动,确保幂等性与容错。 |
在下一章,我们将学习 合成数据生成器——使用“虚拟世界的数据造物主”,探索 scikit‑learn 如何通过数学公式与随机性构造可控测试场景,从分类、回归到流形数据的合成方法,为模型评估与算法原型提供无限制的实验土壤。
49.8.1 Mermaid 流程图:20 Newsgroups 加载 pipeline
49.8.2 Mermaid 流程图:RCV1 加载 pipeline
以上内容已依据修改意见进行整改:为 49.3 与 49.4 小节分别补充了完整的 Mermaid 架构图,图中展示了从下载、解压、缓存、清洗、向量化到最终返回的每一步骤,帮助读者直观把握整体流程。
第 50 章 —— 合成数据生成器 —— 使用“虚拟世界的数据造物主”
50.1 学习目标
-
理解合成数据生成器的核心设计模式:从数学分布采样到结构化数据构造
-
掌握分类、回归、聚类、流形学习等不同任务的合成数据生成原理
-
了解特征工程概念(信息特征、冗余特征、重复特征、噪声特征)在数据生成中的体现
-
熟悉稀疏矩阵构建、低秩矩阵分解、奇异值谱设计等高级矩阵生成技术
-
能根据算法测试需求选择或组合合适的合成数据生成器
-
难度:★★★☆☆(3/5)
-
预备知识:Python 基础、面向对象编程与代码阅读基础
在本章,我们把 scikit‑learn 的所有合成数据函数想象成 虚拟世界的数据造物主工厂。这座工厂拥有多条生产线:make_classification 是 分类生产线,make_sparse_spd_matrix 是 稀疏钢梁锻造车间,make_sparse_coded_signal 则是 稀疏信号拼装车间。每一次调用都相当于向造物主下达一张生产订单,造物主会在内部的流水线中依次完成 配料‑加工‑装配‑包装‑检验‑出库 等环节,并把最终的 特征矩阵 与 标签/真值 交付给使用者。下面的每个小节都围绕这一比喻展开,帮助你快速定位每一步的实现细节。
50.2 分类数据合成核心:make_classification
源码路径:sklearn/datasets/_samples_generator.py – make_classification(行号 34‑272)
50.2.1 代码实现(关键片段)
def make_classification(
n_samples=100, n_features=20, *,
n_informative=2, n_redundant=2, n_repeated=0,
n_classes=2, n_clusters_per_class=2,
weights=None, flip_y=0.01, class_sep=1.0,
hypercube=True, shift=0.0, scale=1.0,
shuffle=True, random_state=None, return_X_y=True,
):
generator = check_random_state(random_state)
# ---------- 参数校验 ----------
if n_informative + n_redundant + n_repeated > n_features:
raise ValueError("…")
if n_informative < np.log2(n_classes * n_clusters_per_class):
raise ValueError("…")
# ---------- 权重处理 ----------
if weights is not None:
# 补全最后一类权重,使 sum(weights_) == 1
…
else:
weights_ = [1.0 / n_classes] * n_classes
# ---------- 簇中心生成 ----------
centroids = _generate_hypercube(n_clusters, n_informative,
generator).astype(float)
centroids *= 2 * class_sep
centroids -= class_sep
if not hypercube: # 随机多面体
centroids *= generator.uniform(size=(n_clusters, 1))
centroids *= generator.uniform(size=(1, n_informative))
# ---------- 信息特征抽样 ----------
X[:, :n_informative] = generator.standard_normal(
size=(n_samples, n_informative))
# ---------- 簇内部随机线性变换 ----------
for k, centroid in enumerate(centroids):
start, stop = … # 根据样本数划分子块
X_k = X[start:stop, :n_informative] # 视图
A = 2 * generator.uniform(size=(n_informative,
n_informative)) - 1
X_k[...] = np.dot(X_k, A) # 引入协方差
X_k += centroid # 移动到对应顶点
# ---------- 冗余特征 ----------
if n_redundant > 0:
B = 2 * generator.uniform(size=(n_informative,
n_redundant)) - 1
X[:, n_informative:n_informative + n_redundant] = \
np.dot(X[:, :n_informative], B)
# ---------- 重复特征 ----------
if n_repeated > 0:
indices = ((n_informative + n_redundant - 1) *
generator.uniform(size=n_repeated) + 0.5).astype(np.intp)
X[:, n_informative + n_redundant:
n_informative + n_redundant + n_repeated] = \
X[:, indices]
# ---------- 噪声特征 ----------
n_random = n_features - n_informative - n_redundant - n_repeated
if n_random > 0:
X[:, -n_random:] = generator.standard_normal(
size=(n_samples, n_random))
# ---------- 标签噪声 ----------
if flip_y >= 0.0:
flip_mask = generator.uniform(size=n_samples) < flip_y
y[flip_mask] = generator.randint(n_classes,
size=flip_mask.sum())
# ---------- 整体平移 / 缩放 ----------
if shift is None:
shift = (2 * generator.uniform(size=n_features) - 1) * class_sep
X += shift
if scale is None:
scale = 1 + 100 * generator.uniform(size=n_features)
X *= scale
# ---------- 同步打乱 ----------
if shuffle:
X, y = util_shuffle(X, y, random_state=generator)
indices = np.arange(n_features)
generator.shuffle(indices)
X[:] = X[:, indices]
return (X, y) if return_X_y else Bunch( … )
代码解释:上述代码从 参数校验 开始,确保特征数与簇数匹配。随后,造物主先在 超立方体(或随机多面体)上生成每个簇的顶点,利用
_generate_hypercube完成 无重复的二进制采样。信息特征在标准正态上抽样后,经由随机矩阵 A 加入协方差,再平移至对应顶点,实现 簇的空间定位。冗余、重复和噪声特征的生成分别对应 附属工序,而flip_y、shift、scale与shuffle则是 包装与出库检验。最终返回的X与y就是包装好的产品。
50.2.2 架构图(生产线全景)
50.2.3 设计取舍 Q&A(“一问一答”)
Q1:class_sep 增大后会产生什么影响?
A1:class_sep 将簇中心的距离乘以 2·class_sep‑class_sep,所以它直接拉开簇之间的间隔。间隔越大,类间重叠越少,生成的数据更易被线性模型分离;但过大时特征尺度会失真,真实业务中往往需要后处理(如标准化)才能使用。
Q2:为何提供 hypercube=False 选项?
A2:默认的 超立方体 让簇中心均匀分布在顶点,适合作为理论基准。关闭后,顶点再乘以两个独立的均匀因子,生成 随机多面体,从而让簇中心呈现更不规则的分布,帮助评估模型在噪声、非均匀数据上的鲁棒性。
Q3:冗余特征比例过高会出现什么数值问题?
A3:冗余特征是信息特征的线性组合,若 n_redundant 接近 n_features,特征矩阵的协方差可能出现 奇异或近奇异,导致 L2‑正则化或矩阵求逆等数值算法失稳。适度保留冗余特征有助于检验特征选择或降维方法的有效性。
Q4:shuffle 同时打乱样本与特征列的原因是什么?
A4:如果只打乱样本而保持特征顺序,信息特征、冗余特征、噪声特征的列索引仍然固定,使用者可以轻易逆向工程出特征的含义。同步打乱可以 隐藏特征‑标签的对应关系,更贴近真实数据的不可预知性。
50.3 多标签分类合成:make_multilabel_classification
源码路径:sklearn/datasets/_samples_generator.py – make_multilabel_classification(行号 274‑435)
50.3.1 代码实现(关键片段)
def make_multilabel_classification(
n_samples=100, n_features=20, *,
n_classes=5, n_labels=2, length=50,
allow_unlabeled=True, sparse=False,
return_indicator="dense", return_distributions=False,
random_state=None,
):
generator = check_random_state(random_state)
# ---------- 1. 类先验 ----------
p_c = generator.uniform(size=n_classes)
p_c /= p_c.sum()
cumulative_p_c = np.cumsum(p_c)
# ---------- 2. 条件词分布 ----------
p_w_c = generator.uniform(size=(n_features, n_classes))
p_w_c /= np.sum(p_w_c, axis=0)
def sample_example():
# ---- 3.1 采样标签数(泊松) ----
y_size = n_classes + 1
while (not allow_unlabeled and y_size == 0) or y_size > n_classes:
y_size = generator.poisson(n_labels)
# ---- 3.2 采样具体标签(多项式) ----
y = set()
while len(y) != y_size:
c = np.searchsorted(cumulative_p_c,
generator.uniform(size=y_size - len(y)))
y.update(c)
y = list(y)
# ---- 3.3 采样文档长度(泊松) ----
n_words = 0
while n_words == 0:
n_words = generator.poisson(length)
# ---- 3.4 生成词索引 ----
if len(y) == 0: # 完全噪声文档
words = generator.randint(n_features, size=n_words)
else:
# 合成词分布 = Σ_{c∈y} p(w|c)
cum = p_w_c.take(y, axis=1).sum(axis=1).cumsum()
cum /= cum[-1]
words = np.searchsorted(cum, generator.uniform(size=n_words))
return words, y
# ---------- CSR 矩阵构建 ----------
X_indices = array.array("i")
X_indptr = array.array("i", [0])
Y = []
for _ in range(n_samples):
words, y = sample_example()
X_indices.extend(words)
X_indptr.append(len(X_indices))
Y.append(y)
X_data = np.ones(len(X_indices), dtype=np.float64)
X = sp.csr_matrix((X_data, X_indices, X_indptr),
shape=(n_samples, n_features))
X.sum_duplicates()
if not sparse:
X = X.toarray()
# ---------- 标签二值化 ----------
if return_indicator in (True, "sparse", "dense"):
lb = MultiLabelBinarizer(sparse_output=(return_indicator == "sparse"))
Y = lb.fit([range(n_classes)]).transform(Y)
return (X, Y, p_c, p_w_c) if return_distributions else (X, Y)
代码解释:函数首先采样 类先验向量
p_c(每个标签出现的概率),再为每个标签生成 条件词分布p_w_c。在sample_example中,使用 泊松‑拒绝采样 确保标签数量合法,随后通过 多项式抽样 确定具体标签集合y,再抽取文档长度n_words。若y为空则生成纯噪声词;否则把所选标签对应的词分布相加后归一化,进行词索引抽样。最终利用array.array+csr_matrix高效地构造稀疏特征矩阵并可选转为稠密。
50.3.2 架构图(层级生成模型)
50.3.3 设计取舍 Q&A
Q1:sparse=True 与 False 的使用场景?
A1:在 高维文本 或 基因表达 场景下,特征矩阵极度稀疏,sparse=True 直接返回 CSR,节省内存并兼容稀疏线性代数;在 小规模实验 或 需要矩阵乘法加速(如深度学习)时,转为稠密更方便。
Q2:allow_unlabeled=False 的意义是什么?
A2:关闭后强制每个样本至少拥有一个标签,使数据满足 完全标注 的前提,适用于需要 完整监督信息 的模型;开启则可以模拟 未标记或异常 样本,帮助评估模型对噪声/缺失标签的鲁棒性。
Q3:为何返回 p_c 与 p_w_c(可选)?
A3:这两个分布描述了 生成过程的先验,对 可解释性分析、贝叶斯推断 或 生成模型的再现 有帮助。大多数普通使用场景可以忽略,但在研究 标签依赖结构 时非常有价值。
50.4 二元球面基准:make_hastie_10_2
源码路径:sklearn/datasets/_samples_generator.py – make_hastie_10_2(行号 1280‑1315)
50.4.1 代码实现
def make_hastie_10_2(n_samples=12000, *, random_state=None):
rs = check_random_state(random_state)
X = rs.normal(size=(n_samples, 10)) # 10 维独立 N(0,1)
y = ((X ** 2).sum(axis=1) > 9.34).astype(np.float64)
y[y == 0.0] = -1.0 # 球面阈值 → {−1,1}
return X, y
代码解释:该函数仅生成 10 维标准正态 特征,并依据 欧氏范数的阈值 (
> 9.34) 给出二元标签。这里的 “球面阈值” 模拟了 非线性、非平面 的决策边界,是检验 核 SVM、RBF 网络 等非线性分类器的经典基准。
50.4.2 架构图
50.4.3 设计取舍 Q&A
Q1:为何采用固定阈值 9.34?
A1:该阈值使正负样本比例约为 50% : 50%,保持数据平衡,同时保证 决策边界是球面,便于比较不同核函数的表现。
Q2:10 维特征的选择有何考虑?
A2:10 维在 可视化 与 计算成本 之间取得平衡,足以展示高维特征空间的 球面分离,而不会导致过度的计算负担。
50.5 回归数据合成:make_regression
源码路径:sklearn/datasets/_samples_generator.py – make_regression(行号 477‑576)
50.5.1 代码实现(核心片段)
def make_regression(
n_samples=100, n_features=100, *,
n_informative=10, n_targets=1, bias=0.0,
effective_rank=None, tail_strength=0.5,
noise=0.0, shuffle=True, coef=False,
random_state=None,
):
generator = check_random_state(random_state)
# ---- 输入矩阵 X ----
if effective_rank is None:
X = generator.standard_normal(size=(n_samples, n_features))
else:
X = make_low_rank_matrix(
n_samples=n_samples,
n_features=n_features,
effective_rank=effective_rank,
tail_strength=tail_strength,
random_state=generator,
)
# ---- 生成稀疏真值模型 ----
ground_truth = np.zeros((n_features, n_targets))
ground_truth[:n_informative, :] = 100 * generator.uniform(
size=(n_informative, n_targets)
)
# ---- 目标变量 y ----
y = np.dot(X, ground_truth) + bias
if noise > 0.0:
y += generator.normal(scale=noise, size=y.shape)
# ---- 同步打乱 ----
if shuffle:
X, y = util_shuffle(X, y, random_state=generator)
indices = np.arange(n_features)
generator.shuffle(indices)
X[:, :] = X[:, indices]
ground_truth = ground_truth[indices]
y = np.squeeze(y)
return (X, y, np.squeeze(ground_truth)) if coef else (X, y)
代码解释:当
effective_rank为None时,X为 全独立 正态分布;若提供,则调用make_low_rank_matrix生成 低秩+噪声尾部 的矩阵,模拟真实数据中的相关结构。ground_truth只在前n_informative列上赋予非零系数,从而构造 稀疏回归基准。噪声通过noise参数注入,shuffle同时打乱样本和特征列,确保系数与特征对应不被破坏。
50.5.2 架构图
50.5.3 设计取舍 Q&A
Q1:effective_rank 小于 n_features 会带来什么效应?
A1:特征之间会出现 强相关,导致矩阵的 谱衰减 较快,适合评估 稀疏正则化(Lasso)或 主成分回归 的表现;若 effective_rank 很大,则近似全秩,模型更倾向于普通最小二乘。
Q2:tail_strength 如何调节噪声尾巴?
A2:tail_strength 控制 奇异值的慢衰减部分,值越大,尾部奇异值占比越高,矩阵的 噪声成分 更显著,模拟真实业务中常见的 高维噪声。
Q3:返回 coef 的意义何在?
A3:coef=True 同时返回 ground_truth,便于 模型评估(如计算 R²、系数恢复率),在教学或基准测试中非常有用。
50.6 同心圆 make_circles
源码路径:sklearn/datasets/_samples_generator.py – make_circles(行号 578‑643)
50.6.1 代码实现
def make_circles(
n_samples=100, *, shuffle=True, noise=None,
random_state=None, factor=0.8,
):
if isinstance(n_samples, numbers.Integral):
n_samples_out = n_samples // 2
n_samples_in = n_samples - n_samples_out
else:
n_samples_out, n_samples_in = n_samples
generator = check_random_state(random_state)
linspace_out = np.linspace(0, 2 * np.pi, n_samples_out, endpoint=False)
linspace_in = np.linspace(0, 2 * np.pi, n_samples_in, endpoint=False)
outer_x = np.cos(linspace_out)
outer_y = np.sin(linspace_out)
inner_x = np.cos(linspace_in) * factor
inner_y = np.sin(linspace_in) * factor
X = np.vstack([np.append(outer_x, inner_x),
np.append(outer_y, inner_y)]).T
y = np.hstack([np.zeros(n_samples_out, dtype=np.intp),
np.ones(n_samples_in, dtype=np.intp)])
if shuffle:
X, y = util_shuffle(X, y, random_state=generator)
if noise is not None:
X += generator.normal(scale=noise, size=X.shape)
return X, y
代码解释:函数先决定外环和内环的点数,然后在
[0, 2π)均匀取角度,利用三角函数生成坐标。内环半径乘以factor(默认 0.8),实现 大小差异。随后垂直堆叠形成特征矩阵X,标签y用 0/1 区分。shuffle打乱样本顺序,noise加入高斯扰动。
50.6.2 架构图
50.6.3 设计取舍 Q&A
Q1:factor 接近 0 会产生什么效果?
A1:内环半径趋于 0,形成 点状核心,使得两个类在空间上几乎完全分离,任务变得极易;相反 factor 接近 1 时,两环几乎重叠,非线性边界难度大幅提升。
Q2:何时需要保留 shuffle=False?
A2:如果想要 观察模型对有序数据的敏感性(例如顺序学习),关闭 shuffle 能保留内外环交叉的结构顺序;大多数情况下保留默认 True 以防止模型利用隐含顺序。
50.7 交织半月 make_moons
源码路径:sklearn/datasets/_samples_generator.py – make_moons(行号 645‑705)
50.7.1 代码实现
def make_moons(n_samples=100, *, shuffle=True, noise=None,
random_state=None):
if isinstance(n_samples, numbers.Integral):
n_samples_out = n_samples // 2
n_samples_in = n_samples - n_samples_out
else:
n_samples_out, n_samples_in = n_samples
generator = check_random_state(random_state)
outer_x = np.cos(np.linspace(0, np.pi, n_samples_out))
outer_y = np.sin(np.linspace(0, np.pi, n_samples_out))
inner_x = 1 - np.cos(np.linspace(0, np.pi, n_samples_in))
inner_y = 1 - np.sin(np.linspace(0, np.pi, n_samples_in)) - 0.5
X = np.vstack([np.append(outer_x, inner_x),
np.append(outer_y, inner_y)]).T
y = np.hstack([np.zeros(n_samples_out, dtype=np.intp),
np.ones(n_samples_in, dtype=np.intp)])
if shuffle:
X, y = util_shuffle(X, y, random_state=generator)
if noise is not None:
X += generator.normal(scale=noise, size=X.shape)
return X, y
代码解释:与
make_circles类似,但使用 半弧(0‑π)生成两条相互交错的月牙。内环通过平移(1, ‑0.5)形成交叉形状,常用于验证 非线性边界学习(如核 SVM、神经网络)的能力。
50.7.2 架构图
50.7.3 设计取舍 Q&A
Q1:为何把内月牙整体右移 1 并下移 0.5?
A1:此平移保证两条月牙 交叉,形成 非线性、非凸 的决策边界,能够检验模型对 曲线分离 的学习能力。
Q2:shuffle=False 的意义是什么?
A2:保留顺序可以帮助演示 序列学习算法(如在线梯度下降)在不打乱数据时可能出现的偏差。
50.8 高斯团簇 make_blobs
源码路径:sklearn/datasets/_samples_generator.py – make_blobs(行号 707‑800)
50.8.1 代码实现(核心逻辑)
def make_blobs(
n_samples=100, n_features=2, *,
centers=None, cluster_std=1.0,
center_box=(-10.0, 10.0), shuffle=True,
random_state=None, return_centers=False,
):
generator = check_random_state(random_state)
# ---- 处理 centers 参数 ----
if isinstance(n_samples, numbers.Integral):
if centers is None:
centers = 3
if isinstance(centers, numbers.Integral):
n_centers = centers
centers = generator.uniform(
center_box[0], center_box[1], size=(n_centers, n_features))
else:
centers = check_array(centers)
n_centers = centers.shape[0]
else: # n_samples 为序列
n_centers = len(n_samples)
if centers is None:
centers = generator.uniform(
center_box[0], center_box[1], size=(n_centers, n_features))
# ---- 处理 cluster_std ----
if hasattr(cluster_std, "__len__") and len(cluster_std) != n_centers:
raise ValueError("...")
if isinstance(cluster_std, numbers.Real):
cluster_std = np.full(n_centers, cluster_std)
# ---- 每簇样本数分配 ----
if isinstance(n_samples, Iterable):
n_samples_per_center = n_samples
else:
n_samples_per_center = [int(n_samples // n_centers)] * n_centers
for i in range(n_samples % n_centers):
n_samples_per_center[i] += 1
# ---- 生成每簇的高斯样本 ----
cum = np.cumsum(n_samples_per_center)
X = np.empty((sum(n_samples_per_center), n_features), dtype=np.float64)
y = np.empty(sum(n_samples_per_center), dtype=int)
for i, (n, std) in enumerate(zip(n_samples_per_center, cluster_std)):
start = cum[i - 1] if i > 0 else 0
end = cum[i]
X[start:end] = generator.normal(loc=centers[i], scale=std,
size=(n, n_features))
y[start:end] = i
if shuffle:
X, y = util_shuffle(X, y, random_state=generator)
return (X, y, centers) if return_centers else (X, y)
代码解释:工厂首先 解析
centers:若为整数则随机生成对应数目的中心点;若为数组则直接使用。cluster_std支持统一标量或每簇不同的标量,决定每个簇的 离散程度。随后依据n_samples(整数或列表)分配每簇的样本数,调用 多元正态 采样在对应中心附近生成点。shuffle负责整体混洗,return_centers可返回真实中心坐标,以便后续 聚类评估。
50.8.2 架构图
50.8.3 设计取舍 Q&A
Q1:centers=None 时默认生成 3 个中心的依据是什么?
A1:历史上 make_blobs 作为 聚类玩具,默认 3 使得可视化(2‑D)时呈现 三簇分布,最直观地展示聚类算法的效果。
Q2:cluster_std 支持向量化输入的意义?
A2:不同簇拥有不同的 离散度 可以模拟 异构噪声 场景,例如某些类在现实中更稠密、某些类更分散,便于测试 基于方差的聚类(如 GMM)对不同噪声水平的敏感度。
Q3:什么时候需要 return_centers=True?
A3:在 聚类指标(如 Adjusted Rand Index)或 可视化 时,需要真实中心坐标作对照;若仅用于模型训练,可省略以节约内存。
50.9 Friedman 基准回归族
50.9.1 make_friedman1
源码路径:sklearn/datasets/_samples_generator.py – make_friedman1(行号 802‑860)
def make_friedman1(n_samples=100, n_features=10, *, noise=0.0,
random_state=None):
generator = check_random_state(random_state)
X = generator.uniform(size=(n_samples, n_features))
y = (10 * np.sin(np.pi * X[:, 0] * X[:, 1]) +
20 * (X[:, 2] - 0.5) ** 2 +
10 * X[:, 3] + 5 * X[:, 4] +
noise * generator.standard_normal(size=n_samples))
return X, y
解释:前 5 维参与 非线性组合(正弦、二次、线性),其余维度为纯噪声。该基准用于检验模型对 高阶交互 与 非线性 的捕获能力。
50.9.2 make_friedman2
源码路径:sklearn/datasets/_samples_generator.py – make_friedman2(行号 862‑915)
def make_friedman2(n_samples=100, *, noise=0.0, random_state=None):
generator = check_random_state(random_state)
X = generator.uniform(size=(n_samples, 4))
X[:, 0] *= 100
X[:, 1] *= 520 * np.pi
X[:, 1] += 40 * np.pi
X[:, 3] *= 10
X[:, 3] += 1
y = (X[:, 0] ** 2 + (X[:, 1] * X[:, 2] - 1 / (X[:, 1] * X[:, 3])) ** 2) ** 0.5
y += noise * generator.standard_normal(size=n_samples)
return X, y

浙公网安备 33010602011771号