使用 MLflow 跟踪服务器进行远程实验跟踪
在本教程中,您将学习如何使用 MLflow Tracking Server 为团队开发设置 MLflow Tracking 环境。
使用 MLflow Tracking Server 进行远程实验跟踪有许多好处:
- 协作: 多个用户可以向相同的端点记录运行,并查询其他用户记录的运行和模型。
- 共享结果:跟踪服务器还提供一个 Tracking UI 端点,团队成员可以轻松查看彼此的结果。
- 集中访问: 跟踪服务器可以作为远程元数据和工件访问的代理运行,从而更容易确保对数据访问的安全性并进行审计。
它是如何工作的?
下图描述了使用远程 MLflow Tracking Server 与 PostgreSQL 和 S3 的架构
您可以在 artifact stores 和 backend stores 文档指南中找到受支持的数据存储列表。
当您开始向 MLflow Tracking Server 记录运行时,会发生以下情况:
-
第1部分 a 和 b:
- MLflow 客户端创建了一个 RestStore 实例并发送 REST API 请求来记录 MLflow 实体
- Tracking Server 会创建一个 SQLAlchemyStore 实例并连接到远程主机,用于将跟踪信息插入到数据库(即指标、参数、标签等)。
-
第1部分 c 和 d:
- 客户端的检索请求会返回来自已配置的 SQLAlchemyStore 表的信息
-
第2部分 a 和 b:
- 用于工件的日志事件由客户端使用
HttpArtifactRepository将文件写入 MLflow Tracking Server - 跟踪服务器随后使用假定角色认证将这些文件写入配置的对象存储位置
- 用于工件的日志事件由客户端使用
-
第2c和2d部分:
- 为用户请求从已配置的后端存储检索工件时,使用与在服务器启动时配置的相同授权认证。
- 工件通过跟踪服务器的
HttpArtifactRepository接口传递给最终用户
开始使用
前言
在实际的生产部署环境中,如上图所示,您会有多台远程主机来运行跟踪服务器和数据库。然而,出于本教程的目的, 我们将仅使用一台机器,在不同端口上运行多个 Docker 容器,以用更简单的评估教程设置来模拟远程环境。我们还将使用 MinIO, 一种与 S3 兼容的对象存储,作为工件存储,这样您就不需要拥有 AWS 账户来运行本教程。
步骤 1 - 获取 MLflow 及额外依赖
MLflow 在 PyPI 上可用。另需安装 pyscopg2 和 boto3 以便使用 Python 访问 PostgreSQL 和 S3。 如果您的系统尚未安装它们,可以通过以下方式安装:
pip install mlflow psycopg2 boto3
步骤 2 - 设置远程数据存储
MLflow Tracking Server 可以与多种数据存储交互,以存储实验和运行数据以及工件。 在本教程中,我们将使用 Docker Compose 启动两个容器,每个容器模拟实际环境中的远程服务器。
- PostgreSQL 数据库作为后端存储。
- MinIO 服务器作为制品存储。
安装 docker 和 docker-compose
这些 Docker 步骤仅为教程目的所需。MLflow 本身完全不依赖于 Docker。
按照官方说明安装 Docker 和 Docker Compose。然后,运行 docker --version 和 docker-compose --version 以确保它们已正确安装。
创建 compose.yaml
创建一个名为 compose.yaml 的文件,其内容如下:
version: "3.7"
services:
# PostgreSQL database
postgres:
image: postgres:latest
environment:
POSTGRES_USER: user
POSTGRES_PASSWORD: password
POSTGRES_DB: mlflowdb
ports:
- 5432:5432
volumes:
- ./postgres-data:/var/lib/postgresql/data
# MinIO server
minio:
image: minio/minio
expose:
- "9000"
ports:
- "9000:9000"
# MinIO Console is available at http://localhost:9001
- "9001:9001"
environment:
MINIO_ROOT_USER: "minio_user"
MINIO_ROOT_PASSWORD: "minio_password"
healthcheck:
test: timeout 5s bash -c ':> /dev/tcp/127.0.0.1/9000' || exit 1
interval: 1s
timeout: 10s
retries: 5
command: server /data --console-address ":9001"
# Create a bucket named "bucket" if it doesn't exist
minio-create-bucket:
image: minio/mc
depends_on:
minio:
condition: service_healthy
entrypoint: >
bash -c "
mc alias set minio http://minio:9000 minio_user minio_password &&
if ! mc ls minio/bucket; then
mc mb minio/bucket
else
echo 'bucket already exists'
fi
"
启动容器
在包含 compose.yaml 文件的相同目录中运行以下命令以启动容器。
这将以后台方式启动 PostgreSQL 和 Minio 服务器的容器,并在 Minio 中创建一个名为 "bucket" 的新存储桶。
docker compose up -d
第3步 - 启动跟踪服务器
在实际环境中,你会有一台远程主机来运行跟踪服务器,但在本教程中我们将仅使用本地机器作为远程主机的模拟替代。
配置访问
为了让跟踪服务器能够访问远程存储,需要为其配置必要的凭据。
export MLFLOW_S3_ENDPOINT_URL=http://localhost:9000 # Replace this with remote storage endpoint e.g. s3://my-bucket in real use cases
export AWS_ACCESS_KEY_ID=minio_user
export AWS_SECRET_ACCESS_KEY=minio_password
您可以在 Supported Storage 中找到有关如何为其他存储配置凭据的说明。
启动跟踪服务器
要指定后端存储和工件存储,可以分别使用 --backend-store-uri 和 --artifacts-store-uri 选项。
mlflow server \
--backend-store-uri postgresql://user:password@localhost:5432/mlflowdb \
--artifacts-destination s3://bucket \
--host 0.0.0.0 \
--port 5000
在实际环境中,将 localhost 替换为数据库服务器的远程主机名或 IP 地址。
步骤 4:将日志记录到跟踪服务器
一旦跟踪服务器运行,你可以通过将 MLflow Tracking URI 设置为该跟踪服务器的 URI 来记录运行。或者,你也可以使用 mlflow.set_tracking_uri() API 来设置跟踪 URI。
export MLFLOW_TRACKING_URI=http://127.0.0.1:5000 # Replace with remote host name or IP address in an actual environment
然后像往常一样使用 MLflow tracking APIs 运行您的代码。以下代码在糖尿病数据集上对 scikit-learn RandomForest 模型进行训练:
import mlflow
from sklearn.model_selection import train_test_split
from sklearn.datasets import load_diabetes
from sklearn.ensemble import RandomForestRegressor
mlflow.autolog()
db = load_diabetes()
X_train, X_test, y_train, y_test = train_test_split(db.data, db.target)
# Create and train models.
rf = RandomForestRegressor(n_estimators=100, max_depth=6, max_features=3)
rf.fit(X_train, y_train)
# Use the model to make predictions on the test dataset.
predictions = rf.predict(X_test)
第 5 步:在 Tracking UI 查看已记录的 Run
我们的伪远程 MLflow 跟踪服务器也在同一端点托管跟踪 UI。在具有远程跟踪服务器的实际部署环境中也是如此。
您可以在浏览器中通过访问 http://127.0.0.1:5000(在实际环境中将其替换为远程主机名或 IP 地址)来访问该 UI。
第 6 步:下载工件
MLflow Tracking Server 还作为工件访问的代理主机。通过诸如 models:/、mlflow-artifacts:/ 等代理 URI 启用对工件的访问,使用户能够访问该位置,而无需管理直接访问的凭证或权限。
import mlflow
model_id = "YOUR_MODEL_ID" # You can find model ID in the Tracking UI
# Download artifact via the tracking server
mlflow_artifact_uri = f"models:/{model_id}"
local_path = mlflow.artifacts.download_artifacts(mlflow_artifact_uri)
# Load the model
model = mlflow.sklearn.load_model(local_path)
接下来做什么?
现在您已经学会如何为远程实验跟踪设置 MLflow Tracking Server! 还有一些更高级的主题您可以探索:
- 跟踪服务器的其他配置:默认情况下,MLflow Tracking Server 同时提供后端存储和 artifact 存储。你也可以配置 Tracking Server 仅提供后端存储或仅提供 artifact 存储,以应对诸如大流量或安全性等不同用例。有关如何为这些用例自定义 Tracking Server,请参见 other use cases。
- 保护跟踪服务器:
--host选项会将服务暴露到所有接口上。如果在生产环境中运行服务器,我们建议不要广泛暴露内置服务器(因为它未经过身份验证且未加密)。阅读 Secure Tracking Server 以了解在生产环境中保护跟踪服务器的最佳实践。