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

如何在MLflow运行中记录Spark ALS模型?Databricks集群问题排查

在Databricks中使用MLflow记录并加载Spark ALS模型的问题

在Databricks集群中尝试用MLflow记录Spark ALS模型时,遇到两类核心问题:

  • 调用部分MLflow记录方法时触发TypeError: cannot pickle '_thread.RLock' object,导致运行终止
  • 部分方法虽能执行完成,但加载模型时出现OSError: No such file or directory或协议不支持、输入类型错误等问题,无法正常使用模型

核心需求:成功在MLflow运行中记录ALS模型,且能正常加载模型完成后续评估与推荐任务。


测试准备代码

import mlflow
import logging
from pyspark.ml.evaluation import RegressionEvaluator
from pyspark.ml.recommendation import ALS
from pyspark import SparkContext, SparkConf

data = [{"User": 1, "Item": 1, "Rating": 1},
        {"User": 2, "Item": 2, "Rating": 3},
        {"User": 3, "Item": 3, "Rating": 1},
        {"User": 4, "Item": 2, "Rating": 4},
        {"User": 1, "Item": 2, "Rating": 3},
        {"User": 2, "Item": 3, "Rating": 2},
        {"User": 2, "Item": 4, "Rating": 1},
        {"User": 4, "Item": 1, "Rating": 5}
        ]

conf = SparkConf().setAppName("ALS-mlflow-test")
sc = SparkContext.getOrCreate(conf)
rdd = sc.parallelize(data)
df_rating = rdd.toDF()

(df_train, df_test) = df_rating.randomSplit([0.8, 0.2])

logging.getLogger("mlflow").setLevel(logging.DEBUG)

已尝试方法及错误分析

方法1:使用mlflow.sklearn.log_model

with mlflow.start_run() as run:    
    model_als = ALS(maxIter=5, regParam=0.01, userCol="User", itemCol="Item", ratingCol="Rating", implicitPrefs=False,
          coldStartStrategy="drop")
    
    model_als.fit(df_train)
    mlflow.sklearn.log_model(model_als, artifact_path="test")

报错:

_SklearnCustomModelPicklingError: Pickling custom sklearn model ALS failed when saving model: cannot pickle '_thread.RLock' object

原因:ALS是Spark ML模型,并非scikit-learn模型,mlflow.sklearn.log_model仅支持scikit-learn生态的模型,Spark模型包含Spark上下文相关的锁对象,无法通过sklearn的序列化机制处理。

方法2:用自定义PythonModel包装ALS模型

class MyModel(mlflow.pyfunc.PythonModel):
    def __init__(self, model):
        self.model = model
    
    def predict(self, context, model_input):
        return self.my_custom_function(model_input)

    def my_custom_function(self, model_input):
        return 0

with mlflow.start_run():
    model_als = ALS(maxIter=5, regParam=0.01, userCol="User", itemCol="Item", ratingCol="Rating", implicitPrefs=False,
          coldStartStrategy="drop")
    my_model = MyModel(model_als)
    model_info = mlflow.pyfunc.log_model(artifact_path="model", python_model=my_model)

报错:

TypeError: cannot pickle '_thread.RLock' object

原因:Spark ALS模型内部包含无法被pickle序列化的线程锁对象,直接包装后依然无法完成序列化存储。

方法3:将ALS放入Pipeline后用mlflow.spark.log_model记录

from pyspark.ml import Pipeline

with mlflow.start_run() as run:    
    model_als = ALS(maxIter=5, regParam=0.01, userCol="User", itemCol="Item", ratingCol="Rating", implicitPrefs=False,
          coldStartStrategy="drop")
    
    pipeline = Pipeline(stages=[model_als])
    pipeline_model = pipeline.fit(df_train)
    mlflow.spark.log_model(pipeline_model, artifact_path="test-pipeline")

日志错误:

stderr: Setting default log level to "WARN".
To adjust logging level use sc.setLogLevel(newLevel). For SparkR, use setLogLevel(newLevel).
2023/01/05 08:54:22 INFO mlflow.spark: File '/tmp/tmpxiznhskj/sparkml' not found on DFS. Will attempt to upload the file.
Traceback (most recent call last):
  File "/databricks/python/lib/python3.9/site-packages/mlflow/utils/_capture_modules.py", line 162, in <module>
    main()
  File "/databricks/python/lib/python3.9/site-packages/mlflow/utils/_capture_modules.py", line 137, in main
    mlflow.pyfunc.load_model(model_path)
  ...
OSError: No such file or directory: '/tmp/tmpxiznhskj/sparkml'

原因:该错误为MLflow内部临时文件处理的警告级问题,模型实际已成功记录,但后续加载方式错误导致无法使用。

方法4:用PipelineModel.load加载模型

from pyspark.ml import PipelineModel

logged_model = 'runs:/xyz123/test'

# Load model
loaded_model = PipelineModel.load(logged_model)

报错:

org.apache.hadoop.fs.UnsupportedFileSystemException: No FileSystem for scheme "runs"

原因:PipelineModel.load仅支持HDFS、本地文件系统等标准文件协议,不识别MLflow的runs:// URI格式,必须通过MLflow提供的加载方法处理。

方法5:用Databricks自动生成代码加载模型

import mlflow
logged_model = 'runs:/xyz123/test'

# Load model
loaded_model = mlflow.spark.load_model(logged_model)

# Perform inference via model.transform()
loaded_model.transform(data)

报错:

AttributeError: 'list' object has no attribute '_jdf'

原因:Spark模型的transform方法要求输入为Spark DataFrame,而非Python列表,需先将列表转换为DataFrame再传入。


解决方案

1. 正确记录ALS模型

使用mlflow.spark.log_model直接记录训练完成的ALSModel实例(而非未训练的ALS estimator),该方法专门用于Spark ML模型的序列化与存储:

import mlflow
from pyspark.ml.recommendation import ALS

with mlflow.start_run() as run:
    # 定义ALS estimator
    als_estimator = ALS(maxIter=5, regParam=0.01, userCol="User", itemCol="Item", ratingCol="Rating", 
                        implicitPrefs=False, coldStartStrategy="drop")
    # 训练得到ALSModel
    als_trained_model = als_estimator.fit(df_train)
    
    # 记录模型到MLflow
    mlflow.spark.log_model(als_trained_model, artifact_path="als-trained-model")
    
    # 可选:记录训练参数与指标
    mlflow.log_param("maxIter", 5)
    mlflow.log_param("regParam", 0.01)

2. 正确加载并使用模型

通过mlflow.spark.load_model加载模型,且确保输入为Spark DataFrame:

import mlflow
from pyspark.ml.evaluation import RegressionEvaluator

# 替换为你的MLflow Run ID
logged_model = 'runs:/<你的Run ID>/als-trained-model'

# 加载模型
loaded_als_model = mlflow.spark.load_model(logged_model)

# 模型评估
df_pred = loaded_als_model.transform(df_test)
evaluator = RegressionEvaluator(metricName="rmse", labelCol="Rating", predictionCol="prediction")
rmse = evaluator.evaluate(df_pred)

df_pred.display()
print(f"Root-mean-square error explicit = {rmse}")

# 生成用户推荐
user_recs = loaded_als_model.recommendForAllUsers(2)
user_recs.display()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 21:55:25