ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

MLflow pmdarima 模型风味完整指南:ARIMA 时间序列模型的保存、记录与 PyFunc 推理

MLflow pmdarima 模型风味完整指南:ARIMA 时间序列模型的保存、记录与 PyFunc 推理 MLflow pmdarima 模型风味完整指南ARIMA 时间序列模型的保存、记录与 PyFunc 推理【免费下载链接】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/mlflowMLflow 的mlflow.pmdarima模型风味flavor为基于pmdarima库训练的单变量 ARIMA 时间序列模型提供了标准的保存、记录与加载接口。本文以 mlflow.pmdarima API 文档 为核心结合 风味实现源码、官方示例 与 模型导出测试完整讲解save_model/log_model/load_model三大核心 API、PyFunc 推理的单行 DataFrame 配置约定以及签名、示例、依赖环境等工程细节帮助你在 MLflow 中落地可复现、可服务化的 ARIMA 预测流程。一、模块定位两种导出格式与一个弃用声明mlflow.pmdarima模块专为单变量univariatepmdarima 模型提供记录与加载能力模型导出为以下两种格式见 模块文档字符串Pmdarima 原生格式使用 pickle 序列化的pmdarima模型实例以model.pmd文件形式存储mlflow.pyfunc格式由该风味自动附加到生成的 MLflow Model 上用于通用的 pyfunc 部署工具与历史预测的批量审计。因此任何通过mlflow.pmdarima.save_model/log_model产生的模型其MLmodel配置中都同时包含pmdarima与python_function两种 flavor——这一点由源码中 pyfunc.add_to_model 与 mlflow_model.add_flavor 的连续调用 保证测试 test_pmdarima_log_model 也断言了pyfunc.FLAVOR_NAME in model_config.flavors。重要前提从源码init.py 第 114-118 行 可以看到该风味在导入时会抛出FutureWarningwarnings.warn( pmdarima flavor is deprecated and will be removed in a future release, FutureWarning, stacklevel2, )即pmdarima 风味已标记为弃用并将在未来版本中移除。在当前仓库版本中它仍可正常工作但新项目在选用时需评估这一维护状态。二、核心 API 全景save_model / log_model / load_model模块公开的三个核心函数另有辅助的get_default_pip_requirements、get_default_conda_env承担了模型生命周期的三个环节。2.1 save_model保存到本地路径save_model将一个已完成拟合的 pmdarimaARIMA模型或Pipeline对象以 pickle 格式保存到本地文件系统源码 L142-L298。其完整参数如下参数类型/默认值说明pmdarima_modelARIMA / Pipeline已fit于时间序列的模型对象pathstr本地目标目录序列化模型model.pmd存放于此conda_envdict/str环境描述可为 conda 环境 dict 或指向 yaml 文件的路径code_pathslist[str]应打包进模型的自定义代码文件路径mlflow_modelModel附加此 flavor 的mlflow.models.Model实例signatureModelSignature输入输出签名不传但提供input_example时自动推断设为False可禁用推断input_example任意输入示例用于签名推断与示例保存pip_requirementsstr/list[str]覆盖默认 pip 依赖的 requirementsextra_pip_requirementsstr/list[str]在默认依赖基础上追加的依赖metadatadict附加到模型上的自定义元数据extra_fileslist[str]额外需要复制进模型目录的文件文档中给出的最小可用示例源码 docstring L196-L221import pandas as pd import mlflow import pmdarima # 指定源数据与模型产物路径 SOURCE_DATA https://raw.githubusercontent.com/facebook/prophet/master/examples/example_retail_sales.csv ARTIFACT_PATH model # 读取数据并重命名字段 sales_data pd.read_csv(SOURCE_DATA) sales_data.rename(columns{y: sales, ds: date}, inplaceTrue) # 划分训练/测试 train_size int(0.8 * len(sales_data)) train sales_data[:train_size] test sales_data[train_size:] with mlflow.start_run(): # 创建模型 model pmdarima.auto_arima(train[sales], seasonalTrue, m12) # 保存到指定路径 mlflow.pmdarima.save_model(model, model)从实现上看save_model内部执行的关键步骤源码 L226-L297包括校验环境参数与保存路径、复制 code paths、保存输入示例、在缺少签名时基于示例推断签名、将模型 pickle 到model.pmd、写入MLmodel、并生成conda.yaml/requirements.txt/python_env.yaml三个环境文件。2.2 log_model记录到当前 runlog_model将 pmdarimaARIMA或Pipeline对象作为当前 run 的 artifact 记录下来源码 L300-L439返回包含模型元数据的ModelInfo实例。它继承了save_model的全部参数并额外支持artifact_path已弃用改用name模型在 run 中的 artifact 路径registered_model_name若指定则同时创建/注册一个模型版本await_registration_for控制等待注册完成的秒数默认 5 分钟传0或None跳过等待name、params、tags、model_type、step、model_id等模型记录参数。带签名与指标记录的标准用法源码 docstring L380-L416import pandas as pd import mlflow from mlflow.models import infer_signature import pmdarima from pmdarima.metrics import smape SOURCE_DATA https://raw.githubusercontent.com/facebook/prophet/master/examples/example_retail_sales.csv ARTIFACT_PATH model sales_data pd.read_csv(SOURCE_DATA) sales_data.rename(columns{y: sales, ds: date}, inplaceTrue) train_size int(0.8 * len(sales_data)) train sales_data[:train_size] test sales_data[train_size:] with mlflow.start_run(): model pmdarima.auto_arima(train[sales], seasonalTrue, m12) # 计算指标 prediction model.predict(n_periodslen(test)) metrics {smape: smape(test[sales], prediction)} # 推断签名 input_sample pd.DataFrame(train[sales]) output_sample pd.DataFrame(model.predict(n_periods5)) signature infer_signature(input_sample, output_sample) # 记录模型 mlflow.pmdarima.log_model(model, nameARTIFACT_PATH, signaturesignature)仓库还提供了可直接运行的 examples/pmdarima/train.py它使用load_wineind数据集训练季节 ARIMAm12记录 aicc/aic/bic/hqic/oob 指标并用RollingForecastCV做滚动回测交叉验证得到 smape/MAE/MSE最后log_modelload_model输出未来 30 期预测。对应的 MLproject 将python train.py声明为入口点可配合mlflow run执行。2.3 load_model加载原生模型load_model从本地文件或任意 run/artifact URI 加载 pmdarimaARIMA模型或Pipeline对象源码 L442-L529。支持的model_uri形式包括本地路径/Users/me/path/to/local/model、相对路径、s3://my_bucket/path/to/model、runs:/mlflow_run_id/run-relative/path/to/model、mlflow-artifacts:/path/to/model。dst_path参数可指定模型产物的本地下载目录须已存在缺省时自动创建。model_info mlflow.pmdarima.log_model( model, nameARTIFACT_PATH, signaturesignature, input_exampleinput_sample.head() ) loaded_model mlflow.pmdarima.load_model(model_info.model_uri) forecast loaded_model.predict(n_periods60) print(fforecast: {forecast})示例输出源码 docstring L512-L519forecast: 234 382452.397246 235 380639.458720 236 359805.611219 ...加载路径同样支持远程 URI——测试 test_pmdarima_load_from_remote_uri_succeeds 验证了从 S3 artifact repository 加载模型的完整流程。三、PyFunc 推理约定单行 DataFrame 配置通过mlflow.pyfunc.load_model加载 pmdarima 模型后其推理接口使用单行single-rowPandas DataFrame配置参数来声明预测需求源码见 _PmdarimaModelWrapper.predict。支持的配置列如下同时见 官方模型文档列名必填说明n_periods是要生成的未来期数从训练集最后一个时间点开始频率沿用训练序列的频率。例如训练数据每小时一个值要预测 3 天则设为72X否外生回归变量exogenous regressors2D 数组仅在 pmdarima 1.8.0 时支持return_conf_int否是否返回置信区间布尔值默认Falsealpha否计算置信区间的显著性水平默认0.05一个典型的 pyfunc 预测配置示例官方文档 L1028-L1046import pmdarima import mlflow import pandas as pd data pmdarima.datasets.load_airpassengers() with mlflow.start_run(): model pmdarima.auto_arima(data, seasonalTrue) mlflow.pmdarima.save_model(model, /tmp/model.pmd) loaded_pyfunc mlflow.pyfunc.load_model(/tmp/model.pmd) prediction_conf pd.DataFrame([{n_periods: 4, return_conf_int: True, alpha: 0.1}]) predictions loaded_pyfunc.predict(prediction_conf)输出为包含yhat、yhat_lower、yhat_upper三列的 DataFrame示例数值Indexyhatyhat_loweryhat_upper0467.573731423.30995511.837511490.494467416.17449564.814442509.138684420.56255597.711173492.554714397.30634587.80309关于输入输出格式有三个从源码与文档中确认的重要约定输入必须恰好为 1 行_PmdarimaModelWrapper.predict在len(dataframe) 1时直接抛出MlflowException错误码INVALID_PARAMETER_VALUE缺失n_periods列或n_periods非整数同样报错源码 L578-L603。输出列随return_conf_int变化为False/None时输出单列[yhat]为True时输出[yhat, yhat_lower, yhat_upper]源码 L640-L649。外生变量的版本门槛若传入X而安装的 pmdarima 版本低于 1.8.0会发出警告且X不会生效源码 L605-L614。测试 test_pmdarima_autoarima_pyfunc_save_and_load 对上述行为做了精确验证以{n_periods: 60, return_conf_int: True, alpha: 0.1}配置调用 pyfunc断言yhat等于原生predict的第一返回值、yhat_lower/yhat_upper与置信区间元组拆解结果一致。四、签名与输入示例注意置信区间模式提供input_example而未显式传signature时save_model/log_model会用示例自动推断签名显式传signatureFalse可关闭推断源码 L236-L240。警告若在原生ARIMA.predict()中以return_conf_intTrue生成置信区间签名将无法正确推断——其元组返回类型不是被识别的签名类型。只有通过 pyfunc 风味推理时infer_signature才能正常工作模块文档 与 官方模型文档 L1087-L1091 均有说明。对应测试见 test_pmdarima_signature_and_example_for_confidence_interval_mode该测试先加载 pyfunc、用{n_periods: 10, return_conf_int: True, alpha: 0.2}配置得到预测再对配置与预测结果推断签名并保存验证。手动推断签名的推荐方式save_model docstring L177-L183from mlflow.models import infer_signature model pmdarima.auto_arima(data) predictions model.predict(n_periods30, return_conf_intFalse) signature infer_signature(data, predictions)五、依赖环境管理pickle 反序列化的安全门槛模型目录会生成conda.yaml、requirements.txt、python_env.yaml三类环境文件默认依赖来自 get_default_pip_requirements即钉定版本的pmdarima并会额外推断模型目录中的 pip 依赖做并集源码 L274-L297。可用conda_env、pip_requirements、extra_pip_requirements三种方式覆盖或追加测试 test_pmdarima_log_model_with_pip_requirements 与 test_pmdarima_log_model_with_extra_pip_requirements 覆盖了文件路径、依赖列表、constraints 文件三种形式。由于原生格式依赖 pickle源码_save_model/_load_modelL532-L549反序列化存在安全风险MLflow 默认禁止pickle 反序列化除非满足以下任一条件环境变量MLFLOW_ALLOW_PICKLE_DESERIALIZATION设为true当前处于 Databricks Runtime 或 Databricks Model Serving 环境。否则load_model会抛出MlflowException源码 L537-L549对应测试 test_load_model_disallows_pickle_deserialization 验证了该行为。这意味着只应加载来自可信来源的 pmdarima 模型。六、更贴近实战的示例完整记录与回测流程examples/pmdarima/train.py 给出了一个比文档示例更完整的实战模板核心流程为加载pmdarima.datasets.load_wineind()数据按train_size150划分训练/测试用auto_arima(..., seasonalTrue, m12)拟合季节模型记录模型参数arima.get_params(deepTrue)与 AIC/AICc/BIC/HQIC/OOB 指标用model_selection.RollingForecastCV(h10, step20, initial60)做滚动回测计算 smape、MAE、MSEmlflow.pmdarima.log_model(pmdarima_modelarima, namemodel, signaturesignature)记录模型mlflow.pmdarima.load_model(model_info.model_uri)重新加载并输出未来 30 期预测。# 关键片段完整代码见 examples/pmdarima/train.py arima auto_arima(train, error_actionignore, traceFalse, suppress_warningsTrue, maxiter5, seasonalTrue, m12) metrics {x: getattr(arima, x)() for x in [aicc, aic, bic, hqic, oob]} cross_validator model_selection.RollingForecastCV(h10, step20, initial60) for x in [smape, mean_absolute_error, mean_squared_error]: metrics[x] calculate_cv_metrics(arima, data, x, cross_validator) model_info mlflow.pmdarima.log_model(pmdarima_modelarima, nameARTIFACT_PATH, signaturesignature) mlflow.log_params(parameters) mlflow.log_metrics(metrics)该示例可用于mlflow run examples/pmdarima直接复现是理解“训练 → 指标 → 记录 → 加载 → 预测”完整闭环的最佳起点。七、快速参考与注意事项兼容对象已拟合的 pmdarimaARIMA或Pipelineauto_arima返回值即属此类。模型二进制model.pmdpickleflavor 名称为pmdarima相关常量见 源码 L106-L109。PyFunc 输入单行 DataFrame至少含n_periods列可选X、return_conf_int、alpha。PyFunc 输出return_conf_intFalse时为[yhat]True时为[yhat, yhat_lower, yhat_upper]。签名注意return_conf_intTrue时原生 API 无法推断签名请用 pyfunc 路径或关闭置信区间。安全pickle 反序列化默认被禁用加载前需设置MLFLOW_ALLOW_PICKLE_DESERIALIZATIONtrue或处于 Databricks 环境并只加载可信模型。维护状态该风味已声明弃用FutureWarning将在未来版本移除请评估项目长期维护计划。延伸阅读mlflow.pmdarima API 参考风味源码实现 mlflow/pmdarima/init.py模型导出测试 tests/pmdarima/test_pmdarima_model_export.py官方示例 examples/pmdarima/train.py模型文档中的 Pmdarima 小节ml-package-versions.yml 中 pmdarima 的测试配置【免费下载链接】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),仅供参考
返回列表