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

自定义网格搜索超参数:多进程运行函数及CUDA报错解决

解决方案建议

1. 修复基础循环逻辑

原代码对inputs的处理存在错误,inputs是包含数千组超参数的大列表,需遍历每组参数;同时核心函数缺失train_size参数,需补充后传入:

def search_best_parameters(dataset, train_size, name_of_model, num_of_hidden_layers, normalized):
    """补充train_size参数,实现模型训练与准确率计算逻辑"""
    # 此处添加基于train_size划分数据集、构建模型、训练、评估的代码
    return accuracy

if __name__ == "__main__":
    results = []  # 用列表存储每组参数与对应准确率,便于后续筛选
    for dataset in datasets:
        inputs = get_search_space()
        for param_set in inputs:
            train_size, name, layers, normalized = param_set
            acc = search_best_parameters(dataset, train_size, name, layers, normalized)
            # 结构化存储结果
            results.append({
                "dataset": dataset,
                "params": {
                    "train_size": train_size,
                    "model_name": name,
                    "hidden_layers": layers,
                    "normalized": normalized
                },
                "accuracy": acc
            })

2. 解决CUDA多进程初始化报错

Colab默认的multiprocessing fork启动方式会导致CUDA上下文冲突,改用spawn启动方式即可避免报错,以下是两种实现方式:

方式一:使用multiprocessing的spawn上下文

import multiprocessing

def worker(args):
    # 将dataset与参数组合并为单个参数传入进程
    dataset, param_set = args
    train_size, name, layers, normalized = param_set
    acc = search_best_parameters(dataset, train_size, name, layers, normalized)
    return {
        "dataset": dataset,
        "params": param_set,
        "accuracy": acc
    }

if __name__ == "__main__":
    # 指定spawn启动方式,规避CUDA重初始化问题
    ctx = multiprocessing.get_context('spawn')
    # 根据GPU显存调整进程数,避免显存溢出(Colab建议1-2个进程)
    with ctx.Pool(processes=2) as pool:
        # 生成所有任务参数:每个数据集对应所有超参数组合
        tasks = []
        for dataset in datasets:
            for param_set in get_search_space():
                tasks.append((dataset, param_set))
        # 并行执行任务
        results = pool.map(worker, tasks)

方式二:使用concurrent.futures.ProcessPoolExecutor

from concurrent.futures import ProcessPoolExecutor

if __name__ == "__main__":
    tasks = []
    for dataset in datasets:
        for param_set in get_search_space():
            tasks.append((dataset, param_set))
    
    # 指定spawn启动方式
    with ProcessPoolExecutor(max_workers=2, mp_context=multiprocessing.get_context('spawn')) as executor:
        results = list(executor.map(worker, tasks))

3. 结果存储与最优参数筛选

用列表存储结果后,可快速筛选最优参数,也可将结果保存到文件长期留存:

# 筛选所有结果中准确率最高的参数组合
best_result = max(results, key=lambda x: x["accuracy"])
print("最优参数组合:", best_result["params"])
print("对应准确率:", best_result["accuracy"])
print("数据集:", best_result["dataset"])

# 将结果保存为JSON文件,方便后续查看分析
import json
with open("grid_search_results.json", "w") as f:
    json.dump(results, f, indent=2)

内容的提问来源于stack exchange,提问作者Satyen Singh

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.22 04:15:09