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,却触发上述报错。
报错原因
- 输入参数获取错误:MLflow会将POST请求的
{"inputs": ...}解析为pandas DataFrame,原代码使用字典的get方法(model_input.get("records"))获取参数,而DataFrame无此方法,导致类型错误。 - 返回值类型不符合要求: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
相关产品推荐
相关产品推荐

