自定义网格搜索超参数:多进程运行函数及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
相关产品推荐
相关产品推荐

