MLflow 数据追踪实战指南:深入解析 mlflow.data 模块的数据集记录与溯源体系
【免费下载链接】mlflowThe open source AI engineering platform for agents, LLMs, and ML models. MLflow enables teams of all sizes to debug, evaluate, monitor, and optimize production-quality AI applications while controlling costs and managing access to models and data.项目地址: https://gitcode.com/GitHub_Trending/ml/mlflow
MLflow 的mlflow.data模块是面向模型训练与评估数据集的统一记录与检索接口,它让团队可以把训练集、验证集等数据资产作为一等公民随 Run 一并落库,并完整保存数据集的名称、摘要哈希、Schema、Profile 与来源信息。本文以官方 API 文档 mlflow.data.rst 为核心骨架,结合仓库源码深入拆解Dataset与DatasetSource两大抽象、各数据框架构造器的参数语义、来源解析注册机制,以及日志、检索、重载来源的完整实战流程。
一、mlflow.data 解决了什么问题
机器学习项目中的可复现性不仅取决于模型代码与超参数,还取决于"模型是在哪份数据上训练出来的"。传统做法是在实验记录里手写一段数据路径备注,既不可检索也不可校验。mlflow.data模块的目标是:
- 将训练/评估数据集以结构化元数据的形式记录到 MLflow Tracking 的 Run 中(通过
mlflow.log_input()); - 支持从多种 Python 数据对象构造数据集,包括 Pandas DataFrame(
mlflow.data.from_pandas())、NumPy 数组(mlflow.data.from_numpy())、Spark DataFrame(mlflow.data.from_spark()/mlflow.data.load_delta())、Polars DataFrame(mlflow.data.from_polars()),以及 Hugging Face Dataset 与 TensorFlow Dataset; - 记录数据集的 name、digest(哈希指纹)、schema、profile 与 source,并支持通过
mlflow.data.get_source()反查并重载数据来源。
整个模块的核心由两个抽象基类构成:Dataset与DatasetSource,其定义分别位于 mlflow/data/dataset.py 与 mlflow/data/dataset_source.py。
二、两个核心抽象:Dataset 与 DatasetSource
2.1 Dataset:数据集元数据的统一载体
Dataset抽象类代表一份用于模型训练或评估的数据集,它承载的元数据包括:
- features / targets / predictions:特征、目标(监督学习标签)与预测列;
- name:数据集名称,如
"iris_data"、"myschema.mycatalog.mytable@v1";未指定时默认返回"dataset"(见 dataset.py); - digest:数据集的唯一哈希指纹,如
"498c7496",用于跨团队比对"同一份数据"(见 dataset.py); - schema:可选的数据集 Schema,如
mlflow.types.Schema; - profile:可选的统计概要,如行数、各列均值/中位数/标准差等;
- source:数据来源,即
DatasetSource实例,例如 S3 目录、Delta Table 或 Web URL。
从源码看,Dataset基类在构造时会先保存 name 与 source,再调用抽象方法_compute_digest()计算摘要(若用户未显式传入 digest),并定义了to_dict()/to_json()序列化接口与_to_mlflow_entity()到mlflow.entities.Dataset的转换逻辑(dataset.py)。这意味着"日志数据集"本质上是把内存中的数据对象转成一份 JSON 元数据,再写入 Tracking 后端。
2.2 DatasetSource:数据来源的抽象与可重载性
DatasetSource抽象类代表数据集的来源,比如 S3 中的文件目录、一个 Delta Table 或一个 Web URL。它定义了四个关键能力(dataset_source.py):
| 方法 | 作用 |
|---|---|
_get_source_type() | 返回来源类型字符串,如"s3"、"delta_table"、"http" |
load() | 加载来源指向的文件/对象:可能把 CSV 从 S3 下载到本地、把 Delta Table 加载为 Spark DataFrame 等,并返回加载结果 |
_can_resolve(raw_source) | 判断本类能否解析某个原始来源对象,如能解析"s3://mybucket/..."但不能解析 Azure Blob 的"wasbs://..." |
_resolve(raw_source) | 由原始来源(URI 字符串、delta 表标识等)构造出具体的DatasetSource实例 |
此外DatasetSource还支持to_dict()/to_json()/from_dict()/from_json()的完整序列化闭环,这正是mlflow.data.get_source()能从 Run 里反序列化出数据源的关键。
值得注意:Dataset 的 features/targets 与 source 并不一定完全一致——如果训练前做过转换或过滤,记录的是处理后的数据特征,而 source 保留的是原始来源。这是 mlflow.data 溯源设计上的一个重要语义。
三、一个完整的端到端实战示例
官方文档给出了一个从构造数据集、日志到 Run、检索元数据、再重载来源的完整例子。下面结合源码逐段拆解(原文见 mlflow.data.rst 中的示例代码块):
import mlflow.data import pandas as pd from mlflow.data.pandas_dataset import PandasDataset # 1. 用 Web URL 的葡萄酒质量数据构造 Pandas DataFrame dataset_source_url = "http://archive.ics.uci.edu/ml/machine-learning-databases/wine-quality/winequality-red.csv" df = pd.read_csv(dataset_source_url) # 2. 从 DataFrame 构造 PandasDataset,并把 Web URL 指定为数据来源 dataset: PandasDataset = mlflow.data.from_pandas(df, source=dataset_source_url) with mlflow.start_run(): # 3. 把数据集记录到 Run,context="training" 表明该数据集用于模型训练 mlflow.log_input(dataset, context="training") # 4. 取回 Run 及其数据集元数据 run = mlflow.get_run(mlflow.last_active_run().info.run_id) dataset_info = run.inputs.dataset_inputs[0].dataset print(f"Dataset name: {dataset_info.name}") print(f"Dataset digest: {dataset_info.digest}") print(f"Dataset profile: {dataset_info.profile}") print(f"Dataset schema: {dataset_info.schema}") # 5. 通过 get_source() 拿到来源并 load() 下载内容到本地文件系统 dataset_source = mlflow.data.get_source(dataset_info) dataset_source.load()这个流程中:
- 第 2 步:
from_pandas()内部会调用resolve_dataset_source(source)把 URL 字符串解析为具体的DatasetSource(此处为HTTPDatasetSource)。若source传None,则会用当前运行上下文标签构造CodeDatasetSource,把"数据集来源是当前代码位置"(如 notebook cell、脚本)记录下来,实现细节见 pandas_dataset.py。 - 第 3 步:
mlflow.log_input()接收 Dataset 与上下文标签("training"、"eval"等),把数据集元数据作为 DatasetInput 写入 Run。 - 第 5 步:
mlflow.data.get_source()支持三种入参:Dataset、mlflow.entities.Dataset或DatasetInput;对实体类型它会依据source_type从 JSON 还原出对应的DatasetSource(见 mlflow/data/init.py)。
四、按数据框架逐一详解构造器
mlflow.data的构造器遵循统一的命名规范:以from_或load_开头,且必须接收可选的name与digest关键字参数。这些约束在 dataset_registry.py 的注册校验逻辑中被强制执行,校验通过后构造器会被动态挂载到mlflow.data命名空间(见 mlflow/data/init.py),因此你既可以写mlflow.data.from_pandas(...),也可以从对应模块直接导入。
4.1 Pandas:from_pandas()/PandasDataset
import mlflow import pandas as pd x = pd.DataFrame( [["tom", 10, 1, 1], ["nick", 15, 0, 1], ["july", 14, 1, 1]], columns=["Name", "Age", "Label", "ModelOutput"], ) dataset = mlflow.data.from_pandas(x, targets="Label", predictions="ModelOutput")参数语义(见 pandas_dataset.py):
df:必填的 Pandas DataFrame;source:来源,可以是 URI、路径字符串或DatasetSource实例;不传则默认为CodeDatasetSource;targets:可选的目标列名,必须存在于 df 中,否则构造时抛出MlflowException;predictions:可选的预测列名,同样必须存在于 df 中;name/digest:可选,未指定时自动生成。
PandasDataset的profile返回{"num_rows": 行数, "num_elements": 元素总数},schema通过_infer_schema()推导为mlflow.types.Schema;digest 由compute_pandas_digest()计算(pandas_dataset.py)。此外它实现了to_pyfunc()与to_evaluation_dataset(),可无缝配合mlflow.evaluate()做模型评估(这也是文档中 RST 指令排除这两个成员的原因——它们服务于评估流程而非元数据记录)。
4.2 NumPy:from_numpy()/NumpyDataset
import mlflow import numpy as np # 基础用法:features + targets x = np.random.uniform(size=[2, 5, 4]) y = np.random.randint(2, size=[2]) dataset = mlflow.data.from_numpy(x, targets=y) # 字典用法:命名特征 x = { "feature_1": np.random.uniform(size=[2, 5, 4]), "feature_2": np.random.uniform(size=[2, 5, 4]), } y = np.random.randint(2, size=[2]) dataset = mlflow.data.from_numpy(x, targets=y)NumpyDataset的 features/targets 既可以是单个np.ndarray,也可以是dict[str, np.ndarray]的命名数组(numpy_dataset.py)。它的profile会记录 features/targets 的shape、size、nbytes,schema使用TensorDatasetSchema(features + targets 两个 TensorSpec)表示,digest 由compute_numpy_digest()计算。
4.3 Spark 与 Delta:from_spark()/load_delta()/SparkDataset
Spark 场景提供了两条路径:
路径 A:从已有的 Spark DataFrame 构造
from pyspark.sql import SparkSession import mlflow spark = SparkSession.builder.getOrCreate() df = spark.read.csv("path/to/data") dataset = mlflow.data.from_spark( df, path="path/to/data", # 或 table_name=...、sql=... targets="label", )from_spark()允许从path、table_name、sql中至多指定一个来描述 DataFrame 的原始来源(三者全不指定则使用CodeDatasetSource),并有三条校验规则(spark_dataset.py):
path、table_name、sql同时指定多于一个时抛错;sql与version不能同时指定;- 若指定
version,对应的 path/table_name 必须确实指向 Delta 表。
当 path/table_name 指向 Delta 表时,from_spark()会自动构造DeltaDatasetSource并附带 Delta 版本号;否则构造普通的SparkDatasetSource。
路径 B:直接按 Delta 表加载
dataset = mlflow.data.load_delta( table_name="my_table", # 或 path="dbfs:/path/to/delta" version=1, # 可选,不指定则自动推断最新版本 targets="label", )load_delta()要求path与table_name恰好指定其中一个;version不指定时会尝试自动获取 Delta 表最新版本;name未指定时默认生成"{table_name}@v{version}"的形式(spark_dataset.py)。
SparkDataset的 digest 计算值得一提:它不对 DataFrame 逐行哈希,而是对 DataFrame 逻辑计划的semanticHash()做 MD5 归一化(Spark 3.1.0+ 直接调用df.semanticHash()),既高效又确定(spark_dataset.py);profile则通过countApprox在 5 秒超时内给出approx_count近似行数(超时估算为 0 时返回"unknown")。
4.4 Hugging Face:from_huggingface()/HuggingFaceDataset
import mlflow from datasets import load_dataset ds = load_dataset("databricks/databricks-dolly-15k", split="train") dataset = mlflow.data.from_huggingface( ds, path="databricks/databricks-dolly-15k", # 必须与 Hub 上的路径一致才能重载 targets=None, revision="v1", )关键约束(huggingface_dataset.py):
ds必须是datasets.Dataset实例,不支持DatasetDict等其他类型,否则抛错;path指定时构造HuggingFaceDatasetSource,其中config_name、split从ds自动提取,data_dir、data_files、revision、trust_remote_code由用户传入,用于将来调用datasets.load_dataset()重载数据;- 若同时传了
source与path,source优先生效,path被忽略(并给出告警); path不传则退化为CodeDatasetSource。
HuggingFaceDataset的 digest 与 schema 推导只取前 10000 行(_MAX_ROWS_FOR_DIGEST_COMPUTATION_AND_SCHEMA_INFERENCE),避免大数据集上的性能问题;profile包含num_rows、dataset_size、size_in_bytes。
4.5 Polars:from_polars()/PolarsDataset
import mlflow import polars as pl x = pl.DataFrame( [["tom", 10, 1, 1], ["nick", 15, 0, 1], ["julie", 14, 1, 1]], schema=["Name", "Age", "Label", "ModelOutput"], ) dataset = mlflow.data.from_polars(x, targets="Label", predictions="ModelOutput")注意版本前提:mlflow.data.polars_dataset要求 polars >= 1.0.0,模块导入时即做版本检查(polars_dataset.py)。Polars 的 Schema 推导有一套专门的类型映射:TYPE_MAP精确映射(如Int64 -> long、String -> string),CLOSE_MAP做近似映射(如Categorical/Enum -> string、Date -> datetime),无法映射的类型则标记为"Unknown"或抛错(polars_dataset.py)。PolarsDataset的 profile 为{"num_rows": height, "num_elements": height * width},digest 通过df.hash_rows().sum()计算。
4.6 TensorFlow:from_tensorflow()/TensorFlowDataset
import tensorflow as tf import mlflow features = tf.random.uniform([100, 4]) targets = tf.random.uniform([100, 1]) dataset = mlflow.data.from_tensorflow(features, targets=targets)TensorFlowDataset的 features/targets 必须是tf.data.Dataset或 TensorFlow Tensor,且二者类型必须一致(Tensor 对 Tensor、Dataset 对 Dataset),否则构造时报错(tensorflow_dataset.py)。其 digest 计算会迭代tf.data.Dataset元素并分桶聚合,schema 同样采用TensorDatasetSchema。
五、Dataset Sources:内置来源与注册优先级
文档底部列出了五类内置 DatasetSource,均可在mlflow.data.sources命名空间下访问(模块导入时由_define_dataset_sources_in_sources_module()动态挂载):
| 来源类 | 典型来源示例 | 说明 |
|---|---|---|
FileSystemDatasetSource | 本地文件系统路径 | 文件系统目录/文件 |
HTTPDatasetSource | Web URL | HTTP(S) 下载型来源,load()会把内容下载到本地临时目录 |
HuggingFaceDatasetSource | Hugging Face Hub 数据集路径 | 保留 config/split/revision 等重载参数 |
DeltaDatasetSource | Delta 表名或路径 + 版本 | 加载为 Spark DataFrame |
SparkDatasetSource | Spark 表名、文件目录或 SQL 语句 | 通过spark.table()/spark.read/spark.sql重载 |
此外仓库中还实现了CodeDatasetSource(把来源标记为代码位置)、UCVolumeDatasetSource(Databricks UC Volume)以及基于 artifact 仓库的通用来源,后者通过 artifact_dataset_sources.py 注册(如 S3、Azure Blob、GCS 等,因此文档示例中的 S3 目录也可以直接作为 source)。
来源解析的优先级由注册顺序决定(dataset_source_registry.py):
- 先注册 artifact 通用来源(优先级最低);
- 再注册
HTTPDatasetSource与外部 entrypoint 来源; - 依次注册 HuggingFace、Spark、Delta、Code、UC Volume、Databricks 评估数据源(优先级更高)。
当用户传入的原始 source 能被多个来源类解析时(如"s3://..."),MLflow 会告警并选择最后注册的匹配类;get_dataset_source_from_json()反序列化时则按相反顺序遍历、以source_type精确匹配。
六、与评估流程的衔接:EvaluationDataset 与 PyFunc 转换
各数据集类都实现了to_pyfunc()与to_evaluation_dataset()两个方法(RST 指令中之所以:exclude-members: to_pyfunc, to_evaluation_dataset,正是为了让 API 文档聚焦于数据记录接口本身):
to_pyfunc():把数据集拆分为 pyfunc inputs 与 outputs(PyFuncInputsOutputs),供mlflow.evaluate()直接喂给模型函数;to_evaluation_dataset():转换为EvaluationDataset,其中可携带 path、feature_names、predictions 等评估所需信息(evaluation_dataset.py)。
例如PandasDataset.to_pyfunc()会在指定targets时把目标列从特征中剥离作为 outputs(pandas_dataset.py);SparkDataset出于驱动内存考虑只取前 10000 行做转换(spark_dataset.py)。
七、扩展机制:插件式注册自定义数据集
mlflow.data的设计允许第三方通过 Python entrypoint 扩展:
mlflow.dataset_constructor:注册自定义数据集构造器(须以from_/load_开头、含可选name/digest参数、返回类型标注为Dataset子类),注册后自动暴露为mlflow.data.from_xxx();mlflow.dataset_source:注册自定义来源类,参与resolve_dataset_source()的统一解析。
注册与校验逻辑见 dataset_registry.py 与 dataset_source_registry.py,这使得 mlflow.data 的溯源体系可以覆盖任意内部数据平台。
八、小结与最佳实践
围绕mlflow.data,几个值得固化的实践:
- 训练前必记录:在
mlflow.start_run()中调用mlflow.log_input(dataset, context="training"),让每次实验自带数据指纹; - 善用 digest 做数据版本比对:digest 基于数据内容计算(Pandas 哈希行列值、Spark 哈希逻辑计划、HF 取前 10000 行),同一份数据在不同实验间 digest 一致,可快速发现"换了数据";
- source 优先选择可重载形式:指定 URL / Delta 表(含版本)/ HF Hub 路径等,而非留空——这样
mlflow.data.get_source(dataset_info).load()才能真实还原数据; - 区分 source 与处理后数据:features/targets 是内存中经转换后的数据,source 是原始位置,二者天然允许不同。
相关源码速查:核心抽象在 mlflow/data/dataset.py 与 mlflow/data/dataset_source.py,构造器注册在 dataset_registry.py,来源解析在 dataset_source_registry.py,各框架实现分布在mlflow/data/pandas_dataset.py、numpy_dataset.py、spark_dataset.py、polars_dataset.py、huggingface_dataset.py、tensorflow_dataset.py中;官方测试用例可参考 tests/data 目录(若存在)以验证各构造器的行为边界。
【免费下载链接】mlflowThe open source AI engineering platform for agents, LLMs, and ML models. MLflow enables teams of all sizes to debug, evaluate, monitor, and optimize production-quality AI applications while controlling costs and managing access to models and data.项目地址: https://gitcode.com/GitHub_Trending/ml/mlflow
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考