news 2026/9/12 8:16:35

MLflow Keras 3 Flavor 完整指南:autolog 自动追踪、模型保存与加载实战

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
MLflow Keras 3 Flavor 完整指南:autolog 自动追踪、模型保存与加载实战

MLflow Keras 3 Flavor 完整指南:autolog 自动追踪、模型保存与加载实战

【免费下载链接】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

本文以 MLflow 仓库中mlflow.kerasPython API 参考文档(docs/api_reference/source/python_api/mlflow.keras.rst)为核心骨架,系统讲解 Keras 3 模型的自动日志记录(autolog)、回调(MlflowCallback)、模型保存(save_model/log_model)与加载(load_model)四大能力。读完本文,你将掌握在 MLflow 中端到端管理 Keras 3 训练实验、注册模型版本并完成推理部署的完整实战方法,并理解其底层实现原理。

一、mlflow.keras模块概览

在 MLflow 中,Keras flavor 由四个子模块组成,分别对应 API 参考中的四个automodule块:

子模块相对路径职责
autologmlflow/keras/autologging.py一行启用 Keras 训练的自动追踪
callbackmlflow/keras/callback.py提供MlflowCallback,手动将指标写入 MLflow
loadmlflow/keras/load.py从 MLflow 加载已保存的 Keras 模型
savemlflow/keras/save.py将 Keras 模型保存/记录到 MLflow

mlflow.keras的入口文件 mlflow/keras/init.py 会根据安装的 Keras 版本自动选择实现路径:

  • keras.__version__主版本小于 3时,mlflow.keras.autologload_modellog_modelsave_model会被重定向到mlflow.tensorflowflavor,以保证旧版本模型的向后兼容加载(对应_load_pyfunc的重定向);
  • Keras 3安装时,才使用本模块独立的autologgingcallbackloadsave实现,并额外暴露MlflowCallbackget_default_pip_requirementsget_default_conda_env等接口,同时保留MLflowCallback作为MlflowCallback的向后兼容别名。

因此,本文所有内容均以Keras 3为前提;如果你仍在使用旧版 Keras(又称 tf-keras),请参考mlflow.tensorflowflavor。

二、一行代码启用自动追踪:mlflow.keras.autolog()

autolog()是 Keras 3 集成中最常用的入口,其核心机制是替换keras.Model.fit方法为 MLflow 提供的定制版本,从而在训练过程中自动记录指标、参数、数据集信息与模型本身。从源码看,这一替换通过safe_patch("keras", keras.Model, "fit", _patched_inference, manage_run=True, ...)实现(见 autologging.py),manage_run=True意味着若当前没有活动的 run,MLflow 会自动创建。

2.1 完整参数说明

参数默认值说明
log_every_epochTrue每个 epoch 结束时记录训练指标
log_every_n_stepsNone若设置,则每n个训练步记录一次指标;log_every_epoch=True时必须为None
log_modelsTruemodel.fit()结束时自动将 Keras 模型记录到 MLflow
log_model_signaturesTrue自动捕获并记录模型签名(输入/输出的张量 shape 与 dtype)
save_exported_modelFalse若为True保存为导出格式(编译后的计算图,适合部署);否则保存为.keras格式(含架构与权重)
log_datasetsTrue记录数据集元数据
log_input_examplesFalse是否记录输入示例
disableFalse若为True,禁用 Keras autologging
exclusiveFalse若为True,自动记录的内容不会写入用户创建的 fluent run
disable_for_unsupported_versionsFalse对未测试/不兼容的 Keras 版本禁用 autologging
silentFalse抑制 autologging 期间 MLflow 的事件日志与警告
registered_model_nameNone设置后,每次训练完成会把模型注册为该名称的新版本(不存在时自动创建)
save_model_kwargsNone透传给keras.Model.save()的额外 kwargs
extra_tagsNone为 autologging 自动创建的每个 run 附加的标签字典

2.2 最小实战示例

import keras import mlflow import numpy as np mlflow.keras.autolog() # 准备一个 2 分类的模拟数据 data = np.random.uniform([8, 28, 28, 3]) label = np.random.randint(2, size=8) model = keras.Sequential([ keras.Input([28, 28, 3]), keras.layers.Flatten(), keras.layers.Dense(2), ]) model.compile( loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True), optimizer=keras.optimizers.Adam(0.001), metrics=[keras.metrics.SparseCategoricalAccuracy()], ) with mlflow.start_run() as run: model.fit(data, label, batch_size=4, epochs=2)

以上代码来自 autologging.py 的官方示例。autolog 会在训练开始前自动推断并记录batch_size参数(见_infer_batch_sizekeras_fit_kwargsx/batch_size的解析逻辑),并记录除selfxycallbacksvalidation_dataverbose之外的所有fit参数。

2.3 自动完成的工作与底层原理

调用autolog()后,_patched_inference(autologging.py)会在每次fit时依次执行:

  1. 记录超参数log_fn_args_as_paramsfit的 kwargs 记录为 run 参数;若设置了batch_size或能从数据集中推断出,则额外记录batch_size参数;
  2. 记录数据集:当log_datasets=True时,通过_log_dataset将 numpy 数组、TensorFlowtf.data.Datasettf.Tensor(x, y)元组数据记录为train/eval数据集(分别由CodeDatasetSource提供来源上下文,详见 autologging.py);validation_data会被记录为eval上下文;
  3. 注入回调:自动向callbacks列表追加一个MlflowCallback,用于按 epoch 或按 step 记录指标(_check_existing_mlflow_callback会检测并拒绝在 autolog 开启时显式再添加MlflowCallback,避免重复记录);
  4. 训练后记录模型fit结束后,若log_models=True,调用_log_keras_model记录模型,此时会通过get_model_signature(mlflow/keras/utils.py)自动推断模型签名——将model.input_shape/model.output_shape中的None维度替换为-1(代表动态 batch 维),并转换为TensorSpec构成ModelSignature

2.4 版本与兼容性约束

需要特别注意的是,autologging仅支持 Keras 3;使用更低版本(tf-keras)时应改用mlflow.tensorflowflavor。autolog 与 Keras 支持的所有后端(TensorFlow、PyTorch、JAX)兼容,但只对model.fit()流程生效——如果你使用自定义训练循环,必须退回到手动日志记录(见下文回调方式)。从 tests/keras/test_autolog.py 的test_custom_autolog_behavior可以看到,save_exported_model=True的测试在非 TensorFlow 后端会被跳过,这印证了导出格式依赖 TensorFlow 的事实。

三、手动追踪训练过程:MlflowCallback

MlflowCallback(mlflow/keras/callback.py)继承自keras.callbacks.Callback,是面向自定义训练流程(如关闭 autolog、自定义回调列表、自定义训练循环)时的手动记录方案。它将模型的优化器参数、架构摘要与训练指标写入当前 MLflow run。

3.1 参数与校验规则

mlflow.keras.MlflowCallback(log_every_epoch=True, log_every_n_steps=None, model_id=None)

构造函数内置了两条严格校验(见 callback.py):

  • log_every_epoch=True时,log_every_n_steps必须为None,否则抛出ValueError
  • log_every_epoch=False时,必须显式指定log_every_n_steps

3.2 四个生命周期钩子

钩子触发时机记录内容
on_train_begin训练开始时将优化器配置写入参数(形如optimizer_learning_rateoptimizer_weight_decay等,key 前缀为optimizer_);将模型架构摘要写入工件文件model_summary.txt(通过log_text
on_epoch_end每个 epoch 结束时log_every_epoch=True,以step=epoch记录该 epoch 的指标
on_batch_end每个 batch 结束时若设置了log_every_n_steps,当optimizer.iterationsn的整数倍时记录指标
on_test_end验证结束时将验证指标以validation_前缀记录(如validation_lossvalidation_sparse_categorical_accuracy

3.3 手动使用示例

import keras import mlflow import numpy as np data = np.random.uniform([8, 28, 28, 3]) label = np.random.randint(2, size=8) model = keras.Sequential([ keras.Input([28, 28, 3]), keras.layers.Flatten(), keras.layers.Dense(2), ]) model.compile( loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True), optimizer=keras.optimizers.Adam(0.001), metrics=[keras.metrics.SparseCategoricalAccuracy()], ) with mlflow.start_run() as run: model.fit( data, label, batch_size=4, epochs=2, callbacks=[mlflow.keras.MlflowCallback()], )

上述示例来自 callback.py 官方文档。tests/keras/test_callback.py 的test_keras_mlflow_callback_log_every_n_steps验证了按步记录时,记录的指标数量应等于optimizer.iterations // log_every_n_steps,说明 step 记录基于优化器迭代计数实现。

四、保存与记录模型:save_model/log_model

4.1save_model:保存到本地文件系统

save_model(model, path, ...)(mlflow/keras/save.py)将 Keras 模型连同签名、conda 环境等元数据保存到本地路径。它在磁盘上生成如下结构:

path/ ├── MLmodel # flavor 元数据(keras 版本、后端、data 路径等) ├── conda.yaml # 默认 conda 环境 ├── python_env.yaml # Python 环境 ├── requirements.txt # pip 依赖 ├── constraints.txt # 约束文件(仅在有约束时生成) └── data/ ├── model.keras # 模型文件(或 model/ 导出目录) └── keras_module.txt # 记录 keras 模块名

主要参数:

参数默认值说明
model必填keras.Model实例
path必填本地保存路径
save_exported_modelFalseTrue保存为导出格式(编译图,适合 serving);False保存为.keras格式
conda_envNoneconda 环境配置
mlflow_modelNone现有的mlflow.models.Model配置对象,为空则新建
signatureNone模型签名ModelSignature
input_exampleNone输入示例
pip_requirementsNonepip 依赖列表(覆盖自动推断)
extra_pip_requirementsNone额外附加的 pip 依赖(与自动推断合并)
save_model_kwargsNone透传给keras.Model.save的 kwargs
metadataNone自定义元数据字典,写入 MLmodel 文件

签名校验是保存流程的重要一环(save.py):若签名缺失会输出警告;若提供签名,则要求输入 schema 至少包含一个字段、所有字段必须是TensorSpec类型、且每个输入的第一维必须为-1(动态 batch 维),否则抛出INVALID_PARAMETER_VALUE错误。

保存格式细节:默认情况下模型以.keras后缀保存(model_path = data/model + ".keras");若目标路径以/dbfs/开头(Databricks 文件系统,其 FUSE 实现不支持随机写入),会先保存到临时文件再shutil.copy2拷贝,以规避写入错误。当save_exported_model=True时,则走_export_keras_model(save.py):它要求签名非空、必须安装 TensorFlow,并通过keras.export.ExportArchivemodel.call包装为名为serve的端点导出。

环境推断:默认 pip 依赖至少包含当前版本的 keras(get_default_pip_requirements返回[_get_pinned_requirement("keras")],save.py);随后通过infer_pip_requirements扫描模型代码推断附加依赖,与默认依赖取并集后写出requirements.txt/conda.yaml/python_env.yaml

import keras import mlflow model = keras.Sequential([ keras.Input([28, 28, 3]), keras.layers.Flatten(), keras.layers.Dense(2), ]) with mlflow.start_run() as run: mlflow.keras.save_model(model, "./model")

4.2log_model:记录到 MLflow 并可选注册

log_model(model, artifact_path=None, ...)save_model的云端版本,底层调用Model.log(flavor=mlflow.keras, ...)(save.py),将模型作为 run 的 artifact 记录到 MLflow 跟踪服务器,并支持:

  • registered_model_name:设置后在模型记录完成后自动创建/注册模型版本(模型不存在时自动创建);
  • await_registration_for:等待模型版本进入READY状态的秒数,默认DEFAULT_AWAIT_MAX_SLEEP_SECONDS(5 分钟),设为0None跳过等待;
  • name/params/tags/model_type/step/model_id:与 MLflow 新式模型记录 API 对齐的进阶参数,artifact_path已标记为 Deprecated(用name替代)。
import keras import mlflow model = keras.Sequential([ keras.Input([28, 28, 3]), keras.layers.Flatten(), keras.layers.Dense(2), ]) with mlflow.start_run() as run: mlflow.keras.log_model(model, name="model")

代码来自 save.py。tests/keras/test_save.py 的test_keras_save_model_exporttest_keras_save_model_non_export分别覆盖了save_exported_model=TrueFalse两条保存路径的加载验证。

五、加载模型并部署:load_model与 PyFunc 集成

5.1load_model:加载为 Keras 模型

load_model(model_uri, dst_path=None, custom_objects=None, load_model_kwargs=None)(mlflow/keras/load.py)支持丰富的 URI 形式:

  • 本地路径:/Users/me/path/to/local/modelrelative/path/to/local/model
  • 对象存储:s3://my_bucket/path/to/model
  • 运行内 artifact:runs:/<mlflow_run_id>/run-relative/path/to/model
  • 模型注册表:models:/<model_name>/<model_version>models:/<model_name>/<stage>

加载流程为:先通过_download_artifact_from_uri下载 artifact,再读取MLmodel文件中的kerasflavor 信息,最后根据save_exported_model标志决定加载方式(load.py):

  • 导出格式:要求安装 TensorFlow,通过tf.saved_model.load加载为可 serving 的计算图;
  • .keras格式:通过keras.saving.load_model加载,支持透传custom_objects(自定义层/激活函数)与load_model_kwargs
import keras import mlflow import numpy as np model = keras.Sequential([ keras.Input([28, 28, 3]), keras.layers.Flatten(), keras.layers.Dense(2), ]) with mlflow.start_run() as run: mlflow.keras.log_model(model) model_url = f"runs:/{run.info.run_id}/model" loaded_model = mlflow.keras.load_model(model_url) # 验证加载后的模型与原始模型输出一致 test_input = np.random.uniform(size=[2, 28, 28, 3]) np.testing.assert_allclose( keras.ops.convert_to_numpy(model(test_input)), loaded_model.predict(test_input), )

5.2 PyFunc 推理:_load_pyfuncKerasModelWrapper

mlflow.keras同时注册了 PyFunc loader(loader_module="mlflow.keras"),因此模型可以统一通过mlflow.pyfunc.load_model加载,并配合mlflow.models部署能力(如mlflow models serve)对外提供 REST 推理服务。

其核心是KerasModelWrapper(load.py)——一个实现了predict(data)的包装类:

  • 输入为pandas.DataFrame时,返回带原索引的DataFrame预测结果;
  • 输入支持np.ndarraylisttupledict,其他类型会抛出INVALID_PARAMETER_VALUE错误;
  • 返回结果统一通过keras.ops.convert_to_numpy转换为 numpy 数组,保证 serving 输出格式稳定;
  • 根据是否导出模型,内部调用model.serve(导出格式)或model.predict.keras格式),由get_model_call_method动态选择。

_load_pyfunc会依次在path/MLmodel与上级目录查找MLmodel文件以兼容不同 artifact 布局(load.py),体现了对旧版 MLflow 保存布局的向后兼容设计。

六、实践建议与注意事项

  1. 版本选择:Keras 3 请使用mlflow.keras;旧版 tf-keras 请使用mlflow.tensorflowflavor,两者 API 名称相同,便于迁移。
  2. 后端一致性:autolog 兼容 TensorFlow、PyTorch、JAX 三种后端,但save_exported_model=True的导出与加载路径依赖 TensorFlow;在非 TensorFlow 后端上训练时,建议保持默认的.keras格式。
  3. 自定义训练循环:autolog 只作用于model.fit()。若编写自定义训练循环,应手动调用log_metrics/log_params,或使用MlflowCallback结合keras的训练回调机制。
  4. autolog 与手动回调互斥:开启 autolog 后不要再向callbacks中显式添加MlflowCallback,否则会抛出异常提示(需先mlflow.keras.autolog(disable=True))。
  5. 签名规范:Keras 3 模型签名要求输入 schema 全部为TensorSpec且第一维为-1(动态 batch 维)。autolog 会自动从model.input_shape推断签名;手动保存时可先构造符合规范的ModelSignature再调用log_model
  6. 依赖与环境:每次记录模型都会生成requirements.txt/conda.yaml/python_env.yaml,默认固定 keras 版本并自动推断附加依赖;生产部署时建议基于这些文件构建运行环境,保证可复现性。

通过 autolog、MlflowCallbacksave_model/log_modelload_model的组合,你可以在 MLflow 上完成从实验追踪、指标记录、模型版本注册到 PyFunc 部署的完整 Keras 3 工作流。相关 API 细节可进一步查阅 docs/api_reference/source/python_api/mlflow.keras.rst 的自动生成文档,以及 docs/docs/classic-ml/deep-learning/keras/index.mdx 的入门指南。

【免费下载链接】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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/12 8:15:03

基于Matlab的手指手掌静脉识别实现与算法详解

简介&#xff1a;面向机器视觉课程创新实践&#xff0c;这份手指手掌静脉识别Matlab工程聚焦手部静脉图像预处理算法实验研究。项目完整覆盖静脉识别链路&#xff0c;对手指和手掌分别进行轮廓分割、感兴趣区域&#xff08;ROI&#xff09;截取、静脉纹理增强与分割&#xff0c…

作者头像 李华
网站建设 2026/9/12 8:11:56

光声峰峰值成像:MATLAB实现与参数调优指南

简介&#xff1a;光声峰峰值成像利用光吸收产生的超声信号重建组织内部光吸收分布&#xff0c;在生物医学光学成像与病变识别中具有实用价值。面向光声成像研究者、生物医学工程相关专业学生以及需要快速构建成像算法的开发者&#xff0c;这份MATLAB资源提供了一套完整的峰峰值…

作者头像 李华
网站建设 2026/9/12 8:11:12

AI模型部署实战:从训练完成到生产上线的5大关键环节

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/12 8:10:49

深入解析Node.js自动化框架OpenClaw与Nanobot架构

1. OpenClaw与Nanobot项目概述 OpenClaw是一个基于Node.js的自动化开发框架&#xff0c;而Nanobot则是其核心组件之一。这两个项目在开发者社区中近期获得了不少关注&#xff0c;特别是在自动化脚本和AI辅助编程领域。我第一次接触OpenClaw是在尝试解决一些重复性编码任务时&am…

作者头像 李华
网站建设 2026/9/12 8:09:07

真实电路中的放大器:定义、分类与工程选型实战指南

1. 这不是教科书里的“放大器”&#xff0c;而是你修电路、调音频、搭传感器时真正会碰上的那个“放大器” “放大器”这三个字&#xff0c;听起来像高中物理课本里那个画着三角形符号、标着“A_v”的抽象概念。但如果你拆过功放机、调过麦克风增益、给单片机接温度传感器、甚至…

作者头像 李华