MLflow 官方示例精读:用 statsmodels 训练 OLS 模型并完成自动日志记录
【免费下载链接】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
本篇技术指南以 examples/statsmodels/README.md 为主线,完整讲解如何在 MLflow 中使用 statsmodels 训练一个 OLS(Ordinary Least Squares,普通最小二乘)回归模型,并借助 MLflow Tracking 自动记录超参数、MSE 指标与训练好的模型。读完本文,你将掌握两种运行方式(直接运行 Python 脚本、以 MLflow Project 方式运行)、--inverse-method参数的底层含义(QR 分解 vs Moore-Penrose 伪逆),以及mlflow.statsmodelsflavor 自动日志(autolog)背后真正记录了哪些内容。
示例概览:一个 4 文件的完整可运行示例
该示例位于仓库的 examples/statsmodels 目录,共包含 4 个文件:
| 文件 | 作用 |
|---|---|
| train.py | 示例主程序:生成合成数据、训练 OLS 模型、启用 autolog 并记录 MSE 指标 |
| MLproject | 将示例包装为 MLflow Project 的元数据文件,声明参数与入口命令 |
| python_env.yaml | 项目运行时依赖:mlflow、statsmodels、scikit-learn |
| README.md | 官方使用说明(本文的骨架) |
python_env.yaml只声明了三项 pip 依赖,其中scikit-learn是为了计算mean_squared_error指标,statsmodels是训练主依赖,mlflow提供 Tracking 与 autolog 能力。
合成数据生成:OLS 的最小可行实验
在 train.py 中,示例用 numpy 构造了一个带二次项和噪声的回归问题:
np.random.seed(9876789) nsamples = 100 x = np.linspace(0, 10, 100) X = np.column_stack((x, x**2)) beta = np.array([1, 0.1, 10]) e = np.random.normal(size=nsamples) X = sm.add_constant(X) y = np.dot(X, beta) + e- 自变量矩阵
X由线性项x、二次项x**2和一列常数项(截距)构成,形状为(100, 3); - 真实系数
beta = [1, 0.1, 10],分别对应截距、一次项与二次项; - 因变量
y = X · beta + e,其中e为标准正态噪声; - 固定随机种子
9876789保证结果可复现——这一数据生成逻辑与仓库测试夹具 tests/statsmodels/model_fixtures.py 中的ols_model()完全一致,测试与示例共用同一套数据,便于对照验证。
随后执行:
ols = sm.OLS(y, X) model = ols.fit(method=args.inverse_method)即用 statsmodels 的sm.OLS拟合最小二乘模型,method参数由命令行传入。
直接运行:--inverse-method参数的两种求逆策略
README 给出的第一种运行方式是直接执行脚本:
python train.py --inverse-method qr--inverse-method控制的是求解最小二乘问题时逆矩阵的计算方式,可选值为qr或pinv(默认),其参数定义见 train.py:
| 取值 | 默认值 | 求解原理 |
|---|---|---|
pinv | ✅ | 使用 Moore-Penrose 伪逆(np.linalg.pinv)求解最小二乘问题 |
qr | ❌ | 使用 QR 分解(np.linalg.qr)求解 |
两种方法的取舍要点:
pinv(伪逆法):数值上更稳健,尤其适合矩阵接近奇异(病态)的情况,因为它基于奇异值分解,能自动处理秩亏矩阵;代价是计算开销略高;qr(QR 分解法):计算效率更高、内存占用更少,适合数据规模较大且矩阵良态的场景;但在处理近乎奇异的矩阵时数值稳定性不如伪逆。
README 明确建议读者两种方法都试一遍,甚至可以省略--inverse-method参数(此时自动回落到默认值pinv)。这正是 MLflow 实验跟踪的典型使用场景:通过多次运行对比不同求解策略下模型的指标表现。由于超参数与指标都会被自动记录,你可以在 MLflow UI 中直接横向对比qr与pinv两批运行的结果。
以 MLflow Project 方式运行
README 提供的第二种运行方式是利用 MLproject 将示例作为 MLflow Project 执行:
mlflow run . -P inverse_method=qrMLproject 中的入口定义如下:
name: statsmodels-example python_env: python_env.yaml entry_points: main: parameters: inverse_method: {type: str, default: 'pinv'} command: | python train.py \ --inverse-method={inverse_method}与直接运行相比,Project 方式有两处差异:
- 参数以
-P形式传入,且使用下划线风格inverse_method,与命令行脚本中的连字符风格--inverse-method由入口定义自动映射; - 环境自动管理:MLflow 会读取
python_env.yaml创建(或复用)Python 环境并安装mlflow、statsmodels、scikit-learn,再执行命令,因此无需预先手动装好全部依赖。
用 MLflow UI 查看实验对比
无论以哪种方式运行,都可以启动 MLflow 追踪服务器查看实验:
mlflow server随后在浏览器中打开默认地址,即可看到每次运行自动记录的:
- 超参数:
fit的入参(如inverse_method、method等),由 autolog 的log_fn_args_as_params机制自动写入; - 指标:autolog 记录的统计指标集合,以及示例代码手动记录的
mse; - 模型工件:以
statsmodelsflavor 保存的训练模型。
源码视角:mlflow.statsmodelsautolog 究竟记录了什么
示例的核心是 train.py 中这一行:
mlflow.statsmodels.autolog()其实现位于 mlflow/statsmodels/init.py,从源码可以确认,调用autolog()后每次fit会自动记录三类内容:
1. 白名单统计指标。源码中的_autolog_metric_allowlist(见 mlflow/statsmodels/init.py)列出了一组回归诊断指标,凡是拟合结果对象上存在且为数值的都会被记录,包括:aic、bic、rsquared、rsquared_adj、fvalue、f_pvalue、ssr、ess、mse_model、mse_resid、mse_total、df_model、df_resid、llf、scale、condition_number、centered_tss、uncentered_tss等共 18 项。测试 tests/statsmodels/test_statsmodels_autolog.py 验证了记录指标集合与白名单完全一致;若个别指标求值抛异常,autolog 会记录一条Failed to autolog metrics警告而不会中断训练。
2. 训练好的模型。每次fit结束后模型会被自动以mlflow.statsmodels.log_model记录为模型工件,同时附带model_summary.txt(由model.summary().as_text()生成)的文本摘要工件。
3. 运行管理。autolog 会自动创建/管理 MLflow run,并支持registered_model_name自动注册模型版本。
示例中mse是手动通过mlflow.log_metrics({"mse": mse})记录的——它不在 autolog 白名单内,需要自行用sklearn.metrics.mean_squared_error计算并记录,这正好演示了「autolog 自动记录 + 手动补充业务指标」的典型组合写法。
进阶:模型的保存、加载与 pyfunc 推理
示例止步于训练与记录,但mlflow.statsmodelsflavor(模块文档见 mlflow/statsmodels/init.py)还提供完整的模型生命周期 API,可直接沿用到你的生产流程:
- 保存:
mlflow.statsmodels.save_model(model, path)将模型序列化为model.statsmodels文件,同时生成MLmodel元数据与requirements.txt、conda.yaml等环境文件(见 mlflow/statsmodels/init.py)。remove_data=True可在保存前清空长度为nobs的原始数据数组以缩小体积——当模型工件超过 100 MB 时,autolog 会提示改用remove_data=True手动记录以降低存储开销; - 加载:
mlflow.statsmodels.load_model("runs:/<run_id>/model")支持本地路径、s3://、runs:/等 URI 直接加载回 statsmodels 的Results对象; - pyfunc 推理:flavor 同时注册了
mlflow.pyfunc加载入口,可通过通用 pyfunc 接口做批量推理与部署。需要说明的是,该 flavor 依赖 pickle 反序列化,因此默认受安全限制:若未处于受信任的 Databricks 环境中,加载时会要求显式设置环境变量MLFLOW_ALLOW_PICKLE_DESERIALIZATION=true才允许反序列化(见 mlflow/statsmodels/init.py)。
可验证性与测试佐证
本示例的正确性由仓库测试体系背书,相关测试文件可直接对照学习:
- tests/statsmodels/model_fixtures.py 提供 13 种 statsmodels 模型夹具(OLS、GLS、WLS、GLM、ARIMA、GEE 等),
ols_model()与示例共享相同的合成数据生成逻辑; - tests/statsmodels/test_statsmodels_autolog.py 覆盖 autolog 的指标白名单、参数记录、摘要工件、模型注册、异常恢复(
test_statsmodels_autolog_works_after_exception)等行为; - tests/statsmodels/test_statsmodels_model_export.py 验证
save_model/load_model的保存加载往返一致性。
快速上手小结
- 进入示例目录后,先以默认方式运行一次:
python train.py(等价于--inverse-method pinv); - 再用 QR 分解运行一次:
python train.py --inverse-method qr; - 或以 Project 方式运行:
mlflow run . -P inverse_method=qr; - 启动
mlflow server打开 UI,对比两次运行的rsquared、aic、mse等指标,直观体会求解策略对拟合结果的影响; - 需要复用模型时,用
mlflow.statsmodels.load_model加载,或借助 pyfunc 接口接入部署流程。
通过这个示例,你可以将「statsmodels 建模 + MLflow 跟踪 + 自动日志」的组合模式直接迁移到自己的回归、时间序列或广义线性模型项目中。
【免费下载链接】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),仅供参考