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

对比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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.05 02:03:02