基于PyTorch与Ray的CNN超参数调优代码报错求助
PyTorch+Ray CNN超参调优代码错误修复与优化方案
一、核心错误分析与修复步骤
1. 卷积层输入通道不匹配
MNIST是单通道灰度数据集,你的Net类中conv1输入通道设为10,会直接导致输入维度不匹配:
# 错误代码 self.conv1 = nn.Conv2d(10, conv1_filter, kernel_size=conv1_kernel) # 修复后 self.conv1 = nn.Conv2d(1, conv1_filter, kernel_size=conv1_kernel)
2. 全连接层输入维度硬编码错误
原代码硬编码conv2_filter *7*7作为全连接层输入,但特征图大小会随卷积核尺寸变化,需动态计算:
# 修复后的Net类forward方法 def forward(self, x): out = self.pool(self.act(self.conv1(x))) out = self.pool(self.act(self.conv2(out))) # 动态展平特征图,自动适配维度 out = out.view(out.size(0), -1) out = self.act(self.lin(out)) out = self.result(out) return out # 同时修改__init__中的线性层初始化,先占位再动态适配 self.lin = nn.Linear(0, dense1_units)
3. 缺失核心训练/测试函数
代码中调用了train()和test()但未定义,需补充:
def train(model, optimizer, train_loader): model.train() criterion = nn.CrossEntropyLoss() for data, target in train_loader: optimizer.zero_grad() loss = criterion(model(data), target) loss.backward() optimizer.step() def test(model, test_loader): model.eval() correct = 0 total = 0 with torch.no_grad(): for data, target in test_loader: output = model(data) _, predicted = torch.max(output.data, 1) total += target.size(0) correct += (predicted == target).sum().item() return correct / total
4. 最佳模型加载参数缺失
test_best_model中实例化Net时未传入必要参数,需从最佳trial的配置中获取:
def test_best_model(results: tune.ResultGrid): best_result = results.get_best_result() best_config = best_result.config # 用最佳配置实例化模型 best_model = Net( best_config["conv1_filter"], best_config["conv1_kernel"], best_config["conv2_filter"], best_config["conv2_kernel"], best_config["dense1_units"] ) # 加载权重并测试 checkpoint_dict = best_result.checkpoint.to_dict() best_model.load_state_dict(checkpoint_dict["model"]) test_loader = get_data_loaders(best_config["batch_size"])[1] test_acc = test(best_model, test_loader) print("最佳模型准确率: ", test_acc)
5. PBT调度器参数错误
PopulationBasedTraining的hyperparam_mutations需指定可变异的超参范围,而非完整搜索空间:
# 错误代码 scheduler = PopulationBasedTraining( time_attr="training_iteration", perturbation_interval=5, hyperparam_mutations = config ) # 修复后 scheduler = PopulationBasedTraining( time_attr="training_iteration", perturbation_interval=5, hyperparam_mutations={ "learning_rate": tune.loguniform(1e-5, 1e-1), "batch_size": tune.choice([2, 4, 8, 16]) } )
6. Checkpoint存储冲突问题
原代码用固定目录my_model存储checkpoint,多trial运行时会互相覆盖,建议用Checkpoint.from_dict直接序列化数据:
# 替换原checkpoint保存逻辑 if step % 5 == 0: checkpoint = Checkpoint.from_dict({ "step": step, "model": model.state_dict() })
二、更简洁的实现方式
1. 利用Ray Tune内置API简化
- 封装数据加载器为函数,避免重复创建
- 使用内置
CombinedStopper替代自定义Stopper,同时实现多条件停止 - 省略物理文件操作,直接用
Checkpoint.from_dict存储权重
2. 核心简化片段
# 统一数据加载函数 def get_data_loaders(batch_size): train_loader = torch.utils.data.DataLoader( MNIST('data', train=True, download=True, transform=transforms.ToTensor()), batch_size=batch_size, shuffle=True ) test_loader = torch.utils.data.DataLoader( MNIST('data', train=False, download=True, transform=transforms.ToTensor()), batch_size=batch_size, shuffle=False ) return train_loader, test_loader # 简化训练函数 def training(config): train_loader, test_loader = get_data_loaders(config["batch_size"]) model = Net( config["conv1_filter"], config["conv1_kernel"], config["conv2_filter"], config["conv2_kernel"], config["dense1_units"] ) optimizer = optim.SGD(model.parameters(), lr=config["learning_rate"]) step = 0 if session.get_checkpoint(): checkpoint_dict = session.get_checkpoint().to_dict() model.load_state_dict(checkpoint_dict["model"]) step = checkpoint_dict["step"] while True: train(model, optimizer, train_loader) acc = test(model, test_loader) checkpoint = Checkpoint.from_dict({ "step": step, "model": model.state_dict() }) if step %5 ==0 else None session.report({"mean_accuracy": acc}, checkpoint=checkpoint) step +=1
3. 内置Stopper替代自定义类
from ray.tune.stopper import CombinedStopper, MaximumIterationStopper, TrialPlateauStopper stopper = CombinedStopper( MaximumIterationStopper(max_iter=5 if args.smoke_test else 100), TrialPlateauStopper(metric="mean_accuracy", mode="max", num_results=3, std=0.001) )
三、完整修复后代码
import argparse import torch import torch.nn as nn import torch.optim as optim from torchvision.datasets import MNIST import torchvision.transforms as transforms from ray import air, tune from ray.air import session from ray.air.checkpoint import Checkpoint from ray.tune.schedulers import PopulationBasedTraining from ray.tune.stopper import CombinedStopper, MaximumIterationStopper, TrialPlateauStopper class Net(nn.Module): def __init__(self, conv1_filter, conv1_kernel, conv2_filter, conv2_kernel, dense1_units): super().__init__() self.conv1 = nn.Conv2d(1, conv1_filter, kernel_size=conv1_kernel) self.conv2 = nn.Conv2d(conv1_filter, conv2_filter, kernel_size=conv2_kernel) self.pool = nn.MaxPool2d(2) self.act = nn.Tanh() self.lin = nn.Linear(0, dense1_units) self.result = nn.Linear(dense1_units, 10) def forward(self, x): out = self.pool(self.act(self.conv1(x))) out = self.pool(self.act(self.conv2(out))) # 动态适配全连接层输入维度 if self.lin.in_features == 0: self.lin = nn.Linear(out.size(1)*out.size(2)*out.size(3), self.lin.out_features).to(out.device) out = out.view(out.size(0), -1) out = self.act(self.lin(out)) out = self.result(out) return out def train(model, optimizer, train_loader): model.train() criterion = nn.CrossEntropyLoss() for data, target in train_loader: optimizer.zero_grad() loss = criterion(model(data), target) loss.backward() optimizer.step() def test(model, test_loader): model.eval() correct = 0 total = 0 with torch.no_grad(): for data, target in test_loader: output = model(data) _, predicted = torch.max(output.data, 1) total += target.size(0) correct += (predicted == target).sum().item() return correct / total def get_data_loaders(batch_size): train_loader = torch.utils.data.DataLoader( MNIST('data', train=True, download=True, transform=transforms.ToTensor()), batch_size=batch_size, shuffle=True ) test_loader = torch.utils.data.DataLoader( MNIST('data', train=False, download=True, transform=transforms.ToTensor()), batch_size=batch_size, shuffle=False ) return train_loader, test_loader def training(config): train_loader, test_loader = get_data_loaders(config["batch_size"]) model = Net( config["conv1_filter"], config["conv1_kernel"], config["conv2_filter"], config["conv2_kernel"], config["dense1_units"] ) optimizer = optim.SGD(model.parameters(), lr=config["learning_rate"]) step = 0 if session.get_checkpoint(): checkpoint_dict = session.get_checkpoint().to_dict() model.load_state_dict(checkpoint_dict["model"]) step = checkpoint_dict["step"] while True: train(model, optimizer, train_loader) acc = test(model, test_loader) checkpoint = Checkpoint.from_dict({ "step": step, "model": model.state_dict() }) if step % 5 == 0 else None session.report({"mean_accuracy": acc}, checkpoint=checkpoint) step += 1 def test_best_model(results: tune.ResultGrid): best_result = results.get_best_result() best_config = best_result.config best_model = Net( best_config["conv1_filter"], best_config["conv1_kernel"], best_config["conv2_filter"], best_config["conv2_kernel"], best_config["dense1_units"] ) checkpoint_dict = best_result.checkpoint.to_dict() best_model.load_state_dict(checkpoint_dict["model"]) test_loader = get_data_loaders(best_config["batch_size"])[1] test_acc = test(best_model, test_loader) print("最佳模型准确率: ", test_acc) if __name__ == "__main__": config = { "batch_size": tune.choice([2, 4, 8, 16]), "conv1_filter": tune.randint(1, 10), "conv1_kernel": tune.randint(3, 5), "conv2_filter": tune.randint(1, 10), "conv2_kernel": tune.randint(3, 5), "dense1_units": tune.qlograndint(5, 8, 2), "learning_rate": tune.loguniform(1e-5, 1e-1) } parser = argparse.ArgumentParser() parser.add_argument("--smoke-test", action="store_true", help="快速测试模式") args, _ = parser.parse_known_args() scheduler = PopulationBasedTraining( time_attr="training_iteration", perturbation_interval=5, hyperparam_mutations={ "learning_rate": tune.loguniform(1e-5, 1e-1), "batch_size": tune.choice([2, 4, 8, 16]) } ) stopper = CombinedStopper( MaximumIterationStopper(max_iter=5 if args.smoke_test else 100), TrialPlateauStopper(metric="mean_accuracy", mode="max", num_results=3, std=0.001) ) tuner = tune.Tuner( training, run_config=air.RunConfig( name="pbt_cnn_mnist", stop=stopper, verbose=1, checkpoint_config=air.CheckpointConfig( checkpoint_score_attribute="mean_accuracy", num_to_keep=4, ), ), tune_config=tune.TuneConfig( scheduler=scheduler, metric="mean_accuracy", mode="max", num_samples=4, ), param_space=config ) results = tuner.fit() test_best_model(results)
内容的提问来源于stack exchange,提问作者김윤도
相关产品推荐
相关产品推荐

