如何将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())
关键代码解释
- 导入warnings模块:用于捕获和处理运行时警告。
- 自定义警告处理器:
warning_handler函数将每个警告的类别和信息存入列表warning_messages。 - 上下文管理器捕获警告:
warnings.catch_warnings(record=True):开启警告记录模式warnings.simplefilter("always"):确保所有类型的警告都被记录,不会被过滤- 替换
warnings.showwarning为自定义处理器,收集所有警告信息
- 结果整理:将收集到的警告用换行符拼接成字符串,若无警告则设为
None,存入DataFrame的warning列。 - 参数适配:将
base_estimator改为estimator适配Sklearn新版本,若需测试弃用警告可改回原参数名。
内容的提问来源于stack exchange,提问作者Bradley Sutliff
相关产品推荐
相关产品推荐

