端到端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
相关产品推荐
相关产品推荐

