加载含自定义Transformer的foundry_ml PySpark模型时stage_transform未注册报错
解决Foundry加载自定义PySpark转换器的stage_transform未注册错误
问题根源
训练阶段仅在训练脚本中注册了自定义FeatureExtractor的stage_transform与序列化器,但加载模型的脚本未执行相同的注册操作。Foundry在加载模型时需要这些注册信息来识别并处理自定义组件,因此会触发No stage_transform registered异常。
解决方案1:在加载脚本中重复注册自定义组件
修改模型加载脚本,添加与训练脚本完全一致的自定义类定义、stage_transform函数、序列化/反序列化器,以及注册逻辑:
# 导入必要的序列化库 import cloudpickle # 导入foundry_ml相关模块 from foundry_ml import Model from foundry_ml.stage.flexible import stage_transform, register_stage_transform_for_class from foundry_ml.stage.serialization import deserializer, serializer, register_serializer_for_class # 导入数据处理工具 from foundry_object.utils import safe_write_data, load_data # 导入PySpark相关类 from pyspark.ml import Transformer # 导入transforms相关模块 from transforms.api import transform, Input, Output # 定义与训练脚本完全一致的FeatureExtractor类 class FeatureExtractor(Transformer): def __init__(self): super(FeatureExtractor, self).__init__() def _transform(self, df): # 保持和训练阶段相同的特征提取逻辑 return df # 定义自定义stage_transform函数 @stage_transform() def stage_transform_wrapper(model, df): return model._transform(df) # 定义反序列化器 @deserializer("serialized_feature_extractor_model.dill", force=True) def deserializer_feature_extractor(filesystem, path): return cloudpickle.loads(load_data(filesystem, path, True), encoding="latin1") # 定义序列化器 @serializer(deserializer_feature_extractor) def serializer_feature_extractor(filesystem, value): path = "serialized_feature_extractor_model.dill" safe_write_data(filesystem, path, cloudpickle.dumps(value), base64_encode=True) return path @transform( input_data=Input("INPUT_PATH"), input_model=Input("MODEL_PATH"), output_data=Output("OUTPUT_PATH"), ) def model_transform(input_data, input_model, output_data): # 加载模型前必须完成自定义组件的注册 register_stage_transform_for_class(FeatureExtractor, stage_transform_wrapper, force=True) register_serializer_for_class(FeatureExtractor, serializer_feature_extractor, force=True) df = input_data.dataframe() model = Model.load(input_model) df = model.transform(df) output_data.write_dataframe(df)
解决方案2:抽离公共模块避免重复代码
将自定义组件和注册逻辑封装到公共Python模块(如custom_transformers.py),在训练和加载脚本中统一导入调用,减少代码冗余:
公共模块custom_transformers.py
import cloudpickle from foundry_ml.stage.flexible import stage_transform, register_stage_transform_for_class from foundry_ml.stage.serialization import deserializer, serializer, register_serializer_for_class from foundry_object.utils import safe_write_data, load_data from pyspark.ml import Transformer class FeatureExtractor(Transformer): def __init__(self): super(FeatureExtractor, self).__init__() def _transform(self, df): # 特征提取逻辑 return df @stage_transform() def stage_transform_wrapper(model, df): return model._transform(df) @deserializer("serialized_feature_extractor_model.dill", force=True) def deserializer_feature_extractor(filesystem, path): return cloudpickle.loads(load_data(filesystem, path, True), encoding="latin1") @serializer(deserializer_feature_extractor) def serializer_feature_extractor(filesystem, value): path = "serialized_feature_extractor_model.dill" safe_write_data(filesystem, path, cloudpickle.dumps(value), base64_encode=True) return path def register_custom_components(): """统一注册自定义组件的辅助函数""" register_stage_transform_for_class(FeatureExtractor, stage_transform_wrapper, force=True) register_serializer_for_class(FeatureExtractor, serializer_feature_extractor, force=True)
训练脚本中调用
from custom_transformers import FeatureExtractor, register_custom_components @transform( output_model=Output("MODEL_PATH"), input_data=Input("INPUT_PATH"), ) def model_training(input_data, output_model): register_custom_components() # 后续训练逻辑(实例化组件、构建Pipeline等)
加载脚本中调用
from custom_transformers import register_custom_components from foundry_ml import Model from transforms.api import transform, Input, Output @transform( input_data=Input("INPUT_PATH"), input_model=Input("MODEL_PATH"), output_data=Output("OUTPUT_PATH"), ) def model_transform(input_data, input_model, output_data): register_custom_components() df = input_data.dataframe() model = Model.load(input_model) df = model.transform(df) output_data.write_dataframe(df)
关键注意事项
- 自定义类的**完全限定名(包路径+类名)**必须在训练和加载环境中完全一致,否则Foundry无法匹配注册的transform逻辑。
- 注册时使用
force=True参数,确保覆盖可能存在的旧注册,避免冲突。 - 序列化与反序列化使用的cloudpickle版本需保持一致,防止反序列化失败。
内容的提问来源于stack exchange,提问作者CtrlAltDefeat
相关产品推荐
相关产品推荐

