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

如何将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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.06 17:12:32