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手动将指标写入 MLflowloadmlflow/keras/load.py从 MLflow 加载已保存的 Keras 模型savemlflow/keras/save.py将 Keras 模型保存/记录到 MLflowmlflow.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_runTrue, ...)实现见 autologging.pymanage_runTrue意味着若当前没有活动的 runMLflow 会自动创建。2.1 完整参数说明参数默认值说明log_every_epochTrue每个 epoch 结束时记录训练指标log_every_n_stepsNone若设置则每n个训练步记录一次指标当log_every_epochTrue时必须为Nonelog_modelsTruemodel.fit()结束时自动将 Keras 模型记录到 MLflowlog_model_signaturesTrue自动捕获并记录模型签名输入/输出的张量 shape 与 dtypesave_exported_modelFalse若为True保存为导出格式编译后的计算图适合部署否则保存为.keras格式含架构与权重log_datasetsTrue记录数据集元数据log_input_examplesFalse是否记录输入示例disableFalse若为True禁用 Keras autologgingexclusiveFalse若为True自动记录的内容不会写入用户创建的 fluent rundisable_for_unsupported_versionsFalse对未测试/不兼容的 Keras 版本禁用 autologgingsilentFalse抑制 autologging 期间 MLflow 的事件日志与警告registered_model_nameNone设置后每次训练完成会把模型注册为该名称的新版本不存在时自动创建save_model_kwargsNone透传给keras.Model.save()的额外 kwargsextra_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, size8) model keras.Sequential([ keras.Input([28, 28, 3]), keras.layers.Flatten(), keras.layers.Dense(2), ]) model.compile( losskeras.losses.SparseCategoricalCrossentropy(from_logitsTrue), optimizerkeras.optimizers.Adam(0.001), metrics[keras.metrics.SparseCategoricalAccuracy()], ) with mlflow.start_run() as run: model.fit(data, label, batch_size4, epochs2)以上代码来自 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_inferenceautologging.py会在每次fit时依次执行记录超参数log_fn_args_as_params将fit的 kwargs 记录为 run 参数若设置了batch_size或能从数据集中推断出则额外记录batch_size参数记录数据集当log_datasetsTrue时通过_log_dataset将 numpy 数组、TensorFlowtf.data.Dataset、tf.Tensor或(x, y)元组数据记录为train/eval数据集分别由CodeDatasetSource提供来源上下文详见 autologging.pyvalidation_data会被记录为eval上下文注入回调自动向callbacks列表追加一个MlflowCallback用于按 epoch 或按 step 记录指标_check_existing_mlflow_callback会检测并拒绝在 autolog 开启时显式再添加MlflowCallback避免重复记录训练后记录模型fit结束后若log_modelsTrue调用_log_keras_model记录模型此时会通过get_model_signaturemlflow/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_modelTrue的测试在非 TensorFlow 后端会被跳过这印证了导出格式依赖 TensorFlow 的事实。三、手动追踪训练过程MlflowCallbackMlflowCallbackmlflow/keras/callback.py继承自keras.callbacks.Callback是面向自定义训练流程如关闭 autolog、自定义回调列表、自定义训练循环时的手动记录方案。它将模型的优化器参数、架构摘要与训练指标写入当前 MLflow run。3.1 参数与校验规则mlflow.keras.MlflowCallback(log_every_epochTrue, log_every_n_stepsNone, model_idNone)构造函数内置了两条严格校验见 callback.pylog_every_epochTrue时log_every_n_steps必须为None否则抛出ValueErrorlog_every_epochFalse时必须显式指定log_every_n_steps。3.2 四个生命周期钩子钩子触发时机记录内容on_train_begin训练开始时将优化器配置写入参数形如optimizer_learning_rate、optimizer_weight_decay等key 前缀为optimizer_将模型架构摘要写入工件文件model_summary.txt通过log_texton_epoch_end每个 epoch 结束时若log_every_epochTrue以stepepoch记录该 epoch 的指标on_batch_end每个 batch 结束时若设置了log_every_n_steps当optimizer.iterations为n的整数倍时记录指标on_test_end验证结束时将验证指标以validation_前缀记录如validation_loss、validation_sparse_categorical_accuracy3.3 手动使用示例import keras import mlflow import numpy as np data np.random.uniform([8, 28, 28, 3]) label np.random.randint(2, size8) model keras.Sequential([ keras.Input([28, 28, 3]), keras.layers.Flatten(), keras.layers.Dense(2), ]) model.compile( losskeras.losses.SparseCategoricalCrossentropy(from_logitsTrue), optimizerkeras.optimizers.Adam(0.001), metrics[keras.metrics.SparseCategoricalAccuracy()], ) with mlflow.start_run() as run: model.fit( data, label, batch_size4, epochs2, 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_model4.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保存为导出格式编译图适合 servingFalse保存为.keras格式conda_envNoneconda 环境配置mlflow_modelNone现有的mlflow.models.Model配置对象为空则新建signatureNone模型签名ModelSignatureinput_exampleNone输入示例pip_requirementsNonepip 依赖列表覆盖自动推断extra_pip_requirementsNone额外附加的 pip 依赖与自动推断合并save_model_kwargsNone透传给keras.Model.save的 kwargsmetadataNone自定义元数据字典写入 MLmodel 文件签名校验是保存流程的重要一环save.py若签名缺失会输出警告若提供签名则要求输入 schema 至少包含一个字段、所有字段必须是TensorSpec类型、且每个输入的第一维必须为-1动态 batch 维否则抛出INVALID_PARAMETER_VALUE错误。保存格式细节默认情况下模型以.keras后缀保存model_path data/model .keras若目标路径以/dbfs/开头Databricks 文件系统其 FUSE 实现不支持随机写入会先保存到临时文件再shutil.copy2拷贝以规避写入错误。当save_exported_modelTrue时则走_export_keras_modelsave.py它要求签名非空、必须安装 TensorFlow并通过keras.export.ExportArchive将model.call包装为名为serve的端点导出。环境推断默认 pip 依赖至少包含当前版本的 kerasget_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_pathNone, ...)是save_model的云端版本底层调用Model.log(flavormlflow.keras, ...)save.py将模型作为 run 的 artifact 记录到 MLflow 跟踪服务器并支持registered_model_name设置后在模型记录完成后自动创建/注册模型版本模型不存在时自动创建await_registration_for等待模型版本进入READY状态的秒数默认DEFAULT_AWAIT_MAX_SLEEP_SECONDS5 分钟设为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, namemodel)代码来自 save.py。tests/keras/test_save.py 的test_keras_save_model_export与test_keras_save_model_non_export分别覆盖了save_exported_modelTrue与False两条保存路径的加载验证。五、加载模型并部署load_model与 PyFunc 集成5.1load_model加载为 Keras 模型load_model(model_uri, dst_pathNone, custom_objectsNone, load_model_kwargsNone)mlflow/keras/load.py支持丰富的 URI 形式本地路径/Users/me/path/to/local/model、relative/path/to/local/model对象存储s3://my_bucket/path/to/model运行内 artifactruns:/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 fruns:/{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与KerasModelWrappermlflow.keras同时注册了 PyFunc loaderloader_modulemlflow.keras因此模型可以统一通过mlflow.pyfunc.load_model加载并配合mlflow.models部署能力如mlflow models serve对外提供 REST 推理服务。其核心是KerasModelWrapperload.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_modelTrue的导出与加载路径依赖 TensorFlow在非 TensorFlow 后端上训练时建议保持默认的.keras格式。自定义训练循环autolog 只作用于model.fit()。若编写自定义训练循环应手动调用log_metrics/log_params或使用MlflowCallback结合keras的训练回调机制。autolog 与手动回调互斥开启 autolog 后不要再向callbacks中显式添加MlflowCallback否则会抛出异常提示需先mlflow.keras.autolog(disableTrue)。签名规范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),仅供参考