Sklearn-源码解析-书-v1-0-二十一-

Sklearn 源码解析(书)v1.0(二十一)

实现细节

  1. 首行携带元信息(样本数、特征数、标签名称),通过手动解析获取这些元数据。

  2. 随后逐行读取数值部分,填充到预先分配好的 np.empty 矩阵中,保持了对大文件的内存友好。

  3. 若提供 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 中加入了 filenamedata_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_namestarget_names,随后分别加载数值矩阵。

46.8.7 统一返回

所有玩具数据集最终返回 Bunch,包含 datatargetfeature_namestarget_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 流程图

graph TD A[传入 container_path] --> B[列出子目录 → 类别列列表] B --> C{categories 参数} C -->|有| D[过滤仅保留指定类别] C -->|无| D D --> E{allowed_extensions} E -->|有| F[仅保留匹配扩展名的文件] E -->|无| G[全部文件] F --> H[构建 target、filenames 列表] G --> H H --> I{shuffle} I -->|True| J[随机打乱顺序] I -->|False| J J --> K{load_content} K -->|True| L[读取文件内容并可选解码] K -->|False| L[仅返回路径] L --> M[返回 Bunch]

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 动手练习

  1. 追踪远程文件下载的完整生命周期

    • 阅读 src/sklearn/datasets/_base.py_fetch_remote()(约第 680‑760 行)

    • 关键点:

      1. 临时文件命名模式 prefix=remote.filename + '.part_' 的并发安全作用。

      2. shutil.move(temp_file_path, file_path) 实现原子写入的原理。

      3. SHA256 校验在下载前(已有文件)和下载后(新文件)的两次检查。

      4. n_retriesdelay 参数如何实现指数退避重试。

    • 思考

      • 为什么使用 NamedTemporaryFile(delete=False) 而不是直接写目标文件?

      • 若下载中断留下 .part_ 临时文件,下次调用会如何处理?

      • URLErrorTimeoutError 之外的异常(如 KeyboardInterrupt)如何被捕获并清理?

  2. 对比三种本地数据装载器的设计差异

    • 阅读 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/ 子包)?

  3. 剖析 Digits 数据集的图像重塑与特征名生成

    • 阅读 load_digits() 实现。

    • 思考

      • flat_data.view() 后直接修改 images.shape 为什么不拷贝内存?这依赖什么条件?

      • 如果去掉 copy=Falsetarget 的内存行为会有什么不同?

      • 特征名列表推导式中 row_idxcol_idx 的嵌套顺序决定了什么?如何验证它与 images 的空间布局一致?

  4. 探究玩具数据集加载器的共性与差异

    • 对比阅读 load_wineload_irisload_breast_cancerload_diabetesload_linnerud

    • 思考

      • load_linnerud 为何不复用 load_csv_data 而是手动读取两个文件?

      • load_diabetesscaled=True 时的 data /= data.shape[0] ** 0.5 有何统计学含义?

      • load_irisload_breast_cancer 返回的 Bunch 多出 filenamedata_module 字段,用途是什么?

  5. 实战 load_files 目录树加载器

    • 假设目录结构如下:
container/
  ├── class_a/
  │   ├── doc1.txt
  │   └── doc2.txt
  └── class_b/
      └── doc3.txt
  • 思考

    • categories=['class_a'] 参数如何过滤目标类别?

    • allowed_extensions=['.txt']encoding='utf-8' 如何协作完成文件筛选与解码?

    • shuffle=Truerandom_state=42 时,返回的 filenamestarget 的对应关系如何保证?

    • 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_framereturn_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 源码地图:函数调用与数据转换流程(对应「源码逐行解析」小节)

flowchart TD A[开始:fetch_california_housing] --> B{缓存文件存在?} B -- 是 --> C[joblib.load 读取 .pkz] B -- 否 --> D{download_if_missing 是 True?} D -- 否 --> E[抛出 OSError:数据未找到] D -- 是 --> F[_fetch_remote 下载 .tgz] F --> G[tarfile 解压压缩包] G --> H[np.loadtxt 读取原始 CSV] H --> I[按 columns_index 重排列顺序] I --> J[joblib.dump 写入 .pkz 缓存(compress=6)] J --> K[删除临时 .tgz 文件] C --> L[拆分目标与特征:target=data[:,0], data=data[:,1:]] K --> L L --> M[特征工程:AveRooms、AveBedrms、AveOccup 除以 households] M --> N[目标除以 10 万:转为 10 万美元单位] N --> O{as_frame 是 True?} O -- 是 --> P[_convert_data_dataframe 生成 pandas 对象] O -- 否 --> Q[保持 ndarray] P --> R[返回 Bunch 包含 frame] Q --> R R --> S[结束:返回数据]

47.5.2 特征工程逻辑图:从原始总和到每户平均值(对应「特征工程」段落)

flowchart LR A[原始数据列] --> B[total_rooms (列索引 2)] A --> C[total_bedrooms (列索引 3)] A --> D[population (列索引 4)] A --> E[households (列索引 5)] B --> F[AveRooms = total_rooms / households] C --> G[AveBedrms = total_bedrooms / households] D --> H[AvgOccup = population / households] E --> F E --> G E --> H F --> I[最终特征:AveRooms] G --> J[最终特征:AveBedrms] H --> K[最终特征:AveOccup]

47.5.3 目标尺度变换图:美元 → 10 万美元(对应「目标尺度缩放」段落)

flowchart A[目标原始值:美元] --> B[除以 100000] B --> C[目标最终值:10 万美元] style A fill:#f9f,stroke:#333 style C fill:#9f9,stroke:#333

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 的缓存相关代码:

  1. California Housing 使用单文件 cal_housing.pkz 存储合并数组,而 Covertype 分离 samplestargets 两个文件,这种设计差异的工程考量是什么?

  2. Covertype 使用 TemporaryDirectory(dir=covtype_dir) + os.rename 实现原子写入,California Housing 直接 joblib.dump 到目标路径,前者解决了什么并发风险?

  3. 两者 compress 参数分别为 6 和 9,压缩级别选择的权衡因素是什么?

47.9.2 分析KDDCUP99的子集构建逻辑与工程妥协

阅读 _kddcup99.pyfetch_kddcup99subset 参数处理逻辑:

  1. subset='SA' 时,为何保留所有正常样本而仅随机抽取 3377 条异常样本?这种极度不平衡构造的基准意义是什么?

  2. subset='SF'/'http'/'smtp' 共享 logged_in==1 筛选与 log(x+0.1) 变换,但列索引硬编码(如 data[:,11]),这种硬编码在工程上为何可接受?

  3. _fetch_brute_kddcup99 中定义的结构化 dtype dt 如何支撑混合类型(int/float/bytes)的精确解析?

47.9.3 探究LFW双任务模式的数据组织差异与缓存复用

对比 _lfw.py_fetch_lfw_people_fetch_lfw_pairs 的实现:

  1. 两者均使用 joblib.Memory(location=lfw_home, compress=6).cache 装饰加载器,这种缓存粒度(函数级)与California Housing的文件级缓存有何异同?

  2. fetch_lfw_people 返回 images (n, H, W) 与 data (n, HW),而 fetch_lfw_pairs 返回 pairs (n, 2, H, W) 与 data (n, 2H*W),这种形状设计如何服务于分类与验证两种任务的典型模型输入?

  3. slice_resize 参数如何正交控制图像预处理管线?默认 slice_(70:195, 78:172) + resize=0.5 的几何含义是什么?

47.9.4 解析物种分布数据集的地理栅格重建机制

阅读 _species_distributions.py 的核心解析流程:

  1. SAMPLES.zipCOVERAGES.zip 均通过 np.load 作为 npz 读取,再用 BytesIO 包装内部文件,这种“嵌套归档”解析模式的优势是什么?

  2. _load_coverageNODATA_value 替换为 -9999 的处理,如何配合 dtype=np.int16 节省内存?若原始 NODATA_value 已为 -9999 会怎样?

  3. 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_ignoreis_row_identifier),以及分类标签的逆映射与 category dtype 统一。

  • 离线测试夹具:了解 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‑ARFFPandas) | 两条装配线:手工工艺(通用,能处理稀疏部件) 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'):当元数据中的 formatsparse_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 架构视图——解析器调度图

flowchart TD A[fetch_openml] --> B[_download_data_to_bunch] B --> C[_load_arff_response] C --> D{parser} D -->|liac-arff| E[_liac_arff_parser] D -->|pandas| F[_pandas_arff_parser] E --> G{output_arrays_type} F --> G G -->|pandas| H[返回 X, y, frame, None] G -->|numpy / sparse| I[返回 X, y, None, categories] H --> J[包装成 Bunch] I --> J

说明fetch_openml 根据元信息决定 parserauto 时会检查 format),随后 _load_arff_response 调用统一入口 load_arff_from_gzip_file,最终的返回结构取决于 output_arrays_typepandasnumpysparse)。


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 流程图——缓存与重试

flowchart TD Start[调用 fetch_openml] --> CheckCache{本地缓存是否存在?} CheckCache -- Yes --> LoadCache[读取缓存文件] CheckCache -- No --> Download[_open_openml_url 下载] Download --> MD5Check{MD5 校验} MD5Check -- ✅ --> Parse[_load_arff_response 解析] MD5Check -- ❌ --> CleanCache[_retry_with_clean_cache 清理并重下] CleanCache --> Download Parse --> Return[Bunch 返回]

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_ignoreis_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...")

代码作用概述:确保模型训练时目标列的 类型统一,防止出现混合 floatobject 导致的 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_namestarget_namesDESCRdetailsurl)封装进 Bunch

return Bunch(
    data=X,
    target=y,
    frame=frame,
    categories=categories,
    feature_names=data_columns,
    target_names=target_columns,
)

代码作用概述:它是从 底层网络流用户可直接使用的 Bunch 的桥梁,负责 校验返回类型选择异常统一处理最终包装

48.6.5 流程图——完整管线

flowchart TD A[fetch_openml] --> B[_download_data_to_bunch] B --> C[_load_arff_response] C --> D{parser} D -->|liac-arff| E[_liac_arff_parser] D -->|pandas| F[_pandas_arff_parser] E --> G[返回 X, y, frame, categories] F --> G G --> H[包装成 Bunch] H --> I[返回给用户]

48.7 OpenML 测试数据夹具 —— 离线测试的“标本库”

  • 位置sklearn/datasets/tests/data/openml/id_<id>/,每个子目录包含 预下载的 .arff.gz对应的 JSON 元信息

  • 离线模式_monkey_patch_webbased_functionsurlopen 替换为读取本地夹具的函数,所有网络请求都被拦截。

  • 覆盖场景

    • 稠密数据 (id_1, id_2, id_3)

    • 稀疏 ARFF (id_61, id_62) → 验证 LIAC‑ARFF 分支

    • 多目标 (id_292, id_561)

    • 边界情况:缺失值、字符串特征、不同 quotechar/escapechar (id_1119, id_1590)

    • 新版兼容性 (id_40589id_40675id_40945id_40966id_42074id_42585)

48.7.1 夹具在测试中的角色

flowchart LR TestSuite[测试套件] --> MockPatch[_monkey_patch_webbased_functions] MockPatch --> LoadFixture[加载本地 ARFF / JSON 夹具] LoadFixture --> fetch_openml[调用 fetch_openml(cache=False)] fetch_openml --> Bunch[返回 Bunch 对象] Bunch --> Assertions[断言检查(shape、dtype、categories 等)]

说明:通过这些本地资源,单元测试在 不依赖网络 的情况下,实现 确定性高速,保证所有边界场景均得到完整覆盖。


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,会发生什么?

Afetch_openml 在解析前会检测 format=="sparse_arff" 并抛出 ValueError,提示用户改用 liac‑arff。这样可以提前拦截不兼容的配置,避免隐藏的运行时错误。


48.9 动手练习

  1. 追踪双解析器的分流与汇合逻辑

    • 查看 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'")

  2. 剖析缓存原子写入与 MD5 校验的协同机制

    • _open_openml_url 使用 TemporaryDirectory + shutil.move 实现 原子写入

    • MD5 在 _load_arff_response下载完成后 校验,若不匹配触发 _retry_with_clean_cache 删除缓存并重下。

  3. 扩展新输出格式(Polars)

    • fetch_openmlStrOptions 中加入 '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 已保留 Categorical dtype。


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_ignoreis_row_identifier 标记的列。 |

| Bunch | 包含 datatargetframecategoriesDESCRdetailsurl 的字典式容器。 |

| 离线测试夹具 | 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_headerstrip_newsgroup_footerstrip_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_filepathCACHE_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)…

逐行注释

  1. 缓存键filebase 加上 remove 参数,确保不同清洗配置对应不同缓存文件,避免混用。

  2. 复用:先调用 fetch_20newsgroups 获得已经清洗好的文本,保持下载与清洗只进行一次。

  3. 缓存检查:若向量化结果已存在,直接 joblib.load,省去重新拟合的成本。

  4. 现场拟合CountVectorizer(dtype=np.int16) 采用 16 位整数压缩计数矩阵,随后保存。

  5. 归一化preprocessing.normalize 原地修改矩阵,避免额外拷贝。

  6. 子集返回:依据 subset 选择训练、测试或全部数据,sp.vstack 用于稀疏矩阵高效拼接。

  7. DataFrame:如果用户需要 pandas 接口,调用内部 _convert_data_dataframe 完成包装。

49.4.5 架构图:20 Newsgroups 加载 Pipeline

flowchart TD A[检查本地缓存] -->|存在| B[读取 .pkz 并解压] A -->|不存在| C[下载 tar.gz] C --> D[解压到临时目录] D --> E[load_files 读取 train / test] E --> F[压缩并写入 .pkz(原子写入)] F --> G[删除临时目录] B --> H[可选文本清洗 (headers / footers / quotes)] H --> I[类别筛选 & 连续标签映射] I --> J[shuffle(可选)] J --> K[返回 Bunch 或 (data, target)]

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,得到对应的排序索引 tu_inverse_permutation(u)u 逆转,使其能够映射回原始位置,最后 t[u_] 即为把 ysample_id_bis 重新排列,以匹配特征矩阵的顺序。

49.5.4 架构图:RCV1 加载 Pipeline

flowchart TD subgraph 下载与特征构建 A1[检查特征缓存] -->|缺失| B1[下载 5 个 .gz 分片] B1 --> C1[流式 load_svmlight_files] C1 --> D1[堆叠特征矩阵 X & 合并 sample_id] D1 --> E1[压缩并保存 X、sample_id] A1 -->|存在| F1[直接加载 X、sample_id] end subgraph 下载与标签构建 A2[检查标签缓存] -->|缺失| B2[下载 topics.qrels.gz] B2 --> C2[解析构建稠密 y 与 sample_id_bis] C2 --> D2[_find_permutation 对齐 y] D2 --> E2[稀疏化 & 压缩保存 y、categories] A2 -->|存在| F2[直接加载 y、categories] end F1 & F2 --> G[根据 subset 切分 (train / test / all)] G --> H[可选 shuffle] H --> I[返回 Bunch 或 (X, y, sample_id)]

49.6 设计中的取舍

在 20 Newsgroups 的文本清洗实现上,作者选择了 纯标准库字符串操作partitionsplit、正则)而非像 BeautifulSoup 那样的通用 HTML 解析器。此举带来的 trade‑off 如下:

  • 优势:零外部依赖、执行速度极快、行为透明、便于审计,尤其在缓存后只执行一次。

  • 劣势:对极端、畸形邮件格式的鲁棒性有限;若数据源结构发生变化,需手动调整正则或分割逻辑。

结论:对于已知且相对固定的文本数据集(如 20 Newsgroups),“恰好够用”的轻量实现更符合工程效率与可维护性的目标;若面向更通用的邮件或 HTML 内容,则应考虑更健壮的解析库。

49.7 动手练习

49.7.1 练习 1:fetch_20newsgroups 细节

  1. 缓存键为何要包含 remove 参数?

    因为不同的文本清洗组合会产生不同的实际内容。如果缓存键不区分 remove,则在一次清洗后缓存的结果会在后续请求中被错误复用,导致模型看到与预期不一致的文本。将 remove 作为键的一部分保证每种清洗配置都有独立缓存,避免交叉污染。

  2. categories 筛选后为何使用 np.searchsorted 重新编码 target

    原始 target 是全局的类别索引(0‑19),当只保留子集 categories 时,这些索引会出现间隙。np.searchsorted(labels, data.target) 将剩余标签重新映射为连续的整数序列(0, 1, …),确保后续模型训练时标签空间紧凑且不出现空洞。

49.7.2 练习 2:fetch_rcv1 细节

  1. 分片下载与 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

  2. 标签稠密‑稀疏构建过程

    • 初始化全零稠密矩阵 ynp.uint8)以及 sample_id_bis 用于记录文档 ID 的顺序。

    • 逐行解析 .qrels.gz,为每个文档‑主题对在稠密矩阵中设置 1

    • 完成后通过 sp.csr_matrix 将稠密矩�压缩为 CSR 稀疏格式,再使用 joblib.dump(..., compress=9) 缓存。

  3. sample_id_bissample_id 顺序不一致的现实场景

    • 特征文件下载顺序(即时间戳)排列,保证训练样本在前。

    • 标签文件 则是按 文档 ID(内部编号)排序,可能与特征的时间顺序不匹配。此不一致来源于原始数据集的历史生成方式,需要通过置换对齐。

  4. subset='train'shuffle=True 同时指定时的先后顺序

    • 首先依据 subsetXysample_id 进行 时间序列切分(训练集为前 23 149 条)。

    • shuffle=True,随后在切分后的子集上执行随机打乱(shuffle_),这一步会破坏时间顺序,但仅影响已划定的训练或测试集合。

49.7.3 练习 3:向量化与缓存策略比较

  1. 两种策略的优缺点

    • 20 Newsgroups(现场拟合 + 缓存)

      • 优点:灵活,可随 removestop_wordsngram_range 等参数自由组合;首次运行后缓存即可重复使用。

      • 缺点:首次向量化需要遍历全部文本,计算量随文本规模线性增长;不适合极大数据集(数十 GB)会导致内存压力。

    • RCV1(预计算稀疏向量)

      • 优点:特征矩阵已在服务器端完成向量化,下载后直接使用,极大降低本地计算成本,适合 百万级样本

      • 缺点:缺乏自定义空间;如果需要不同的特征抽取方式(如 TF‑IDF、子词),只能重新下载或自行重算。

  2. 若在 fetch_20newsgroups_vectorized 中支持自定义 TfidfVectorizer,缓存键的设计建议

    将关键的向量化参数(如 use_idfnormsublinear_tfngram_rangemax_dfmin_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

flowchart TD A[检查本地缓存] -->|存在| B[读取 .pkz] A -->|不存在| C[下载 tarball] C --> D[解压到临时目录] D --> E[load_files 读取 train/test] E --> F[压缩并写入 .pkz] F --> G[删除临时目录] B --> H[可选文本清洗 (headers/footers/quotes)] H --> I[类别筛选 & 重编码] I --> J[shuffle(可选)] J --> K[返回 Bunch 或 (data, target)]

49.8.2 Mermaid 流程图:RCV1 加载 pipeline

flowchart TD A1[检查特征缓存] -->|缺失| B1[下载 5 个 .gz 分片] B1 --> C1[流式 load_svmlight_files] C1 --> D1[堆叠特征矩阵 X & 合并 sample_id] D1 --> E1[压缩并保存 X、sample_id] A1 -->|存在| F1[直接加载 X、sample_id] A2[检查标签缓存] -->|缺失| B2[下载 topics.qrels.gz] B2 --> C2[解析构建稠密 y 与 sample_id_bis] C2 --> D2[_find_permutation 对齐 y] D2 --> E2[稀疏化 & 压缩保存 y、categories] A2 -->|存在| F2[直接加载 y、categories] F1 & F2 --> G[根据 subset 切分 (train/test/all)] G --> H[可选 shuffle] H --> I[返回 Bunch 或 (X, y, sample_id)]

以上内容已依据修改意见进行整改:为 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.pymake_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_yshiftscaleshuffle 则是 包装与出库检验。最终返回的 Xy 就是包装好的产品。

50.2.2 架构图(生产线全景)

flowchart TD A[参数校验] --> B[权重补全] B --> C[生成超立方体顶点] C --> D{hypercube ?} D -->|True| E[直接使用顶点] D -->|False| F[随机多面体变形] E & F --> G[抽样信息特征 N(0,1)] G --> H[簇内部矩阵 A 施加协方差] H --> I[平移至簇中心] I --> J[生成冗余特征 B] J --> K[复制重复特征] K --> L[填充噪声特征] L --> M[标签噪声 flip_y] M --> N[整体 shift/scale] N --> O[shuffle 样本 & 特征] O --> P[返回 (X,y) / Bunch]

50.2.3 设计取舍 Q&A(“一问一答”)

Q1:class_sep 增大后会产生什么影响?

A1class_sep 将簇中心的距离乘以 2·class_sepclass_sep,所以它直接拉开簇之间的间隔。间隔越大,类间重叠越少,生成的数据更易被线性模型分离;但过大时特征尺度会失真,真实业务中往往需要后处理(如标准化)才能使用。

Q2:为何提供 hypercube=False 选项?

A2:默认的 超立方体 让簇中心均匀分布在顶点,适合作为理论基准。关闭后,顶点再乘以两个独立的均匀因子,生成 随机多面体,从而让簇中心呈现更不规则的分布,帮助评估模型在噪声、非均匀数据上的鲁棒性。

Q3:冗余特征比例过高会出现什么数值问题?

A3:冗余特征是信息特征的线性组合,若 n_redundant 接近 n_features,特征矩阵的协方差可能出现 奇异或近奇异,导致 L2‑正则化或矩阵求逆等数值算法失稳。适度保留冗余特征有助于检验特征选择或降维方法的有效性。

Q4:shuffle 同时打乱样本与特征列的原因是什么?

A4:如果只打乱样本而保持特征顺序,信息特征、冗余特征、噪声特征的列索引仍然固定,使用者可以轻易逆向工程出特征的含义。同步打乱可以 隐藏特征‑标签的对应关系,更贴近真实数据的不可预知性。


50.3 多标签分类合成:make_multilabel_classification

源码路径sklearn/datasets/_samples_generator.pymake_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 架构图(层级生成模型)

flowchart TD A[采样类先验 p_c] --> B[采样条件词分布 p_w_c] B --> C[循环 sample_example() for each sample] C --> D[泊松抽样标签数 n_labels] D --> E[多项式抽样标签集合 y] C --> F[泊松抽样文档长度 length] F --> G[若 y 为空 → 生成噪声词] G --> H[否则 → 合成词分布 Σ p(w|c)] H --> I[抽样词索引,填充到 CSR] I --> J[构建稀疏矩阵 X] J --> K[MultiLabelBinarizer 二值化 Y] K --> L{返回} L -->|X,Y| M[返回稠密/稀疏 X 与 Y] L -->|+分布| N[返回 (X,Y,p_c,p_w_c)]

50.3.3 设计取舍 Q&A

Q1:sparse=TrueFalse 的使用场景?

A1:在 高维文本基因表达 场景下,特征矩阵极度稀疏,sparse=True 直接返回 CSR,节省内存并兼容稀疏线性代数;在 小规模实验需要矩阵乘法加速(如深度学习)时,转为稠密更方便。

Q2:allow_unlabeled=False 的意义是什么?

A2:关闭后强制每个样本至少拥有一个标签,使数据满足 完全标注 的前提,适用于需要 完整监督信息 的模型;开启则可以模拟 未标记或异常 样本,帮助评估模型对噪声/缺失标签的鲁棒性。

Q3:为何返回 p_cp_w_c(可选)?

A3:这两个分布描述了 生成过程的先验,对 可解释性分析贝叶斯推断生成模型的再现 有帮助。大多数普通使用场景可以忽略,但在研究 标签依赖结构 时非常有价值。


50.4 二元球面基准:make_hastie_10_2

源码路径sklearn/datasets/_samples_generator.pymake_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) 给出二元标签。这里的 “球面阈值” 模拟了 非线性、非平面 的决策边界,是检验 核 SVMRBF 网络 等非线性分类器的经典基准。

50.4.2 架构图

flowchart TD A[采样 10 维 N(0,1)] --> B[计算每行的平方和] B --> C[阈值 9.34 判别 → y = 1 / -1] C --> D[返回 X, y]

50.4.3 设计取舍 Q&A

Q1:为何采用固定阈值 9.34?

A1:该阈值使正负样本比例约为 50% : 50%,保持数据平衡,同时保证 决策边界是球面,便于比较不同核函数的表现。

Q2:10 维特征的选择有何考虑?

A2:10 维在 可视化计算成本 之间取得平衡,足以展示高维特征空间的 球面分离,而不会导致过度的计算负担。


50.5 回归数据合成:make_regression

源码路径sklearn/datasets/_samples_generator.pymake_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_rankNone 时,X全独立 正态分布;若提供,则调用 make_low_rank_matrix 生成 低秩+噪声尾部 的矩阵,模拟真实数据中的相关结构。ground_truth 只在前 n_informative 列上赋予非零系数,从而构造 稀疏回归基准。噪声通过 noise 参数注入,shuffle 同时打乱样本和特征列,确保系数与特征对应不被破坏。

50.5.2 架构图

flowchart TD A[读取参数] --> B{effective_rank?} B -->|None| C[生成标准正态 X] B -->|指定| D[调用 make_low_rank_matrix → 低秩 X] C & D --> E[构造稀疏 ground_truth] E --> F[计算 y = X·ground_truth + bias] F --> G{noise>0?} G -->|Yes| H[加高斯噪声] G -->|No| I[保持 y 不变] H & I --> J[shuffle 同步打乱] J --> K{coef?} K -->|True| L[返回 (X, y, ground_truth)] K -->|False| M[返回 (X, y)]

50.5.3 设计取舍 Q&A

Q1:effective_rank 小于 n_features 会带来什么效应?

A1:特征之间会出现 强相关,导致矩阵的 谱衰减 较快,适合评估 稀疏正则化(Lasso)或 主成分回归 的表现;若 effective_rank 很大,则近似全秩,模型更倾向于普通最小二乘。

Q2:tail_strength 如何调节噪声尾巴?

A2tail_strength 控制 奇异值的慢衰减部分,值越大,尾部奇异值占比越高,矩阵的 噪声成分 更显著,模拟真实业务中常见的 高维噪声

Q3:返回 coef 的意义何在?

A3coef=True 同时返回 ground_truth,便于 模型评估(如计算 、系数恢复率),在教学或基准测试中非常有用。


50.6 同心圆 make_circles

源码路径sklearn/datasets/_samples_generator.pymake_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 架构图

flowchart TD A[计算外环/内环点数] --> B[均匀采样角度] B --> C[计算坐标并对内环乘以 factor] C --> D[堆叠坐标得到 X] D --> E[生成标签 y (0/1)] E --> F{shuffle?} F -->|Yes| G[随机置换 X, y] F -->|No| G G --> H{noise?} H -->|Yes| I[加高斯噪声] H -->|No| I I --> J[返回 X, y]

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.pymake_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 架构图

flowchart TD A[分配外/内样本数] --> B[生成上半圆坐标] B --> C[生成下半圆(平移)坐标] C --> D[堆叠得到 X] D --> E[标签 y = 0/1] E --> F{shuffle?} F -->|Yes| G[打乱样本] F -->|No| G G --> H{noise?} H -->|Yes| I[添加高斯噪声] H -->|No| I I --> J[返回 X, y]

50.7.3 设计取舍 Q&A

Q1:为何把内月牙整体右移 1 并下移 0.5?

A1:此平移保证两条月牙 交叉,形成 非线性、非凸 的决策边界,能够检验模型对 曲线分离 的学习能力。

Q2:shuffle=False 的意义是什么?

A2:保留顺序可以帮助演示 序列学习算法(如在线梯度下降)在不打乱数据时可能出现的偏差。


50.8 高斯团簇 make_blobs

源码路径sklearn/datasets/_samples_generator.pymake_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 架构图

flowchart TD A[解析 centers 参数] --> B[生成或使用中心坐标] B --> C[解析 cluster_std] C --> D[分配每簇样本数] D --> E[循环:多元正态采样每簇] E --> F[拼接所有簇得到 X, y] F --> G{shuffle?} G -->|Yes| H[随机置换样本顺序] G -->|No| H H --> I{return_centers?} I -->|Yes| J[返回 (X, y, centers)] I -->|No| K[返回 (X, y)]

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.pymake_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.pymake_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
posted @ 2026-09-04 04:08  绝不原创的飞龙  阅读(3)  评论(0)    收藏  举报