端到端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()
错误根源总结
- 构造函数参数名与调用时的关键字参数不匹配
- 构造函数未正确使用传入的配置对象
- 静态方法缺少
@staticmethod装饰器 - 调用了类中不存在的
train()方法,实际应调用initiate_model_training()
内容的提问来源于stack exchange,提问作者Md. Ehsanul Haque Kanan
相关产品推荐
相关产品推荐

