自定义Transformer构建的Spark ML Pipeline模型加载推断失败求助
问题分析与解决方案
核心原因
Spark ML的模型加载流程中,会先调用自定义Transformer的无参构造方法创建实例,再通过MLReadable的逻辑恢复持久化的状态。如果你的CategoricalCleaner的__init__方法定义了必填参数,加载时就会因缺少参数触发报错。
解决方案要点
- 为自定义Transformer添加无参构造方法,将原有必填参数设为可选(赋予默认值)
- 正确实现
MLWritable的write()方法,持久化拟合后的状态(比如清洗规则、映射表等)和配置参数 - 正确实现
MLReadable的read()和defaultReaders()方法,读取并恢复实例状态
可运行示例代码
自定义CategoricalCleaner实现
from pyspark.ml import Estimator, Transformer from pyspark.ml.param.shared import HasInputCol, HasOutputCol from pyspark.ml.util import MLWritable, MLReadable, DefaultParamsWriter, DefaultParamsReader from pyspark.sql import DataFrame from pyspark.sql.functions import col, when class CategoricalCleaner(Estimator, Transformer, HasInputCol, HasOutputCol, MLWritable, MLReadable): # 无参构造方法,所有参数设为可选默认值 def __init__(self, inputCol=None, outputCol=None, replace_map=None): super().__init__() self.replace_map = replace_map if replace_map is not None else {} # 初始化父类参数 if inputCol is not None: self.setInputCol(inputCol) if outputCol is not None: self.setOutputCol(outputCol) # Estimator的拟合逻辑:生成清洗映射表 def _fit(self, dataset: DataFrame) -> Transformer: # 示例逻辑:统计高频类别,将低频替换为"OTHER" freq_df = dataset.groupBy(self.getInputCol()).count() total = dataset.count() threshold = 0.01 # 可通过HyperOpt调优的参数 valid_cats = freq_df.filter(col("count") / total > threshold).select(self.getInputCol()).rdd.flatMap(lambda x: x).collect() replace_map = {cat: cat for cat in valid_cats} # 低频类别映射为"OTHER" all_cats = dataset.select(self.getInputCol()).distinct().rdd.flatMap(lambda x: x).collect() for cat in all_cats: if cat not in replace_map: replace_map[cat] = "OTHER" # 返回拟合后的Transformer实例 return CategoricalCleaner( inputCol=self.getInputCol(), outputCol=self.getOutputCol(), replace_map=replace_map ) # Transformer的转换逻辑 def _transform(self, dataset: DataFrame) -> DataFrame: input_col = self.getInputCol() output_col = self.getOutputCol() # 构建when条件链 cond = when(col(input_col).isin(self.replace_map.keys()), col(input_col)) for old_val, new_val in self.replace_map.items(): if new_val != old_val: cond = cond.when(col(input_col) == old_val, new_val) cond = cond.otherwise("OTHER") return dataset.withColumn(output_col, cond) # MLWritable实现:持久化参数和状态 def write(self): return DefaultParamsWriter(self) # MLReadable实现:读取并恢复实例 @classmethod def read(cls): return DefaultParamsReader(cls) @classmethod def defaultReaders(cls): return {"CategoricalCleaner": cls}
模型训练、保存与加载测试
from pyspark.sql import SparkSession from pyspark.ml import Pipeline import mlflow import mlflow.spark # 初始化SparkSession spark = SparkSession.builder.appName("CustomTransformerTest").getOrCreate() # 生成测试数据 data = spark.createDataFrame([ ("A",), ("A",), ("B",), ("C",), ("D",), ("D",), ("D",), ("E",) ], ["category"]) # 构建Pipeline cleaner = CategoricalCleaner(inputCol="category", outputCol="cleaned_category") pipeline = Pipeline(stages=[cleaner]) # 训练模型 model = pipeline.fit(data) # 本地保存模型(MLflow或PipelineModel.write都可) model.write().overwrite().save("./categorical_cleaner_pipeline") # MLflow保存示例 with mlflow.start_run(): mlflow.spark.log_model(model, "categorical_cleaner_model") # 加载模型并推断 loaded_model = Pipeline.load("./categorical_cleaner_pipeline") test_data = spark.createDataFrame([("A",), ("F",), ("B",)], ["category"]) loaded_model.transform(test_data).show() # MLflow加载示例(替换<run_id>为实际运行ID) loaded_mlflow_model = mlflow.spark.load_model("./mlruns/0/<run_id>/artifacts/categorical_cleaner_model") loaded_mlflow_model.transform(test_data).show()
关键注意事项
- 所有需要持久化的状态(比如
replace_map)必须在__init__中定义,且会被DefaultParamsWriter自动序列化 - 不要在
__init__中加入依赖外部未持久化数据的逻辑,加载时这些数据还未恢复 - 生产环境建议通过继承
Params定义自定义参数,而非直接使用实例属性,更符合Spark ML的参数管理规范
内容的提问来源于stack exchange,提问作者Sumit
相关产品推荐
相关产品推荐

