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

如何将Sklearn GridSearchCV运行警告保存至DataFrame?

解决Sklearn GridSearchCV运行警告捕获并存入DataFrame的问题

需求说明

使用Sklearn的GridSearchCV针对不同数据集优化AdaBoost分类器参数,需将数据集名称、best_params_、best_score_以及运行过程中产生的各类警告(如ConvergenceWarning、弃用警告等)存入同一个DataFrame,且无需通过文件读写中转实现。

问题分析

原代码未对警告做捕获处理,导致DataFrame的warning列始终为NA。示例中使用AdaBoostClassifier(base_estimator=RandomForestClassifier())会触发弃用警告(Sklearn 1.2+版本中base_estimator已被弃用,需改用estimator参数),但这些警告未被收集。

解决方案

利用Python标准库warnings的上下文管理器,在每个数据集的模型训练过程中捕获所有警告,将警告信息整理后存入DataFrame对应列。

修改后的完整代码

from sklearn.model_selection import GridSearchCV, StratifiedKFold
from sklearn.ensemble import AdaBoostClassifier, RandomForestClassifier
import numpy as np
import tqdm as tq
import pandas as pd
from sklearn.preprocessing import StandardScaler
import warnings

# 初始化结果DataFrame,增加数据集名称列
df_params = pd.DataFrame(columns=['dataset_name', 'learning_rate', 'n_estimators', 'accuracy', 'warning'])
# 修正参数名:Sklearn 1.2+中base_estimator已弃用,改用estimator(若需测试弃用警告可改回base_estimator)
abc = AdaBoostClassifier(estimator=RandomForestClassifier())

parameters = {'n_estimators':[5,10],
              'learning_rate':[0.01,0.2]}

# 定义数据集及对应名称
datasets = [
    ('dataset_a', np.random.random((50, 3))),
    ('dataset_b', np.random.random((70, 3))),
    ('dataset_c', np.random.random((50, 5)))
]

for i, (name, data) in tq.tqdm(enumerate(datasets)):
    X = data
    sc = StandardScaler()
    X = sc.fit_transform(X)
    y = ['foo', 'bar'] * int(len(X)/2)
    
    skf = StratifiedKFold(n_splits=5, shuffle=True, random_state=None)
    clf = GridSearchCV(abc, parameters, cv=skf, scoring='accuracy', n_jobs=-1)
    
    # 捕获当前数据集训练过程中的所有警告
    warning_messages = []
    def warning_handler(message, category, filename, lineno, file=None, line=None):
        warning_messages.append(f"{category.__name__}: {message}")
    
    with warnings.catch_warnings(record=True) as w:
        warnings.simplefilter("always")  # 强制记录所有警告,不做过滤
        warnings.showwarning = warning_handler  # 替换默认警告处理器,收集警告信息
        clf.fit(X, y)
    
    # 整理结果字典
    dict_best_params = clf.best_params_.copy()
    dict_best_params['dataset_name'] = name
    dict_best_params['accuracy'] = clf.best_score_
    dict_best_params['warning'] = '\n'.join(warning_messages) if warning_messages else None
    
    # 合并到结果DataFrame
    best_params = pd.DataFrame(dict_best_params, index=[i])
    df_params = pd.concat([df_params, best_params], ignore_index=True)

print(df_params.head())

关键代码解释

  1. 导入warnings模块:用于捕获和处理运行时警告。
  2. 自定义警告处理器:warning_handler函数将每个警告的类别和信息存入列表warning_messages。
  3. 上下文管理器捕获警告:
    • warnings.catch_warnings(record=True):开启警告记录模式
    • warnings.simplefilter("always"):确保所有类型的警告都被记录,不会被过滤
    • 替换warnings.showwarning为自定义处理器,收集所有警告信息
  4. 结果整理:将收集到的警告用换行符拼接成字符串,若无警告则设为None,存入DataFrame的warning列。
  5. 参数适配:将base_estimator改为estimator适配Sklearn新版本,若需测试弃用警告可改回原参数名。

内容的提问来源于stack exchange,提问作者Bradley Sutliff

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.19 21:15:35