8.配置文件管理
第 8 章 — 配置文件管理(YAML)
第 7 章把日志(logging + TensorBoard)收拾利落,但命令行参数堆到了十几个:
--epochs --batch-size --lr --lr-scheduler --data-dir --num-workers --device --ckpt-dir --patience --log-dir ...,一行命令又长又难对比,换个数据集就得重拼一遍,也没法把"一组实验配置"当文件存档。本章引入 YAML 配置文件:把不常变的超参数按模块写进config/config.yaml,命令行只保留--config / --lr / --epochs / --batch-size / --resume五个常用覆盖项——"一份写好的菜单 + 临时口头加两笔菜"的组合。
一、本章要解决的问题
- 之前(第 7 章):14 个超参数全部挂在
argparse上,参数越多命令越长;跑不同实验要重新敲一遍命令行,还容易手滑打错参数。 - 现在:绝大部分超参数(数据路径、批次大小、学习率、调度器、早停、日志目录等)写进
config/config.yaml,按data / train / model / checkpoint / log模块分组;命令行只留 5 个:--config(选哪份菜单)+--lr / --epochs / --batch-size(临时改三个高频变量)+--resume(断点续训)。 - 判断标准:
python train.py什么都不传也能跑(配置全在 YAML 里);python train.py --lr 1e-4 --epochs 50只改想改的;启动日志打印"最终生效配置",一眼看清每个值到底来自命令行还是配置文件。
二、核心概念速览
下面 6 个概念是本章代码的"生词"。先花 5 分钟读完,再看代码会轻松很多。
1. 配置文件 vs 命令行参数(点菜单 vs 口头点菜)
把一次训练实验想象成去餐厅吃饭。命令行参数像"口头点菜":方便快捷,但一桌十几道菜全靠嘴报,既累又容易漏。配置文件像"菜单/点餐单":把整套菜写在一张纸上,服务员照单上菜,改菜谱只要改纸、不用重新喊一遍。所以本章的策略是:完整配置写进 YAML,临时改动用命令行覆盖——日常"照单上菜",偶尔"这盘少放盐"。
2. YAML 语法三要素
YAML 的入门只需要记三件事:缩进表达层级——就像 Word 的目录,缩进越深表示层级越深,data: 下面缩进的 root:、batch_size: 都是它的子项(用空格,别用 Tab);键值对——key: value 的形式,冒号后面必须有一个空格,就像填表格"姓名:张三";列表——以 - item 开头的行,就像购物清单每项前面打个减号。记住这三条,配置文件就看得懂了。注意:YAML 里 # 开头是注释,写了中文注释也完全没问题(文件本身就是实验的可读记录)。
3. yaml.safe_load 与 yaml.load 的区别(安全)
yaml.load 不加限制地解析 YAML 里的各种标签,如果配置文件被篡改,解析过程可能执行任意代码——相当于"收到陌生邮件里的附件,还双击运行了它"。yaml.safe_load 只解析普通数据(字符串、数字、列表、字典),永不执行代码,是官方推荐的安全做法。我们的配置里只有纯数据,用 safe_load 完全够用,所以记住一个原则:永远用 safe_load。
4. 为什么用 Config(dict) + getattr 实现点访问
yaml.safe_load 读出来是嵌套的 dict,访问要写成 cfg["data"]["batch_size"],又长又容易手滑。让 Config 继承 dict,再定义 __getattr__:当属性访问找不到时,自动去 dict 里按 key 取值。于是 cfg.data.batch_size 等价于 cfg["data"]["batch_size"],代码干净得像在直接读配置文件本身。生活类比:平时家里有什么东西直接用(dict 原生查找),家里缺东西才"开门出去找"(__getattr__ 触发兜底),找到就带回、找不到就报 AttributeError。
5. 配置覆盖优先级
命令行 > 配置文件 > 代码默认值。这就像点餐的三层规则:餐厅固定的菜谱(代码默认值)是底线;厨房按菜单备菜(配置文件);你临时口头加一句"这盘少放盐"(命令行覆盖),以你的要求为准。实现上非常朴素:命令行参数设 default=None,None 表示"用户没传";配置文件加载完之后,逐个检查命令行参数,不是 None 才覆盖配置里的对应项。这样 --lr 没传时保持配置文件里的 train.lr,传了才改。
6. argparse 只保留关键覆盖项的权衡
第 7 章把 14 个参数全放命令行,命令又长又难对比。本章把低频变量(数据路径、调度器参数、worker 数、patience、日志目录等)沉进配置文件,命令行只留 5 个高频项。这是一个明确 trade-off:--lr/--epochs/--batch-size 是实验对比最常动的"旋钮",留在命令行方便快速调;其余属于"换环境才动"的项,写进 YAML 沉淀成可复现的实验记录。划分没有绝对标准(团队习惯说了算),但"配置文件为主、命令行做增量覆盖"是业界主流做法。
三、解决思路
- 配置文件先行:新增
config/config.yaml,顶部device单独放,其余按data / train / model / checkpoint / log分组,每个键一行注释说明它对应第 7 章的哪个--参数,让配置"自带说明书"。 - 点访问包装:
Config(dict)+__getattr__把嵌套 dict 变成可点访问的配置对象,_to_config递归转换(dict 换 Config、list 里的元素也一并处理)。 - 命令行只留覆盖项:argparse 保留
--config / --lr / --epochs / --batch-size / --resume五个参数,且四个覆盖项的default=None;cfg加载完成后逐项判空覆盖。 - 其余沿用:logging(RichHandler + FileHandler)、TensorBoard、tqdm、checkpoint / 早停 / 恢复逻辑与第 7 章一致,只把
args.xxx换成cfg.xxx.yyy;启动时用yaml.safe_dump(cfg)打印"最终生效配置",方便核对每一项的来源。
trade-off:参数从命令行移进配置文件后,命令变短、实验可复现(YAML 文件本身就是存档);代价是多学一层 YAML 语法,且"配置写错了不会立刻报错、只会静默跑歪"。因此本章在训练前打印最终生效配置,把"跑歪"的风险尽早暴露出来。
四、代码变更
相对第 7 章的改动:
+ import yaml # 新增:读取 YAML 配置文件
import os
import logging
import argparse
...
from torch.utils.tensorboard import SummaryWriter
from torchvision import datasets, transforms
from collections import defaultdict
+ from tqdm import tqdm # 沿用前几章:训练/验证进度条
+ from rich.logging import RichHandler # 沿用前几章:控制台彩色日志
# 0. 命令行参数:从 14 个精简到 5 个
- parser.add_argument("--epochs", type=int, default=30, ...)
- parser.add_argument("--batch-size", type=int, default=64, ...)
- parser.add_argument("--lr", type=float, default=1e-3, ...)
- parser.add_argument("--lr-scheduler", ...)
- parser.add_argument("--lr-step-size", ...)
- parser.add_argument("--lr-gamma", ...)
- parser.add_argument("--data-dir", ...)
- parser.add_argument("--num-workers", ...)
- parser.add_argument("--device", ...)
- parser.add_argument("--ckpt-dir", ...)
- parser.add_argument("--patience", ...)
- parser.add_argument("--log-dir", ...)
+ parser.add_argument("--config", type=str, default="config/config.yaml",
+ help="YAML 配置文件路径")
+ parser.add_argument("--lr", type=float, default=None, ...) # 覆盖 train.lr
+ parser.add_argument("--epochs", type=int, default=None, ...) # 覆盖 train.epochs
+ parser.add_argument("--batch-size", type=int, default=None, ...) # 覆盖 data.batch_size
+ parser.add_argument("--resume", type=str, default=None, ...)
+ # 新增:Config(dict) + __getattr__ 实现 cfg.train.lr 点访问
+ class Config(dict):
+ def __getattr__(self, key):
+ try:
+ return self[key]
+ except KeyError as e:
+ raise AttributeError(key) from e
+ def _to_config(obj): ... # 递归把 dict/list 全部换成 Config
+ def load_config(path):
+ with open(path, "r", encoding="utf-8") as f:
+ return _to_config(yaml.safe_load(f))
+ cfg = load_config(args.config) # 全局唯一的"生效配置"
+ # 命令行覆盖配置:default=None 判空(没传就保持配置文件的值)
+ if args.lr is not None:
+ cfg.train.lr = args.lr
+ if args.epochs is not None:
+ cfg.train.epochs = args.epochs
+ if args.batch_size is not None:
+ cfg.data.batch_size = args.batch_size
# 1. 日志工具:控制台 handler 换成 Rich
- console = logging.StreamHandler()
- console.setFormatter(fmt)
+ console = RichHandler(rich_tracebacks=True, markup=True)
...
- logger = setup_logger(args.log_dir)
- writer = SummaryWriter(log_dir=args.log_dir)
+ logger = setup_logger(cfg.log.dir)
+ writer = SummaryWriter(log_dir=cfg.log.dir)
# 2. 设备
- if args.device == "auto":
+ if cfg.device == "auto":
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
else:
- device = torch.device(args.device)
+ device = torch.device(cfg.device)
# 3. 模型
- model = SimpleCNN().to(device)
+ model = SimpleCNN(num_classes=cfg.model.num_classes).to(device)
# 4. 数据:变换抽成函数(第 9 章在此加入数据增强),根目录/验证集比例进配置
- transform = transforms.Compose([...])
+ def build_transforms(cfg, train=False):
+ """按配置构建训练/验证两套变换(本章先用基础版,第 9 章再加入增强)。"""
+ ...
- train_dataset = datasets.CIFAR10(root=args.data_dir, train=True, ...)
- test_dataset = datasets.CIFAR10(root=args.data_dir, train=False, ...)
- train_dataset, val_dataset = random_split(
- train_dataset, [45000, len(train_dataset) - 45000])
+ train_dataset = datasets.CIFAR10(root=cfg.data.root, train=True, ...)
+ test_dataset = datasets.CIFAR10(root=cfg.data.root, train=False, ...)
+ val_size = int(len(train_dataset) * cfg.data.val_ratio) # 比例式切分
+ train_dataset, val_dataset = random_split(
+ train_dataset, [len(train_dataset) - val_size, val_size])
...
- train_loader = DataLoader(train_dataset, batch_size=args.batch_size,
- shuffle=True, num_workers=args.num_workers)
+ train_loader = DataLoader(train_dataset, batch_size=cfg.data.batch_size,
+ shuffle=True, num_workers=cfg.data.num_workers)
...(val/test loader 同理)
# 5. 优化器 + 调度器
- optimizer = optim.Adam(model.parameters(), lr=args.lr)
+ optimizer = optim.Adam(model.parameters(), lr=cfg.train.lr)
if ...:
- optimizer, step_size=args.lr_step_size, gamma=args.lr_gamma)
+ optimizer, step_size=cfg.train.lr_step_size, gamma=cfg.train.lr_gamma)
...
# 6. 训练/验证函数:循环加 tqdm 进度条
- for images, labels in loader:
+ for images, labels in tqdm(loader, desc="Train", leave=False):
...
# 7. 恢复逻辑(不变,--resume 仍走命令行)
...
# 8. 主训练循环
- early_stopping = EarlyStopping(patience=args.patience, verbose=True)
+ early_stopping = EarlyStopping(patience=cfg.train.patience, verbose=True)
...
- logger.info("训练配置:")
- for k, v in vars(args).items():
- logger.info(f" {k:12s} = {v}")
+ logger.info("最终生效配置:")
+ logger.info(yaml.safe_dump(cfg, allow_unicode=True, sort_keys=False))
...
- for epoch in range(start_epoch, args.epochs + 1):
+ for epoch in range(start_epoch, cfg.train.epochs + 1):
...
- save_checkpoint(state, os.path.join(args.ckpt_dir, "last.pth"))
+ save_checkpoint(state, os.path.join(cfg.checkpoint.dir, "last.pth"))
...五、完整代码
新建两个文件:先创建 config/config.yaml(配置),再创建 train.py(训练脚本)。
创建 config/config.yaml:
# CIFAR-10 图像分类实验配置
# YAML 语法三要素速查:
# 1) 缩进表达层级:data: 下面缩进的 root:/batch_size: 都是它的子项(用空格,别用 Tab);
# 2) 键值对:key: value,冒号后必须有一个空格;
# 3) 列表:以 "- " 开头的行(本章暂未用到,练习里会加)。
device: auto # auto / cuda / cpu
data:
root: ./data # 数据集存放目录(对应第 7 章的 --data-dir)
batch_size: 64 # 每个 batch 的样本数(对应 --batch-size)
num_workers: 2 # DataLoader 数据加载进程数(对应 --num-workers)
val_ratio: 0.1 # 从训练集切 10% 做验证集,比写死 45000 更通用
train:
epochs: 30 # 训练轮数上限(对应 --epochs)
lr: 0.001 # 初始学习率(对应 --lr)
lr_scheduler: step # 学习率调度:step / cosine
lr_step_size: 15 # StepLR:每 15 轮衰减一次(对应 --lr-step-size)
lr_gamma: 0.1 # StepLR 衰减系数(对应 --lr-gamma)
patience: 7 # 早停:验证集连续 7 轮无提升则停止(对应 --patience)
model:
num_classes: 10 # CIFAR-10 一共 10 个类别
checkpoint:
dir: ./checkpoints # 断点保存目录,last.pth / best.pth 都会存到这里
log:
dir: ./runs/cifar10 # TensorBoard 事件文件与 train.log 的输出目录(对应 --log-dir)创建 train.py:
import os # 文件/目录操作:建目录、拼路径
import logging # 日志:控制台 + 文件双输出(第 7 章引入)
import argparse # 命令行参数:本章只保留少量"覆盖项"
import yaml # 读取/打印 YAML 配置文件
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader, random_split
from torch.utils.tensorboard import SummaryWriter # TensorBoard(第 7 章引入)
from torchvision import datasets, transforms
from collections import defaultdict
from tqdm import tqdm # 训练/验证进度条(前几章已引入,沿用)
from rich.logging import RichHandler # 控制台彩色日志 handler(前几章已引入,沿用)
# ----------------------------------------------------------------------------
# 0. 命令行参数:只保留 5 个,其余全部交给配置文件
# ----------------------------------------------------------------------------
def parse_args():
# --config 决定"读哪份菜单";--lr/--epochs/--batch-size 是"临时改菜";
# --resume 负责断点续训。注意三个覆盖项的 default 都是 None:
# None = "用户没传",这样我们才能区分"没传"和"传了 0"。
parser = argparse.ArgumentParser(description="CIFAR-10 图像分类训练")
parser.add_argument("--config", type=str, default="config/config.yaml",
help="YAML 配置文件路径")
parser.add_argument("--lr", type=float, default=None, help="覆盖配置文件中的 train.lr")
parser.add_argument("--epochs", type=int, default=None, help="覆盖配置文件中的 train.epochs")
parser.add_argument("--batch-size", type=int, default=None,
help="覆盖配置文件中的 data.batch_size")
parser.add_argument("--resume", type=str, default=None,
help="从指定 checkpoint 恢复训练")
return parser.parse_args()
args = parse_args()
# ----------------------------------------------------------------------------
# 0.5 配置读取:把 YAML 变成"可以点访问"的 Config 对象
# ----------------------------------------------------------------------------
class Config(dict):
"""dict 的薄包装:cfg.train.lr 等价于 cfg["train"]["lr"]。"""
def __getattr__(self, key):
# __getattr__ 是"属性不存在时才调用的钩子":
# 平时直接找属性,找不到就下来这里,帮我们从 dict 里按 key 取值。
try:
return self[key]
except KeyError as e:
raise AttributeError(key) from e
def _to_config(obj):
# 递归:dict 全部换成 Config,list 里的元素也一并处理
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):
# yaml.safe_load 把 YAML 读成嵌套 dict,再用 _to_config 套上点访问能力
with open(path, "r", encoding="utf-8") as f:
return _to_config(yaml.safe_load(f))
cfg = load_config(args.config)
# 命令行显式传入的参数覆盖配置文件(None = 未传入,保持配置文件值)
# 覆盖优先级:命令行 > 配置文件 > 代码默认值,这里实现"命令行 > 配置文件"
if args.lr is not None:
cfg.train.lr = args.lr
if args.epochs is not None:
cfg.train.epochs = args.epochs
if args.batch_size is not None:
cfg.data.batch_size = args.batch_size
# ----------------------------------------------------------------------------
# 1. 日志工具:RichHandler(控制台) + FileHandler(train.log) + TensorBoard
# ----------------------------------------------------------------------------
def setup_logger(log_dir):
os.makedirs(log_dir, exist_ok=True) # 目录不存在就自动创建
logger = logging.getLogger("train")
logger.handlers.clear() # 防止重复 setup 时叠加 handler
logger.setLevel(logging.INFO)
console = RichHandler(rich_tracebacks=True, markup=True) # 控制台:Rich 彩色
console.setLevel(logging.INFO)
logger.addHandler(console)
file_handler = logging.FileHandler(os.path.join(log_dir, "train.log"),
encoding="utf-8") # 文件:普通文本
file_handler.setFormatter(logging.Formatter(
"%(asctime)s | %(levelname)s | %(message)s", datefmt="%Y-%m-%d %H:%M:%S"))
logger.addHandler(file_handler)
return logger
logger = setup_logger(cfg.log.dir) # 日志目录来自配置(log.dir)
writer = SummaryWriter(log_dir=cfg.log.dir) # TensorBoard 事件文件目录
# ----------------------------------------------------------------------------
# 2. 设备
# ----------------------------------------------------------------------------
if cfg.device == "auto":
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
else:
device = torch.device(cfg.device)
logger.info(f"Using device: {device}")
# ----------------------------------------------------------------------------
# 3. 模型:与第 7 章相同,只有 num_classes 改为从配置读取
# ----------------------------------------------------------------------------
class SimpleCNN(nn.Module):
def __init__(self, num_classes=10):
super().__init__()
self.features = nn.Sequential(
nn.Conv2d(3, 32, kernel_size=3, padding=1),
nn.BatchNorm2d(32),
nn.ReLU(inplace=True),
nn.MaxPool2d(2),
nn.Conv2d(32, 64, kernel_size=3, padding=1),
nn.BatchNorm2d(64),
nn.ReLU(inplace=True),
nn.MaxPool2d(2),
nn.Conv2d(64, 128, kernel_size=3, padding=1),
nn.BatchNorm2d(128),
nn.ReLU(inplace=True),
nn.MaxPool2d(2),
)
self.classifier = nn.Sequential(
nn.AdaptiveAvgPool2d(1),
nn.Flatten(),
nn.Linear(128, num_classes),
)
def forward(self, x):
return self.classifier(self.features(x))
def build_transforms(cfg, train=False):
"""按配置构建训练/验证两套变换(本章先用基础版,第 9 章再加入增强)。"""
transforms_list = [
transforms.ToTensor(),
transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)),
]
return transforms.Compose(transforms_list)
train_dataset = datasets.CIFAR10(root=cfg.data.root, train=True, download=True,
transform=build_transforms(cfg, train=True))
test_dataset = datasets.CIFAR10(root=cfg.data.root, train=False, download=True,
transform=build_transforms(cfg, train=False))
# 验证集大小由比例 val_ratio 决定(0.1 -> 5000 张),比写死 45000 更通用
val_size = int(len(train_dataset) * cfg.data.val_ratio)
train_dataset, val_dataset = random_split(
train_dataset, [len(train_dataset) - val_size, val_size]
)
train_loader = DataLoader(train_dataset, batch_size=cfg.data.batch_size,
shuffle=True, num_workers=cfg.data.num_workers)
val_loader = DataLoader(val_dataset, batch_size=cfg.data.batch_size,
shuffle=False, num_workers=cfg.data.num_workers)
test_loader = DataLoader(test_dataset, batch_size=cfg.data.batch_size,
shuffle=False, num_workers=cfg.data.num_workers)
# ----------------------------------------------------------------------------
# 4. 损失函数 + 优化器 + 调度器(超参数全部来自 cfg.train.*)
# ----------------------------------------------------------------------------
model = SimpleCNN(num_classes=cfg.model.num_classes).to(device)
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=cfg.train.lr)
if cfg.train.lr_scheduler == "step":
scheduler = optim.lr_scheduler.StepLR(
optimizer, step_size=cfg.train.lr_step_size, gamma=cfg.train.lr_gamma)
elif cfg.train.lr_scheduler == "cosine":
scheduler = optim.lr_scheduler.CosineAnnealingLR(
optimizer, T_max=cfg.train.epochs)
else:
scheduler = None
# ----------------------------------------------------------------------------
# 5. 工具函数与类(与第 7 章一致,仅循环里加了 tqdm 进度条)
# ----------------------------------------------------------------------------
def train_one_epoch(model, loader, criterion, optimizer, device):
model.train() # 训练模式:启用 BN 的 batch 统计
total_loss, correct, total = 0.0, 0, 0
for images, labels in tqdm(loader, desc="Train", leave=False): # tqdm 进度条
images, labels = images.to(device), labels.to(device)
optimizer.zero_grad() # 清零梯度(不清会累加)
outputs = model(images)
loss = criterion(outputs, labels)
loss.backward()
optimizer.step()
total_loss += loss.item() * images.size(0)
correct += (outputs.argmax(dim=1) == labels).sum().item()
total += images.size(0)
return total_loss / total, correct / total
def validate(model, loader, criterion, device):
model.eval() # 评估模式:BN 用全局统计量
total_loss, correct, total = 0.0, 0, 0
with torch.no_grad(): # 验证不计算梯度,省显存
for images, labels in tqdm(loader, desc="Val", leave=False):
images, labels = images.to(device), labels.to(device)
outputs = model(images)
loss = criterion(outputs, labels)
total_loss += loss.item() * images.size(0)
correct += (outputs.argmax(dim=1) == labels).sum().item()
total += images.size(0)
return total_loss / total, correct / total
def save_checkpoint(state, path):
os.makedirs(os.path.dirname(path), exist_ok=True)
torch.save(state, path)
logger.info(f"[Checkpoint] 已保存到 {path}")
class EarlyStopping:
def __init__(self, patience=7, delta=0.0, verbose=False):
self.patience = patience
self.delta = delta
self.verbose = verbose
self.best_score = None
self.counter = 0
self.early_stop = False
def __call__(self, val_acc):
score = val_acc
if self.best_score is None:
self.best_score = score
elif score < self.best_score + self.delta:
self.counter += 1
if self.verbose:
logger.info(f"[EarlyStopping] 连续 {self.counter}/{self.patience} 轮未提升")
if self.counter >= self.patience:
self.early_stop = True
else:
self.best_score = score
self.counter = 0
# ----------------------------------------------------------------------------
# 6. 断点续训(--resume 仍走命令行)
# ----------------------------------------------------------------------------
metrics = defaultdict(list)
start_epoch, best_acc = 1, 0.0
if args.resume:
ckpt = torch.load(args.resume, map_location=device, weights_only=False)
model.load_state_dict(ckpt["model_state"])
optimizer.load_state_dict(ckpt["optimizer_state"])
if scheduler is not None and ckpt.get("scheduler_state") is not None:
scheduler.load_state_dict(ckpt["scheduler_state"])
start_epoch = ckpt["epoch"] + 1
best_acc = ckpt.get("best_acc", 0.0)
metrics = defaultdict(list, ckpt.get("metrics", {}))
logger.info(f"[Resume] 从 epoch {ckpt['epoch']} 恢复,历史最优 val_acc={best_acc:.4f}")
early_stopping = EarlyStopping(patience=cfg.train.patience, verbose=True)
# 启动时打印"最终生效配置":核对每个值到底来自命令行还是配置文件
logger.info("=" * 60)
logger.info("最终生效配置:")
logger.info(yaml.safe_dump(cfg, allow_unicode=True, sort_keys=False))
logger.info("=" * 60)
# ----------------------------------------------------------------------------
# 7. 主训练循环(结构同第 7 章,变量全部换成 cfg.*)
# ----------------------------------------------------------------------------
for epoch in range(start_epoch, cfg.train.epochs + 1):
train_loss, train_acc = train_one_epoch(model, train_loader, criterion, optimizer, device)
val_loss, val_acc = validate(model, val_loader, criterion, device)
if scheduler is not None:
scheduler.step()
current_lr = optimizer.param_groups[0]["lr"]
metrics["train_loss"].append(train_loss)
metrics["train_acc"].append(train_acc)
metrics["val_loss"].append(val_loss)
metrics["val_acc"].append(val_acc)
logger.info(f"Epoch {epoch:02d}/{cfg.train.epochs} | "
f"Train loss {train_loss:.4f} acc {train_acc:.4f} | "
f"Val loss {val_loss:.4f} acc {val_acc:.4f} | "
f"lr {current_lr:.2e}")
writer.add_scalar("train/loss", train_loss, epoch)
writer.add_scalar("train/acc", train_acc, epoch)
writer.add_scalar("val/loss", val_loss, epoch)
writer.add_scalar("val/acc", val_acc, epoch)
writer.add_scalar("lr", current_lr, epoch)
state = {
"model_state": model.state_dict(),
"optimizer_state": optimizer.state_dict(),
"scheduler_state": scheduler.state_dict() if scheduler else None,
"epoch": epoch,
"best_acc": best_acc,
"metrics": dict(metrics),
}
save_checkpoint(state, os.path.join(cfg.checkpoint.dir, "last.pth"))
if val_acc > best_acc:
best_acc = val_acc
state["best_acc"] = best_acc
save_checkpoint(state, os.path.join(cfg.checkpoint.dir, "best.pth"))
early_stopping(val_acc)
if early_stopping.early_stop:
logger.info(f"[EarlyStopping] 连续 {cfg.train.patience} 轮无提升,在 epoch {epoch} 停止")
break
test_loss, test_acc = validate(model, test_loader, criterion, device)
logger.info(f"[Test] loss {test_loss:.4f} acc {test_acc:.4f}")
writer.close()运行与查看:
# 什么都不传:全部参数取 config/config.yaml
python train.py
# 只改两个高频旋钮:其余仍取配置文件(启动日志里能看出覆盖生效)
python train.py --lr 1e-4 --epochs 50两种跑法启动时都会打印"最终生效配置"(用 yaml.safe_dump 原样输出,含中文注释),跑完用 tensorboard --logdir runs 照常看曲线,--resume 照常断点续训——这些和上一章完全一样,只是参数的"默认值"从代码挪到了 YAML。
六、本章小结
- 学到了什么
- YAML 三要素:缩进表达层级、键值对(冒号后要有空格)、列表(
- item);#注释让配置文件自带说明。 - 配置读取链路:
yaml.safe_load(安全解析)→Config(dict)+__getattr__(点访问),cfg.data.batch_size等价于cfg["data"]["batch_size"]。 - 覆盖优先级:命令行 > 配置文件 > 代码默认值,实现技巧是
default=None+ 判空,None即"没传、保持配置文件值"。 - argparse 只保留高频覆盖项:不常变的参数沉入配置文件,命令变短、实验可复现(YAML 文件本身就是一次实验的存档)。
- YAML 三要素:缩进表达层级、键值对(冒号后要有空格)、列表(
- 常见坑
- YAML 缩进必须用空格、不能混用 Tab;
key: value冒号后必须有一个空格,否则解析直接报错。 - 忘用
yaml.safe_load而用yaml.load:除非有特殊需求,一律safe_load,防代码注入。 - 覆盖项没判空就赋值:会把配置里的值覆盖成
None,训练静默跑歪。 - 打印配置记得
allow_unicode=True,否则中文路径/注释会被转义成一堆\uXXXX。 - 配置里
num_classes: 10是字符串"10"还是数字10要分清——YAML 里不带引号的数字就是数字,与nn.Linear等接口要类型匹配。
- YAML 缩进必须用空格、不能混用 Tab;
- 下一章预告:
build_transforms(cfg, train)已经留好了"训练/验证两套变换"的接口,但训练集目前还没做数据增强——第 9 章将加入随机裁剪/翻转等增强,并用配置项控制开关,方便对比"有增强 vs 无增强"。
七、动手练习
- 验证覆盖优先级:把
config/config.yaml里的train.epochs改成 5,分别运行python train.py(不带--epochs)和python train.py --epochs 3,观察启动日志打印的"最终生效配置",确认"命令行 > 配置文件"成立。 - 新增一个配置项:在
config.yaml的train:块下加weight_decay: 1e-4,然后把train.py里的optim.Adam(model.parameters(), lr=cfg.train.lr)改成optim.Adam(model.parameters(), lr=cfg.train.lr, weight_decay=cfg.train.weight_decay),体会"加一个超参数"的完整流程。 - 故意制造 YAML 报错:把
data:下面的root: ./data缩进改成 Tab(或删掉冒号后的空格),运行python train.py看报错信息,记住这类语法错误长什么样。 - 对比点访问与 dict 访问:把代码里的
cfg.train.lr、cfg.data.batch_size全部改成cfg["train"]["lr"]、cfg["data"]["batch_size"]跑一遍,说出两种写法的优缺点(代码可读性 vs 灵活性)。 - 加一个列表配置:在
config.yaml里加train: { milestones: [10, 20] },尝试用yaml.safe_load读出来并打印type(cfg.train.milestones),验证列表在_to_config里的处理方式,然后想想它能为第 9 章的学习率策略(如 MultiStepLR)提供什么。
