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

SDV CTGAN模型MLflow部署报错:序列化输入DataFrame不兼容

问题与解决方案

问题详情

报错信息:

模型评估出错。请验证序列化的输入DataFrame是否与模型推理兼容。

使用SDV CTGAN模型,自定义MLflow模型包装类代码如下:

class Model_Wrapper(mlflow.pyfunc.PythonModel):

    def __init__(self,):
      self.model = None

    def load_context(self,context):
        self.model=mlflow.pyfunc.load_model(context.artifacts["Original_Model"])

    def predict(self, context, model_input):
       ss = self.model.sample(int(model_input.get("records")[0]))
       return ss.to_json()

调用方式:通过POST请求调用invocations接口,输入格式为{"inputs":{"records":[2]}},输入符合MLflow规范,但期望输出为DataFrame,却触发上述报错。

报错原因

  1. 输入参数获取错误:MLflow会将POST请求的{"inputs": ...}解析为pandas DataFrame,原代码使用字典的get方法(model_input.get("records"))获取参数,而DataFrame无此方法,导致类型错误。
  2. 返回值类型不符合要求:MLflow的pyfunc模型要求predict方法返回pandas DataFrame、numpy数组、列表或字典,原代码返回JSON字符串,违反兼容性要求。

修正后的代码

class Model_Wrapper(mlflow.pyfunc.PythonModel):

    def __init__(self,):
      self.model = None

    def load_context(self,context):
        self.model=mlflow.pyfunc.load_model(context.artifacts["Original_Model"])

    def predict(self, context, model_input):
        # 从DataFrame中提取生成记录数
        num_records = int(model_input["records"].iloc[0])
        # 生成CTGAN样本
        generated_samples = self.model.sample(num_records)
        # 直接返回DataFrame,符合MLflow要求与预期输出
        return generated_samples

说明

  • 改用model_input["records"].iloc[0]从DataFrame中获取单个参数值,适配MLflow的输入解析规则。
  • 直接返回生成的DataFrame,既满足MLflow的模型兼容性要求,也符合期望输出为DataFrame的需求。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.19 19:37:14