16.自动超参数搜索
第 16 章 — 自动超参数搜索(Optuna)
第 15 章交付的模板工程已经足够"能打":配置驱动、一键训练、自动存档、早停、多卡……所有配置项都可以在
config.yaml里手改。但一个致命的问题还在:"lr 设成 0.001 还是 0.0003?batch_size 用 64 还是 128?调度器用 step 还是 cosine?"这些超参数到底怎么"找"出来,还是靠猜、靠经验、靠碰运气。 本章引入 Optuna,把"猜超参"这件事变成机器自动做的实验,让你从"调参玄学"里解放出来。
一、本章要解决的问题
- 之前:第 15 章完成了"修改
config.yaml即可切换实验",但config.yaml里那些超参的取值本身从哪来?仍然是靠经验、靠手试,而且每试一组都要等一次完整训练跑完(30 轮,可能几十分钟到几小时)。 - 现在:用 Optuna 自动做超参搜索——给定每个超参的取值范围,让算法自己决定"下一组试什么",用尽量少的尝试找到尽量好的组合。
- 判断标准:
python optuna_search.py --trials 20跑完,能得到一组明显优于"拍脑袋默认值"的超参,并且结果写进optuna_best.txt,可以直接填回config.yaml精调。
二、核心概念速览
下面 8 个概念是本章代码的全部"生词"。先花 5 分钟读完,再看代码会轻松很多。
1. 超参数 vs 模型参数
模型参数是训练过程中"学"出来的——卷积核的权重、偏置,训练结束的那一刻才知道它们是多少。超参数是训练开始前就要定的"旋钮"——学习率、batch_size、调度器类型、weight_decay。类比学做菜:盐放多少、火候多大是"参数",靠一次次练习调好;而"用铁锅还是不粘锅、用花生油还是黄油"是"超参数",开火之前就得定下来。
2. 为什么调参难:组合爆炸
超参数很少是独立的,它们互相影响(lr 太大配大 batch 可能发散,配小 batch 可能还行)。如果每个超参取 10 个候选值、共有 6 个超参,那就是 10⁶ = 100 万种组合;每组再完整训练 30 轮,就是天文数字。这就是"菜谱"类比:调料少的时候,排列组合可以挨个试;调料一多,把整本菜谱的组合全部做一遍就完全不可能了,必须靠聪明的"挑着试"。
3. 网格搜索与随机搜索
- 网格搜索(Grid Search):把每个超参的候选值做笛卡尔积,穷举所有组合。低维(2~3 个超参)还行,维度一高就撞上组合爆炸。
- 随机搜索(Random Search):每次从取值范围里随机抽一组。看起来"乱试",但研究(Bergstra & Bengio, 2012)表明:同样的尝试次数下,随机搜索往往不输甚至优于网格搜索——因为它在每个超参的"数值分布"上覆盖得更广,而不是死板地按格子铺。 类比:一个把整本菜谱按排列组合全部做一遍,一个闭着眼睛随机翻几页做几道——后者往往更早碰到好吃的菜。
4. Optuna 的贝叶斯优化 TPE
Optuna 默认用的不是随机搜索,而是贝叶斯优化,具体算法叫 TPE(Tree-structured Parzen Estimator,树状结构帕尔岑估计器)。核心思想:根据已经跑过的几组 trial 的结果,构建一个"哪里更可能出高分"的猜测模型,下一次就在高分区域附近重点采样;跑得越多,猜测越准。类比一个越来越懂你口味的厨师:尝过前几道菜之后,下次下料会往"你喜欢的味道"那个方向靠,而不是永远乱放调料。
5. trial 与 study
- trial(一次尝试):一组超参 + 用它们训练 SEARCH_EPOCHS 轮 + 得到一个分数(本章是验证集准确率)。
- study(一组尝试):装下所有 trial 的"账本",帮你记住历史记录、自动挑出最优的 trial(
study.best_params/study.best_value)。 类比:trial 是"做一道菜并打分",study 是"整场试菜活动"——活动主办方(Optuna)根据前面几道菜的得分决定下一道怎么做。
6. trial.suggest_ 采样接口*
Optuna 怎么知道"超参该从什么范围里取"?靠 trial.suggest_* 系列接口,你在目标函数里告诉它每个超参的取值范围:
suggest_float("lr", 1e-4, 1e-2, log=True):连续数值;log=True表示对数尺度采样。为什么要 log?因为学习率跨多个数量级(1e-4 到 1e-2),在线性轴上采样时 [0.008, 0.01] 这样的小区间会被挤到几乎没样本;而对数轴上每个数量级分到的"机会"一样多——对 lr 这种量级敏感的旋钮,log 采样才是真正的"均匀"。suggest_categorical("batch_size", [32, 64, 128]):从给定列表里挑一个(离散选项)。suggest_int:整数采样(本代码没用,但搜索轮数、层数等整数超参时常用)。 类比:Optuna 像发了一张"超参问卷",你在每个"取值范围"里填好选项,它负责替你决定每个 trial 到底填哪个值,而且填得越来越聪明。
7. 为什么搜索阶段用小 epoch 预算
本章 SEARCH_EPOCHS = 5:每组超参只训 5 轮,而不是完整训练的 30 轮。理由:超参的相对好坏在训练早期就能看个大概——lr 太大可能前几轮就发散,batch 太小可能收敛明显偏慢,5 轮足够把"明显不行"的淘汰掉。省下来的时间用来试更多组组合。类比海选:先让选手唱 30 秒片段筛掉明显不行的,决赛再让入围者唱完整首歌——没有人会要求海选就唱 4 分钟。
8. pruning 剪枝(选读)
再进一步省时间的思路:如果某组超参训练到一半,分数已经明显落后于历史最好成绩,就**提前砍掉(prune)**这个 trial,把算力让给更可能出成绩的组合。Optuna 支持在训练中调用 trial.report(中间分数, epoch) 上报中间指标,再配合 MedianPruner 等剪枝器自动砍掉落后 trial。类比厨艺比赛:前两轮就明显烧糊的选手,不用等成品出炉就能被淘汰。
三、解决思路
- 复用第 15 章的组件:
build_dataloaders/build_model/set_seed/setup_logger都是现成的工厂,搜索脚本直接 import,只写一个轻量搜索训练循环(十几行),不引入第 13 章的Trainer——搜索要的是"快"和"换超参重来",Trainer 里的存档、TensorBoard、早停都是搜索阶段用不上的重量。 - 定义搜索空间:用
trial.suggest_*声明 4 个超参——lr(对数均匀 1e-41e-2)、batch_size(32/64/128)、lr_scheduler(step/cosine)、weight_decay(01e-3)。Optuna 会自己决定每次 trial 的取值。 - 定义目标函数:
train_and_eval(cfg, trial, device)对一个 trial 的超参训练SEARCH_EPOCHS轮,返回验证集准确率。这个返回值就是 Optuna 要最大化的"分数"。 - 让 Optuna 自动做实验:
optuna.create_study(direction="maximize")+study.optimize(..., n_trials=20),中间过程完全交给 TPE——哪些组合试过了、下一组试什么,都不用你操心。 - 结果落盘:把最优超参写进
optuna_best.txt,之后手动填回config.yaml用完整训练精调,形成"粗筛 → 精调"的闭环。
trade-off:搜索阶段用小 epoch + 小规模探索,得到的是"有希望的超参区间"而不是最终冠军;用 5 轮粗筛出来的最优超参跑完整 30 轮时,分数可能跟粗筛时略有出入。这是"用精度换时间"的必然取舍——先快速圈定好区域,再在好区域里精调。
四、代码变更
相对第 15 章的改动(只需要新增一个文件,其余一律不动):
project/
├── config/
│ └── config.yaml # 无需修改:搜索脚本自动读取并按 trial 覆盖超参
├── dataset/datasets.py # 无需修改:复用 build_dataloaders
├── models/classifier.py # 无需修改:复用 build_model
├── engine/trainer.py # 无需修改:搜索用轻量循环,不依赖 Trainer
├── utils/
│ ├── __init__.py
│ ├── logger.py
│ ├── checkpoint.py
│ ├── metrics.py
│ ├── early_stopping.py
│ ├── seed.py
│ ├── dist.py
│ └── eval_tools.py
+ ├── optuna_search.py # 新增:本章唯一的新文件(全部代码见下)
├── train.py
├── predict.py
├── analyze.py
├── README.md
└── requirements.txt # 追加一行依赖requirements.txt 的改动:
tensorboard>=2.13
matplotlib>=3.7
+ optuna>=3.2为什么要"不依赖 Trainer 自己写循环"?因为
Trainer是为完整训练设计的:存档、恢复、早停、TensorBoard、混合精度……搜索阶段每 5 轮就要换一组超参重来,这些功能全部用不上,还会拖慢节奏。搜索脚本只保留"前向 → 反向 → step → 算验证集准确率"的最小循环,这是"粗筛"阶段该有的轻量。
五、完整代码
创建 optuna_search.py:
运行示例(首次需安装 Optuna):
pip install optuna
python optuna_search.py --trials 20搜索过程会在终端滚动显示每个 trial 的进度条(tqdm);跑完后日志会打印三行关键信息:
最优 trial: 14
最优 val_acc: 0.7213
最优超参: {'lr': 0.0016, 'batch_size': 64, 'lr_scheduler': 'cosine', 'weight_decay': 0.0002}如何查看 study.best_params:study.best_params 就是一个普通 Python 字典(形如 {'lr': 0.0016, 'batch_size': 64, ...}),脚本里已经用 logger.info 打印,并逐行写进了 optuna_best.txt,直接打开文件即可。注意:本章默认的 create_study 没传 storage,study 存在内存里,进程退出就没了;想让结果可事后加载、可断点续搜,把 create_study 加上 storage="sqlite:///optuna.db",之后就能这样查:
import optuna
study = optuna.load_study(study_name="cifar10_search",
storage="sqlite:///optuna.db")
print(study.best_params) # 最优超参字典
print(study.best_value) # 最优分数
print(study.trials_dataframe()) # 全部 trial 的历史记录完整代码(复用第 15 章的 build_dataloaders / build_model 组件,只新增本文件):
"""Optuna 超参数搜索:复用工程的 build_* 组件,自动搜索最优超参。
用法:python optuna_search.py --trials 20
"""
import argparse # 命令行参数:--config / --trials
import logging # 日志模块(setup_logger 基于它实现)
import torch
import torch.nn as nn
import torch.optim as optim
import optuna # 超参搜索库(本章的主角)
import yaml # 解析 config.yaml
from tqdm import tqdm # 训练进度条
# 复用第 15 章工程的组件:数据/模型/种子/日志工厂,一行不改
from dataset.datasets import build_dataloaders
from models.classifier import build_model
from utils.seed import set_seed, worker_init_fn
from utils.logger import setup_logger
SEARCH_EPOCHS = 5 # 搜索阶段每组超参只训 5 轮(小预算,先粗筛)
class Config(dict):
# 让 config 支持 cfg.train.lr 这种"点语法"(普通 dict 只能 cfg["train"]["lr"])
def __getattr__(self, key):
try:
return self[key]
except KeyError as e:
raise AttributeError(key) from e
def _to_config(obj):
# 递归把 yaml 读出来的 dict/list 全部转成 Config,深层配置也能点语法访问
if isinstance(obj, dict):
return Config({k: _to_config(v) for k, v in obj.items()})
if isinstance(obj, list):
return [_to_config(v) for v in obj]
return obj
def load_config(path):
# 读入第 15 章的 config.yaml,返回 Config 对象
with open(path, "r", encoding="utf-8") as f:
return _to_config(yaml.safe_load(f))
def train_and_eval(cfg, trial, device):
"""对一组超参训练 SEARCH_EPOCHS 轮,返回验证集准确率(Optuna 最大化目标)。"""
# 从 trial 采样超参:Optuna 自动决定下一组试什么
lr = trial.suggest_float("lr", 1e-4, 1e-2, log=True) # 对数均匀采样
batch_size = trial.suggest_categorical("batch_size", [32, 64, 128])
lr_scheduler = trial.suggest_categorical("lr_scheduler", ["step", "cosine"])
weight_decay = trial.suggest_float("weight_decay", 0.0, 1e-3)
# 把采样到的超参写回 cfg,让 build_* 工厂"感知"本次 trial 的取值
cfg.train.lr = lr
cfg.data.batch_size = batch_size
cfg.train.lr_scheduler = lr_scheduler
# 复用工程的组件:数据与模型工厂不变,只有超参在变
train_loader, val_loader, test_loader = build_dataloaders(
cfg, worker_init_fn=worker_init_fn)
model = build_model(cfg).to(device)
criterion = nn.CrossEntropyLoss()
# 注意:weight_decay 直接传给优化器即可,无需写回 cfg
optimizer = optim.Adam(model.parameters(), lr=lr, weight_decay=weight_decay)
# 按采样到的调度器类型创建对应 scheduler(沿用第 4 章的两种选择)
if lr_scheduler == "step":
scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=3, gamma=0.5)
else:
scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=SEARCH_EPOCHS)
# 轻量训练循环:不依赖第 13 章的 Trainer,够用即可
for epoch in range(SEARCH_EPOCHS):
model.train()
for images, labels in tqdm(train_loader,
desc=f"Trial {trial.number} E{epoch + 1}",
leave=False):
images, labels = images.to(device), labels.to(device)
optimizer.zero_grad()
outputs = model(images)
loss = criterion(outputs, labels)
loss.backward()
optimizer.step()
scheduler.step()
# 用验证集准确率作为 trial 的分数
model.eval()
correct, total = 0, 0
with torch.no_grad():
for images, labels in val_loader:
images, labels = images.to(device), labels.to(device)
outputs = model(images)
correct += (outputs.argmax(dim=1) == labels).sum().item()
total += images.size(0)
return correct / total
def main():
parser = argparse.ArgumentParser(description="Optuna 超参搜索")
parser.add_argument("--config", type=str, default="config/config.yaml")
parser.add_argument("--trials", type=int, default=20, help="搜索多少组超参")
args = parser.parse_args()
cfg = load_config(args.config)
set_seed(cfg.seed) # 每组 trial 起点一致,排除随机干扰
logger = setup_logger(cfg.log.dir)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
logger.info(f"Using device: {device}")
# direction="maximize":目标(val_acc)越大越好
study = optuna.create_study(direction="maximize",
study_name="cifar10_search")
# optimize 会按需调用 train_and_eval,跑满 n_trials 组超参
study.optimize(lambda trial: train_and_eval(cfg, trial, device),
n_trials=args.trials)
# 搜索结束:打印最优 trial 的编号、分数与超参
logger.info(f"最优 trial: {study.best_trial.number}")
logger.info(f"最优 val_acc: {study.best_value:.4f}")
logger.info(f"最优超参: {study.best_params}")
# 把结果写进文件,方便手动填入 config.yaml 精调
with open("optuna_best.txt", "w", encoding="utf-8") as f:
f.write(f"val_acc={study.best_value:.4f}\n")
for k, v in study.best_params.items():
f.write(f"{k}={v}\n")
if __name__ == "__main__":
main()六、本章小结
- 学到了什么
- 超参数(训练前定的旋钮)vs 模型参数(训练学出来的):调参难在组合爆炸,靠手试不现实。
- Optuna 用贝叶斯优化(TPE)代替穷举:先跑几组,再根据历史结果"猜"更优的方向,越试越聪明。
- 搜索的完整工作流:
suggest_*声明搜索空间 → 目标函数返回分数 →create_study(direction=...)→study.optimize(n_trials=N)→ 读study.best_params。 - 工程实践:搜索阶段用小 epoch 预算(5 轮粗筛)省时间;
log=True让跨数量级的 lr 采样更均匀;set_seed让每组 trial 起点一致,排除随机噪声。
- 常见坑
- 忘记固定种子:不
set_seed,同一个 trial 换个随机种子结果可能天差地别,Optuna 会被噪声误导。 - 搜索预算当精调:5 轮粗筛出的"最优"只代表搜索阶段的排名,一定要用完整训练验证后再决定最终配置。
- 搜索空间定太宽/太窄:太宽浪费时间,太窄找不到更好解;先用小
--trials快速试水再扩大范围。 - study 不持久化:默认
create_study的结果在内存里,进程结束就丢;记得把结果落盘(optuna_best.txt)或传storage参数。
- 忘记固定种子:不
七、动手练习
- 扩大搜索空间:给
train_and_eval增加一个trial.suggest_int("num_workers", 2, 8)(搜索 DataLoader 加载进程数),跑--trials 20,观察最优值与总耗时的变化。 - 持久化 study:把
create_study改成optuna.create_study(direction="maximize", study_name="cifar10_search", storage="sqlite:///optuna.db"),再写个 5 行小脚本用optuna.load_study加载并打印study.trials_dataframe(),看看每个 trial 都试了什么。 - 对比采样器:把
create_study的采样器换成optuna.samplers.RandomSampler(),固定同一随机种子各跑--trials 20,对比 TPE 与随机搜索谁的最优 val_acc 更高——体会"越试越聪明"的差距。 - 闭环精调:把
optuna_best.txt里的超参填回config.yaml,用第 13 章的python train.py完整训练 30 轮,与第 15 章默认配置的最终测试集准确率对比,验证"粗筛 → 精调"闭环。 - 进阶(可选)· 实现 pruning:在
train_and_eval每个 epoch 的验证后调用trial.report(val_acc, epoch)和trial.should_prune()(配合optuna.exceptions.TrialPruned),并给create_study加pruner=optuna.pruners.MedianPruner(),比较开/关剪枝的总耗时。
八、16 章全览表
至此整套教程收官。下表回顾 16 章各自的标题、引入的库与核心功能,既是索引,也是"将来遇到问题回哪一章查"的地图:
| 章 | 标题 | 引入的库 | 核心功能 |
|---|---|---|---|
| 1 | 最小可运行实现 | torch / torchvision | 跑通训练+验证循环,看清训练"骨架" |
| 2 | 训练/验证/测试三段划分与指标管理 | random_split / 内置容器 | 泛化能力评估、指标历史记录 |
| 3 | 命令行参数管理 | argparse | 超参数命令行化,可追溯可覆盖 |
| 4 | 学习率调度器 | torch.optim.lr_scheduler | 动态调整学习率,突破收敛瓶颈 |
| 5 | 模型保存与恢复 | torch.save / torch.load | last.pth / best.pth / --resume 断点续训 |
| 6 | 早停机制 | 自定义 EarlyStopping | 连续无提升自动停止,防过拟合省时间 |
| 7 | 日志系统与 TensorBoard | logging / torch.utils.tensorboard | 日志落盘、标量曲线可视化 |
| 8 | 配置文件管理 | pyyaml | config.yaml 成为"唯一事实来源" |
| 9 | 数据增强与自定义 Transform | torchvision.transforms | 训练集增强,缓解过拟合 |
| 10 | 混合精度训练 | torch.amp | FP16 前向 + FP32 反向,提速省显存 |
| 11 | 代码模块化拆分 | 标准库 os | 单一职责封装 + 工厂函数,工程化结构 |
| 12 | 可复现性保障 | random / numpy / torch | 固定全部随机种子 |
| 13 | 多 GPU 与分布式训练 | torch.distributed | DataParallel / DDP,多卡并行 |
| 14 | 测试推理与可视化 | matplotlib | predict.py 推理、混淆矩阵、样例图 |
| 15 | 最终模板整合与文档 | —(整合) | README、全开关配置表、交付验收 |
| 16 | 自动超参数搜索 | Optuna | 用 Optuna 自动做超参搜索 |
