Swin Transformer实战:从零开始搭建图像分类模型(PyTorch版)

Swin Transformer实战:从零搭建图像分类模型(PyTorch版)

如果你已经对Vision Transformer(ViT)有所了解,并且正在寻找一种既能处理高分辨率图像,又能将计算复杂度控制在合理范围内的视觉Transformer架构,那么Swin Transformer很可能就是你工具箱里缺失的那块拼图。我第一次接触Swin Transformer是在处理一个遥感图像分类项目时,当时面对数千张高分辨率卫星影像,标准的ViT模型在显存和速度上都显得力不从心。Swin Transformer提出的层级化设计滑动窗口注意力机制,巧妙地解决了这个问题,让我在保持模型性能的同时,显著提升了训练和推理效率。

这篇文章面向有一定PyTorch和深度学习基础的开发者,特别是那些希望将前沿的Transformer架构真正落地到图像分类任务中的朋友。我们将完全从零开始,手把手地构建一个完整的Swin Transformer图像分类模型。整个过程不仅仅是“调包”,我会带你深入理解其核心模块的设计动机,并分享在实际编码、训练和调试过程中积累的宝贵经验。你会发现,从数据加载到模型部署,每一个环节都有值得注意的细节。

1. 环境准备与数据预处理

在开始构建模型之前,确保你的开发环境已经就绪。我强烈建议使用Anaconda来管理Python环境,它能有效避免不同项目间的依赖冲突。

# 创建并激活一个新的conda环境
conda create -n swin_torch python=3.8
conda activate swin_torch

# 安装PyTorch(请根据你的CUDA版本选择对应命令,这里以CUDA 11.3为例)
pip install torch==1.12.1+cu113 torchvision==0.13.1+cu113 --extra-index-url https://download.pytorch.org/whl/cu113

# 安装其他必要的库
pip install timm  # 一个非常棒的PyTorch图像模型库,我们将参考其Swin实现
pip install opencv-python pillow matplotlib scikit-learn pandas

提示:timm库(PyTorch Image Models)是Ross Wightman维护的一个宝藏库,它包含了大量预训练模型的高质量实现。即使我们打算从零构建,参考它的代码结构也是极佳的学习方式。

接下来是数据预处理。一个鲁棒的数据管道是成功训练模型的一半。我们以经典的CIFAR-10数据集为例,但这里的流程可以轻松迁移到你的自定义数据集。

import torch
from torchvision import datasets, transforms
from torch.utils.data import DataLoader

def build_dataloaders(data_dir='./data', batch_size=32):
    """
    构建训练和验证数据加载器。
    注意:Swin Transformer原始论文输入为224x224,但我们可以根据任务调整。
    """
    # 定义数据增强和归一化
    # 对于小数据集(如CIFAR-10),增强尤为重要
    train_transform = transforms.Compose([
        transforms.RandomResizedCrop(224),  # 随机裁剪并缩放到224x224
        transforms.RandomHorizontalFlip(p=0.5),
        transforms.RandomRotation(degrees=15),
        transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2),
        transforms.ToTensor(),
        transforms.Normalize(mean=[0.485, 0.456, 0.406],
                             std=[0.229, 0.224, 0.225])  # ImageNet统计值,通用性较好
    ])

    val_transform = transforms.Compose([
        transforms.Resize(256),  # 验证时先等比缩放
        transforms.CenterCrop(224),  # 再从中心裁剪
        transforms.ToTensor(),
        transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
    ])

    # 加载数据集
    train_dataset = datasets.CIFAR10(root=data_dir, train=True,
                                     download=True, transform=train_transform)
    val_dataset = datasets.CIFAR10(root=data_dir, train=False,
                                   download=True, transform=val_transform)

    # 创建数据加载器
    train_loader = DataLoader(train_dataset, batch_size=batch_size,
                              shuffle=True, num_workers=4, pin_memory=True)
    val_loader = DataLoader(val_dataset, batch_size=batch_size,
                            shuffle=False, num_workers=4, pin_memory=True)

    return train_loader, val_loader, train_dataset.classes

这里有几个关键点需要注意:

  • 图像尺寸:Swin Transformer的原始设计输入是224x224,因为其窗口划分(Window Partition)下采样(Patch Merging) 模块对输入尺寸有特定要求(通常需要能被patch_sizewindow_size整除)。如果你使用其他尺寸,需要相应调整模型参数。
  • 数据增强:对于图像分类,适当的数据增强是防止过拟合、提升模型泛化能力的利器。除了上面用到的,根据具体任务还可以考虑CutMixMixUpRandAugment等更高级的策略。
  • 归一化参数:我们使用了ImageNet的均值和标准差。如果你的数据集与ImageNet差异很大(例如医学图像、卫星图像),强烈建议计算自己数据集的统计值并替换它们,这通常能带来小幅度的性能提升。

2. 深入理解Swin Transformer的核心模块

在动手编码之前,我们必须吃透Swin Transformer的几个核心设计思想。这能帮助你在调试模型时,清楚地知道每一层输入输出的变化,而不是一个“黑盒”。

2.1 从Patch Embedding到层级结构

与ViT将图像直接切割成固定大小的patch并展平不同,Swin Transformer更像CNN,采用了层级化(Hierarchical) 的特征图构建方式。这个过程始于Patch Embedding

你可以把它想象成一个步长(stride)和卷积核大小(kernel_size)都为patch_size(通常是4)的卷积层。对于一个224x224x3的输入图像,经过patch_size=4的嵌入后,会得到56x56xC的特征图(C是嵌入维度,Swin-Tiny中是96)。这相当于把图像分成了56x564x4的像素块,每个块被映射成一个C维的向量。

import torch.nn as nn

class PatchEmbed(nn.Module):
    """ 将图像分割成不重叠的块并进行嵌入。"""
    def __init__(self, img_size=224, patch_size=4, in_chans=3, embed_dim=96, norm_layer=None):
        super().__init__()
        self.img_size = (img_size, img_size)
        self.patch_size = (patch_size, patch_size)
        self.grid_size = (img_size // patch_size, img_size // patch_size)
        self.num_patches = self.grid_size[0] * self.grid_size[1]

        # 核心就是一个卷积层
        self.proj = nn.Conv2d(in_chans, embed_dim,
                              kernel_size=patch_size, stride=patch_size)
        self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity()

    def forward(self, x):
        B, C, H, W = x.shape
        # 确保输入尺寸符合预期
        assert H == self.img_size[0] and W == self.img_size[1], \
            f"Input image size ({H}*{W}) doesn't match model ({self.img_size[0]}*{self.img_size[1]})."
        # 卷积操作: (B, 3, 224, 224) -> (B, 96, 56, 56)
        x = self.proj(x)
        x = x.flatten(2).transpose(1, 2)  # (B, 96, 56, 56) -> (B, 56*56, 96)
        x = self.norm(x)
        return x

2.2 窗口注意力与滑动窗口注意力

这是Swin Transformer的灵魂,也是其计算效率远超标准ViT自注意力的关键。标准自注意力在序列长度N(即patch数量)上的计算复杂度是O(N²)。对于56x56=3136patch,这个计算量是巨大的。

窗口多头自注意力(W-MSA) 的巧妙之处在于,它将特征图划分成一个个不重叠的、固定大小(如7x7)的窗口,只在每个窗口内部计算自注意力。这样,计算复杂度就从全局的O(N²)降为了O(M² * N/M²) = O(M² * (N/M²)),其中M是窗口大小。当N很大时,这带来了近乎线性的复杂度增长。

但W-MSA有个明显缺陷:窗口之间没有信息交互。为了解决这个问题,滑动窗口多头自注意力(SW-MSA) 被引入。它在下一个Transformer Block中,将窗口向右下角滑动窗口大小//2个像素,然后重新划分窗口。这样,同一个位置在不同的Block中,会和不同的邻居patch进行交互,从而实现了跨窗口的信息流通。

然而,滑动窗口带来了新的问题:新窗口可能包含来自原图中不相邻的区域。Swin Transformer使用了一个精妙的掩码(Mask)机制,在计算注意力时,通过给不属于同一原始区域的patch对加上一个极大的负偏置(如-100),使得经过Softmax后其注意力权重趋近于0,从而屏蔽无效的交互。

def create_mask(window_size, shift_size, H, W):
    """
    生成SW-MSA所需的注意力掩码。
    这是一个简化示例,用于理解原理。
    """
    Hp = int(np.ceil(H / window_size)) * window_size
    Wp = int(np.ceil(W / window_size)) * window_size
    img_mask = torch.zeros((1, Hp, Wp, 1))
    h_slices = (slice(0, -window_size),
                slice(-window_size, -shift_size),
                slice(-shift_size, None))
    w_slices = (slice(0, -window_size),
                slice(-window_size, -shift_size),
                slice(-shift_size, None))
    cnt = 0
    for h in h_slices:
        for w in w_slices:
            img_mask[:, h, w, :] = cnt
            cnt += 1

    mask_windows = window_partition(img_mask, window_size)
    mask_windows = mask_windows.view(-1, window_size * window_size)
    attn_mask = mask_windows.unsqueeze(1) - mask_windows.unsqueeze(2)
    attn_mask = attn_mask.masked_fill(attn_mask != 0, float(-100.0)).masked_fill(attn_mask == 0, float(0.0))
    return attn_mask

2.3 下采样与相对位置偏置

Patch Merging 是Swin Transformer层级结构的下采样层,作用类似于CNN中的池化或步长卷积。它将相邻的2x2patch的特征拼接起来,然后通过一个线性层将通道数翻倍(例如从C变为2C),同时空间尺寸减半(H, W -> H/2, W/2)。这既增大了感受野,又构建了类似CNN的金字塔特征。

相对位置偏置(Relative Position Bias) 是Swin Transformer在注意力计算中加入的一项改进。与ViT使用绝对位置编码不同,Swin Transformer认为,在自注意力中,patch之间的相对位置关系比绝对位置更重要。它为注意力分数矩阵中的每个元素(i, j)添加了一个可学习的偏置项B(i, j),这个偏置只依赖于patch ipatch j的相对位置(如上、下、左、右偏移量)。实践表明,这种简单的编码方式比绝对位置编码或复杂的相对位置编码公式效果更好,且计算开销小。

下表总结了Swin Transformer(Tiny版本)四个Stage的典型配置:

Stage输入尺寸 (H, W, C)Swin Transformer Block 数量输出尺寸 (H, W, C)窗口大小多头注意力头数
156, 56, 96256, 56, 9673
256, 56, 96228, 28, 19276
328, 28, 192614, 14, 384712
414, 14, 38427, 7, 768724

3. 动手搭建Swin Transformer模型

理解了原理,现在我们可以用PyTorch将它们组装起来。我们将采用模块化的方式,从最基础的窗口注意力模块开始构建。

3.1 构建窗口注意力模块

首先实现最核心的WindowAttention模块,它封装了带相对位置偏置的窗口内自注意力计算。

import torch.nn.functional as F
from einops import rearrange  # 可以使用einops库让张量操作更清晰

class WindowAttention(nn.Module):
    """ 基于相对位置偏置的窗口多头自注意力。"""
    def __init__(self, dim, window_size, num_heads, qkv_bias=True, attn_drop=0., proj_drop=0.):
        super().__init__()
        self.dim = dim
        self.window_size = window_size  # (Wh, Ww)
        self.num_heads = num_heads
        head_dim = dim // num_heads
        self.scale = head_dim ** -0.5

        # 相对位置偏置表:一个可学习的参数,大小为 (num_relative_distance, num_heads)
        # 对于窗口大小M,相对位置距离有(2M-1)*(2M-1)种
        self.relative_position_bias_table = nn.Parameter(
            torch.zeros((2 * window_size[0] - 1) * (2 * window_size[1] - 1), num_heads))

        # 生成相对位置索引(固定值,无需学习)
        coords_h = torch.arange(self.window_size[0])
        coords_w = torch.arange(self.window_size[1])
        coords = torch.stack(torch.meshgrid([coords_h, coords_w], indexing='ij'))  # 2, Wh, Ww
        coords_flatten = torch.flatten(coords, 1)  # 2, Wh*Ww
        relative_coords = coords_flatten[:, :, None] - coords_flatten[:, None, :]  # 2, Wh*Ww, Wh*Ww
        relative_coords = relative_coords.permute(1, 2, 0).contiguous()  # Wh*Ww, Wh*Ww, 2
        relative_coords[:, :, 0] += self.window_size[0] - 1
        relative_coords[:, :, 1] += self.window_size[1] - 1
        relative_coords[:, :, 0] *= 2 * self.window_size[1] - 1
        relative_position_index = relative_coords.sum(-1)  # Wh*Ww, Wh*Ww
        self.register_buffer("relative_position_index", relative_position_index)

        self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
        self.attn_drop = nn.Dropout(attn_drop)
        self.proj = nn.Linear(dim, dim)
        self.proj_drop = nn.Dropout(proj_drop)

        nn.init.trunc_normal_(self.relative_position_bias_table, std=.02)

    def forward(self, x, mask=None):
        """
        Args:
            x: 输入特征,形状为 (num_windows*B, N, C),其中N=窗口内patch数(如7*7=49)
            mask: (可选) SW-MSA中使用的注意力掩码,形状为 (nW, N, N)
        """
        B_, N, C = x.shape
        # 生成q, k, v
        qkv = self.qkv(x).reshape(B_, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
        q, k, v = qkv[0], qkv[1], qkv[2]  # 每个形状: (B_, num_heads, N, head_dim)

        q = q * self.scale
        attn = (q @ k.transpose(-2, -1))  # (B_, num_heads, N, N)

        # 添加相对位置偏置
        relative_position_bias = self.relative_position_bias_table[self.relative_position_index.view(-1)].view(
            self.window_size[0] * self.window_size[1], self.window_size[0] * self.window_size[1], -1)  # Wh*Ww, Wh*Ww, nH
        relative_position_bias = relative_position_bias.permute(2, 0, 1).contiguous()  # nH, Wh*Ww, Wh*Ww
        attn = attn + relative_position_bias.unsqueeze(0)

        if mask is not None:
            nW = mask.shape[0]
            attn = attn.view(B_ // nW, nW, self.num_heads, N, N) + mask.unsqueeze(1).unsqueeze(0)
            attn = attn.view(-1, self.num_heads, N, N)

        attn = F.softmax(attn, dim=-1)
        attn = self.attn_drop(attn)

        x = (attn @ v).transpose(1, 2).reshape(B_, N, C)
        x = self.proj(x)
        x = self.proj_drop(x)
        return x

3.2 组装Swin Transformer Block与Basic Layer

一个Swin Transformer Block由两个连续的子Block组成,分别使用W-MSA和SW-MSA。每个子Block都遵循“LayerNorm -> Attention -> 残差连接 -> LayerNorm -> MLP -> 残差连接”的结构。

class SwinTransformerBlock(nn.Module):
    """ Swin Transformer Block.
    包含两个连续的子Block,分别使用W-MSA和SW-MSA。
    """
    def __init__(self, dim, input_resolution, num_heads, window_size=7, shift_size=0,
                 mlp_ratio=4., qkv_bias=True, drop=0., attn_drop=0., drop_path=0.):
        super().__init__()
        self.dim = dim
        self.input_resolution = input_resolution
        self.num_heads = num_heads
        self.window_size = window_size
        self.shift_size = shift_size
        self.mlp_ratio = mlp_ratio

        # 确保shift_size小于window_size
        if self.shift_size > 0:
            assert 0 <= self.shift_size < self.window_size, "shift_size must in 0-window_size"

        self.norm1 = nn.LayerNorm(dim)
        self.attn = WindowAttention(
            dim, window_size=(self.window_size, self.window_size), num_heads=num_heads,
            qkv_bias=qkv_bias, attn_drop=attn_drop, proj_drop=drop)

        self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()
        self.norm2 = nn.LayerNorm(dim)
        mlp_hidden_dim = int(dim * mlp_ratio)
        self.mlp = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, drop=drop)

        # 如果使用滑动窗口,创建注意力掩码
        if self.shift_size > 0:
            H, W = self.input_resolution
            img_mask = torch.zeros((1, H, W, 1))
            h_slices = (slice(0, -self.window_size),
                        slice(-self.window_size, -self.shift_size),
                        slice(-self.shift_size, None))
            w_slices = (slice(0, -self.window_size),
                        slice(-self.window_size, -self.shift_size),
                        slice(-self.shift_size, None))
            cnt = 0
            for h in h_slices:
                for w in w_slices:
                    img_mask[:, h, w, :] = cnt
                    cnt += 1

            mask_windows = window_partition(img_mask, self.window_size)  # nW, window_size, window_size, 1
            mask_windows = mask_windows.view(-1, self.window_size * self.window_size)
            attn_mask = mask_windows.unsqueeze(1) - mask_windows.unsqueeze(2)
            attn_mask = attn_mask.masked_fill(attn_mask != 0, float(-100.0)).masked_fill(attn_mask == 0, float(0.0))
        else:
            attn_mask = None

        self.register_buffer("attn_mask", attn_mask)

    def forward(self, x):
        H, W = self.input_resolution
        B, L, C = x.shape
        assert L == H * W, "input feature has wrong size"

        shortcut = x
        x = self.norm1(x)
        x = x.view(B, H, W, C)

        # 循环移位(cyclic shift)以实现滑动窗口
        if self.shift_size > 0:
            shifted_x = torch.roll(x, shifts=(-self.shift_size, -self.shift_size), dims=(1, 2))
        else:
            shifted_x = x

        # 划分窗口
        x_windows = window_partition(shifted_x, self.window_size)  # nW*B, window_size, window_size, C
        x_windows = x_windows.view(-1, self.window_size * self.window_size, C)  # nW*B, N, C

        # W-MSA/SW-MSA
        attn_windows = self.attn(x_windows, mask=self.attn_mask)  # nW*B, N, C

        # 合并窗口
        attn_windows = attn_windows.view(-1, self.window_size, self.window_size, C)
        shifted_x = window_reverse(attn_windows, self.window_size, H, W)  # B H' W' C

        # 反向循环移位
        if self.shift_size > 0:
            x = torch.roll(shifted_x, shifts=(self.shift_size, self.shift_size), dims=(1, 2))
        else:
            x = shifted_x
        x = x.view(B, H * W, C)

        # 第一个残差连接
        x = shortcut + self.drop_path(x)

        # MLP部分
        x = x + self.drop_path(self.mlp(self.norm2(x)))

        return x

Basic Layer则是由多个SwinTransformerBlock和一个可选的Patch Merging下采样层组成,构成了一个完整的Stage。

class BasicLayer(nn.Module):
    """ 一个Swin Transformer Stage,包含多个Block和一个可选的Patch Merging。"""
    def __init__(self, dim, input_resolution, depth, num_heads, window_size,
                 mlp_ratio=4., qkv_bias=True, drop=0., attn_drop=0.,
                 drop_path=0., downsample=None):
        super().__init__()
        self.dim = dim
        self.input_resolution = input_resolution
        self.depth = depth

        # 构建Block
        self.blocks = nn.ModuleList([
            SwinTransformerBlock(dim=dim, input_resolution=input_resolution,
                                 num_heads=num_heads, window_size=window_size,
                                 shift_size=0 if (i % 2 == 0) else window_size // 2, # 交替使用W-MSA和SW-MSA
                                 mlp_ratio=mlp_ratio,
                                 qkv_bias=qkv_bias,
                                 drop=drop, attn_drop=attn_drop,
                                 drop_path=drop_path[i] if isinstance(drop_path, list) else drop_path)
            for i in range(depth)])

        # 下采样层
        if downsample is not None:
            self.downsample = downsample(input_resolution, dim=dim)
        else:
            self.downsample = None

    def forward(self, x):
        for blk in self.blocks:
            x = blk(x)
        if self.downsample is not None:
            x = self.downsample(x)
        return x

3.3 整合成完整的Swin Transformer

最后,我们将Patch Embedding、多个Basic Layer和一个分类头组合起来,形成完整的Swin Transformer模型。

class SwinTransformer(nn.Module):
    """ 完整的Swin Transformer模型。"""
    def __init__(self, img_size=224, patch_size=4, in_chans=3, num_classes=1000,
                 embed_dim=96, depths=[2, 2, 6, 2], num_heads=[3, 6, 12, 24],
                 window_size=7, mlp_ratio=4., qkv_bias=True,
                 drop_rate=0., attn_drop_rate=0., drop_path_rate=0.1):
        super().__init__()

        self.num_classes = num_classes
        self.num_layers = len(depths)
        self.embed_dim = embed_dim
        self.num_features = int(embed_dim * 2 ** (self.num_layers - 1))

        # 1. Patch Embedding
        self.patch_embed = PatchEmbed(img_size=img_size, patch_size=patch_size,
                                      in_chans=in_chans, embed_dim=embed_dim)
        patches_resolution = self.patch_embed.grid_size
        self.patches_resolution = patches_resolution

        # 2. 绝对位置编码(可选,Swin中有时会加入一个可学习的绝对位置编码)
        self.pos_drop = nn.Dropout(p=drop_rate)

        # 随机深度衰减(stochastic depth)
        dpr = [x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))]

        # 3. 构建各个Stage
        self.layers = nn.ModuleList()
        for i_layer in range(self.num_layers):
            layer = BasicLayer(dim=int(embed_dim * 2 ** i_layer),
                               input_resolution=(patches_resolution[0] // (2 ** i_layer),
                                                 patches_resolution[1] // (2 ** i_layer)),
                               depth=depths[i_layer],
                               num_heads=num_heads[i_layer],
                               window_size=window_size,
                               mlp_ratio=mlp_ratio,
                               qkv_bias=qkv_bias,
                               drop=drop_rate, attn_drop=attn_drop_rate,
                               drop_path=dpr[sum(depths[:i_layer]):sum(depths[:i_layer + 1])],
                               downsample=PatchMerging if (i_layer < self.num_layers - 1) else None)
            self.layers.append(layer)

        # 4. 最后的归一化层和分类头
        self.norm = nn.LayerNorm(self.num_features)
        self.avgpool = nn.AdaptiveAvgPool1d(1)
        self.head = nn.Linear(self.num_features, num_classes) if num_classes > 0 else nn.Identity()

        self.apply(self._init_weights)

    def _init_weights(self, m):
        if isinstance(m, nn.Linear):
            nn.init.trunc_normal_(m.weight, std=.02)
            if isinstance(m, nn.Linear) and m.bias is not None:
                nn.init.constant_(m.bias, 0)
        elif isinstance(m, nn.LayerNorm):
            nn.init.constant_(m.bias, 0)
            nn.init.constant_(m.weight, 1.0)

    def forward_features(self, x):
        x = self.patch_embed(x)
        x = self.pos_drop(x)

        for layer in self.layers:
            x = layer(x)

        x = self.norm(x)  # B, L, C
        x = self.avgpool(x.transpose(1, 2))  # B, C, 1
        x = torch.flatten(x, 1)
        return x

    def forward(self, x):
        x = self.forward_features(x)
        x = self.head(x)
        return x

至此,一个完整的Swin Transformer模型就搭建完成了。你可以通过修改depthsnum_heads等参数来实例化不同规模的模型,如Swin-Tiny、Swin-Small等。

4. 模型训练、调优与部署实战

模型搭建只是第一步,如何高效地训练它并应用到实际项目中,才是真正的挑战。这里分享几个我在实践中总结的关键点。

4.1 训练策略与技巧

直接从头开始训练一个Transformer模型,尤其是在中等规模的数据集上,很容易过拟合或难以收敛。以下策略能显著提升训练效果:

  • 使用预训练权重:如果可能,永远从预训练模型开始微调。你可以在Hugging Face Hub或timm库中找到在ImageNet-1K或ImageNet-22K上预训练好的Swin Transformer权重。这能为你提供一个强大的特征提取器起点。

    import timm
    model = timm.create_model('swin_tiny_patch4_window7_224', pretrained=True, num_classes=10) # 将分类头改为10类
    
  • 学习率与优化器:对于Transformer类模型,AdamW优化器配合余弦退火或带热重启的余弦退火学习率调度器是黄金组合。初始学习率可以设得小一些,例如3e-45e-4

    from torch.optim import AdamW
    from torch.optim.lr_scheduler import CosineAnnealingLR
    
    optimizer = AdamW(model.parameters(), lr=5e-4, weight_decay=0.05)
    scheduler = CosineAnnealingLR(optimizer, T_max=epochs)
    
  • 标签平滑与混合精度训练:标签平滑(Label Smoothing)可以减轻模型对训练标签的过度自信,提升泛化能力。混合精度训练(AMP)则能大幅减少显存占用并加速训练,对于Swin这类大模型尤其有用。

    criterion = nn.CrossEntropyLoss(label_smoothing=0.1)
    scaler = torch.cuda.amp.GradScaler() # 混合精度训练
    
  • 梯度裁剪:Transformer模型训练时梯度可能较大,使用梯度裁剪可以稳定训练过程。

    torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
    

4.2 常见问题与调试

在训练你自己的Swin模型时,可能会遇到以下问题:

  1. Loss为NaN或爆炸

    • 检查输入数据:确保数据归一化正确,像素值在合理范围(如[0,1]或[-1,1])。
    • 降低学习率:这是最常见的原因。
    • 启用梯度裁剪
    • 检查模型初始化:确保我们自定义的_init_weights函数被正确应用。
  2. 验证集准确率远低于训练集(过拟合)

    • 增强数据增强:尝试更激进的数据增强策略,如RandAugmentAutoAugment
    • 增加正则化:提高weight_decay,或尝试Stochastic Depth(已在代码中通过drop_path实现)。
    • 使用更小的模型:对于你的数据集,Swin-Tiny可能已经足够强大。
  3. 训练速度慢

    • 使用混合精度训练:如前所述,能显著加速。
    • 增大批量大小:在显存允许的范围内,更大的批量大小通常能更稳定地训练。
    • 使用更快的优化器:可以尝试LAMB优化器,它对大批量训练有优化。

4.3 模型推理与部署

训练完成后,你需要将模型部署到生产环境。这里有几个方向:

  • ONNX导出:PyTorch模型可以方便地导出为ONNX格式,以便在其他推理引擎(如TensorRT, OpenVINO)上运行。
    torch.onnx.export(model, dummy_input, "swin_transformer.onnx",
                      input_names=['input'], output_names=['output'],
                      dynamic_axes={'input': {0: 'batch_size'}, 'output': {0: 'batch_size'}})
    
  • TorchScript:如果你需要在不依赖Python环境的情况下部署模型,可以考虑使用TorchScript。
  • 使用推理优化库torch.jit.scripttorch.jit.trace或专门的推理优化库如Torch-TensorRT,可以对模型进行图优化和算子融合,提升推理速度。

最后,别忘了用训练好的模型在测试集或真实数据上跑一跑,看看效果。我习惯在验证集上保存性能最好的模型,然后在一个完全独立的测试集上进行最终评估,这能更真实地反映模型的泛化能力。整个从零搭建、训练到部署的过程,虽然会遇到各种“坑”,但每一步的解决都会让你对模型的理解更深一层。当你看到自己亲手搭建的Swin Transformer在任务上取得不错的效果时,那种成就感是直接用现成模型无法比拟的。

评论
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包
实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值