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

如何在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为例,实际流程是:
    1. 所有超参组合先训9/(3^1)=3轮
    2. 根据验证指标排序,保留top 1/3的组合
    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>

三、关键注意事项

  1. 必须记录验证指标:Hyperband完全依赖验证集指标(比如valid_acc)筛选最优组合,若只log训练指标或未匹配配置中的指标名称,Wandb无法判断淘汰逻辑,所有任务都会训满max_iter,这大概率是你任务停不下来的核心原因。
  2. 必须加入早停判断:训练循环里一定要加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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.13 10:40:49