You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在MLflow中记录含KerasClassifier的Sklearn Pipeline并解决pickle RLock报错

问题根因

TypeError: can't pickle _thread.RLock objects报错的核心原因是:

KerasClassifier内部持有TensorFlow运行时生成的线程锁资源,这类非Python原生可序列化对象无法被MLflow默认用于sklearn模型序列化的pickle/cloudpickle工具处理。不接入MLflow时不需要序列化整个Pipeline实例,因此可以正常运行。
你看到的steps截断警告是MLflow自动日志在记录pipeline参数时的截断行为,和报错无直接关联。

解决方案
  • 方案1:关闭MLflow sklearn自动日志的自动模型记录功能,手动分开记录预处理组件和Keras模型
    首先修改自动日志配置,关掉自动存模型逻辑避免自动序列化整个pipeline:

    mlflow.sklearn.autolog(log_models=False)
    

    训练完成后分开记录两个组件:

    # 单独保存预处理组件
    preprocessor = clf.named_steps["preprocessor"]
    mlflow.sklearn.log_model(sk_model=preprocessor, artifact_path="preprocessor")
    
    # 单独保存Keras模型
    keras_model = clf.named_steps["estimator"].model
    mlflow.keras.log_model(keras_model, artifact_path="keras_model", conda_env=conda_env)
    

    如果需要后续一键调用推理,可以自定义MLflow pyfunc包装类把预处理和Keras模型的调用逻辑封装在一起。

  • 方案2:调整序列化配置
    如果你必须存储完整的Pipeline对象,可以在调用mlflow.sklearn.log_model时切换序列化工具,部分场景下可以解决锁对象序列化问题:

    mlflow.sklearn.log_model(
        sk_model=clf, 
        artifact_path="model", 
        signature=signature, 
        conda_env=conda_env,
        # 优先尝试cloudpickle,不生效可替换为SERIALIZATION_FORMAT_JOBLIB
        serialization_format=mlflow.sklearn.SERIALIZATION_FORMAT_CLOUDPICKLE
    )
    
  • 方案3:替换废弃的KerasClassifier实现
    你当前使用的from tensorflow.keras.wrappers.scikit_learn import KerasClassifier是TensorFlow 2.x中已废弃的接口,存在原生序列化缺陷。可以替换为scikeras库提供的KerasClassifier实现,该实现针对sklearn pipeline和序列化场景做了专门适配,不存在RLock对象无法序列化的问题。
    替换后仅需修改导入和初始化逻辑即可:

    # 替换导入
    from scikeras.wrappers import KerasClassifier
    # 初始化时参数直接透传给模型创建函数
    classfier = KerasClassifier(model=create_model, **search_space)
    
验证方法

修复后可以先手动执行序列化操作验证问题是否解决,无报错再调用MLflow存模型即可:

import pickle
pickle.dumps(clf)

内容的提问来源于stack exchange,提问作者mas

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.10.05 03:36:04