对比ML模型性能的函数如何自动选择最优模型
修正后完整可运行代码
from sklearn.model_selection import train_test_split from sklearn.metrics import accuracy_score # 可选,用于更规范的准确率计算 def compare_models(X,y,model1,model2,test_size=0.15,val_size=0.15,random_state=0): # 首先拆分数据集为训练集和测试集,使用传入的参数而非硬编码 X_train_full,X_test,y_train_full,y_test = train_test_split(X, y, random_state=random_state, test_size=test_size) # 再将训练集拆分为训练子集和验证集 X_train,X_val,y_train,y_val = train_test_split(X_train_full,y_train_full,random_state=random_state, test_size=val_size) # 使用验证集对比两个模型的性能 model1.fit(X_train,y_train) val_preds_model1 = model1.predict(X_val) model2.fit(X_train,y_train) val_preds_model2 = model2.predict(X_val) # 计算每个模型的验证准确率 acc_val_model1 = accuracy_score(y_val, val_preds_model1) acc_val_model2 = accuracy_score(y_val, val_preds_model2) # 核心逻辑:选择验证准确率更高的模型作为最优模型 if acc_val_model1 > acc_val_model2: best_model = model1 print(f"选中模型1,验证准确率为{acc_val_model1:.4f},高于模型2的{acc_val_model2:.4f}") elif acc_val_model2 > acc_val_model1: best_model = model2 print(f"选中模型2,验证准确率为{acc_val_model2:.4f},高于模型1的{acc_val_model1:.4f}") else: # 平局时默认选择模型1,可根据需求调整 best_model = model1 print(f"两个模型验证准确率相同,均为{acc_val_model1:.4f},默认选择模型1") # 将选中的最优模型在训练+验证合并集上重新训练 best_model.fit(X_train_full,y_train_full) # 在测试集上评估模型性能 preds_test = best_model.predict(X_test) acc_test = accuracy_score(y_test, preds_test) # 可根据需要调整返回内容,比如新增返回最优模型、验证准确率等 return { "test_accuracy": acc_test, "best_model": best_model, "model1_val_accuracy": acc_val_model1, "model2_val_accuracy": acc_val_model2, "test_predictions": preds_test }
核心修改说明
- 补全了最优模型选择逻辑:通过两个模型的验证准确率比较,动态赋值
best_model变量替换原代码中的XXX,同时兼容准确率平局的边界情况 - 修复了原代码参数未生效问题:原函数定义的
test_size/val_size/random_state参数未实际使用,全部替换为传入参数,同时设置默认值保证调用兼容性 - 优化了准确率计算逻辑:使用
sklearn.metrics.accuracy_score替代手写计算逻辑,更规范且不易出错,如果你希望保留手写逻辑也可直接替换回来 - 丰富了返回结果:除了测试准确率外,还返回最优模型实例、两个模型的验证准确率、测试集预测结果,方便后续调用和分析
调用示例
from sklearn.ensemble import RandomForestClassifier from sklearn.svm import SVC # 示例用的两个待对比模型 rf = RandomForestClassifier() svc = SVC() # 假设你已经预处理好特征矩阵X和标签y result = compare_models(X, y, model1=rf, model2=svc, test_size=0.15, val_size=0.15, random_state=42) # 打印测试集准确率 print(f"最优模型测试集准确率:{result['test_accuracy']:.4f}")
内容的提问来源于stack exchange,提问作者diesmiling
相关产品推荐
相关产品推荐

