实战猫狗大战:可复用的PyTorch项目架构
本文你将学到 :如何把一个深度学习项目组织成"配置 / 数据 / 模型 / 工具 / 主程序"五块分离的可复用骨架;用一个 Dataset 类同时管好训练集、验证集、测试集三种数据(含数据增强的差异化处理);BasicModule 如何给所有模型统一加上save/load能力;用 getattr + 字符串动态切换模型、彻底告别if-else;fire库三行代码搞定命令行接口;完整的train/val/test三段式流程(含混淆矩阵、loss平滑统计、visdom实时监控、"损失不降就衰减学习率"策略);以及ipdb动态调试训练中程序的实用技巧。案例是Kaggle经典比赛"Dogs vs. Cats",全部代码来自陈云《深度学习框架PyTorch:入门与实践(第2版)》第9章配套仓库(PyTorch 1.8)。
文章目录
- 实战猫狗大战:可复用的PyTorch项目架构
-
- 一、任务与目标:Kaggle猫狗二分类
- 二、项目骨架:五块分离的目录结构
- 三、`init.py`:让目录变成包
- 四、config.py:所有可配置项集中管理
- 五、data/dataset.py:一个类管三种数据集
- 六、models/:BasicModule与动态模型选择
-
- [6.1 BasicModule:给所有模型统一加save/load](#6.1 BasicModule:给所有模型统一加save/load)
- [6.2 两个模型:微调的SqueezeNet与手写的ResNet34](#6.2 两个模型:微调的SqueezeNet与手写的ResNet34)
- [6.3 字符串选模型:getattr的妙用](#6.3 字符串选模型:getattr的妙用)
- 七、utils/visualize.py:封装visdom做训练监控
- [八、main.py:fire命令行 + 训练/验证/测试](#八、main.py:fire命令行 + 训练/验证/测试)
-
- [8.1 fire:三行代码的命令行接口](#8.1 fire:三行代码的命令行接口)
- [8.2 训练函数:五步流程](#8.2 训练函数:五步流程)
- [8.3 验证:eval模式切换是重点](#8.3 验证:eval模式切换是重点)
- [8.4 测试:输出提交文件](#8.4 测试:输出提交文件)
- [8.5 帮助函数与实际使用](#8.5 帮助函数与实际使用)
- 九、实用调试技巧:训练到一半也能改学习率
- 总结

一、任务与目标:Kaggle猫狗二分类
"Dogs vs. Cats"是Kaggle上的入门经典:训练集25000张图片混放在一个文件夹里,文件名自带标签,格式为 <category>.<num>.jpg,比如 cat.10000.jpg、dog.100.jpg;测试集12500张,命名只有编号,如 1000.jpg。任务是训练一个二分类器,对测试集每张图输出"是狗的概率",提交CSV:
id,label
10001,0.889
10002,0.01
...
任务本身不难,用预训练模型微调就能得到不错的成绩。这一篇的重点不在调模型,而在工程:做实验往往要改十几次参数、换几种网络、跑无数轮训练,如果代码是一锅粥,每改一处都心惊胆战。一套合理的项目结构能让你把精力花在实验本身上。
一个训练项目通常要覆盖五件事:模型定义、数据加载、训练与验证、过程可视化、测试推理。同时代码要满足三个要求:高度可配置 (改参数不改代码)、结构清晰 (各归其位)、注释完善(别人能看懂)。下面这套结构就是围绕这些目标设计的。
二、项目骨架:五块分离的目录结构
├── checkpoints/ # 训练好的模型权重,程序崩了可以恢复
├── data/
│ ├── __init__.py
│ └── dataset.py # Dataset定义、数据预处理
├── models/
│ ├── __init__.py
│ ├── basic_module.py # nn.Module的增强封装
│ ├── squeezenet.py # 模型一:SqueezeNet(微调)
│ └── resnet34.py # 模型二:ResNet34(从零实现)
├── utils/
│ ├── __init__.py
│ └── visualize.py # visdom可视化封装
├── config.py # 所有可配置项 + 默认值
├── main.py # 入口:train/test/help
├── requirements.txt
└── README.md
职责划分一句话说清:config.py管参数,data管数据,models管网络,utils管杂活,main.py只做流程编排。新增一个模型就加一个文件,新增一个数据集也是加一个文件,互不干扰。这套结构从图像分类到GAN都能套用,是笔者见过复用价值最高的组织方式之一。
#mermaid-svg-h6rzEPgzwPzDH2o3{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-h6rzEPgzwPzDH2o3 .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-h6rzEPgzwPzDH2o3 .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-h6rzEPgzwPzDH2o3 .error-icon{fill:#552222;}#mermaid-svg-h6rzEPgzwPzDH2o3 .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-h6rzEPgzwPzDH2o3 .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-h6rzEPgzwPzDH2o3 .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-h6rzEPgzwPzDH2o3 .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-h6rzEPgzwPzDH2o3 .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-h6rzEPgzwPzDH2o3 .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-h6rzEPgzwPzDH2o3 .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-h6rzEPgzwPzDH2o3 .marker{fill:#333333;stroke:#333333;}#mermaid-svg-h6rzEPgzwPzDH2o3 .marker.cross{stroke:#333333;}#mermaid-svg-h6rzEPgzwPzDH2o3 svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-h6rzEPgzwPzDH2o3 p{margin:0;}#mermaid-svg-h6rzEPgzwPzDH2o3 .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-h6rzEPgzwPzDH2o3 .cluster-label text{fill:#333;}#mermaid-svg-h6rzEPgzwPzDH2o3 .cluster-label span{color:#333;}#mermaid-svg-h6rzEPgzwPzDH2o3 .cluster-label span p{background-color:transparent;}#mermaid-svg-h6rzEPgzwPzDH2o3 .label text,#mermaid-svg-h6rzEPgzwPzDH2o3 span{fill:#333;color:#333;}#mermaid-svg-h6rzEPgzwPzDH2o3 .node rect,#mermaid-svg-h6rzEPgzwPzDH2o3 .node circle,#mermaid-svg-h6rzEPgzwPzDH2o3 .node ellipse,#mermaid-svg-h6rzEPgzwPzDH2o3 .node polygon,#mermaid-svg-h6rzEPgzwPzDH2o3 .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-h6rzEPgzwPzDH2o3 .rough-node .label text,#mermaid-svg-h6rzEPgzwPzDH2o3 .node .label text,#mermaid-svg-h6rzEPgzwPzDH2o3 .image-shape .label,#mermaid-svg-h6rzEPgzwPzDH2o3 .icon-shape .label{text-anchor:middle;}#mermaid-svg-h6rzEPgzwPzDH2o3 .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-h6rzEPgzwPzDH2o3 .rough-node .label,#mermaid-svg-h6rzEPgzwPzDH2o3 .node .label,#mermaid-svg-h6rzEPgzwPzDH2o3 .image-shape .label,#mermaid-svg-h6rzEPgzwPzDH2o3 .icon-shape .label{text-align:center;}#mermaid-svg-h6rzEPgzwPzDH2o3 .node.clickable{cursor:pointer;}#mermaid-svg-h6rzEPgzwPzDH2o3 .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-h6rzEPgzwPzDH2o3 .arrowheadPath{fill:#333333;}#mermaid-svg-h6rzEPgzwPzDH2o3 .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-h6rzEPgzwPzDH2o3 .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-h6rzEPgzwPzDH2o3 .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-h6rzEPgzwPzDH2o3 .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-h6rzEPgzwPzDH2o3 .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-h6rzEPgzwPzDH2o3 .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-h6rzEPgzwPzDH2o3 .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-h6rzEPgzwPzDH2o3 .cluster text{fill:#333;}#mermaid-svg-h6rzEPgzwPzDH2o3 .cluster span{color:#333;}#mermaid-svg-h6rzEPgzwPzDH2o3 div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-h6rzEPgzwPzDH2o3 .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-h6rzEPgzwPzDH2o3 rect.text{fill:none;stroke-width:0;}#mermaid-svg-h6rzEPgzwPzDH2o3 .icon-shape,#mermaid-svg-h6rzEPgzwPzDH2o3 .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-h6rzEPgzwPzDH2o3 .icon-shape p,#mermaid-svg-h6rzEPgzwPzDH2o3 .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-h6rzEPgzwPzDH2o3 .icon-shape .label rect,#mermaid-svg-h6rzEPgzwPzDH2o3 .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-h6rzEPgzwPzDH2o3 .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-h6rzEPgzwPzDH2o3 .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-h6rzEPgzwPzDH2o3 :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} config.py
默认参数
main.py
流程编排
命令行参数
fire解析
data/dataset.py
DogCat
models/
ResNet34等
utils/visualize.py
Visualizer
checkpoints/
模型权重
result.csv
三、__init__.py:让目录变成包
结构里几乎每个子目录都有 __init__.py------有它,目录才是Python包,别的模块才能从中import。它可以是空文件,但也可以帮我们缩短导入路径。比如在 data/__init__.py 里写:
python
from .dataset import DogCat
那么 main.py 里就可以直接 from data import DogCat,而不必写全 from data.dataset import DogCat。models包同理,这个小技巧是后面"字符串选模型"的基础。
四、config.py:所有可配置项集中管理
调参最忌讳参数散落在代码各处。把它们全部收进一个类,给足默认值和注释:
python
# config.py
import warnings
import torch as t
class DefaultConfig(object):
env = 'default' # visdom环境名
vis_port = 8097 # visdom端口
model = 'SqueezeNet' # 使用的模型,名字必须与models/__init__.py中一致
train_data_root = './data/train/' # 训练集路径
test_data_root = './data/test/' # 测试集路径
load_model_path = None # 预训练权重路径,None表示不加载
batch_size = 32 # batch大小
use_gpu = True # 是否用GPU
num_workers = 4 # DataLoader进程数
print_freq = 20 # 每N个batch打印/画图一次
debug_file = '/tmp/debug' # 该文件存在则进入调试模式(后文详解)
result_file = 'result.csv' # 测试结果输出
max_epoch = 10
lr = 0.001 # 初始学习率
lr_decay = 0.5 # 损失不降时 lr = lr * lr_decay
weight_decay = 0e-5 # 权重衰减
def _parse(self, kwargs):
"""根据字典kwargs更新配置参数"""
for k, v in kwargs.items():
if not hasattr(self, k):
warnings.warn("Warning: opt has not attribut %s" % k)
setattr(self, k, v)
# 根据use_gpu统一决定device,后续代码只认opt.device
opt.device = t.device('cuda') if opt.use_gpu else t.device('cpu')
print('user config:')
for k, v in self.__class__.__dict__.items():
if not k.startswith('_'):
print(k, getattr(self, k))
opt = DefaultConfig()
_parse 是点睛之笔:命令行传进来的参数以字典形式覆盖默认值,拼错参数名会立刻warning提示,最后把生效配置完整打印一遍------每次实验日志开头都留有一份参数快照,事后对比实验时非常有用。
有人会问:为什么不用标准库的 argparse?当然可以,但每加一个参数都要写一行 parser.add_argument('--lr', type=float, default=0.001, help='...'),冗长且在Jupyter/IPython里很难交互调试。一个纯Python类直观得多,这是作者的个人偏好,供参考。
五、data/dataset.py:一个类管三种数据集
Kaggle只给了训练集和测试集,实践中还要从训练集切出验证集来监控过拟合。三种数据的差异有两处:划分方式 不同、预处理 不同(训练集要数据增强,验证/测试集不能加随机性)。与其写三个Dataset类,不如用一个 mode 参数区分:
python
# data/dataset.py
import os
from PIL import Image
from torch.utils import data
from torchvision import transforms as T
class DogCat(data.Dataset):
def __init__(self, root, transforms=None, mode=None):
"""
获取所有图片路径,并按 train/val/test 划分数据
mode ∈ ["train", "test", "val"]
"""
assert mode in ["train", "test", "val"]
self.mode = mode
imgs = [os.path.join(root, img) for img in os.listdir(root)]
# 测试集文件名形如 1000.jpg,训练集形如 cat.1000.jpg,排序键不同
if self.mode == "test":
imgs = sorted(imgs, key=lambda x: int(x.split('.')[-2].split('/')[-1]))
else:
imgs = sorted(imgs, key=lambda x: int(x.split('.')[-2]))
imgs_num = len(imgs)
# 训练:验证 = 7:3
if self.mode == "test": self.imgs = imgs
if self.mode == "train": self.imgs = imgs[:int(0.7 * imgs_num)]
if self.mode == "val": self.imgs = imgs[int(0.7 * imgs_num):]
if transforms is None:
# ImageNet统计量,用预训练模型必须配这组均值方差
normalize = T.Normalize(mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225])
if self.mode == "test" or self.mode == "val":
# 验证/测试:确定性变换,保证结果可复现
self.transforms = T.Compose([
T.Resize(224),
T.CenterCrop(224),
T.ToTensor(),
normalize
])
else:
# 训练:随机裁剪+随机水平翻转做数据增强
self.transforms = T.Compose([
T.Resize(256),
T.RandomResizedCrop(224),
T.RandomHorizontalFlip(),
T.ToTensor(),
normalize
])
def __getitem__(self, index):
"""
返回一张图片的数据;
训练/验证集:label为 1(dog) / 0(cat)
测试集:label为图片编号,如 1000.jpg 返回 1000(提交CSV要用)
"""
img_path = self.imgs[index]
if self.mode == "test":
label = int(self.imgs[index].split('.')[-2].split('/')[-1])
else:
label = 1 if 'dog' in img_path.split('/')[-1] else 0
data = Image.open(img_path)
data = self.transforms(data)
return data, label
def __len__(self):
return len(self.imgs)
几个设计点:
- 划分前先排序 。
os.listdir返回顺序不确定,不排序的话,两次运行切出来的训练/验证集可能不同,验证指标就没有可比性了。 - 费时操作放
__getitem__。读图、解码这些IO都放在这里,配合DataLoader的num_workers多进程并行,主进程几乎不等数据。 - 测试集的label是编号 。因为提交文件需要
id,这样dataloader直接吐出来,测试代码不用再解析文件名。 - 数据清洗这类重活建议单独写脚本预处理,不要塞进Dataset里。
使用方式与前面章节讲的一致:
python
train_data = DogCat(opt.train_data_root, mode="train")
train_dataloader = DataLoader(train_data, opt.batch_size,
shuffle=True, num_workers=opt.num_workers)
for ii, (data, label) in enumerate(train_dataloader):
...
六、models/:BasicModule与动态模型选择
6.1 BasicModule:给所有模型统一加save/load
模型多了以后,保存、加载权重是每个模型都要的功能。与其在每个模型里复制粘贴,不如封装一个基类:
python
# models/basic_module.py
import torch as t
import time
class BasicModule(t.nn.Module):
"""封装nn.Module,提供save和load两个方法"""
def __init__(self):
super(BasicModule, self).__init__()
self.model_name = str(type(self)) # 默认模型名
def load(self, path):
"""加载指定路径的模型权重"""
self.load_state_dict(t.load(path))
def save(self, name=None):
"""保存模型,默认文件名 = 模型名 + 时间,如 resnet34_0710_23_57_29.pth"""
if name is None:
prefix = 'checkpoints/' + self.model_name + '_'
name = time.strftime(prefix + '%m%d_%H_%M_%S.pth')
t.save(self.state_dict(), name)
return name
def get_optimizer(self, lr, weight_decay):
return t.optim.Adam(self.parameters(), lr=lr, weight_decay=weight_decay)
save 用"模型名+时间戳"自动命名,权重永远不会互相覆盖,回溯任何一次实验都有据可查。get_optimizer 也放进了基类------默认优化全部参数,子类可以覆写它,这就是下面SqueezeNet微调的关键。
6.2 两个模型:微调的SqueezeNet与手写的ResNet34
SqueezeNet直接拿torchvision的预训练模型改分类头:
python
# models/squeezenet.py
from torchvision.models import squeezenet1_1
from models.basic_module import BasicModule
from torch import nn
from torch.optim import Adam
class SqueezeNet(BasicModule):
def __init__(self, num_classes=2):
super(SqueezeNet, self).__init__()
self.model_name = 'squeezenet'
self.model = squeezenet1_1(pretrained=True)
# 预训练模型是1000分类,换成2分类的分类头
self.model.num_classes = num_classes
self.model.classifier = nn.Sequential(
nn.Dropout(p=0.5),
nn.Conv2d(512, num_classes, 1),
nn.ReLU(inplace=True),
nn.AvgPool2d(13, stride=1)
)
def forward(self, x):
return self.model(x)
def get_optimizer(self, lr, weight_decay):
# 微调策略:只训练新换的分类头,冻结前面的特征提取部分
return Adam(self.model.classifier.parameters(), lr,
weight_decay=weight_decay)
注意覆写的 get_optimizer:只把 classifier 的参数交给优化器,前面ImageNet学来的特征提取层保持不动。这是小数据集微调最稳妥的起手式。
ResNet34则演示了从零搭网络的推荐姿势------子模块化 + 函数生成重复结构:
python
# models/resnet34.py
from .basic_module import BasicModule
from torch import nn
from torch.nn import functional as F
class ResidualBlock(nn.Module):
"""子module:残差块"""
def __init__(self, inchannel, outchannel, stride=1, shortcut=None):
super(ResidualBlock, self).__init__()
self.left = nn.Sequential(
nn.Conv2d(inchannel, outchannel, 3, stride, 1, bias=False),
nn.BatchNorm2d(outchannel),
nn.ReLU(inplace=True),
nn.Conv2d(outchannel, outchannel, 3, 1, 1, bias=False),
nn.BatchNorm2d(outchannel))
self.right = shortcut # 维度不一致时用1x1卷积调整
def forward(self, x):
out = self.left(x)
residual = x if self.right is None else self.right(x)
out += residual # 残差连接
return F.relu(out)
class ResNet34(BasicModule):
"""主module:ResNet34 = 前处理 + 4个layer + 全连接
每个layer由多个ResidualBlock组成,用_make_layer函数生成"""
def __init__(self, num_classes=2):
super(ResNet34, self).__init__()
self.model_name = 'resnet34'
self.pre = nn.Sequential(
nn.Conv2d(3, 64, 7, 2, 3, bias=False),
nn.BatchNorm2d(64),
nn.ReLU(inplace=True),
nn.MaxPool2d(3, 2, 1))
# 4个layer分别含3、4、6、3个残差块
self.layer1 = self._make_layer(64, 64, 3, 1, is_shortcut=False)
self.layer2 = self._make_layer(64, 128, 4, 2)
self.layer3 = self._make_layer(128, 256, 6, 2)
self.layer4 = self._make_layer(256, 512, 3, 2)
self.fc = nn.Linear(512, num_classes)
def _make_layer(self, inchannel, outchannel, block_num, stride, is_shortcut=True):
if is_shortcut:
shortcut = nn.Sequential(
nn.Conv2d(inchannel, outchannel, 1, stride, bias=False),
nn.BatchNorm2d(outchannel))
else:
shortcut = None
layers = [ResidualBlock(inchannel, outchannel, stride, shortcut)]
for i in range(1, block_num):
layers.append(ResidualBlock(outchannel, outchannel))
return nn.Sequential(*layers)
def forward(self, x):
x = self.pre(x)
x = self.layer1(x)
x = self.layer2(x)
x = self.layer3(x)
x = self.layer4(x)
x = F.avg_pool2d(x, 7)
x = x.view(x.size(0), -1)
return self.fc(x)
模型定义的三条经验:尽量用 nn.Sequential;常用结构封装成子module(如ResidualBlock);重复且有规律的结构用函数生成(如 _make_layer)。
6.3 字符串选模型:getattr的妙用
在 models/__init__.py 里注册模型:
python
from .squeezenet import SqueezeNet
from .resnet34 import ResNet34
主程序里就可以用字符串动态实例化:
python
import models
model = getattr(models, opt.model)() # opt.model = 'ResNet34' 或 'SqueezeNet'
这一行是整套架构里复用价值最高的技巧:换模型只需改命令行参数 --model='ResNet34',主程序零改动 。以后新增模型,只要在models目录加文件并在 __init__.py 里加一行import,无需任何if-else。
七、utils/visualize.py:封装visdom做训练监控
visdom是PyTorch生态常用的可视化工具(先 python -m visdom.server 启动服务)。原生接口每次画点都要传窗口名、坐标等一堆参数,封装一层用起来舒服得多:
python
# utils/visualize.py(节选)
import visdom
import time
import numpy as np
class Visualizer(object):
"""封装visdom的基本操作,仍可通过self.vis.function调用原生接口"""
def __init__(self, env='default', **kwargs):
self.vis = visdom.Visdom(env=env, **kwargs)
self.index = {} # 记录每条曲线画到第几个点,充当横坐标
self.log_text = ''
def plot(self, name, y, **kwargs):
"""self.plot('loss', 1.00) ------ 自动追加到名为name的曲线"""
x = self.index.get(name, 0)
self.vis.line(Y=np.array([y]), X=np.array([x]), win=name,
opts=dict(title=name),
update=None if x == 0 else 'append', **kwargs)
self.index[name] = x + 1
def img(self, name, img_, **kwargs):
"""self.img('input_img', t.Tensor(3, 64, 64))"""
self.vis.images(img_.cpu().numpy(), win=name,
opts=dict(title=name), **kwargs)
def log(self, info, win='log_text'):
"""带时间戳的文字日志,如 self.log({'loss':1, 'lr':0.0001})"""
self.log_text += ('[{time}] {info} <br>'.format(
time=time.strftime('%m%d_%H%M%S'), info=info))
self.vis.text(self.log_text, win)
def __getattr__(self, name):
# 未定义的方法直接转发给原生visdom
return getattr(self.vis, name)
封装的核心是 plot:内部用 self.index 字典维护每条曲线的横坐标,调用方只管 vis.plot('loss', loss值),一行代码追加一个点。__getattr__ 兜底转发保证封装不损失任何原生功能。
八、main.py:fire命令行 + 训练/验证/测试
8.1 fire:三行代码的命令行接口
Google开源的fire库能把任意Python函数直接暴露成命令行命令:
python
# example.py
import fire
def add(x, y):
return x + y
if __name__ == '__main__':
fire.Fire()
bash
python example.py add 1 2 # 执行 add(1, 2)
python example.py add --x=1 --y=2
于是 main.py 的骨架就是四个函数:train(训练)、val(验证,内部调用)、test(推理)、help(帮助),末尾一句 fire.Fire()。执行 python main.py train --lr=0.01 时,fire会调用 train(lr=0.01),kwargs 再交给 opt._parse 覆盖默认配置------命令行、配置文件、程序逻辑就这样串起来了。
8.2 训练函数:五步流程
python
# main.py
from config import opt
import os
import torch as t
import models
from data.dataset import DogCat
from torch.utils.data import DataLoader
from torchnet import meter
from utils.visualize import Visualizer
from tqdm import tqdm
def train(**kwargs):
opt._parse(kwargs) # 命令行参数覆盖默认配置
vis = Visualizer(opt.env, port=opt.vis_port)
# step1: 模型------字符串动态实例化
model = getattr(models, opt.model)()
if opt.load_model_path:
model.load(opt.load_model_path)
model.to(opt.device)
# step2: 数据
train_data = DogCat(opt.train_data_root, mode="train")
val_data = DogCat(opt.train_data_root, mode="val")
train_dataloader = DataLoader(train_data, opt.batch_size,
shuffle=True, num_workers=opt.num_workers)
val_dataloader = DataLoader(val_data, opt.batch_size,
shuffle=False, num_workers=opt.num_workers)
# step3: 损失函数和优化器(优化器由模型自己决定,微调模型只优化分类头)
criterion = t.nn.CrossEntropyLoss()
lr = opt.lr
optimizer = model.get_optimizer(lr, opt.weight_decay)
# step4: 统计指标------平滑损失 + 混淆矩阵
loss_meter = meter.AverageValueMeter()
confusion_matrix = meter.ConfusionMeter(2)
previous_loss = 1e10
# step5: 训练主循环
for epoch in range(opt.max_epoch):
loss_meter.reset()
confusion_matrix.reset()
for ii, (data, label) in tqdm(enumerate(train_dataloader)):
input = data.to(opt.device)
target = label.to(opt.device)
optimizer.zero_grad()
score = model(input)
loss = criterion(score, target)
loss.backward()
optimizer.step()
# 更新统计指标;detach一下更安全,避免统计代码挂进计算图
loss_meter.add(loss.item())
confusion_matrix.add(score.detach(), target.detach())
if (ii + 1) % opt.print_freq == 0:
vis.plot('loss', loss_meter.value()[0])
# 存在debug标识文件则进入调试模式(见第九节)
if os.path.exists(opt.debug_file):
import ipdb
ipdb.set_trace()
model.save() # 每个epoch存一次权重
# 验证 + 可视化
val_cm, val_accuracy = val(model, val_dataloader)
vis.plot('val_accuracy', val_accuracy)
vis.log("epoch:{epoch},lr:{lr},loss:{loss},train_cm:{train_cm},val_cm:{val_cm}"
.format(epoch=epoch, loss=loss_meter.value()[0],
val_cm=str(val_cm.value()),
train_cm=str(confusion_matrix.value()), lr=lr))
# 损失不再下降则衰减学习率
if loss_meter.value()[0] > previous_loss:
lr = lr * opt.lr_decay
# 直接改param_group,不重建优化器,保留momentum等状态
for param_group in optimizer.param_groups:
param_group['lr'] = lr
previous_loss = loss_meter.value()[0]
几处值得展开:
- meter工具 :来自torchnet。
AverageValueMeter统计一个epoch内损失的均值方差,比看单个batch的loss曲线平滑得多;ConfusionMeter统计混淆矩阵------比如50张狗有35张判对、15张误判成猫。样本比例不均衡时,准确率会骗人,混淆矩阵不会。 - 学习率衰减策略 :epoch结束时若平均损失比上一轮还高,就把学习率乘0.5。修改方式是直接遍历
optimizer.param_groups改lr字段,而不是新建optimizer------后者会把Adam积累的动量信息清零。 - 每个epoch都save :配合时间戳文件名,训练异常中断随时可以
--load-model-path恢复。
8.3 验证:eval模式切换是重点
python
@t.no_grad()
def val(model, dataloader):
"""计算模型在验证集上的准确率等信息"""
model.eval() # 切验证模式:影响BatchNorm/Dropout行为
confusion_matrix = meter.ConfusionMeter(2)
for ii, (val_input, label) in tqdm(enumerate(dataloader)):
val_input = val_input.to(opt.device)
score = model(val_input)
confusion_matrix.add(score.detach().squeeze(), label.type(t.LongTensor))
model.train() # 用完切回训练模式!
cm_value = confusion_matrix.value()
accuracy = 100. * (cm_value[0][0] + cm_value[1][1]) / (cm_value.sum())
return confusion_matrix, accuracy
两个必须成对出现的动作:进来 model.eval(),出去 model.train()。BatchNorm在eval模式用全局统计量、Dropout停止随机失活,忘记切换是新手最常见的"验证指标诡异"根源。整个函数再套一个 @t.no_grad() 装饰器,不记录计算图,省显存也提速。
8.4 测试:输出提交文件
python
@t.no_grad()
def test(**kwargs):
opt._parse(kwargs)
# 加载模型
model = getattr(models, opt.model)().eval()
if opt.load_model_path:
model.load(opt.load_model_path)
model.to(opt.device)
# 加载测试数据
test_data = DogCat(opt.test_data_root, mode="test")
test_dataloader = DataLoader(test_data, batch_size=opt.batch_size,
shuffle=False, num_workers=opt.num_workers)
results = []
for ii, (data, path) in tqdm(enumerate(test_dataloader)):
input = data.to(opt.device)
score = model(input)
# softmax后取第1列 = 属于狗的概率
probability = t.nn.functional.softmax(score, dim=1)[:, 1].detach().tolist()
batch_results = [(path_.item(), probability_)
for path_, probability_ in zip(path, probability)]
results += batch_results
write_csv(results, opt.result_file)
return results
def write_csv(results, file_name):
import csv
with open(file_name, 'w') as f:
writer = csv.writer(f)
writer.writerow(['id', 'label'])
writer.writerows(results)
注意 softmax(score, dim=1)[:, 1]:模型输出的是两个类的原始分数,softmax归一化后第1列就是"是狗的概率",正好是Kaggle要的格式。
8.5 帮助函数与实际使用
help 函数用标准库 inspect.getsource 直接打印config类的源码------配置项改了,帮助信息自动同步,不用手工维护两份:
python
def help():
"""打印帮助信息:python main.py help"""
print("""
usage : python main.py <function> [--args=value]
<function> := train | test | help
example:
python main.py train --env='env0701' --lr=0.01
python main.py test --dataset='path/to/dataset/root/'
avaiable args:""")
from inspect import getsource
print(getsource(opt.__class__))
日常使用就是三条命令(fire会自动把 --train-data-root 转成 train_data_root):
bash
# 训练
python main.py train --train-data-root=data/train/ --lr=0.005 \
--batch-size=32 --model='ResNet34' --max-epoch=20
# 测试
python main.py test --test-data-root=data/test1 --model='ResNet34' \
--load-model-path='checkpoints/resnet34_00_23_05.pth'
# 帮助
python main.py help
九、实用调试技巧:训练到一半也能改学习率
训练代码里埋了一个不起眼却极好用的机关:
python
if os.path.exists(opt.debug_file):
import ipdb
ipdb.set_trace()
程序每隔 print_freq 个batch检查一次 /tmp/debug 文件是否存在。想调试时在终端执行 touch /tmp/debug,训练程序会在下一个检查点自动挂起,进入ipdb交互界面。此时你可以:
- 查看/修改学习率:
opt.lr = 0.001,再循环改optimizer.param_groups,不重启程序就完成调参; - 手动
model.save()存一份权重; - 用
s命令步进到model(input)内部,逐层查看输出的均值方差,定位数值异常出在哪一层; - 调完
rm /tmp/debug,输入c继续训练;想安全退出就输入quit------比Ctrl+C更稳妥,能保证DataLoader的多进程正确释放资源。
另外两条排错经验:模型训了几小时快结束时因为小bug崩了?如果是在IPython里用 %run 跑的,立刻 %debug 进入事后调试,现场执行 model.save() 抢救权重。遇到 CUDNN_STATUS_BAD_PARAM 这类天书报错,先把模型和数据挪回CPU跑一遍(model.cpu()(input.cpu())),CPU路径的报错信息友好得多------常见根源无非三类:类型不匹配(CrossEntropyLoss的target必须是LongTensor)、数据忘了搬到GPU(多个module放进list不会被 .cuda() 转移,要用 nn.ModuleList)、张量形状不对。
总结
要点回顾:
- 项目五块分离:config.py(参数)、data/(数据)、models/(网络)、utils/(工具)、main.py(编排),checkpoints/存权重。
- 一个Dataset类用
mode参数管理train/val/test:划分前先排序保证可复现,训练集加数据增强,验证/测试集只做确定性变换。 BasicModule统一封装save(时间戳命名)/load/get_optimizer;微调模型覆写get_optimizer只训练分类头。getattr(models, opt.model)()用字符串选模型,新增模型不改主程序。- fire三行代码搞定命令行,
opt._parse(kwargs)完成参数覆盖并打印配置快照。 - 训练循环五步:模型→数据→损失/优化器→meter统计→循环训练;验证必须
model.eval()/model.train()成对出现;损失不降则原地衰减学习率。 - debug_file + ipdb实现"运行中调试":不重启就能改学习率、存模型、逐层检查网络输出。
下一篇预告:《GAN实战:用DCGAN生成动漫头像》------把本篇这套工程骨架直接套到生成对抗网络上,从GAN博弈原理讲到转置卷积生成器、判别器的逐行实现,以及生成器/判别器交替训练的种种技巧。