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

端到端ML项目Model Trainer报错:TypeError: __init__()收到意外参数'config'

问题解决步骤

1. 修正ModelTrainer类的构造函数

你的ModelTrainer构造函数接收的参数是model_trainer_config,但调用时用了关键字参数config=,导致参数名不匹配;同时构造函数内部没有使用传入的配置,反而重新实例化了ModelTrainerConfig(),这两个问题一起引发了报错。修改构造函数:

class ModelTrainer:
    def __init__(self, model_trainer_config):
        # 使用传入的配置对象,而非重新创建
        self.model_trainer_config = model_trainer_config

2. 修正调用代码的参数传递

调用ModelTrainer时,去掉不匹配的关键字参数名,同时注意类中不存在train()方法,需调用实际定义的initiate_model_training():

try:
    config = ConfigurationManager()
    model_trainer_config = config.get_model_trainer_config()
    # 直接传入参数,无需指定关键字config=
    model_trainer = ModelTrainer(model_trainer_config)
    # 需先加载训练测试数据(需补充从artifact读取数据的逻辑)
    model_trainer.initiate_model_training(X_train, X_test, y_train, y_test)
except Exception as e:
    raise e

3. 修复类中静态方法的装饰器缺失问题

save_obj和evaluate_model是静态方法,但未添加@staticmethod装饰器,直接调用会引发错误,补充装饰器:

class ModelTrainer:
    def __init__(self, model_trainer_config):
        self.model_trainer_config = model_trainer_config

    @staticmethod
    def save_obj(file_path, obj):
        try:
            dir_path = os.path.dirname(file_path)
            os.makedirs(dir_path, exist_ok=True)
            with open(file_path, 'wb') as file_obj:
                joblib.dump(obj, file_obj, compress=('gzip'))
        except Exception as e:
            logger.info('Error occured in utils save_obj')
            raise e

    @staticmethod
    def evaluate_model(X_train, y_train, X_test, y_test, models):
        try:
            report = {}
            for i in range(len(models)):
                model = list(models.values())[i]
                model.fit(X_train,y_train)
                y_test_pred = model.predict(X_test)
                test_model_score = r2_score(y_test,y_test_pred)
                report[list(models.keys())[i]] = test_model_score
            return report
        except Exception as e:
            logger.info('Exception occured during model training')
            raise e

    # 其余方法保持不变

4. 补充训练数据加载逻辑

initiate_model_training需要传入X_train、X_test、y_train、y_test,需从之前的数据处理产物中读取,示例代码:

# 在调用initiate_model_training前添加
import pandas as pd

# 假设数据保存在data_transformation的artifact目录下
X_train = pd.read_csv('artifacts/data_transformation/train_features.csv')
y_train = pd.read_csv('artifacts/data_transformation/train_target.csv').values.ravel()
X_test = pd.read_csv('artifacts/data_transformation/test_features.csv')
y_test = pd.read_csv('artifacts/data_transformation/test_target.csv').values.ravel()

错误根源总结

  1. 构造函数参数名与调用时的关键字参数不匹配
  2. 构造函数未正确使用传入的配置对象
  3. 静态方法缺少@staticmethod装饰器
  4. 调用了类中不存在的train()方法,实际应调用initiate_model_training()

内容的提问来源于stack exchange,提问作者Md. Ehsanul Haque Kanan

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.04 13:07:08