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块:
| 子模块 | 相对路径 | 职责 |
|---|---|---|
autolog | mlflow/keras/autologging.py | 一行启用 Keras 训练的自动追踪 |
callback | mlflow/keras/callback.py | 提供MlflowCallback,手动将指标写入 MLflow |
load | mlflow/keras/load.py | 从 MLflow 加载已保存的 Keras 模型 |
save | mlflow/keras/save.py | 将 Keras 模型保存/记录到 MLflow |
mlflow.keras的入口文件 mlflow/keras/init.py 会根据安装的 Keras 版本自动选择实现路径:
- 当
keras.__version__主版本小于 3时,mlflow.keras.autolog、load_model、log_model、save_model会被重定向到mlflow.tensorflowflavor,以保证旧版本模型的向后兼容加载(对应_load_pyfunc的重定向); - 当Keras 3安装时,才使用本模块独立的
autologging、callback、load、save实现,并额外暴露MlflowCallback、get_default_pip_requirements、get_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_epoch | True | 每个 epoch 结束时记录训练指标 |
log_every_n_steps | None | 若设置,则每n个训练步记录一次指标;当log_every_epoch=True时必须为None |
log_models | True | model.fit()结束时自动将 Keras 模型记录到 MLflow |
log_model_signatures | True | 自动捕获并记录模型签名(输入/输出的张量 shape 与 dtype) |
save_exported_model | False | 若为True保存为导出格式(编译后的计算图,适合部署);否则保存为.keras格式(含架构与权重) |
log_datasets | True | 记录数据集元数据 |
log_input_examples | False | 是否记录输入示例 |
disable | False | 若为True,禁用 Keras autologging |
exclusive | False | 若为True,自动记录的内容不会写入用户创建的 fluent run |
disable_for_unsupported_versions | False | 对未测试/不兼容的 Keras 版本禁用 autologging |
silent | False | 抑制 autologging 期间 MLflow 的事件日志与警告 |
registered_model_name | None | 设置后,每次训练完成会把模型注册为该名称的新版本(不存在时自动创建) |
save_model_kwargs | None | 透传给keras.Model.save()的额外 kwargs |
extra_tags | None | 为 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_size对keras_fit_kwargs中x/batch_size的解析逻辑),并记录除self、x、y、callbacks、validation_data、verbose之外的所有fit参数。
2.3 自动完成的工作与底层原理
调用autolog()后,_patched_inference(autologging.py)会在每次fit时依次执行:
- 记录超参数:
log_fn_args_as_params将fit的 kwargs 记录为 run 参数;若设置了batch_size或能从数据集中推断出,则额外记录batch_size参数; - 记录数据集:当
log_datasets=True时,通过_log_dataset将 numpy 数组、TensorFlowtf.data.Dataset、tf.Tensor或(x, y)元组数据记录为train/eval数据集(分别由CodeDatasetSource提供来源上下文,详见 autologging.py);validation_data会被记录为eval上下文; - 注入回调:自动向
callbacks列表追加一个MlflowCallback,用于按 epoch 或按 step 记录指标(_check_existing_mlflow_callback会检测并拒绝在 autolog 开启时显式再添加MlflowCallback,避免重复记录); - 训练后记录模型:
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_rate、optimizer_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.iterations为n的整数倍时记录指标 |
on_test_end | 验证结束时 | 将验证指标以validation_前缀记录(如validation_loss、validation_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_model | False | True保存为导出格式(编译图,适合 serving);False保存为.keras格式 |
conda_env | None | conda 环境配置 |
mlflow_model | None | 现有的mlflow.models.Model配置对象,为空则新建 |
signature | None | 模型签名ModelSignature |
input_example | None | 输入示例 |
pip_requirements | None | pip 依赖列表(覆盖自动推断) |
extra_pip_requirements | None | 额外附加的 pip 依赖(与自动推断合并) |
save_model_kwargs | None | 透传给keras.Model.save的 kwargs |
metadata | None | 自定义元数据字典,写入 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.ExportArchive将model.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 分钟),设为0或None跳过等待;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_export与test_keras_save_model_non_export分别覆盖了save_exported_model=True与False两条保存路径的加载验证。
五、加载模型并部署: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/model、relative/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_pyfunc与KerasModelWrapper
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.ndarray、list、tuple、dict,其他类型会抛出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 保存布局的向后兼容设计。
六、实践建议与注意事项
- 版本选择:Keras 3 请使用
mlflow.keras;旧版 tf-keras 请使用mlflow.tensorflowflavor,两者 API 名称相同,便于迁移。 - 后端一致性:autolog 兼容 TensorFlow、PyTorch、JAX 三种后端,但
save_exported_model=True的导出与加载路径依赖 TensorFlow;在非 TensorFlow 后端上训练时,建议保持默认的.keras格式。 - 自定义训练循环:autolog 只作用于
model.fit()。若编写自定义训练循环,应手动调用log_metrics/log_params,或使用MlflowCallback结合keras的训练回调机制。 - autolog 与手动回调互斥:开启 autolog 后不要再向
callbacks中显式添加MlflowCallback,否则会抛出异常提示(需先mlflow.keras.autolog(disable=True))。 - 签名规范:Keras 3 模型签名要求输入 schema 全部为
TensorSpec且第一维为-1(动态 batch 维)。autolog 会自动从model.input_shape推断签名;手动保存时可先构造符合规范的ModelSignature再调用log_model。 - 依赖与环境:每次记录模型都会生成
requirements.txt/conda.yaml/python_env.yaml,默认固定 keras 版本并自动推断附加依赖;生产部署时建议基于这些文件构建运行环境,保证可复现性。
通过 autolog、MlflowCallback、save_model/log_model与load_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),仅供参考