跳到主要内容

MLflow 数据集追踪

mlflow.data 模块是贯穿机器学习模型开发工作流程的数据集管理综合解决方案。它支持您对训练、验证和评估过程中使用的数据集进行追踪、版本控制和管理,从而提供从原始数据到模型预测的完整血缘关系。

为什么数据集追踪很重要

数据集追踪对于可复现的机器学习至关重要,并提供以下几个关键优势:

  • 数据血缘:追踪从原始数据源到模型输入的完整历程
  • 可复现性:确保实验能够使用相同的数据集进行复现
  • 版本控制:随着数据集的演变,对其不同版本进行管理
  • 协作:在团队间共享数据集及其元数据
  • 评估集成:与 MLflow 的评估功能无缝集成
  • 生产环境监控:追踪生产推理和评估中使用的数据集

核心组件

MLflow 的数据集追踪围绕两个主要抽象概念展开

Dataset (数据集)

Dataset 抽象是一个元数据追踪对象,用于保存有关已记录数据集的全面信息。存储在 Dataset 对象中的信息包括:

核心属性

  • 名称 (Name):数据集的描述性标识符(若未指定,默认为 "dataset")
  • 摘要 (Digest):用于识别数据集的唯一哈希值/指纹(自动计算)
  • 来源 (Source):包含原始数据位置血缘信息的 DatasetSource
  • 模式 (Schema):可选的数据集模式(具体实现可能不同,例如 MLflow Schema)
  • 配置 (Profile):可选的汇总统计信息(具体实现可能不同,例如行数、列统计信息)

支持的数据集类型

特殊数据集类型

  • EvaluationDataset - 内部数据集类型,专用于 mlflow.models.evaluate() 模型评估工作流

DatasetSource (数据源)

DatasetSource 组件提供了与数据原始来源的关联血缘,无论是文件 URL、S3 存储桶、数据库表还是任何其他数据源。这确保了您始终可以追溯到数据的原始来源。

可以使用 mlflow.data.get_source() API 获取 DatasetSource,该 API 接受 DatasetDatasetEntityDatasetInput 的实例。

快速入门:基础数据集追踪

以下是如何开始使用基础数据集追踪的方法

python
import mlflow.data
import pandas as pd

# Load your data
dataset_source_url = (
"https://raw.githubusercontent.com/mlflow/mlflow/master/tests/datasets/winequality-white.csv"
)
raw_data = pd.read_csv(dataset_source_url, delimiter=";")

# Create a Dataset object
dataset = mlflow.data.from_pandas(
raw_data, source=dataset_source_url, name="wine-quality-white", targets="quality"
)

# Log the dataset to an MLflow run
with mlflow.start_run():
mlflow.log_input(dataset, context="training")

# Your training code here
# model = train_model(raw_data)
# mlflow.sklearn.log_model(model, "model")

数据集信息和元数据

创建数据集时,MLflow 会自动捕获丰富的元数据

python
# Access dataset metadata
print(f"Dataset name: {dataset.name}") # Defaults to "dataset" if not specified
print(f"Dataset digest: {dataset.digest}") # Unique hash identifier (computed automatically)
print(f"Dataset source: {dataset.source}") # DatasetSource object
print(f"Dataset profile: {dataset.profile}") # Optional: implementation-specific statistics
print(f"Dataset schema: {dataset.schema}") # Optional: implementation-specific schema

示例输出

text
Dataset name: wine-quality-white
Dataset digest: 2a1e42c4
Dataset profile: {"num_rows": 4898, "num_elements": 58776}
Dataset schema: {"mlflow_colspec": [
{"type": "double", "name": "fixed acidity"},
{"type": "double", "name": "volatile acidity"},
...
{"type": "long", "name": "quality"}
]}
Dataset source: <DatasetSource object>
数据集属性

profileschema 属性是特定于实现的,可能会因数据集类型(PandasDataset、SparkDataset 等)而异。某些数据集类型对于这些属性可能返回 None

数据源和血缘

MLflow 支持来自各种来源的数据集

python
# From local file
local_dataset = mlflow.data.from_pandas(df, source="/path/to/local/file.csv", name="local-data")

# From cloud storage
s3_dataset = mlflow.data.from_pandas(df, source="s3://bucket/data.parquet", name="s3-data")

# From database
db_dataset = mlflow.data.from_pandas(df, source="postgresql://user:pass@host/db", name="db-data")

# From URL
url_dataset = mlflow.data.from_pandas(df, source="https://example.com/data.csv", name="web-data")

MLflow UI 中的数据集追踪

当您将数据集记录到 MLflow 运行中时,它们会显示在带有全面元数据的 MLflow UI 中。您可以直接在界面中查看数据集信息、模式和血缘。

Dataset in MLflow UI

UI 显示内容:

  • 数据集名称和摘要
  • 包含列类型的模式信息
  • 概况统计(行数等)
  • 来源血缘信息
  • 使用该数据集的上下文

与 MLflow Evaluate 集成

MLflow 数据集最强大的功能之一是它与 MLflow 评估功能的无缝集成。使用 mlflow.models.evaluate() 时,MLflow 会在内部自动将各种数据类型转换为 EvaluationDataset 对象。

EvaluationDataset

MLflow 在使用 mlflow.models.evaluate() 时使用内部的 EvaluationDataset 类。此数据集类型是根据您的输入数据自动创建的,并为评估工作流提供了专门优化的哈希和元数据追踪。

直接将数据集用于 MLflow evaluate

python
import mlflow
from sklearn.ensemble import RandomForestClassifier
from sklearn.model_selection import train_test_split

# Prepare data and train model
data = pd.read_csv("classification_data.csv")
X = data.drop("target", axis=1)
y = data["target"]
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.3, random_state=42)

model = RandomForestClassifier()
model.fit(X_train, y_train)

# Create evaluation dataset
eval_data = X_test.copy()
eval_data["target"] = y_test

eval_dataset = mlflow.data.from_pandas(eval_data, targets="target", name="evaluation-set")

with mlflow.start_run():
# Log model
mlflow.sklearn.log_model(model, name="model", input_example=X_test)

# Evaluate using the dataset
result = mlflow.models.evaluate(
model="runs:/{}/model".format(mlflow.active_run().info.run_id),
data=eval_dataset,
model_type="classifier",
)

print(f"Accuracy: {result.metrics['accuracy_score']:.3f}")

MLflow Evaluate 集成示例

这是一个完整的示例,展示了数据集如何与 MLflow 的评估功能集成

Dataset Evaluation in MLflow UI

评估运行展示了数据集、模型、指标和评估工件(如混淆矩阵)是如何一起被记录的,从而提供了评估过程的完整视图。

高级数据集管理

追踪随演变产生的数据集版本

python
def create_versioned_dataset(data, version, base_name="customer-data"):
"""Create a versioned dataset with metadata."""

dataset = mlflow.data.from_pandas(
data,
source=f"data_pipeline_v{version}",
name=f"{base_name}-v{version}",
targets="target",
)

with mlflow.start_run(run_name=f"Dataset_Version_{version}"):
mlflow.log_input(dataset, context="versioning")

# Log version metadata
mlflow.log_params({
"dataset_version": version,
"data_size": len(data),
"features_count": len(data.columns) - 1,
"target_distribution": data["target"].value_counts().to_dict(),
})

# Log data quality metrics
mlflow.log_metrics({
"missing_values_pct": (data.isnull().sum().sum() / data.size) * 100,
"duplicate_rows": data.duplicated().sum(),
"target_balance": data["target"].std(),
})

return dataset


# Create multiple versions
v1_dataset = create_versioned_dataset(data_v1, "1.0")
v2_dataset = create_versioned_dataset(data_v2, "2.0")
v3_dataset = create_versioned_dataset(data_v3, "3.0")

生产环境使用案例

监控生产环境批量预测中使用的数据集

python
def monitor_batch_predictions(batch_data, model_version, date):
"""Monitor production batch prediction datasets."""

# Create dataset for batch predictions
batch_dataset = mlflow.data.from_pandas(
batch_data,
source=f"production_batch_{date}",
name=f"batch_predictions_{date}",
targets="true_label" if "true_label" in batch_data.columns else None,
predictions="prediction" if "prediction" in batch_data.columns else None,
)

with mlflow.start_run(run_name=f"Batch_Monitor_{date}"):
mlflow.log_input(batch_dataset, context="production_batch")

# Log production metadata
mlflow.log_params({
"batch_date": date,
"model_version": model_version,
"batch_size": len(batch_data),
"has_ground_truth": "true_label" in batch_data.columns,
})

# Monitor prediction distribution
if "prediction" in batch_data.columns:
pred_metrics = {
"prediction_mean": batch_data["prediction"].mean(),
"prediction_std": batch_data["prediction"].std(),
"unique_predictions": batch_data["prediction"].nunique(),
}
mlflow.log_metrics(pred_metrics)

# Evaluate if ground truth is available
if all(col in batch_data.columns for col in ["prediction", "true_label"]):
result = mlflow.models.evaluate(data=batch_dataset, model_type="classifier")
print(f"Batch accuracy: {result.metrics.get('accuracy_score', 'N/A')}")

return batch_dataset


# Usage
batch_dataset = monitor_batch_predictions(daily_batch_data, "v2.1", "2024-01-15")

最佳实践

使用 MLflow 数据集时,请遵循以下最佳实践:

数据质量:在记录数据集之前,请务必验证数据质量。检查缺失值、重复项和数据类型。

命名规范:为数据集使用一致的描述性名称,并包含版本信息和上下文。

来源文档:始终指定有意义的源 URL 或标识符,以便能够追溯到原始数据。

上下文规范:记录数据集时使用清晰的上下文标签(例如 "training"、"validation"、"evaluation"、"production")。

元数据记录:包括有关数据采集、预处理步骤和数据特征的相关元数据。

版本控制:显式追踪数据集版本,特别是在数据预处理或采集方法发生变化时。

摘要计算:不同数据集类型的计算方式不同:

  • 标准数据集:基于数据内容和结构。
  • MetaDataset:仅基于元数据(名称、来源、模式)—— 不对实际数据进行哈希。
  • EvaluationDataset:针对大数据集使用采样行进行优化哈希。

来源灵活性:DatasetSource 支持多种来源类型,包括 HTTP URL、文件路径、数据库连接和云存储位置。

评估集成:在设计数据集时考虑到评估需求,明确指定目标列和预测列。

核心优势

MLflow 数据集追踪为 ML 团队提供了几个关键优势:

可复现性:确保实验能够使用相同的数据集进行复现,即使数据源发生变化。

血缘追踪:保持从来源到模型预测的完整数据血缘,从而实现更好的调试和合规性。

协作:通过统一的接口在团队成员之间共享数据集及其元数据。

评估集成:与 MLflow 的评估功能无缝集成,进行全面的模型评估。

生产环境监控:追踪生产系统中用于性能监控和数据偏移检测的数据集。

质量保证:自动捕获数据质量指标并监控随时间的变化。

无论您是追踪训练数据集、管理评估数据还是监控生产批量预测,MLflow 的数据集追踪功能都为您可靠且可复现的机器学习工作流奠定了基础。