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

使用mlflow.pyfunc保存/加载模型时传递额外artifacts及predict报错问题

问题:自定义MLflow PyFunc模型保存加载后predict参数报错

问题背景

此前在Stack Overflow咨询过mlflow.pyfunc.PythonModel与mlflow.pyfunc.PyFuncModel的区别,得到了明确解答。现在遇到扩展问题:自定义了一个带有fit和predict方法的类,fit计算并返回参数字典,predict依赖该字典。本地直接运行时(按之前方案传入None到predict)正常,但将模型保存后重新加载,调用m2.predict(None, df, d)时报错:'predict() takes 2 positional arguments but 4 were given'。

错误原因

加载后的m2是PyFuncModel实例,它的predict方法签名和你自定义的PythonModel的predict方法完全不同:

  • 自定义PythonModel的predict你定义为predict(self, context, X, d, y=None),支持多参数
  • 但PyFuncModel的predict方法仅接受最多2个位置参数:第一个是输入数据,第二个是可选的params字典(用于传递额外参数)。当你调用m2.predict(None, df, d)时,三个位置参数会被全部传入PyFuncModel的predict,超出其参数数量限制,因此触发报错。

另外,MLflow保存PythonModel时,只会序列化实例的属性状态,不会保留外部传递的参数。你之前将fit返回的d作为外部参数传给predict,加载后的模型实例无法直接获取这个外部变量,这也是问题根源之一。

解决方案

将fit得到的参数字典d存储为模型实例的属性,让predict直接使用实例内部的参数,而不是依赖外部传递。这样模型保存时,参数会被序列化到模型文件中,加载后可直接调用predict。

修改后的代码示例

1. 调整自定义模型类

import pandas as pd
import mlflow.pyfunc

# 测试数据
data = {'col1': [1, 2], 'col2': [3, 4]}
df = pd.DataFrame(data=data)

# 修改后的模型类
class PredictSpeciality(mlflow.pyfunc.PythonModel):
    def fit(self):
        print('fit')
        self.d = {'mult': 2}  # 将参数保存为实例属性
    
    def predict(self, context, X, y=None):  # 移除外部传入的d参数
        print('predict')
        X['pred'] = X['col1'] * self.d['mult']
        return X

2. 本地运行与保存加载

# 本地运行
m = PredictSpeciality()
m.fit()
m.predict(None, df)

# 保存模型
mlflow.pyfunc.save_model(path="temp_model", python_model=m)

# 加载模型并调用predict
m2 = mlflow.pyfunc.load_model("temp_model")
m2.predict(df)  # 仅需传入输入数据,无需额外参数

3. 若需动态传递额外参数的处理

如果确实需要在predict时动态传递参数(而非依赖fit时的固定参数),可以利用PyFuncModel.predict的params参数:

# 修改predict方法支持params参数
class PredictSpeciality(mlflow.pyfunc.PythonModel):
    def fit(self):
        print('fit')
        self.default_d = {'mult': 2}
    
    def predict(self, context, X, params=None, y=None):
        print('predict')
        # 使用传入的params,若无则用默认值
        d = params if params is not None else self.default_d
        X['pred'] = X['col1'] * d['mult']
        return X

# 加载后调用时传入params
m2 = mlflow.pyfunc.load_model("temp_model")
m2.predict(df, params={'mult': 3})  # 动态传递参数

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.15 00:35:25