MLflow 官方示例精读:用 statsmodels 训练 OLS 模型并完成自动日志记录
2026/9/12 16:00:35 网站建设 项目流程

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项目运行时依赖:mlflowstatsmodelsscikit-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控制的是求解最小二乘问题时逆矩阵的计算方式,可选值为qrpinv(默认),其参数定义见 train.py:

取值默认值求解原理
pinv使用 Moore-Penrose 伪逆(np.linalg.pinv)求解最小二乘问题
qr使用 QR 分解(np.linalg.qr)求解

两种方法的取舍要点:

  • pinv(伪逆法):数值上更稳健,尤其适合矩阵接近奇异(病态)的情况,因为它基于奇异值分解,能自动处理秩亏矩阵;代价是计算开销略高;
  • qr(QR 分解法):计算效率更高、内存占用更少,适合数据规模较大且矩阵良态的场景;但在处理近乎奇异的矩阵时数值稳定性不如伪逆。

README 明确建议读者两种方法都试一遍,甚至可以省略--inverse-method参数(此时自动回落到默认值pinv)。这正是 MLflow 实验跟踪的典型使用场景:通过多次运行对比不同求解策略下模型的指标表现。由于超参数与指标都会被自动记录,你可以在 MLflow UI 中直接横向对比qrpinv两批运行的结果。

以 MLflow Project 方式运行

README 提供的第二种运行方式是利用 MLproject 将示例作为 MLflow Project 执行:

mlflow run . -P inverse_method=qr

MLproject 中的入口定义如下:

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 方式有两处差异:

  1. 参数以-P形式传入,且使用下划线风格inverse_method,与命令行脚本中的连字符风格--inverse-method由入口定义自动映射;
  2. 环境自动管理:MLflow 会读取python_env.yaml创建(或复用)Python 环境并安装mlflowstatsmodelsscikit-learn,再执行命令,因此无需预先手动装好全部依赖。

用 MLflow UI 查看实验对比

无论以哪种方式运行,都可以启动 MLflow 追踪服务器查看实验:

mlflow server

随后在浏览器中打开默认地址,即可看到每次运行自动记录的:

  • 超参数fit的入参(如inverse_methodmethod等),由 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)列出了一组回归诊断指标,凡是拟合结果对象上存在且为数值的都会被记录,包括:aicbicrsquaredrsquared_adjfvaluef_pvaluessressmse_modelmse_residmse_totaldf_modeldf_residllfscalecondition_numbercentered_tssuncentered_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.txtconda.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的保存加载往返一致性。

快速上手小结

  1. 进入示例目录后,先以默认方式运行一次:python train.py(等价于--inverse-method pinv);
  2. 再用 QR 分解运行一次:python train.py --inverse-method qr
  3. 或以 Project 方式运行:mlflow run . -P inverse_method=qr
  4. 启动mlflow server打开 UI,对比两次运行的rsquaredaicmse等指标,直观体会求解策略对拟合结果的影响;
  5. 需要复用模型时,用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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询