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

加载含自定义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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 23:44:55