如何在Wandb Sweep中使用Early Terminate?参数与实践解惑
Wandb Sweep + Hyperband 早停实操指南
一、核心参数大白话解释
你对resource的理解没错,它就是训练的epoch数。其他参数直接对应Hyperband的筛选逻辑,用人话拆解:
max_iter:单个超参数组合能拿到的最大训练epoch数(比如你设9,就是最多训9轮)min_iter:单个组合的最低训练门槛,低于这个轮数的组合不会被早停(比如设1,就是至少训1轮才会评估是否淘汰)eta:淘汰比例的倒数,默认是3。意思是每轮筛选后,只保留1/eta比例的最优组合(比如eta=3,就留top 33%,剩下的直接终止)s:筛选“阶段组”的数量,s越大,筛选轮次越多,但整体计算量也会上升。以你设的s=1、max_iter=9、eta=3为例,实际流程是:- 所有超参组合先训
9/(3^1)=3轮 - 根据验证指标排序,保留top 1/3的组合
- 这些保留的组合继续训练,直到达到max_iter=9轮
- 所有超参组合先训
二、极简验证代码
1. Sweep配置文件(sweep_config.yaml)
program: train.py method: hyperband # 必须指定用来筛选的指标,要和训练代码里log的名称完全一致 metric: name: valid_acc goal: maximize # 准确率需要最大化,若用损失则设minimize parameters: lr: distribution: uniform min: 0.001 max: 0.01 batch_size: values: [32, 64] early_terminate: type: hyperband s: 1 eta: 3 min_iter: 1 max_iter: 9
2. 训练代码(train.py)
import wandb import torch import torch.nn as nn from torch.utils.data import DataLoader from torchvision.datasets import MNIST from torchvision.transforms import ToTensor def train(): # 初始化wandb,自动加载sweep分配的超参数 run = wandb.init() config = run.config # 加载数据集,必须分训练集和验证集 train_data = MNIST(root="./data", train=True, download=True, transform=ToTensor()) val_data = MNIST(root="./data", train=False, download=True, transform=ToTensor()) train_loader = DataLoader(train_data, batch_size=config.batch_size, shuffle=True) val_loader = DataLoader(val_data, batch_size=config.batch_size, shuffle=False) # 定义简单模型 model = nn.Sequential( nn.Flatten(), nn.Linear(28*28, 128), nn.ReLU(), nn.Linear(128, 10) ) loss_fn = nn.CrossEntropyLoss() optimizer = torch.optim.Adam(model.parameters(), lr=config.lr) # 训练循环,核心是加入早停判断 for epoch in range(config.max_iter): # 训练步骤 model.train() train_loss = 0.0 for X, y in train_loader: optimizer.zero_grad() pred = model(X) loss = loss_fn(pred, y) loss.backward() optimizer.step() train_loss += loss.item() # 验证步骤(必须做,用来生成筛选指标) model.eval() val_acc = 0.0 with torch.no_grad(): for X, y in val_loader: pred = model(X) val_acc += (pred.argmax(1) == y).type(torch.float).sum().item() val_acc /= len(val_data) # 记录指标,必须和sweep配置里的metric.name一致 wandb.log({ "train_loss": train_loss/len(train_loader), "valid_acc": val_acc, "epoch": epoch+1 }) # 关键:检查是否被Hyperband标记停止,是的话直接终止训练 if wandb.run.should_stop(): print(f"Epoch {epoch+1}: 被Hyperband早停") break if __name__ == "__main__": train()
3. 运行方式
在终端执行以下命令:
# 初始化sweep wandb sweep sweep_config.yaml # 启动agent执行调优,替换成你的sweep ID wandb agent <你的sweep ID>
三、关键注意事项
- 必须记录验证指标:Hyperband完全依赖验证集指标(比如
valid_acc)筛选最优组合,若只log训练指标或未匹配配置中的指标名称,Wandb无法判断淘汰逻辑,所有任务都会训满max_iter,这大概率是你任务停不下来的核心原因。 - 必须加入早停判断:训练循环里一定要加
if wandb.run.should_stop(): break,否则即使Wandb标记任务停止,代码仍会继续跑满max_iter。
四、任务无法停止的排查步骤
针对你设s=1、max_iter=9但任务不停的问题,按以下顺序排查:
- 检查训练代码里是否log了
valid_acc,且名称和sweep配置里的metric.name完全一致 - 确认训练循环里加入了
wandb.run.should_stop()的判断并执行break - 检查超参数组合数量:如果组合数太少(比如只有2个),按eta=3的比例,top 1/3会向上取整为1个,但如果Wandb判断所有组合都值得保留,就会让它们都训到max_iter
- 确认
s和max_iter的匹配性:max_iter最好是eta^(s+1)的整数倍(比如s=1、eta=3时,3^(1+1)=9,刚好符合你的设置,这部分没问题)
内容的提问来源于stack exchange,提问作者zlq147
相关产品推荐
相关产品推荐

