
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),僅供參考