• 👉 声明:本文为学习记录性文章,参考「365天深度学习训练营」相关内容整理。
  • 🍨 本文为🔗365天深度学习训练营 中的学习记录博客
  • 🍖 原作者:K同学啊


前言

本篇是我训练营的第8次学习,主要目标是使用 PyTorch 完成儿童胸部 X 光片肺炎识别,并重点学习经典卷积神经网络 ResNet-34。上一周(P7)学习了手动搭建 VGG-16 网络框架实现马铃薯病害识别,本周进一步学习ResNet(残差网络)

VGG-16 的信息主要沿着一条主干网络逐层向后传播;ResNet-34 除了卷积主路径,还增加了一条快捷连接,让输入可以跨过若干卷积层直接与主路径结果相加。残差连接可以让信息和梯度更顺畅地在深层网络中传播,缓解网络加深后训练困难以及准确率退化的问题。

ResNet-34 标准结构

ResNet-34 由 33 个卷积层 + 1 个全连接层组成,共 34 层可训练层。它使用BasicBlock作为基本构建单元,每个 BasicBlock 包含两个 3×3 卷积层和一个残差连接(shortcut)。网络分为 5 个阶段:conv1(7×7 大卷积核下采样)、conv2_x、conv3_x、conv4_x、conv5_x(每个阶段由多个 BasicBlock 堆叠),最后接全局平均池化和全连接层。

本周使用的是 儿童胸部 X 光肺炎数据集,包含 8530 张胸部 X 光片,任务是根据 X 光片判断患者是正常(NORMAL)还是患有肺炎(PNEUMONIA)

P8 学习目标

  1. 跑通 ResNet34 代码,并完成胸部 X 光片二分类训练。
  2. 根据官方代码的结构输出 + 代码结构图,手动搭建 ResNet34 算法网络(见四、方案二

感谢 K同学啊 老师的教学,以及 ChatGPT 和 Kimi。


一、准备工作

1. 设置运行设备:GPU 或 CPU

import torch
import torch.nn as nn
import torchvision.transforms as transforms
import torchvision
from torchvision import transforms, datasets
import os, PIL, pathlib, warnings

warnings.filterwarnings("ignore")             # 忽略警告信息

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
device
device(type='cuda')

本周我成功配置了 AMD GPU 加速环境。我的电脑是AMD GPU,在前七周的学习中,我一直使用 CPU 版本的 PyTorch 进行训练,本周我根据 AMD 官方文档 成功安装了 ROCm 驱动和 GPU 版本的 PyTorch,具体步骤如下:

  1. 安装 AMD ROCm 驱动:按照官方文档安装适用于 Radeon 显卡的 ROCm for Windows 驱动;
  2. 创建新的 shelter 环境:新建了一个独立的 Python 虚拟环境(命名为 shelter),避免与之前的 CPU 环境冲突;
  3. 安装 GPU 版 PyTorch:在 shelter 环境中安装了支持 ROCm 的 PyTorch 版本。

测试了第7周的可以运行,但本周的内容会有报错,借助AI修改了一下


2. 关于肺炎 X 光数据集

本周使用的是儿童胸部 X 光肺炎数据集(Chest X-Ray Images),包含 8530 张胸部 X 光片,用于检测肺炎。文件夹结构如下:

data/
├── NORMAL/          # 正常胸部X光片
└── PNEUMONIA/       # 肺炎胸部X光片

3. 导入本地数据集

import os, PIL, random, pathlib

data_dir = './data/'
data_dir = pathlib.Path(data_dir)

data_paths  = list(data_dir.glob('*'))
classeNames = [str(path).split("\\")[1] for path in data_paths]
classeNames

P3-P8 周的数据加载方式完全一致,都是使用 pathlib.Path + datasets.ImageFolder 来读取本地图片。P7 周是马铃薯病害 3 分类,本周是肺炎 X 光 2 分类。


4. 数据预处理:transforms.Compose()

# 关于transforms.Compose的更多介绍可以参考:https://blog.csdn.net/qq_38251616/article/details/124878863
train_transforms = transforms.Compose([
    transforms.Resize([224, 224]),  # 将输入图片resize成统一尺寸
    transforms.ToTensor(),          # 将PIL Image或numpy.ndarray转换为tensor,并归一化到[0,1]之间
    transforms.Normalize(           # 标准化处理-->转换为标准正太分布(高斯分布),使模型更容易收敛
        mean=[0.485, 0.456, 0.406], 
        std=[0.229, 0.224, 0.225])  # 其中 mean=[0.485,0.456,0.406]与std=[0.229,0.224,0.225] 从数据集中随机抽样计算得到的。
])

test_transform = transforms.Compose([
    transforms.Resize([224, 224]),  # 将输入图片resize成统一尺寸
    transforms.ToTensor(),          # 将PIL Image或numpy.ndarray转换为tensor,并归一化到[0,1]之间
    transforms.Normalize(           # 标准化处理-->转换为标准正太分布(高斯分布),使模型更容易收敛
        mean=[0.485, 0.456, 0.406], 
        std=[0.229, 0.224, 0.225])  # 其中 mean=[0.485,0.456,0.406]与std=[0.229,0.224,0.225] 从数据集中随机抽样计算得到的。
])

total_data = datasets.ImageFolder("./data/", transform=train_transforms)
total_data
device(type='cuda')
Dataset ImageFolder
    Number of datapoints: 8530
    Root location: ./data/
    StandardTransform
Transform: Compose(
               Resize(size=[224, 224], interpolation=bilinear, max_size=None, antialias=True)
               ToTensor()
               Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
           )

Resize((224, 224)) 把所有原始图片统一调整为 224 × 224。ResNet-34 可以处理其他尺寸,但 torchvision 的 ImageNet 预训练权重通常以 224 × 224 图片作为标准输入,因此本周继续沿用这一尺寸。
ToTensor() 将 PIL 图片转换为 PyTorch Tensor,并把像素值从 0-255 缩放到 0-1。单张图片转换后的 shape 为[3, 224, 224]
mean=[0.485, 0.456, 0.406]std =[0.229, 0.224, 0.225]这是 ImageNet 数据集 RGB 三通道的常用均值和标准差。由于官方 ResNet-34 预训练权重是在 ImageNet 上训练的,使用相同的标准化方式可以让本周图片的数值分布更接近预训练阶段的输入。
本周的数据量是8530张,比P7周(2152张)大了很多,更多的数据对于迁移学习来说非常有利,因为预训练模型需要足够的数据来微调最后一层分类器。

X光片是灰度图像,为什么要用3通道的mean/std标准化?

虽然X光片本身是单通道灰度图像,但ImageFolder加载后会自动转换为3通道RGB格式(复制灰度值到3个通道)。ResNet-34的预训练模型是在ImageNet的3通道彩色图上训练的,所以保持3通道输入是必要的。使用ImageNet的mean/std标准化可以让X光数据的分布与预训练模型的训练数据分布更接近。


5. 使用 ImageFolder 自动生成标签

total_data.class_to_idx
{'NORMAL': 0, 'PNEUMONIA': 1}

ImageFolder 根据文件夹名称自动分配标签,NORMAL 对应索引 0PNEUMONIA 对应索引 1。这是一个二分类问题,最终模型输出 2 个类别的概率分数。


6. 划分训练集和测试集

train_size = int(0.8 * len(total_data))
test_size  = len(total_data) - train_size
train_dataset, test_dataset = torch.utils.data.random_split(total_data, [train_size, test_size])
train_dataset, test_dataset
(<torch.utils.data.dataset.Subset at 0x234507bb680>,
 <torch.utils.data.dataset.Subset at 0x234507bb560>)

这一段的意思是

  • train_size = int(0.8 * len(total_data)):训练集大小为总数据量的 80%。总数据量是 8530,所以训练集大小为 int(0.8 × 8530) = 6824
  • test_size = len(total_data) - train_size:测试集大小为剩余的 20%,即 8530 - 6824 = 1706
  • 使用 random_split 将数据集随机打乱后按 8:2 划分。

7. 创建 DataLoader 数据加载器

batch_size = 64

train_dl = torch.utils.data.DataLoader(train_dataset,
                                       batch_size=batch_size,
                                       shuffle=True,
                                       num_workers=4)
test_dl = torch.utils.data.DataLoader(test_dataset,
                                      batch_size=batch_size,
                                      shuffle=False,
                                      num_workers=4)

本周的 batch_size 从 32 增大到了 64,原因有二:

  1. 数据集更大(8530张 vs 2152张),使用更大的 batch 可以更好地利用 GPU 并行计算能力;
  2. ResNet-34 使用全局平均池化代替了大全连接层,参数量更少,内存占用更小,可以容纳更大的 batch。

num_workers=4 表示使用 4 个子进程加载数据,比前几周的 1 个更快,因为数据量大了,数据加载可能成为瓶颈。

shuffle=False 在测试集上设置为 False,保证测试结果的确定性。


8. 查看一个 batch 的数据格式

for X, y in test_dl:
    print("Shape of X [N, C, H, W]: ", X.shape)
    print("Shape of y: ", y.shape, y.dtype)
    break
Shape of X [N, C, H, W]:  torch.Size([64, 3, 224, 224])
Shape of y:  torch.Size([64]) torch.int64

输入 shape 可以拆开理解:

torch.Size([64, 3, 224, 224])
             ↑   ↑    ↑    ↑
             N   C    H    W
             │   │    │    └── 宽度:224 像素
             │   │    └─────── 高度:224 像素
             │   └──────────── 通道数:3(X光灰度图自动转3通道)
             └──────────────── batch_size,一批 64 张图片

二、ResNet-34 模型

1. 什么是 ResNet?

**ResNet(Residual Network,残差网络)是由微软研究院的何恺明等人在 2015 年提出的深度卷积神经网络架构。ResNet 在当年的 ImageNet 图像识别竞赛中取得了冠军,并引入了革命性的残差连接(Skip Connection)**概念,使得训练超过 100 层的深层网络成为可能。

2. 为什么需要 ResNet?

在较浅的神经网络中,继续增加卷积层通常可以增强特征提取能力。但当网络越来越深时,并不一定能得到更好的结果,反而可能出现:

  1. 梯度在传播过程中越来越弱,前面的层很难得到有效更新;
  2. 网络层数增加后训练误差反而升高;
  3. 更深的网络理论表达能力更强,却比浅层网络更难优化。

这种现象不能简单理解为过拟合,因为有时连训练集上的误差都会变差,它通常被称为退化问题(Degradation Problem)

ResNet 的解决思路不是让每一组卷积层直接学习完整映射 H(x),而是让它学习输入与目标之间的差值,也就是残差:

F(x) = H(x) - x

因此最终输出为:

H(x) = F(x) + x

其中:

  • x:残差块的原始输入;
  • F(x):两层卷积学习到的变化;
  • F(x) + x:把主路径结果与原输入相加;
  • 相加后再经过 ReLU 得到残差块输出。

如果当前最合适的映射接近恒等映射,网络只需要让 F(x) 接近 0,就可以保留输入信息。这通常比让多层卷积重新学习完整的 H(x)=x 更容易。


3. ResNet-34 的整体结构

输入:[3, 224, 224]
↓
Conv1:7×7, 64通道, stride=2
输出:[64, 112, 112]
↓
BatchNorm + ReLU
↓
MaxPool:3×3, stride=2
输出:[64, 56, 56]
↓
layer1:BasicBlock × 3,输出通道64
输出:[64, 56, 56]
↓
layer2:BasicBlock × 4,输出通道128,首块stride=2
输出:[128, 28, 28]
↓
layer3:BasicBlock × 6,输出通道256,首块stride=2
输出:[256, 14, 14]
↓
layer4:BasicBlock × 3,输出通道512,首块stride=2
输出:[512, 7, 7]
↓
AdaptiveAvgPool2d((1, 1))
输出:[512, 1, 1]
↓
Flatten
输出:[512]
↓
Linear(512, 2)
输出:2个类别分数

4. ResNet-34 的 shape 变化

阶段 核心操作 输出 shape(不写 batch)
输入 原始 RGB 图片 [3, 224, 224]
conv1 7×7, stride=2, padding=3 [64, 112, 112]
maxpool 3×3, stride=2, padding=1 [64, 56, 56]
layer1 3 个 BasicBlock [64, 56, 56]
layer2 4 个 BasicBlock,首块下采样 [128, 28, 28]
layer3 6 个 BasicBlock,首块下采样 [256, 14, 14]
layer4 3 个 BasicBlock,首块下采样 [512, 7, 7]
avgpool 自适应平均池化 [512, 1, 1]
flatten 展平 [512]
fc 二分类 [2]

可以看到,随着网络向后传播:

空间尺寸:224 → 112 → 56 → 28 → 14 → 7 → 1
通道数量:3 → 64 → 128 → 256 → 512

空间尺寸逐渐缩小,通道数逐渐增加。前面的层保留较多位置细节,后面的层使用更多通道表示更抽象的高级特征。


三、调用官方 ResNet-34 模型(方案一)

1. 调用模型

from torchvision.models import resnet34

# 加载预训练模型,并且对模型进行微调
model = resnet34(pretrained = True).to(device) # 加载预训练的resnet34模型

for param in model.parameters():
    param.requires_grad = False # 冻结模型的参数,这样子在训练的时候只训练最后一层的参数

# 修改模型的fc层,即(fc): Linear(in_features=512, out_features=2, bias=True)
model.fc = nn.Linear(512,len(classeNames)) # 修改resnet34模型中最后一层全连接层,输出目标类别个数
model.to(device)  
model
Downloading: "https://download.pytorch.org/models/resnet34-b627a593.pth" to C:\Users\admin/.cache\torch\hub\checkpoints\resnet34-b627a593.pth
100%|██████████| 83.3M/83.3M [00:12<00:00, 7.04MB/s]
ResNet(
  (conv1): Conv2d(3, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)
  (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
  (relu): ReLU(inplace=True)
  (maxpool): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False)
  (layer1): Sequential(
    (0): BasicBlock(
      (conv1): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
      (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
      (relu): ReLU(inplace=True)
      (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
      (bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    (1): BasicBlock(
      (conv1): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
      (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
      (relu): ReLU(inplace=True)
      (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
      (bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    (2): BasicBlock(
      (conv1): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
      (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
      (relu): ReLU(inplace=True)
      (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
      (bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
  )
  (layer2): Sequential(
    (0): BasicBlock(
      (conv1): Conv2d(64, 128, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1), bias=False)
      (bn1): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
      (relu): ReLU(inplace=True)
      (conv2): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
      (bn2): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
      (downsample): Sequential(
        (0): Conv2d(64, 128, kernel_size=(1, 1), stride=(2, 2), bias=False)
        (1): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
      )
    )
    (1): BasicBlock(
      (conv1): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
      (bn1): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
      (relu): ReLU(inplace=True)
      (conv2): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
      (bn2): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    (2): BasicBlock(
      (conv1): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
      (bn1): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
      (relu): ReLU(inplace=True)
      (conv2): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
      (bn2): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    (3): BasicBlock(
      (conv1): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
      (bn1): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
      (relu): ReLU(inplace=True)
      (conv2): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
      (bn2): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
  )
  (layer3): Sequential(
    (0): BasicBlock(
      (conv1): Conv2d(128, 256, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1), bias=False)
      (bn1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
      (relu): ReLU(inplace=True)
      (conv2): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
      (bn2): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
      (downsample): Sequential(
        (0): Conv2d(128, 256, kernel_size=(1, 1), stride=(2, 2), bias=False)
        (1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
      )
    )
    (1): BasicBlock(
      (conv1): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
      (bn1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
      (relu): ReLU(inplace=True)
      (conv2): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
      (bn2): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    (2): BasicBlock(
      (conv1): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
      (bn1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
      (relu): ReLU(inplace=True)
      (conv2): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
      (bn2): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    (3): BasicBlock(
      (conv1): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
      (bn1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
      (relu): ReLU(inplace=True)
      (conv2): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
      (bn2): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    (4): BasicBlock(
      (conv1): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
      (bn1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
      (relu): ReLU(inplace=True)
      (conv2): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
      (bn2): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    (5): BasicBlock(
      (conv1): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
      (bn1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
      (relu): ReLU(inplace=True)
      (conv2): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
      (bn2): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
  )
  (layer4): Sequential(
    (0): BasicBlock(
      (conv1): Conv2d(256, 512, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1), bias=False)
      (bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
      (relu): ReLU(inplace=True)
      (conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
      (bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
      (downsample): Sequential(
        (0): Conv2d(256, 512, kernel_size=(1, 1), stride=(2, 2), bias=False)
        (1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
      )
    )
    (1): BasicBlock(
      (conv1): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
      (bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
      (relu): ReLU(inplace=True)
      (conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
      (bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    (2): BasicBlock(
      (conv1): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
      (bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
      (relu): ReLU(inplace=True)
      (conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
      (bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
  )
  (avgpool): AdaptiveAvgPool2d(output_size=(1, 1))
  (fc): Linear(in_features=512, out_features=2, bias=True)
)

这段代码的核心步骤解析:

步骤 1:加载预训练模型

resnet34(pretrained=True):加载在 ImageNet 数据集上预训练好的 ResNet-34 模型。

步骤 2:冻结模型参数

param.requires_grad = False冻结所有参数,在训练时不计算这些参数的梯度,也不更新这些参数。预训练模型已经在 ImageNet 上学习了边缘、纹理、轮廓和较复杂物体特征,本周只需要重新训练最后的分类层。

步骤 3:修改最后一层全连接层

ResNet-34 原来的最后一层是 Linear(512, 1000),输出 1000 个类别。我们的任务只有 2 个类别,所以需要把 model.fc 替换为 Linear(512, 2)

ResNet-34 与 VGG-16 的一个重要区别是,ResNet-34 使用 AdaptiveAvgPool2d 将特征图压缩到 1×1,然后接一个全连接层,因此不需要像 VGG-16 那样展平大量参数。这也是 ResNet-34 参数量远小于 VGG-16 的主要原因。


2. 查看模型详情

原本使用 torchsummary.summary(model, (3, 224, 224)) 查看模型结构时,在 AMD GPU 环境下出现了 miopenStatusUnknownError,GPT告诉我Windows 上的 ROCm PyTorch,而 AMD 也注明 Windows 版本尚未支持完整的 ROCm 软件栈,所以部分第三方工具触发 MIOpen 运算时可能出现兼容问题,因此改用更新的 torchinfo 工具,并暂时将模型放到 CPU 上完成结构统计。torchinfo 的 input_size 需要写出完整的输入 shape,即 (batch_size, channel, height, width),所以这里设置为 (1, 3, 224, 224)。查看完成后,再使用 model.to(device) 将模型放回原来的运行设备,不影响后续使用 GPU 训练。

# 统计模型参数量以及其他指标
from torchinfo import summary

model = model.cpu()
summary(model, input_size=(1, 3, 224, 224), device="cpu")
==========================================================================================
Layer (type:depth-idx)                   Output Shape              Param #
==========================================================================================
ResNet                                   [1, 2]                    --
├─Conv2d: 1-1                            [1, 64, 112, 112]         (9,408)
├─BatchNorm2d: 1-2                       [1, 64, 112, 112]         (128)
├─ReLU: 1-3                              [1, 64, 112, 112]         --
├─MaxPool2d: 1-4                         [1, 64, 56, 56]           --
├─Sequential: 1-5                        [1, 64, 56, 56]           --
│    └─BasicBlock: 2-1                   [1, 64, 56, 56]           --
│    │    └─Conv2d: 3-1                  [1, 64, 56, 56]           (36,864)
│    │    └─BatchNorm2d: 3-2             [1, 64, 56, 56]           (128)
│    │    └─ReLU: 3-3                    [1, 64, 56, 56]           --
│    │    └─Conv2d: 3-4                  [1, 64, 56, 56]           (36,864)
│    │    └─BatchNorm2d: 3-5             [1, 64, 56, 56]           (128)
│    │    └─ReLU: 3-6                    [1, 64, 56, 56]           --
│    └─BasicBlock: 2-2                   [1, 64, 56, 56]           --
│    │    └─Conv2d: 3-7                  [1, 64, 56, 56]           (36,864)
│    │    └─BatchNorm2d: 3-8             [1, 64, 56, 56]           (128)
│    │    └─ReLU: 3-9                    [1, 64, 56, 56]           --
│    │    └─Conv2d: 3-10                 [1, 64, 56, 56]           (36,864)
│    │    └─BatchNorm2d: 3-11            [1, 64, 56, 56]           (128)
│    │    └─ReLU: 3-12                   [1, 64, 56, 56]           --
│    └─BasicBlock: 2-3                   [1, 64, 56, 56]           --
│    │    └─Conv2d: 3-13                 [1, 64, 56, 56]           (36,864)
│    │    └─BatchNorm2d: 3-14            [1, 64, 56, 56]           (128)
│    │    └─ReLU: 3-15                   [1, 64, 56, 56]           --
│    │    └─Conv2d: 3-16                 [1, 64, 56, 56]           (36,864)
│    │    └─BatchNorm2d: 3-17            [1, 64, 56, 56]           (128)
│    │    └─ReLU: 3-18                   [1, 64, 56, 56]           --
├─Sequential: 1-6                        [1, 128, 28, 28]          --
│    └─BasicBlock: 2-4                   [1, 128, 28, 28]          --
│    │    └─Conv2d: 3-19                 [1, 128, 28, 28]          (73,728)
│    │    └─BatchNorm2d: 3-20            [1, 128, 28, 28]          (256)
│    │    └─ReLU: 3-21                   [1, 128, 28, 28]          --
│    │    └─Conv2d: 3-22                 [1, 128, 28, 28]          (147,456)
│    │    └─BatchNorm2d: 3-23            [1, 128, 28, 28]          (256)
│    │    └─Sequential: 3-24             [1, 128, 28, 28]          (8,448)
│    │    └─ReLU: 3-25                   [1, 128, 28, 28]          --
│    └─BasicBlock: 2-5                   [1, 128, 28, 28]          --
│    │    └─Conv2d: 3-26                 [1, 128, 28, 28]          (147,456)
│    │    └─BatchNorm2d: 3-27            [1, 128, 28, 28]          (256)
│    │    └─ReLU: 3-28                   [1, 128, 28, 28]          --
│    │    └─Conv2d: 3-29                 [1, 128, 28, 28]          (147,456)
│    │    └─BatchNorm2d: 3-30            [1, 128, 28, 28]          (256)
│    │    └─ReLU: 3-31                   [1, 128, 28, 28]          --
│    └─BasicBlock: 2-6                   [1, 128, 28, 28]          --
│    │    └─Conv2d: 3-32                 [1, 128, 28, 28]          (147,456)
│    │    └─BatchNorm2d: 3-33            [1, 128, 28, 28]          (256)
│    │    └─ReLU: 3-34                   [1, 128, 28, 28]          --
│    │    └─Conv2d: 3-35                 [1, 128, 28, 28]          (147,456)
│    │    └─BatchNorm2d: 3-36            [1, 128, 28, 28]          (256)
│    │    └─ReLU: 3-37                   [1, 128, 28, 28]          --
│    └─BasicBlock: 2-7                   [1, 128, 28, 28]          --
│    │    └─Conv2d: 3-38                 [1, 128, 28, 28]          (147,456)
│    │    └─BatchNorm2d: 3-39            [1, 128, 28, 28]          (256)
│    │    └─ReLU: 3-40                   [1, 128, 28, 28]          --
│    │    └─Conv2d: 3-41                 [1, 128, 28, 28]          (147,456)
│    │    └─BatchNorm2d: 3-42            [1, 128, 28, 28]          (256)
│    │    └─ReLU: 3-43                   [1, 128, 28, 28]          --
├─Sequential: 1-7                        [1, 256, 14, 14]          --
│    └─BasicBlock: 2-8                   [1, 256, 14, 14]          --
│    │    └─Conv2d: 3-44                 [1, 256, 14, 14]          (294,912)
│    │    └─BatchNorm2d: 3-45            [1, 256, 14, 14]          (512)
│    │    └─ReLU: 3-46                   [1, 256, 14, 14]          --
│    │    └─Conv2d: 3-47                 [1, 256, 14, 14]          (589,824)
│    │    └─BatchNorm2d: 3-48            [1, 256, 14, 14]          (512)
│    │    └─Sequential: 3-49             [1, 256, 14, 14]          (33,280)
│    │    └─ReLU: 3-50                   [1, 256, 14, 14]          --
│    └─BasicBlock: 2-9                   [1, 256, 14, 14]          --
│    │    └─Conv2d: 3-51                 [1, 256, 14, 14]          (589,824)
│    │    └─BatchNorm2d: 3-52            [1, 256, 14, 14]          (512)
│    │    └─ReLU: 3-53                   [1, 256, 14, 14]          --
│    │    └─Conv2d: 3-54                 [1, 256, 14, 14]          (589,824)
│    │    └─BatchNorm2d: 3-55            [1, 256, 14, 14]          (512)
│    │    └─ReLU: 3-56                   [1, 256, 14, 14]          --
│    └─BasicBlock: 2-10                  [1, 256, 14, 14]          --
│    │    └─Conv2d: 3-57                 [1, 256, 14, 14]          (589,824)
│    │    └─BatchNorm2d: 3-58            [1, 256, 14, 14]          (512)
│    │    └─ReLU: 3-59                   [1, 256, 14, 14]          --
│    │    └─Conv2d: 3-60                 [1, 256, 14, 14]          (589,824)
│    │    └─BatchNorm2d: 3-61            [1, 256, 14, 14]          (512)
│    │    └─ReLU: 3-62                   [1, 256, 14, 14]          --
│    └─BasicBlock: 2-11                  [1, 256, 14, 14]          --
│    │    └─Conv2d: 3-63                 [1, 256, 14, 14]          (589,824)
│    │    └─BatchNorm2d: 3-64            [1, 256, 14, 14]          (512)
│    │    └─ReLU: 3-65                   [1, 256, 14, 14]          --
│    │    └─Conv2d: 3-66                 [1, 256, 14, 14]          (589,824)
│    │    └─BatchNorm2d: 3-67            [1, 256, 14, 14]          (512)
│    │    └─ReLU: 3-68                   [1, 256, 14, 14]          --
│    └─BasicBlock: 2-12                  [1, 256, 14, 14]          --
│    │    └─Conv2d: 3-69                 [1, 256, 14, 14]          (589,824)
│    │    └─BatchNorm2d: 3-70            [1, 256, 14, 14]          (512)
│    │    └─ReLU: 3-71                   [1, 256, 14, 14]          --
│    │    └─Conv2d: 3-72                 [1, 256, 14, 14]          (589,824)
│    │    └─BatchNorm2d: 3-73            [1, 256, 14, 14]          (512)
│    │    └─ReLU: 3-74                   [1, 256, 14, 14]          --
│    └─BasicBlock: 2-13                  [1, 256, 14, 14]          --
│    │    └─Conv2d: 3-75                 [1, 256, 14, 14]          (589,824)
│    │    └─BatchNorm2d: 3-76            [1, 256, 14, 14]          (512)
│    │    └─ReLU: 3-77                   [1, 256, 14, 14]          --
│    │    └─Conv2d: 3-78                 [1, 256, 14, 14]          (589,824)
│    │    └─BatchNorm2d: 3-79            [1, 256, 14, 14]          (512)
│    │    └─ReLU: 3-80                   [1, 256, 14, 14]          --
├─Sequential: 1-8                        [1, 512, 7, 7]            --
│    └─BasicBlock: 2-14                  [1, 512, 7, 7]            --
│    │    └─Conv2d: 3-81                 [1, 512, 7, 7]            (1,179,648)
│    │    └─BatchNorm2d: 3-82            [1, 512, 7, 7]            (1,024)
│    │    └─ReLU: 3-83                   [1, 512, 7, 7]            --
│    │    └─Conv2d: 3-84                 [1, 512, 7, 7]            (2,359,296)
│    │    └─BatchNorm2d: 3-85            [1, 512, 7, 7]            (1,024)
│    │    └─Sequential: 3-86             [1, 512, 7, 7]            (132,096)
│    │    └─ReLU: 3-87                   [1, 512, 7, 7]            --
│    └─BasicBlock: 2-15                  [1, 512, 7, 7]            --
│    │    └─Conv2d: 3-88                 [1, 512, 7, 7]            (2,359,296)
│    │    └─BatchNorm2d: 3-89            [1, 512, 7, 7]            (1,024)
│    │    └─ReLU: 3-90                   [1, 512, 7, 7]            --
│    │    └─Conv2d: 3-91                 [1, 512, 7, 7]            (2,359,296)
│    │    └─BatchNorm2d: 3-92            [1, 512, 7, 7]            (1,024)
│    │    └─ReLU: 3-93                   [1, 512, 7, 7]            --
│    └─BasicBlock: 2-16                  [1, 512, 7, 7]            --
│    │    └─Conv2d: 3-94                 [1, 512, 7, 7]            (2,359,296)
│    │    └─BatchNorm2d: 3-95            [1, 512, 7, 7]            (1,024)
│    │    └─ReLU: 3-96                   [1, 512, 7, 7]            --
│    │    └─Conv2d: 3-97                 [1, 512, 7, 7]            (2,359,296)
│    │    └─BatchNorm2d: 3-98            [1, 512, 7, 7]            (1,024)
│    │    └─ReLU: 3-99                   [1, 512, 7, 7]            --
├─AdaptiveAvgPool2d: 1-9                 [1, 512, 1, 1]            --
├─Linear: 1-10                           [1, 2]                    1,026
==========================================================================================
Total params: 21,285,698
Trainable params: 1,026
Non-trainable params: 21,284,672
Total mult-adds (Units.GIGABYTES): 3.66
==========================================================================================
Input size (MB): 0.60
Forward/backward pass size (MB): 59.81
Params size (MB): 85.14
Estimated Total Size (MB): 145.55
==========================================================================================
model = model.to(device)

使用 torchinfo 在 CPU 上对调用的官方 ResNet-34 模型进行结构统计。输入模拟数据的 shape 为 [1, 3, 224, 224],模型最终输出 shape 为 [1, 2],对应 NORMAL 和 PNEUMONIA 两个类别。模型共有 21,285,698 个参数,其中 21,284,672 个参数来自已经冻结的预训练特征提取网络,只有最后一个全连接层的 1,026 个参数参与训练。最后一层为 Linear(512, 2),其参数量为 512×2+2=1,026。输出中带括号的参数量表示这些层已被冻结,不会在训练过程中更新;最后的全连接层参数没有括号,说明它仍然可以训练。这说明预训练模型的参数冻结和分类层替换均已成功完成。

模块 参数量 说明
conv1 + bn1 + maxpool 9,536 初始下采样层
layer1(3个BasicBlock,64通道) 110,976 conv2_x
layer2(4个BasicBlock,128通道) 885,120 conv3_x,含1个下采样Block
layer3(6个BasicBlock,256通道) 8,518,656 conv4_x,含1个下采样Block
layer4(3个BasicBlock,512通道) 11,759,424 conv5_x,含1个下采样Block
avgpool 0 无参数
fc(全连接层) 1,026 Linear(512, 2)
合计 21,285,698

总参数量:约 2100 万,仅为 VGG-16(1.34 亿)的 16%
可训练参数:仅 1,026 个(全连接层),因为我们冻结了所有预训练层
不可训练参数:21,284,672 个(预训练的特征提取层)


3. 训练模型

3.1 编写训练函数

# 训练循环
def train(dataloader, model, loss_fn, optimizer):
    size = len(dataloader.dataset)
    num_batches = len(dataloader)

    train_loss, train_acc = 0, 0
    
    for X, y in dataloader:
        X, y = X.to(device), y.to(device)
        
        pred = model(X)
        loss = loss_fn(pred, y)
        
        optimizer.zero_grad()
        loss.backward()
        optimizer.step()
        
        train_acc  += (pred.argmax(1) == y).type(torch.float).sum().item()
        train_loss += loss.item()
            
    train_acc  /= size
    train_loss /= num_batches

    return train_acc, train_loss

3.2 编写测试函数

def test(dataloader, model, loss_fn):
    size        = len(dataloader.dataset)
    num_batches = len(dataloader)
    test_loss, test_acc = 0, 0
    
    with torch.no_grad():
        for X, y in dataloader:
            X, y = X.to(device), y.to(device)
            
            y_pred = model(X)
            loss   = loss_fn(y_pred, y)
            
            test_loss += loss.item()
            test_acc  += (y_pred.argmax(1) == y).type(torch.float).sum().item()

    test_acc  /= size
    test_loss /= num_batches

    return test_acc, test_loss

3.3 正式训练

# 确保模型位于 AMD GPU
model = model.to(device)

# 冻结 ResNet-34 主体,只训练最后的 fc 层
for param in model.parameters():
    param.requires_grad = False

for param in model.fc.parameters():
    param.requires_grad = True


# 优化器只接收最后一层的参数
optimizer = torch.optim.Adam(model.fc.parameters(), lr=1e-4)
loss_fn = nn.CrossEntropyLoss()

epochs = 10
train_loss = []
train_acc = []
test_loss = []
test_acc = []

best_acc = 0

# 用来保存最佳模型参数
best_model_state = None

for epoch in range(epochs):

    # 不使用 model.train()
    # 让前面冻结的 BatchNorm 保持评估模式
    model.eval()

    # fc 是 Linear 层,仍然可以正常计算梯度和更新参数
    model.fc.train()

    epoch_train_acc, epoch_train_loss = train(
        train_dl,
        model,
        loss_fn,
        optimizer
    )

    # 测试阶段
    model.eval()

    epoch_test_acc, epoch_test_loss = test(
        test_dl,
        model,
        loss_fn
    )

    # 保存最佳模型参数到 CPU,避免复制整个 GPU 模型
    if epoch_test_acc > best_acc:
        best_acc = epoch_test_acc

        best_model_state = {
            name: param.detach().cpu().clone()
            for name, param in model.state_dict().items()
        }

    train_acc.append(epoch_train_acc)
    train_loss.append(epoch_train_loss)
    test_acc.append(epoch_test_acc)
    test_loss.append(epoch_test_loss)

    # 获取当前学习率
    lr = optimizer.param_groups[0]["lr"]

    template = (
        "Epoch:{:2d}, "
        "Train_acc:{:.1f}%, Train_loss:{:.3f}, "
        "Test_acc:{:.1f}%, Test_loss:{:.3f}, "
        "Lr:{:.2E}, Best_acc:{:.1f}%"
    )

    print(
        template.format(
            epoch + 1,
            epoch_train_acc * 100,
            epoch_train_loss,
            epoch_test_acc * 100,
            epoch_test_loss,
            lr,
            best_acc * 100
        )
    )


# 保存最佳模型参数
PATH = "./best_model.pth"
torch.save(best_model_state, PATH)

print("Done")
print("最佳测试准确率:{:.1f}%".format(best_acc * 100))

由于我的设备使用的是 AMD GPU,原训练代码中的 model.train() 会使 ResNet-34 主体网络中的 BatchNorm 层进入训练模式,从而触发 miopenStatusUnknownError。因此,查询GPT后将本次训练修改为冻结预训练模型的主体参数,只训练最后的全连接层 fc;训练时使用 model.eval() 保持前面 BatchNorm 层为评估模式,再用 model.fc.train() 训练分类层。同时,优化器只接收 model.fc.parameters(),并将最佳模型参数复制到 CPU 后保存,避免在 AMD GPU 上直接复制整个模型。这样既保留了 GPU 加速,也减少了 MIOpen 兼容问题。

Epoch: 1, Train_acc:73.4%, Train_loss:0.578, Test_acc:87.0%, Test_loss:0.452, Lr:1.00E-04, Best_acc:87.0%
Epoch: 2, Train_acc:89.8%, Train_loss:0.386, Test_acc:90.0%, Test_loss:0.339, Lr:1.00E-04, Best_acc:90.0%
Epoch: 3, Train_acc:91.8%, Train_loss:0.305, Test_acc:90.9%, Test_loss:0.283, Lr:1.00E-04, Best_acc:90.9%
Epoch: 4, Train_acc:92.5%, Train_loss:0.262, Test_acc:91.9%, Test_loss:0.250, Lr:1.00E-04, Best_acc:91.9%
Epoch: 5, Train_acc:93.2%, Train_loss:0.233, Test_acc:92.7%, Test_loss:0.226, Lr:1.00E-04, Best_acc:92.7%
Epoch: 6, Train_acc:93.7%, Train_loss:0.213, Test_acc:93.0%, Test_loss:0.210, Lr:1.00E-04, Best_acc:93.0%
Epoch: 7, Train_acc:94.0%, Train_loss:0.197, Test_acc:93.4%, Test_loss:0.196, Lr:1.00E-04, Best_acc:93.4%
Epoch: 8, Train_acc:94.3%, Train_loss:0.185, Test_acc:94.0%, Test_loss:0.185, Lr:1.00E-04, Best_acc:94.0%
Epoch: 9, Train_acc:94.6%, Train_loss:0.176, Test_acc:94.3%, Test_loss:0.176, Lr:1.00E-04, Best_acc:94.3%
Epoch:10, Train_acc:94.5%, Train_loss:0.168, Test_acc:94.5%, Test_loss:0.169, Lr:1.00E-04, Best_acc:94.5%
Done
最佳测试准确率:94.5%

训练结果分析:
初始阶段(Epoch 1): 训练准确率 73.4%,测试准确率 87.0%。预训练模型在第一个 epoch 就能达到很高的准确率,ResNet-34 在 ImageNet 上学到的通用视觉特征可以很好地迁移到 X 光片分类上;
稳定收敛(Epoch 1-10): 训练准确率从 73.4% 稳步提升到 94.5%,测试准确率从 87.0% 提升到 94.5%。整个训练过程非常平滑,没有出现 P7 周那样的过拟合波动;
训练-测试一致性: 两条曲线几乎完全贴合(Epoch 10 训练 94.5% vs 测试 94.5%),差距极小。


4. 结果可视化

import matplotlib.pyplot as plt
import warnings
warnings.filterwarnings("ignore")
plt.rcParams['font.sans-serif']    = ['SimHei']
plt.rcParams['axes.unicode_minus'] = False
plt.rcParams['figure.dpi']         = 100

from datetime import datetime
current_time = datetime.now()

epochs_range = range(epochs)

plt.figure(figsize=(12, 3))
plt.subplot(1, 2, 1)

plt.plot(epochs_range, train_acc, label='Training Accuracy')
plt.plot(epochs_range, test_acc, label='Test Accuracy')
plt.legend(loc='lower right')
plt.title('Training and Validation Accuracy')
plt.xlabel(current_time)

plt.subplot(1, 2, 2)
plt.plot(epochs_range, train_loss, label='Training Loss')
plt.plot(epochs_range, test_loss, label='Test Loss')
plt.legend(loc='upper right')
plt.title('Training and Validation Loss')
plt.show()

这是我的结果(调用官方模型)
在这里插入图片描述


5. 模型评估

修改训练代码后,保存的是最佳模型的参数字典 best_model_state,而不是完整的 best_model 模型对象,因此评估时不能再直接使用 best_model.eval()。需要先通过 torch.load() 读取保存的参数,再使用 model.load_state_dict() 将最佳参数加载回原来的 ResNet-34 模型,随后调用 model.eval() 切换到评估模式,并在测试集上计算准确率和损失。

best_model_official.eval()
test_acc_official, test_loss_official = test(test_dl, best_model_official, loss_fn)
test_acc_official, test_loss_official
(0.9449003516998827, 0.168533221439079)

6. 指定图片进行预测

由于训练后保存的是最佳模型的参数字典,而不是单独的完整模型对象,所以单张图片预测前,需要先使用 model.load_state_dict() 将最佳参数加载回 ResNet-34 模型,再将模型放到运行设备并切换为评估模式。预测时使用 torch.inference_mode() 关闭梯度计算,降低计算和显存开销;同时使用 .item() 将 GPU 上的预测索引转换为普通整数,再根据类别列表获取最终预测类别。

from PIL import Image
import matplotlib.pyplot as plt

# 类别名称
classes = total_data.classes

# 加载训练过程中保存的最佳模型参数
model.load_state_dict(
    torch.load("./best_model.pth", map_location=device, weights_only=True)
)

model = model.to(device)
model.eval()


def predict_one_image(image_path, model, transform, classes):
    # 读取图片并转换为RGB格式
    test_img = Image.open(image_path).convert("RGB")

    # 显示原始图片
    plt.imshow(test_img)
    plt.axis("off")
    plt.show()

    # 图片预处理,并增加batch维度
    img = transform(test_img).unsqueeze(0).to(device)

    # 预测时不计算梯度
    with torch.inference_mode():
        output = model(img)
        pred_index = output.argmax(dim=1).item()

    pred_class = classes[pred_index]
    print(f"预测结果是:{pred_class}")


predict_one_image(
    image_path="./data/PNEUMONIA/person1_bacteria_1.jpeg",
    model=model,
    transform=test_transform,
    classes=classes
)

在这里插入图片描述

预测结果是:PNEUMONIA

四、手动搭建 ResNet-34 模型(方案二)

ResNet-34 的核心不是简单地连续堆叠 34 个卷积层,而是先定义一个包含两层 3×3 卷积的 BasicBlock,再按照 [3,4,6,3] 的数量组成四个残差阶段。

模块 组成 输出 shape(单张图)
输入 RGB X 光片 [3,224,224]
网络开头 7×7 Conv + BN + ReLU + MaxPool [64,56,56]
layer1 BasicBlock × 3 [64,56,56]
layer2 BasicBlock × 4,首块下采样 [128,28,28]
layer3 BasicBlock × 6,首块下采样 [256,14,14]
layer4 BasicBlock × 3,首块下采样 [512,7,7]
avgpool 自适应平均池化 [512,1,1]
fc Linear(512,2) [2]

ResNet-34 中的 34 层按照主干带权重层计算:

开头 1 个卷积层
+ 16 个 BasicBlock × 每块 2 个卷积层
+ 最后 1 个全连接层
= 1 + 32 + 1
= 34 层

其中:

3 + 4 + 6 + 3 = 16 个 BasicBlock

1. 定义 BasicBlock(基本残差块)

class BasicBlock(nn.Module):
    """ResNet-34 的基本残差块,包含两个 3x3 卷积层和一个残差连接"""
    expansion = 1

    def __init__(self, in_channels, out_channels, stride=1, downsample=None):
        super(BasicBlock, self).__init__()
        # 第一个卷积层,可能带有下采样(stride=2)
        self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, 
                               stride=stride, padding=1, bias=False)
        self.bn1 = nn.BatchNorm2d(out_channels)
        self.relu = nn.ReLU(inplace=True)
        
        # 第二个卷积层,stride 始终为 1
        self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3,
                               stride=1, padding=1, bias=False)
        self.bn2 = nn.BatchNorm2d(out_channels)
        
        # 残差连接(捷径)
        self.downsample = downsample

    def forward(self, x):
        identity = x

        out = self.conv1(x)
        out = self.bn1(out)
        out = self.relu(out)
        out = self.conv2(out)
        out = self.bn2(out)

        if self.downsample is not None:
            identity = self.downsample(x)

        out += identity
        out = self.relu(out)

        return out

BasicBlock 的关键逻辑

  1. 主路径x → conv1 → bn1 → relu → conv2 → bn2 → out
  2. 捷径路径x → downsample(如果需要) → identity
  3. 残差相加out = out + identity
  4. 最终激活out = relu(out)

什么时候需要 downsample?

stride > 1(特征图尺寸变小)或者 in_channels != out_channels(通道数改变)时,输入 x x x 和卷积输出无法直接相加,需要对 x x x 进行 1×1 卷积下采样来匹配尺寸和通道数。


2. 定义 ResNet-34 网络

class ResNet34(nn.Module):
    def __init__(self, num_classes=2):
        super(ResNet34, self).__init__()

        # ---------- conv1:初始卷积层 ----------
        self.conv1 = nn.Conv2d(3, 64, kernel_size=7, stride=2, padding=3, bias=False)
        self.bn1 = nn.BatchNorm2d(64)
        self.relu = nn.ReLU(inplace=True)
        self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)

        # ---------- layer1 (conv2_x):3 个 BasicBlock,通道数 64 ----------
        self.layer1 = self._make_layer(64, 64, 3, stride=1)

        # ---------- layer2 (conv3_x):4 个 BasicBlock,通道数 128 ----------
        self.layer2 = self._make_layer(64, 128, 4, stride=2)

        # ---------- layer3 (conv4_x):6 个 BasicBlock,通道数 256 ----------
        self.layer3 = self._make_layer(128, 256, 6, stride=2)

        # ---------- layer4 (conv5_x):3 个 BasicBlock,通道数 512 ----------
        self.layer4 = self._make_layer(256, 512, 3, stride=2)

        # ---------- 全局平均池化 + 全连接层 ----------
        self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
        self.fc = nn.Linear(512, num_classes)

        # ---------- 权重初始化 ----------
        for m in self.modules():
            if isinstance(m, nn.Conv2d):
                nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')
            elif isinstance(m, nn.BatchNorm2d):
                nn.init.constant_(m.weight, 1)
                nn.init.constant_(m.bias, 0)

    def _make_layer(self, in_channels, out_channels, num_blocks, stride=1):
        """创建由多个 BasicBlock 组成的层"""
        downsample = None

        if stride != 1 or in_channels != out_channels:
            downsample = nn.Sequential(
                nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=stride, bias=False),
                nn.BatchNorm2d(out_channels),
            )

        layers = []
        layers.append(BasicBlock(in_channels, out_channels, stride, downsample))
        for _ in range(1, num_blocks):
            layers.append(BasicBlock(out_channels, out_channels))

        return nn.Sequential(*layers)

    def forward(self, x):
        x = self.conv1(x)
        x = self.bn1(x)
        x = self.relu(x)
        x = self.maxpool(x)

        x = self.layer1(x)
        x = self.layer2(x)
        x = self.layer3(x)
        x = self.layer4(x)

        x = self.avgpool(x)
        x = torch.flatten(x, 1)
        x = self.fc(x)

        return x


device = "cuda" if torch.cuda.is_available() else "cpu"
print("Using {} device".format(device))

model = ResNet34(num_classes=2).to(device)
model

手动搭建 ResNet-34 的完整结构解析:

输入图片:[3, 224, 224]
↓
conv1(初始卷积层)
  ├─ Conv2d(3→64, 7×7, stride=2, padding=3)    输出: [64, 112, 112]
  ├─ BN + ReLU
  └─ MaxPool2d(3×3, stride=2, padding=1)        输出: [64, 56, 56]
↓
layer1(conv2_x):3 个 BasicBlock,通道数 64
  ├─ Block1: Conv(64→64, 3×3) → BN → ReLU → Conv(64→64, 3×3) → BN + skip(64→64)  输出: [64, 56, 56]
  ├─ Block2: Conv(64→64, 3×3) → BN → ReLU → Conv(64→64, 3×3) → BN + skip(64→64)  输出: [64, 56, 56]
  └─ Block3: Conv(64→64, 3×3) → BN → ReLU → Conv(64→64, 3×3) → BN + skip(64→64)  输出: [64, 56, 56]
↓
layer2(conv3_x):4 个 BasicBlock,通道数 128
  ├─ Block1: Conv(64→128, 3×3, s=2) → BN → ReLU → Conv(128→128, 3×3) → BN + skip(64→128, 1×1, s=2)  输出: [128, 28, 28]
  ├─ Block2-4: Conv(128→128, 3×3) → BN → ReLU → Conv(128→128, 3×3) → BN + skip(128→128)  输出: [128, 28, 28]
↓
layer3(conv4_x):6 个 BasicBlock,通道数 256
  ├─ Block1: Conv(128→256, 3×3, s=2) → BN → ReLU → Conv(256→256, 3×3) → BN + skip(128→256, 1×1, s=2)  输出: [256, 14, 14]
  ├─ Block2-6: Conv(256→256, 3×3) → BN → ReLU → Conv(256→256, 3×3) → BN + skip(256→256)  输出: [256, 14, 14]
↓
layer4(conv5_x):3 个 BasicBlock,通道数 512
  ├─ Block1: Conv(256→512, 3×3, s=2) → BN → ReLU → Conv(512→512, 3×3) → BN + skip(256→512, 1×1, s=2)  输出: [512, 7, 7]
  ├─ Block2-3: Conv(512→512, 3×3) → BN → ReLU → Conv(512→512, 3×3) → BN + skip(512→512)  输出: [512, 7, 7]
↓
avgpool: AdaptiveAvgPool2d(1, 1)                输出: [512, 1, 1]
flatten:                                         输出: [512]
fc: Linear(512, 2)                              输出: [2]

根据官方 ResNet-34 的结构输出,参考上一周的教程思路,通过AI辅助搭建了由初始卷积层、四组残差层、全局平均池化层和全连接层组成的网络。四组残差层分别包含 3、4、6、3 个 BasicBlock,与标准 ResNet-34 结构一致。layer2、layer3 和 layer4 的第一个 BasicBlock 使用步长为 2 的卷积缩小特征图尺寸,并通过 1×1 卷积构建 downsample 分支,使主分支与捷径分支的 shape 保持一致。模型输入 [3,224,224] 的图片后,特征图依次变化为 64×56×56、128×28×28、256×14×14、512×7×7,经过全局平均池化和展平后进入 Linear(512,2),最终输出正常与肺炎两个类别。手动模型与官方模型结构相同,但手动模型使用随机初始化权重,默认全部参数参与训练,而官方模型使用预训练权重并只训练最后的分类层。


3. 查看模型详情

# 统计模型参数量以及其他指标
from torchinfo import summary

model = model.cpu()
summary(model, input_size=(1, 3, 224, 224), device="cpu")
==========================================================================================
Layer (type:depth-idx)                   Output Shape              Param #
==========================================================================================
ResNet34                                 [1, 2]                    --
├─Conv2d: 1-1                            [1, 64, 112, 112]         9,408
├─BatchNorm2d: 1-2                       [1, 64, 112, 112]         128
├─ReLU: 1-3                              [1, 64, 112, 112]         --
├─MaxPool2d: 1-4                         [1, 64, 56, 56]           --
├─Sequential: 1-5                        [1, 64, 56, 56]           --
│    └─BasicBlock: 2-1                   [1, 64, 56, 56]           --
│    │    └─Conv2d: 3-1                  [1, 64, 56, 56]           36,864
│    │    └─BatchNorm2d: 3-2             [1, 64, 56, 56]           128
│    │    └─ReLU: 3-3                    [1, 64, 56, 56]           --
│    │    └─Conv2d: 3-4                  [1, 64, 56, 56]           36,864
│    │    └─BatchNorm2d: 3-5             [1, 64, 56, 56]           128
│    │    └─ReLU: 3-6                    [1, 64, 56, 56]           --
│    └─BasicBlock: 2-2                   [1, 64, 56, 56]           --
│    │    └─Conv2d: 3-7                  [1, 64, 56, 56]           36,864
│    │    └─BatchNorm2d: 3-8             [1, 64, 56, 56]           128
│    │    └─ReLU: 3-9                    [1, 64, 56, 56]           --
│    │    └─Conv2d: 3-10                 [1, 64, 56, 56]           36,864
│    │    └─BatchNorm2d: 3-11            [1, 64, 56, 56]           128
│    │    └─ReLU: 3-12                   [1, 64, 56, 56]           --
│    └─BasicBlock: 2-3                   [1, 64, 56, 56]           --
│    │    └─Conv2d: 3-13                 [1, 64, 56, 56]           36,864
│    │    └─BatchNorm2d: 3-14            [1, 64, 56, 56]           128
│    │    └─ReLU: 3-15                   [1, 64, 56, 56]           --
│    │    └─Conv2d: 3-16                 [1, 64, 56, 56]           36,864
│    │    └─BatchNorm2d: 3-17            [1, 64, 56, 56]           128
│    │    └─ReLU: 3-18                   [1, 64, 56, 56]           --
├─Sequential: 1-6                        [1, 128, 28, 28]          --
│    └─BasicBlock: 2-4                   [1, 128, 28, 28]          --
│    │    └─Conv2d: 3-19                 [1, 128, 28, 28]          73,728
│    │    └─BatchNorm2d: 3-20            [1, 128, 28, 28]          256
│    │    └─ReLU: 3-21                   [1, 128, 28, 28]          --
│    │    └─Conv2d: 3-22                 [1, 128, 28, 28]          147,456
│    │    └─BatchNorm2d: 3-23            [1, 128, 28, 28]          256
│    │    └─Sequential: 3-24             [1, 128, 28, 28]          8,448
│    │    └─ReLU: 3-25                   [1, 128, 28, 28]          --
│    └─BasicBlock: 2-5                   [1, 128, 28, 28]          --
│    │    └─Conv2d: 3-26                 [1, 128, 28, 28]          147,456
│    │    └─BatchNorm2d: 3-27            [1, 128, 28, 28]          256
│    │    └─ReLU: 3-28                   [1, 128, 28, 28]          --
│    │    └─Conv2d: 3-29                 [1, 128, 28, 28]          147,456
│    │    └─BatchNorm2d: 3-30            [1, 128, 28, 28]          256
│    │    └─ReLU: 3-31                   [1, 128, 28, 28]          --
│    └─BasicBlock: 2-6                   [1, 128, 28, 28]          --
│    │    └─Conv2d: 3-32                 [1, 128, 28, 28]          147,456
│    │    └─BatchNorm2d: 3-33            [1, 128, 28, 28]          256
│    │    └─ReLU: 3-34                   [1, 128, 28, 28]          --
│    │    └─Conv2d: 3-35                 [1, 128, 28, 28]          147,456
│    │    └─BatchNorm2d: 3-36            [1, 128, 28, 28]          256
│    │    └─ReLU: 3-37                   [1, 128, 28, 28]          --
│    └─BasicBlock: 2-7                   [1, 128, 28, 28]          --
│    │    └─Conv2d: 3-38                 [1, 128, 28, 28]          147,456
│    │    └─BatchNorm2d: 3-39            [1, 128, 28, 28]          256
│    │    └─ReLU: 3-40                   [1, 128, 28, 28]          --
│    │    └─Conv2d: 3-41                 [1, 128, 28, 28]          147,456
│    │    └─BatchNorm2d: 3-42            [1, 128, 28, 28]          256
│    │    └─ReLU: 3-43                   [1, 128, 28, 28]          --
├─Sequential: 1-7                        [1, 256, 14, 14]          --
│    └─BasicBlock: 2-8                   [1, 256, 14, 14]          --
│    │    └─Conv2d: 3-44                 [1, 256, 14, 14]          294,912
│    │    └─BatchNorm2d: 3-45            [1, 256, 14, 14]          512
│    │    └─ReLU: 3-46                   [1, 256, 14, 14]          --
│    │    └─Conv2d: 3-47                 [1, 256, 14, 14]          589,824
│    │    └─BatchNorm2d: 3-48            [1, 256, 14, 14]          512
│    │    └─Sequential: 3-49             [1, 256, 14, 14]          33,280
│    │    └─ReLU: 3-50                   [1, 256, 14, 14]          --
│    └─BasicBlock: 2-9                   [1, 256, 14, 14]          --
│    │    └─Conv2d: 3-51                 [1, 256, 14, 14]          589,824
│    │    └─BatchNorm2d: 3-52            [1, 256, 14, 14]          512
│    │    └─ReLU: 3-53                   [1, 256, 14, 14]          --
│    │    └─Conv2d: 3-54                 [1, 256, 14, 14]          589,824
│    │    └─BatchNorm2d: 3-55            [1, 256, 14, 14]          512
│    │    └─ReLU: 3-56                   [1, 256, 14, 14]          --
│    └─BasicBlock: 2-10                  [1, 256, 14, 14]          --
│    │    └─Conv2d: 3-57                 [1, 256, 14, 14]          589,824
│    │    └─BatchNorm2d: 3-58            [1, 256, 14, 14]          512
│    │    └─ReLU: 3-59                   [1, 256, 14, 14]          --
│    │    └─Conv2d: 3-60                 [1, 256, 14, 14]          589,824
│    │    └─BatchNorm2d: 3-61            [1, 256, 14, 14]          512
│    │    └─ReLU: 3-62                   [1, 256, 14, 14]          --
│    └─BasicBlock: 2-11                  [1, 256, 14, 14]          --
│    │    └─Conv2d: 3-63                 [1, 256, 14, 14]          589,824
│    │    └─BatchNorm2d: 3-64            [1, 256, 14, 14]          512
│    │    └─ReLU: 3-65                   [1, 256, 14, 14]          --
│    │    └─Conv2d: 3-66                 [1, 256, 14, 14]          589,824
│    │    └─BatchNorm2d: 3-67            [1, 256, 14, 14]          512
│    │    └─ReLU: 3-68                   [1, 256, 14, 14]          --
│    └─BasicBlock: 2-12                  [1, 256, 14, 14]          --
│    │    └─Conv2d: 3-69                 [1, 256, 14, 14]          589,824
│    │    └─BatchNorm2d: 3-70            [1, 256, 14, 14]          512
│    │    └─ReLU: 3-71                   [1, 256, 14, 14]          --
│    │    └─Conv2d: 3-72                 [1, 256, 14, 14]          589,824
│    │    └─BatchNorm2d: 3-73            [1, 256, 14, 14]          512
│    │    └─ReLU: 3-74                   [1, 256, 14, 14]          --
│    └─BasicBlock: 2-13                  [1, 256, 14, 14]          --
│    │    └─Conv2d: 3-75                 [1, 256, 14, 14]          589,824
│    │    └─BatchNorm2d: 3-76            [1, 256, 14, 14]          512
│    │    └─ReLU: 3-77                   [1, 256, 14, 14]          --
│    │    └─Conv2d: 3-78                 [1, 256, 14, 14]          589,824
│    │    └─BatchNorm2d: 3-79            [1, 256, 14, 14]          512
│    │    └─ReLU: 3-80                   [1, 256, 14, 14]          --
├─Sequential: 1-8                        [1, 512, 7, 7]            --
│    └─BasicBlock: 2-14                  [1, 512, 7, 7]            --
│    │    └─Conv2d: 3-81                 [1, 512, 7, 7]            1,179,648
│    │    └─BatchNorm2d: 3-82            [1, 512, 7, 7]            1,024
│    │    └─ReLU: 3-83                   [1, 512, 7, 7]            --
│    │    └─Conv2d: 3-84                 [1, 512, 7, 7]            2,359,296
│    │    └─BatchNorm2d: 3-85            [1, 512, 7, 7]            1,024
│    │    └─Sequential: 3-86             [1, 512, 7, 7]            132,096
│    │    └─ReLU: 3-87                   [1, 512, 7, 7]            --
│    └─BasicBlock: 2-15                  [1, 512, 7, 7]            --
│    │    └─Conv2d: 3-88                 [1, 512, 7, 7]            2,359,296
│    │    └─BatchNorm2d: 3-89            [1, 512, 7, 7]            1,024
│    │    └─ReLU: 3-90                   [1, 512, 7, 7]            --
│    │    └─Conv2d: 3-91                 [1, 512, 7, 7]            2,359,296
│    │    └─BatchNorm2d: 3-92            [1, 512, 7, 7]            1,024
│    │    └─ReLU: 3-93                   [1, 512, 7, 7]            --
│    └─BasicBlock: 2-16                  [1, 512, 7, 7]            --
│    │    └─Conv2d: 3-94                 [1, 512, 7, 7]            2,359,296
│    │    └─BatchNorm2d: 3-95            [1, 512, 7, 7]            1,024
│    │    └─ReLU: 3-96                   [1, 512, 7, 7]            --
│    │    └─Conv2d: 3-97                 [1, 512, 7, 7]            2,359,296
│    │    └─BatchNorm2d: 3-98            [1, 512, 7, 7]            1,024
│    │    └─ReLU: 3-99                   [1, 512, 7, 7]            --
├─AdaptiveAvgPool2d: 1-9                 [1, 512, 1, 1]            --
├─Linear: 1-10                           [1, 2]                    1,026
==========================================================================================
Total params: 21,285,698
Trainable params: 21,285,698
Non-trainable params: 0
Total mult-adds (Units.GIGABYTES): 3.66
==========================================================================================
Input size (MB): 0.60
Forward/backward pass size (MB): 59.81
Params size (MB): 85.14
Estimated Total Size (MB): 145.55
==========================================================================================

手动搭建的 ResNet-34 参数量分析

模块 参数量 说明
conv1 + bn1 + maxpool 9,536 初始下采样层
layer1(3个BasicBlock,64通道) 110,976 conv2_x
layer2(4个BasicBlock,128通道) 885,120 conv3_x
layer3(6个BasicBlock,256通道) 8,518,656 conv4_x
layer4(3个BasicBlock,512通道) 11,759,424 conv5_x
avgpool 0 无参数
fc(全连接层) 1,026 Linear(512, 2)
合计 21,285,698 全部为可训练参数

与方案一(调用官方)结构完全相同,但所有参数都是随机初始化的,没有预训练权重。

model = model.to(device)

4. 训练模型

4.1 训练函数和测试函数

与方案一相同,此处省略。直接跳到训练代码。

4.2 正式训练

import copy

optimizer  = torch.optim.Adam(model.parameters(), lr=1e-4)
loss_fn    = nn.CrossEntropyLoss()

epochs     = 10
train_loss = []
train_acc  = []
test_loss  = []
test_acc   = []

best_acc = 0

# ========== 新增:禁用 MIOpen/cuDNN 后端,规避 AMD GPU 的 batch_norm 错误 ==========
torch.backends.cudnn.enabled = False
# ==============================================================================

for epoch in range(epochs):
    
    model.train()
    epoch_train_acc, epoch_train_loss = train(train_dl, model, loss_fn, optimizer)
    
    model.eval()
    epoch_test_acc, epoch_test_loss = test(test_dl, model, loss_fn)
    
    if epoch_test_acc > best_acc:
        best_acc   = epoch_test_acc
        best_model = copy.deepcopy(model)
    
    train_acc.append(epoch_train_acc)
    train_loss.append(epoch_train_loss)
    test_acc.append(epoch_test_acc)
    test_loss.append(epoch_test_loss)
    
    lr = optimizer.state_dict()['param_groups'][0]['lr']
    
    template = ('Epoch:{:2d}, Train_acc:{:.1f}%, Train_loss:{:.3f}, Test_acc:{:.1f}%, Test_loss:{:.3f}, Lr:{:.2E}')
    print(template.format(epoch+1, epoch_train_acc*100, epoch_train_loss, 
                          epoch_test_acc*100, epoch_test_loss, lr))

PATH = './best_model.pth'
torch.save(best_model.state_dict(), PATH)

print('Done')
Epoch: 1, Train_acc:92.7%, Train_loss:0.181, Test_acc:82.4%, Test_loss:0.758, Lr:1.00E-04
Epoch: 2, Train_acc:96.8%, Train_loss:0.092, Test_acc:94.1%, Test_loss:0.172, Lr:1.00E-04
Epoch: 3, Train_acc:97.2%, Train_loss:0.070, Test_acc:93.4%, Test_loss:0.207, Lr:1.00E-04
Epoch: 4, Train_acc:97.7%, Train_loss:0.063, Test_acc:96.8%, Test_loss:0.095, Lr:1.00E-04
Epoch: 5, Train_acc:99.1%, Train_loss:0.027, Test_acc:97.0%, Test_loss:0.119, Lr:1.00E-04
Epoch: 6, Train_acc:99.0%, Train_loss:0.026, Test_acc:92.2%, Test_loss:0.303, Lr:1.00E-04
Epoch: 7, Train_acc:98.8%, Train_loss:0.031, Test_acc:96.5%, Test_loss:0.125, Lr:1.00E-04
Epoch: 8, Train_acc:99.5%, Train_loss:0.014, Test_acc:96.9%, Test_loss:0.164, Lr:1.00E-04
Epoch: 9, Train_acc:99.7%, Train_loss:0.009, Test_acc:97.1%, Test_loss:0.158, Lr:1.00E-04
Epoch:10, Train_acc:99.4%, Train_loss:0.014, Test_acc:97.1%, Test_loss:0.117, Lr:1.00E-04
Done

手动搭建的 ResNet-34 没有加载预训练权重,模型中的卷积层、BatchNorm 层和全连接层均为随机初始化,因此本次使用 model.parameters() 创建优化器,使约 2128 万个参数全部参与训练。由于在 AMD GPU 环境下,模型进入训练模式后,BatchNorm2d 调用 MIOpen 时出现了 miopenStatusUnknownError,因此在训练前设置 torch.backends.cudnn.enabled = False,关闭相关优化后端,使 BatchNorm 等运算采用兼容性更好的计算方式。这样仍然可以使用 AMD GPU 训练完整的手动 ResNet-34,但训练速度可能比开启 MIOpen 优化时更慢。每轮训练时先使用 model.train() 开启训练模式并更新模型参数,测试时使用 model.eval() 切换到评估模式;当测试集准确率高于之前的最佳结果时,使用 copy.deepcopy(model) 保存当前最佳模型,训练结束后再将最佳模型的参数保存为 best_model.pth,用于后续模型评估和单张图片预测。


5. 结果可视化

import matplotlib.pyplot as plt
import warnings
warnings.filterwarnings("ignore")
plt.rcParams['font.sans-serif']    = ['SimHei']
plt.rcParams['axes.unicode_minus'] = False
plt.rcParams['figure.dpi']         = 100

from datetime import datetime
current_time = datetime.now()

epochs_range = range(epochs)

plt.figure(figsize=(12, 3))
plt.subplot(1, 2, 1)

plt.plot(epochs_range, train_acc, label='Training Accuracy')
plt.plot(epochs_range, test_acc, label='Test Accuracy')
plt.legend(loc='lower right')
plt.title('Training and Validation Accuracy')
plt.xlabel(current_time)

plt.subplot(1, 2, 2)
plt.plot(epochs_range, train_loss, label='Training Loss')
plt.plot(epochs_range, test_loss, label='Test Loss')
plt.legend(loc='upper right')
plt.title('Training and Validation Loss')
plt.show()

这是我的结果(手动搭建模型)
在这里插入图片描述


6. 模型评估

# 加载保存的最佳模型参数
model.load_state_dict(torch.load(PATH, map_location=device))
model = model.to(device)

# 进入评估模式
model.eval()

# 在测试集上评估最佳模型
epoch_test_acc, epoch_test_loss = test(
    test_dl,
    model,
    loss_fn
)

epoch_test_acc, epoch_test_loss
(0.9712778429073857, 0.1576850487143491)

五、总结

本周学习的是 PyTorch 入门第 P8 周:调用并手动搭建 ResNet-34 网络实现肺炎 X 光片分类。

本周最重要的收获

  1. ResNet-34 的核心创新——残差连接(Skip Connection):通过引入"捷径路径"让网络学习输入和输出之间的残差(差异),而不是从零学习完整映射。这解决了深层网络的退化问题,使得训练 34 层甚至 100+ 层的网络成为可能。

  2. 两种模型构建方式

    • 方案一(调用官方预训练):冻结特征提取层,只微调最后一层。训练快、效果好,是实际应用中的推荐方案。
    • 方案二(手动搭建):从零搭建完整网络,没有预训练权重。用于学习网络结构和锻炼代码能力。
  3. BasicBlock 的手动搭建:每个 BasicBlock 包含两个 3×3 卷积层(+BN + ReLU)和一个残差连接。当输入输出尺寸/通道不一致时,使用 1×1 卷积进行下采样。

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐