如何将GridSearchCV的参数值拆分后与评分导出为CSV文件
解决网格搜索参数CSV拆分问题
我懂你现在的困扰——网格搜索结果里的参数字典被一股脑塞进CSV的同一列,查看和分析都特别麻烦。咱们只需要把字典里的每个参数单独拆成一列就搞定啦!
方法一:针对固定参数的直接提取
如果你的param_grid里的参数是固定的(比如就batch_size和epochs),可以直接从每个参数字典里提取对应的值,和评分结果一起写入不同列:
import csv # 假设你的参数网格定义 param_grid=dict(batch_size=[40, 60, 80], epochs=[1000, 2000]) grid = GridSearchCV(estimator=model, param_grid=param_grid, n_jobs=-1, cv=3) grid_result = grid.fit(x_train, y_train, validation_data=(x_test, y_test), callbacks=[es]) # 提取结果数据 means = grid_result.cv_results_['mean_test_score'] stds = grid_result.cv_results_['std_test_score'] params = grid_result.cv_results_['params'] # 注意变量名统一,避免覆盖 exportfile='/Users/test.csv' with open(exportfile, 'w', newline='') as file: writer = csv.writer(file) # 先写入清晰的表头 writer.writerow(['mean_test_score', 'std_test_score', 'batch_size', 'epochs']) # 遍历每一组结果,拆分参数写入 for mean, stdev, param in zip(means, stds, params): batch_size_val = param['batch_size'] epochs_val = param['epochs'] writer.writerow([mean, stdev, batch_size_val, epochs_val])
这样生成的CSV里,每个参数都会单独占一列,再也不会是挤在一起的字典格式了。
方法二:动态适配任意参数网格(推荐)
如果以后你可能会给param_grid添加更多参数(比如learning_rate、dropout_rate),可以用动态方法自动适配,不用每次修改代码:
import csv # 示例:带更多参数的网格 param_grid=dict(batch_size=[40, 60], epochs=[1000, 2000], learning_rate=[0.001, 0.01]) grid = GridSearchCV(estimator=model, param_grid=param_grid, n_jobs=-1, cv=3) grid_result = grid.fit(x_train, y_train, validation_data=(x_test, y_test), callbacks=[es]) means = grid_result.cv_results_['mean_test_score'] stds = grid_result.cv_results_['std_test_score'] params = grid_result.cv_results_['params'] exportfile='/Users/test.csv' with open(exportfile, 'w', newline='') as file: writer = csv.writer(file) # 动态生成表头:评分项 + 所有参数名 header = ['mean_test_score', 'std_test_score'] + list(param_grid.keys()) writer.writerow(header) # 按表头顺序提取每个参数的值,组合成一行写入 for mean, stdev, param in zip(means, stds, params): param_values = [param[key] for key in param_grid.keys()] writer.writerow([mean, stdev] + param_values)
这个方法会自动读取param_grid里的所有参数名作为表头,然后按顺序提取每个参数的值,不管你以后加多少参数,代码都不用改,非常灵活。
小提醒
注意你原代码里的变量名问题:原代码里写的param = grid_result.cv_results_['params'],后面循环又用param作为变量,这会导致变量覆盖,容易出错,所以我改成了params来存储所有参数列表。
内容的提问来源于stack exchange,提问作者D.Luipers
相关产品推荐
相关产品推荐

