《深度学习框架PyTorch入门与实践》系列:11-实战猫狗大战之可复用的PyTorch项目架构

实战猫狗大战:可复用的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)。

文章目录

一、任务与目标:Kaggle猫狗二分类

"Dogs vs. Cats"是Kaggle上的入门经典:训练集25000张图片混放在一个文件夹里,文件名自带标签,格式为 <category>.<num>.jpg,比如 cat.10000.jpgdog.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_groupslr 字段,而不是新建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)、张量形状不对。

总结

要点回顾:

  1. 项目五块分离:config.py(参数)、data/(数据)、models/(网络)、utils/(工具)、main.py(编排),checkpoints/存权重。
  2. 一个Dataset类用 mode 参数管理train/val/test:划分前先排序保证可复现,训练集加数据增强,验证/测试集只做确定性变换。
  3. BasicModule 统一封装save(时间戳命名)/load/get_optimizer;微调模型覆写 get_optimizer 只训练分类头。
  4. getattr(models, opt.model)() 用字符串选模型,新增模型不改主程序。
  5. fire三行代码搞定命令行,opt._parse(kwargs) 完成参数覆盖并打印配置快照。
  6. 训练循环五步:模型→数据→损失/优化器→meter统计→循环训练;验证必须 model.eval()/model.train() 成对出现;损失不降则原地衰减学习率。
  7. debug_file + ipdb实现"运行中调试":不重启就能改学习率、存模型、逐层检查网络输出。

下一篇预告:《GAN实战:用DCGAN生成动漫头像》------把本篇这套工程骨架直接套到生成对抗网络上,从GAN博弈原理讲到转置卷积生成器、判别器的逐行实现,以及生成器/判别器交替训练的种种技巧。

相关推荐
Bruce_Liuxiaowei1 小时前
从零到可运行:基于 Vue3 + FastAPI + DeepSeek-V3 的 AI 英语单词学习系统全栈实战
人工智能·python·学习·fastapi·全栈·智能体
kaixin_啊啊1 小时前
香精近红外总体步骤概览
人工智能·matlab·近红外
u0103055271 小时前
昇腾Model-Agent端云协同架构解析
人工智能
IT爱学堂1 小时前
尚硅谷Java+AI大模型应用开发革新版本 2025年3月
java·开发语言·人工智能
W_326001 小时前
Python-OpenCV图像像素与通道:通道拆分合并、深浅拷贝
图像处理·人工智能·python·opencv·机器学习
CIO_Alliance1 小时前
AI基础系列(1)| 向量、矩阵、张量在AI中分别扮演什么角色?
大数据·人工智能·线性代数·ai·矩阵·企业cio联盟·企业级ai化转型
通问AI1 小时前
人形机器人量产技术笔记:从5500台出货看规模化路径
人工智能
无凭2 小时前
字节Agent框架 DeerFlow 的可观测性(二):运行中的中间事件是怎样到达前端的?
人工智能
XLYcmy2 小时前
小红书 算法一面 四
深度学习·llm·sft·强化学习·多模态·grpo·奖励函数