Cleanlab 输入校验机制深度解析:cleanlab.internal.validation 模块源码导读
【免费下载链接】cleanlabCleanlab's open-source library is the standard>项目地址: https://gitcode.com/GitHub_Trending/cl/cleanlab
导读:cleanlab 作为以"数据为中心"的标准开源库,几乎所有公开 API(如
cleanlab.classification、cleanlab.count、cleanlab.filter以及 Datalab 内部模块)在计算标签噪声、置信度或异常值之前,都要先经过一层严格的输入合法性检查。本文以仓库中的 validation.rst 文档为骨架,结合其对应的 validation.py 源码与 test_validation.py 测试用例,系统拆解 cleanlab 的输入校验协议:标签格式、预测概率矩阵、特征张量索引能力、多标签嵌套结构等,帮助你在自定义数据管道或调用 cleanlab API 时一次性通过校验、避免踩坑。
模块定位:cleanlab 所有算法入口的第一道防线
cleanlab.internal.validation位于 cleanlab 的内部(internal)命名空间,其模块 docstring 只有一句话:"Checks to ensure valid inputs for various methods."(确保各种方法收到的输入是合法的)。它不直接向用户暴露业务功能,而是作为基础设施被上层模块反复调用。从源码结构看,该模块提供 6 个可独立调用的函数:
| 函数 | 职责 |
|---|---|
assert_valid_inputs | 顶层总入口:一次性校验 X、y、pred_probs 三者格式与相互一致性 |
assert_valid_class_labels | 校验单标签 y 为 0 起始整数的一维数组 |
assert_nonempty_input | 校验 X 非 None |
assert_indexing_works | 校验 X 支持列表式索引(含 pandas、PyTorch Dataset、稀疏矩阵特例) |
labels_to_array | 将 list / np.ndarray / pd.Series / pd.DataFrame 统一转换为 1D numpy 数组 |
labels_to_list_multilabel | 将多标签 y 规范化为嵌套 list |
配套的类型别名定义在 typing.py:LabelLike = Union[list, np.ndarray, pd.Series, pd.DataFrame],DatasetLike = Any(即只要行为像数据集即可)。这决定了整个校验协议的设计取向——cleanlab 对数据容器类型是宽松的,但对取值语义是苛刻的。
顶层校验入口:assert_valid_inputs
assert_valid_inputs是模块的核心函数,签名如下:
def assert_valid_inputs( X: DatasetLike, y: LabelLike, pred_probs: Optional[np.ndarray] = None, multi_label: bool = False, allow_missing_classes: bool = True, allow_one_class: bool = False, ) -> None它在一次调用内完成四类检查:标签容器类型检查、X 非空与长度对齐检查、pred_probs 形状与取值检查、多标签结构检查。下面逐一展开。
1. 标签容器类型检查
if not isinstance(y, (list, np.ndarray, np.generic, pd.Series, pd.DataFrame)): raise TypeError("labels should be a numpy array or pandas Series.")y 只接受 list、numpy 数组、pandas Series/DataFrame 四类容器。注意源码中np.generic也被纳入,说明 numpy 标量包装类型同样被兼容。若传入其他类型(如元组 tuple、字典 dict),会直接抛出TypeError。
2. 单标签模式下的类别合法性检查
当multi_label=False时,y 首先经过labels_to_array转为 1D 数组,再交给assert_valid_class_labels做深入检查(详见下一节)。这里的关键参数是:
allow_missing_classes=True(默认):允许数据中缺失部分类别编号,即np.unique(y)可以不等于[0,1,...,K-1]。cleanlab 很多算法能容忍"类别编号不连续"的数据。allow_one_class=False(默认):要求 y 中至少包含 2 个类别,否则抛ValueError。因为单类数据无法做标签噪声估计等任务。allow_missing_classes=False:强制要求 0..K-1 每个类别在数据中都出现,否则抛TypeError,并给出提示 "cleanlab requires zero-indexed integer labels (0,1,2,..,K-1)"。
3. X 与 y 的配对检查(pred_probs 缺省时必须执行)
allow_empty_X = True if pred_probs is None: allow_empty_X = False这是一个容易被忽视的设计:当只提供 X 和 y(pred_probs 为 None)时,X 不允许为空;而当 pred_probs 已提供时,X 可以是 None(因为很多 cleanlab 算法只需 pred_probs 即可工作)。逻辑分支如下:
if not allow_empty_X: assert_nonempty_input(X) try: num_examples = len(X) len_supported = True except: len_supported = False ...如果len(X)不可用,退而尝试X.shape[0];两者都不支持则抛TypeError("Data features X must support either: len(X) or X.shape[0]")。拿到样本数后校验num_examples != len(y),不一致时抛出带具体数字的ValueError(如 "X is length 100 and labels is length 99"),这是实际调试中最常见的报错之一。最后调用assert_indexing_works(X, length_X=num_examples)验证 X 可以被索引(详见第 5 节)。
4. pred_probs 校验
pred_probs 的检查覆盖四个维度:
if not isinstance(pred_probs, (np.ndarray, np.generic)): raise TypeError("pred_probs must be a numpy array.") if len(pred_probs) != len(y): raise ValueError("pred_probs and labels must have same length.") if len(pred_probs.shape) != 2: raise ValueError("pred_probs array must have shape: num_examples x num_classes.")- 类型:必须是 numpy 数组,pandas DataFrame 不满足要求;
- 长度:与 y 等长;
- 形状:必须是二维的
num_examples x num_classes,一维或三维均报错; - 列数下限:根据标签中的最大类别索引动态判定。单标签模式取
highest_class = max(y) + 1;多标签模式取所有非空样本标签子列表的最大值加一:
highest_class = max([max(y_i) for y_i in y if len(y_i) != 0]) + 1 if pred_probs.shape[1] < highest_class: raise ValueError( f"pred_probs must have at least {highest_class} columns, based on the largest class index which appears in labels." )- 取值域:要求所有概率值在
[0, 1]内,但允许浮点误差。这里引用了 constants.py 中定义的FLOATING_POINT_COMPARISON = 1e-6:
if (np.min(pred_probs) < 0 - FLOATING_POINT_COMPARISON) or ( np.max(pred_probs) > 1 + FLOATING_POINT_COMPARISON ): raise ValueError("Values in pred_probs must be between 0 and 1.")也就是说概率值落在[-1e-6, 1+1e-6]区间内都被视为合法,避免因 softmax/sigmoid 的浮点舍入误差误杀正常输入。
- 冗余提醒:当 X 与 pred_probs 同时提供时,会发出
warnings.warn("When X and pred_probs are both provided, the former may be ignored."),提示用户 X 可能不会被算法使用——这是 cleanlab 面向"只需 pred_probs 即可运行"工作流的明确信号。
5. 索引能力检查:assert_indexing_works
cleanlab 大量算法需要按索引子集访问数据(例如置信度剪枝、交叉验证折叠),因此必须确认 X 支持"列表式索引"。assert_indexing_works的检查顺序体现兼容性优先级:
if isinstance(X, (pd.DataFrame, pd.Series)): _ = X.iloc[idx] # 优先 pandas 的 iloc 位置索引 ... import torch if isinstance(X, torch.utils.data.Dataset): _ = torch.utils.data.Subset(X, idx) # PyTorch Dataset 走 Subset 包装 ... _ = X[idx] # 兜底:原生 __getitem__ 列表索引三条路径中任一条成功即通过;全部失败时抛出TypeError,并给出两条可选的修复建议:X[index_list]或 pandas 的X.iloc[index_list]。源码注释特别说明,length_X是可选参数,原因是稀疏矩阵不支持len(X)——当 X 为 scipy 稀疏矩阵时,调用方(如assert_valid_inputs)会把从X.shape[0]得到的长度传进来,从而让本函数对稀疏矩阵同样适用。这一细节保证了 cleanlab 可以处理大规模稀疏特征。
标签规范化工具:labels_to_array 与 labels_to_list_multilabel
这两个函数负责把"用户友好的多种标签容器"统一成 cleanlab 内部约定格式,并在此过程中暴露格式错误。
labels_to_array:统一为 1D numpy 数组
转换逻辑非常直接:
if isinstance(y, pd.Series): return y.to_numpy() elif isinstance(y, pd.DataFrame): y_arr = y.values if y_arr.shape[1] != 1: raise ValueError("labels must be one dimensional.") return y_arr.flatten() else: return np.asarray(y)- pandas Series 直接
to_numpy(); - pandas DataFrame 要求恰好一列,多列时抛
ValueError("labels must be one dimensional.")——对应测试 test_validation.py 中pd.DataFrame({"a": [0, 1], "b": [2, 3]})触发报错的用例; - list、numpy 数组等其余类型走
np.asarray兜底,转换失败时抛ValueError提示需能转成 1D 数组。
测试 test_validation.py 用参数化组合[["a","b","a"], [0,1,2]]×[list, np.array, pd.Series, pd.DataFrame]验证了所有容器都能正确归一为 numpy 数组且值保持不变。
labels_to_list_multilabel:多标签的严格嵌套结构
多标签模式下,y 的合法形态是"每个样本一个类别列表"的嵌套 list:
if not isinstance(y, list): raise ValueError("Unsupported Label format") if not all(isinstance(x, list) for x in y): raise ValueError("Each element in list of labels must be a list.")只支持外层与内层均为 list 的结构。注意它不做长度与 pred_probs 的匹配检查(那由assert_valid_inputs的multi_label=True分支完成),只负责结构规范化,体现职责分离。
assert_valid_class_labels:单标签语义的硬性规则
该函数把 y 的语义约束表达得非常明确:"a 1D numpy array where labels are zero-based integers (not multi-label)",逐条检查如下:
if y.ndim != 1: raise ValueError("Labels must be 1D numpy array.") if any([isinstance(label, str) for label in y]): raise ValueError("Labels cannot be strings, ...") if not np.equal(np.mod(y, 1), 0).all(): raise ValueError("Labels must be zero-indexed integers ...") if min(y) < 0: raise ValueError("Labels must be positive integers ...")四个维度依次是:维度(必须 1D)、非字符串(如["a","b","a"]这类标签会被明确拒绝,理由是 cleanlab 依赖类别索引而非类别名)、整型(通过y % 1 == 0判定,浮点 0.0/1.0 因取模为 0 仍可通过)、非负(从 0 开始)。随后检查类别数:allow_one_class=False时要求至少 2 个类别;allow_missing_classes=False时要求np.unique(y)严格等于np.arange(K)。
值得注意的组合语义:当allow_missing_classes=True(默认)且allow_one_class=False时,诸如[0, 0, 0](单类)会被拒绝,而[0, 0, 2](缺失类别 1)会被接受。这两者的组合是 cleanlab 各 API 面向不同算法需求动态调整的:例如cleanlab.count的多数方法默认allow_missing_classes=True,而某些要求全覆盖的估计器则收紧为False。
模块的实际调用链:校验如何保护算法核心
通过源码检索可以确认,该模块被 cleanlab 各核心模块广泛引用,作为算法执行前的公共前置检查。典型调用路径包括:
- classification.py:
CleanLearning.fit在训练前调用assert_valid_inputs(X, labels, pred_probs)(classification.py#L481),交叉验证阶段同样校验(classification.py#L741); - count.py:导入
assert_valid_inputs, labels_to_array,在estimate_py_noise_matrices(count.py#L124)中传入X=None只校验 pred_probs 与标签,在 count.py#L980 处同时校验 X; - datalab/internal/data.py:Datalab 数据加载时用
labels_to_array归一化标签,空 Datalab 场景下以labels_to_array([])初始化空标签(data.py#L263); - datalab/internal/issue_manager/label.py:标签问题检测器在
run前调用assert_valid_inputs(X=None, y=self.datalab.labels, pred_probs=pred_probs)(label.py#L273); - 此外 rank.py、filter.py、outlier.py、multiannotator.py 以及回归、目标检测等子模块均有引用。
一个值得注意的模式:凡是只需 pred_probs 的算法,调用方统一传X=None,从而自动跳过"X 非空"检查,只走 pred_probs 与标签的校验分支——这与assert_valid_inputs中allow_empty_X的动态逻辑严格对应,是"按需校验"设计的典型体现。
使用建议:如何让数据一次性通过 cleanlab 校验
基于以上源码行为,给实际使用者几条可直接落地的建议:
- 标签用 0 起始整数:将类别映射为
0,1,...,K-1的整数数组;字符串标签会直接抛错,需要自行做编码映射。 - 保持 y 为一维:pandas DataFrame 标签务必只有一列;多列 DataFrame 传入
labels_to_array会抛 "labels must be one dimensional"。 - pred_probs 必须是二维 numpy 数组:形状为
(n_samples, K),行数与 y 对齐,列数至少为max(y)+1(多标签时基于最大类别索引),值域落在[-1e-6, 1+1e-6]。不要传 pandas DataFrame 或一维概率向量。 - 只算置信度/找噪声标签时可以不传 X:cleanlab 多数算法只需 pred_probs 即可运行;传
X=None可跳过特征检查,同时避免 "X 会被忽略" 的告警。 - 特征 X 需支持列表式索引:普通 numpy 数组、pandas DataFrame(用 iloc)、PyTorch Dataset(用 Subset)均可;若使用自定义数据结构,请确保实现
__getitem__或提供len()/.shape[0],稀疏矩阵则依赖调用方传入length_X参数。 - 注意类别数下限:默认要求标签中至少出现 2 个类别(
allow_one_class=False);单类数据要么换用支持该场景的 API,要么显式调整对应参数。
延伸阅读
- 校验模块的 RST 文档:validation.rst
- 校验实现源码:validation.py
- 类型别名定义:typing.py
- 浮点比较阈值:constants.py
- 测试用例:test_validation.py
- 校验被调用的典型入口:classification.py、count.py、filter.py、rank.py
【免费下载链接】cleanlabCleanlab's open-source library is the standard>项目地址: https://gitcode.com/GitHub_Trending/cl/cleanlab
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考