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

端到端ML项目Model Trainer报错:TypeError缺少4个必填位置参数

问题分析与解决方案

核心问题

你调用initiate_model_training()时未传入它定义时要求的4个参数(X_train, X_test, y_train, y_test),同时类中的辅助方法未正确声明为静态方法,不符合Python类方法规范。


分步修复

1. 加载预处理数据并传入训练方法

在04_model_trainer.ipynb中,调用initiate_model_training前需先加载数据转换阶段保存的训练/测试集。假设数据以numpy格式存储在artifacts/data_transformation目录下,修改代码如下:

import numpy as np

try:
    config = ConfigurationManager()
    model_trainer_config = config.get_model_trainer_config()
    model_trainer = ModelTrainer(model_trainer_config)
    
    # 加载预处理后的训练测试数据
    train_data = np.load('artifacts/data_transformation/train.npz')
    X_train = train_data['X_train']
    y_train = train_data['y_train']
    
    test_data = np.load('artifacts/data_transformation/test.npz')
    X_test = test_data['X_test']
    y_test = test_data['y_test']
    
    # 传入参数调用训练方法
    model_trainer.initiate_model_training(X_train, X_test, y_train, y_test)
except Exception as e:
    raise e

注意:请根据你的项目实际数据存储路径调整文件路径。

2. 修正ModelTrainer类的方法定义

将save_obj和evaluate_model声明为静态方法(因为它们不依赖类实例变量),避免潜在的调用错误:

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    

    def initiate_model_training(self, X_train, X_test, y_train, y_test):
        try:
            logger.info('Splitting ')

            models={
            'LinearRegression':LinearRegression(),
            'Lasso':Lasso(),
            'Ridge':Ridge(),
            'Elasticnet':ElasticNet(),
            'RandomForestRegressor': RandomForestRegressor(),
            'GradientBoostRegressor' : GradientBoostingRegressor(),
            "AdaBoost" : AdaBoostRegressor(),
            'DecisionTreeRegressor' : DecisionTreeRegressor(),
            "SupportVectorRegressor" : SVR(),
            "KNN" : KNeighborsRegressor()
            }

            model_report:dict = ModelTrainer.evaluate_model(X_train,y_train, X_test, y_test, models)
            print(model_report)
            print("\n====================================================================================")
            logger.info(f'Model Report : {model_report}')

            best_model_score = max(sorted(model_report.values()))
            best_model_name = list(model_report.keys())[
                list(model_report.values()).index(best_model_score)
            ]
            best_model = models[best_model_name]

            print(f"Best Model Found, Model Name :{best_model_name}, R2-score: {best_model_score}")
            print("\n====================================================================================")
            logger.info(f"Best Model Found, Model name: {best_model_name}, R2-score: {best_model_score}")
            logger.info(f"{best_model.feature_names_in_}")
            
            ModelTrainer.save_obj(
                file_path = self.model_trainer_config.trained_model_file_path,
                obj = best_model
            )

        except Exception as e:
            logger.info('Exception occured at model training')
            raise e

额外修正:修复了initiate_model_training中日志消息的拼写错误(model trianing改为model training),并修正了GradientBoostRegressor()的键名多余括号问题。


内容的提问来源于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 04:10:59