ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

TensorFlow实现ViT与Swin Transformer图像分类实战解析

TensorFlow实现ViT与Swin Transformer图像分类实战解析 自从2017年《Attention Is All You Need》提出Transformer架构以来NLP领域几乎被它彻底重构。而最近这两三年Transformer也开始大举“入侵”计算机视觉——从ViTVision Transformer到Swin Transformer这套基于自注意力机制的架构正在蚕食CNN曾经牢牢占据的图像处理版图。我这次用TensorFlow从零实现了ViT和Swin-T在图像分类任务上的完整流程包含数据预处理、模型搭建、训练调参与推理部署。这篇文章不是简单贴代码而是把我踩过的坑、对两个模型核心设计的理解、以及为什么Swin-T在工业场景往往比ViT更好用都拆开讲清楚。如果你正在用TensorFlow做图像分类或者刚接触Transformer想找个视觉方向的切入点这篇内容可以直接照着跑。1. 整体设计与思路拆解1.1 为什么视觉领域也要用Transformer先说一个很多人困惑的问题CNN已经做得够好了为什么非要折腾TransformerCNN的本质是局部感受野加滑动窗口通过堆叠卷积层来逐步扩大感知范围。这种方式有两个天然的局限第一远距离依赖关系需要非常深的网络才能建模而且信息在层层传递中会衰减第二卷积核的归纳偏置太强模型天然认为邻近像素关系更紧密这个假设在某些任务上反而成了天花板。Transformer的自注意力机制则完全不一样。它在计算每个位置时直接和序列中所有其他位置做相似度计算一步到位建立全局依赖。你在图像里放一只鸟鸟嘴和鸟尾可能隔着几十个像素CNN可能要十几层卷积才能把这两块信息关联起来Transformer在第一层就能看到全图。我当初决定用TensorFlow做这个项目也是因为TensorFlow的Keras API在模型搭建上足够灵活自定义Layer和Model都很顺手而且TF 2.x的tf.data管线配合GPU训练数据吞吐比PyTorch的DataLoader更容易调优。如果你手头是旧版TensorFlow建议直接上2.10以上版本。1.2 ViT和Swin-T的设计目标差异ViT和Swin-T虽然都叫Transformer但设计哲学差异非常大。ViT的做法是最直接的把图像切成一堆16x16的小patch拉平后当作token序列直接送进标准Transformer Encoder。它几乎完全抛弃了CNN的归纳偏置全靠数据量和算力硬扛。这就导致ViT特别“挑食”——在ImageNet这种千万级数据集上它能打过同规模的ResNet但如果你只有几万张图的小数据集ViT的训练效果往往不如ResNet。原因很简单自注意力没有局部性的先验模型需要从零学习“相邻像素通常相关”这件事。Swin-T则是Transformer向CNN“妥协”的产物。它的核心思想是层级化加局部性先在窗口内做自注意力窗口默认7x7再通过Shifted Window让信息跨窗口流动。这就像把CNN的局部感受野和Transformer的全局建模能力做了个折中。Swin-T还引入了类似FPN的层级结构输出不同分辨率的特征图这让它做检测、分割这类密集预测任务时比ViT自然得多。如果你只是做图像分类两者差距不算悬殊但如果后面要扩展做目标检测或者语义分割Swin-T的骨架优势会非常明显。我在项目里两个模型都实现了目的就是把这两种设计思路的差距直观地展示出来。1.3 技术选型与项目环境规划这个项目的技术栈如下Python 3.9推荐3.10TensorFlow 2.10GPU版CUDA 11.2搭配cuDNN 8.1TensorFlow Addons主要用里面的AdamW优化器和tfa.optimizers.MultiOptimizer做分层学习率数据集采用CIFAR-100有100个类、每类600张图规模适中既能体现Transformer的能力训练时间又可控单卡训练NVIDIA RTX 3090 24G显存这里要重点提醒一个坑TensorFlow 2.10是最后一个原生支持Windows GPU的版本2.11之后Windows用户只能走WSL 2。如果主力机是Windows老老实实装2.10如果是Linux服务器装新版倒无所谓。我一开始在Windows上装了TF 2.12折腾了半天发现GPU根本用不上换回2.10才解决。硬件方面ViT-Base在CIFAR-100上跑一个epoch大约需要60到80秒3090Swin-T稍慢一点大概90到110秒。batch size建议调到64再大梯度更新不稳定学习率需要同步调整。2. 核心原理详解与TensorFlow实现2.1 ViT的图片分块与Embedding实现ViT的第一个关键操作是Patch Embedding。假设输入图片是224x224x3patch_size16那么图片会被切成(224/16)^2196个patch每个patch大小16x16x3768维。这个Patch Embedding在实现上有个小技巧不用真的把图片切碎再逐个处理而是直接用一个大步长的卷积核搞定。我用TensorFlow的tf.keras.layers.Conv2D来实现这一步卷积核大小等于patch_size步长也等于patch_size输出通道数就是Embedding的维度。这种方式等价于全连接层的线性投影但计算效率高得多而且梯度传播更稳定。代码如下class PatchEmbed(layers.Layer): def __init__(self, patch_size16, embed_dim768): super(PatchEmbed, self).__init__() self.patch_size patch_size self.embed_dim embed_dim self.proj layers.Conv2D( filtersembed_dim, kernel_sizepatch_size, stridespatch_size, paddingvalid ) self.norm layers.LayerNormalization(epsilon1e-6) def call(self, images): # images shape: [B, H, W, C] patches self.proj(images) # [B, H/p, W/p, embed_dim] batch_size tf.shape(patches)[0] h patches.shape[1] w patches.shape[2] patches tf.reshape(patches, [batch_size, h * w, self.embed_dim]) return self.norm(patches)注意我加了一个LayerNormalization。这个细节在原始ViT论文里是没有的但后续很多改进工作比如DeiT都发现在Patch Embedding后加一层LN能显著稳定训练。尤其是当patch_size比较小、embed_dim比较大的时候不加LN经常出现前几个step的loss剧烈抖动。2.2 位置编码与可学习参数设计ViT的另一关键设计是位置编码。自注意力机制本身是不具备位置信息的你打乱token的顺序注意力计算结果完全不变所以必须把位置信息编码进去。ViT用的方式很简单初始化一个可学习的Positional Embedding矩阵形状是[num_tokens 1, embed_dim]通过训练不断更新。这里的1是class token。ViT借鉴了BERT的[CLS]设计在序列最前面额外加一个可学习的token。这个token不承载任何图像信息它的作用是和所有patch token计算注意力最终拿它的输出接分类头。这样做的好处是把全局信息汇聚到一个固定位置上分类头输入维度恒定不受patch数量影响。class ViT(layers.Layer): def __init__(self, num_patches196, embed_dim768): super(ViT, self).__init__() self.num_patches num_patches self.embed_dim embed_dim self.cls_token self.add_weight( namecls_token, shape[1, 1, embed_dim], initializertf.initializers.TruncatedNormal(stddev0.02), trainableTrue ) self.pos_embed self.add_weight( namepos_embed, shape[1, num_patches 1, embed_dim], initializertf.initializers.TruncatedNormal(stddev0.02), trainableTrue ) def call(self, patches): batch_size tf.shape(patches)[0] cls_tokens tf.tile(self.cls_token, [batch_size, 1, 1]) x tf.concat([cls_tokens, patches], axis1) x x self.pos_embed return x初始化标准差用0.02这个数值是BERT等预训练模型的通行做法我试过用0.1初始化pos_embed训练早期loss下降明显变慢。原因在于位置编码初始值太大会干扰patch token本身的信息表达自注意力在前期会把大量权重分配到位置相似性上而不是内容相似性上。2.3 Swin-T的窗口注意力与位移窗口机制Swin-T最核心的两个创新是Window Attention窗口注意力和Shifted Window Attention位移窗口注意力。窗口注意力的思路是把特征图切成固定大小的窗口通常是7x7只在窗口内部做自注意力。举个例子输入56x56的特征图切成7x7的窗口一共得到8x864个窗口每个窗口内有49个token。注意力计算的复杂度从全局的O(N^2)降到窗口内的O(M^2)其中M49是常数。这让Swin-T可以处理高分辨率输入而ViT在分辨率提升时计算量成平方增长。只在窗口内做注意力显然不够窗口与窗口之间没有信息交流感受野被锁死了。Swin-T的解法很巧妙在下一层Transformer Block中把窗口整体向右下角偏移shift这样原本位于不同窗口的token在新的窗口划分下就“混”到一起了。循环往复信息就能够在不同窗口之间流动。代码实现上这个shift操作有点讲究。直接对特征图做tf.roll实现位移然后重新切窗口计算完注意力后再把窗口还原最后用tf.roll反向移回去。为了避免移位后产生的非规则窗口比如最左上角那个窗口只有1/4大小需要生成一个Mask矩阵把不相关位置的注意力分数置为负无穷。这个Mask的生成是Swin-T实现里最容易出错的地方后面实操部分我会专门讲。2.4 Swin-T的层级结构与patch mergingSwin-T的另一个重要设计是层级化结构。它借鉴了CNN特征金字塔的思路通过Patch Merging层逐步降低空间分辨率、增加通道数。整个过程分四个StageStage 1输入224x224的图patch size是4得到56x56分辨率、96通道的特征图Stage 2Patch Merging把2x2区域的token合并成一个分辨率变成28x28通道数翻倍到192Stage 3继续合并分辨率14x14通道数384Stage 4分辨率7x7通道数768Patch Merging的实现非常直接把2x2邻域的4个token在通道维拼接起来得到4C维然后通过一个线性层压缩到2C维。这里面也加了一个LN。从直觉上理解这个过程相当于不断把空间信息压缩成语义信息和CNN里stride2的卷积十分相似区别在于Patch Merging没有可学习的卷积核纯粹是信息重组加线性变换。这种层级结构给Swin-T带来了一个ViT没有的优势——多尺度特征输出。目标检测里的FPN、语义分割里的U-Net都需要多尺度特征图Swin-T天然就能提供。我在做这个项目的时候是把Swin-T的Stage 3和Stage 4的输出分别接了全局池化和分类头发现比只用最后一层特征效果稳定得多。3. 实操过程与核心环节实现3.1 CIFAR-100数据集的加载与增强策略我先用CIFAR-100做快速验证因为这个数据集小训练一个epoch时间短适合调试模型结构和超参数。验证跑通了再扩展到ImageNet子集或者自定义数据集。CIFAR-100原始图片是32x32直接送进224x224的ViT肯定不行。这里有两个方案一是用插值算法把32x32放大到224x224信息丢失非常严重效果很差二是修改patch_size把32x32的图切成4x4的patchpatch数量是64个序列长度大幅缩短计算开销小很多。我采用的是第二种方案。测试下来patch_size4、embed_dim256、6个Transformer Encoder的轻量级ViT在CIFAR-100上能跑到82%左右的Top-1准确率用RandAugment增强训练时间大约15分钟3090。数据增强我用了完整的组合拳def build_augmenter(): return tf.keras.Sequential([ layers.RandomCrop(32, 32), layers.RandomFlip(horizontal), layers.RandomRotation(0.15), layers.RandomContrast(0.2), layers.RandomZoom(0.1), ])这里有个容易被忽视的点RandomCrop之前必须先用Resizing把图片放大到40x40否则随机裁剪没有意义。CIFAR-100的图片本身就只有32x32不放大的话裁剪区域只能很小增强效果有限。我通常放大到36x36或者40x40再裁剪回32x32。3.2 ViT模型的完整实现与训练配置我把ViT的核心模块拆成三个部分Multi-Head Self-Attention、MLP Block、Transformer Encoder。注意力机制的实现要注意tf.matmul的维度处理。Q、K、V都是从输入线性变换得到的形状是[B, num_tokens, embed_dim]。先reshape成[B, num_tokens, num_heads, head_dim]然后转置成[B, num_heads, num_tokens, head_dim]才能在多头之间并行计算。注意力分数除以sqrt(head_dim)做缩放这一步必不可少。embed_dim768时head_dim64如果省略缩放点积结果动辄几十上百softmax直接饱和梯度非常小模型几乎不学习。class MultiHeadSelfAttention(layers.Layer): def __init__(self, embed_dim768, num_heads12, dropout0.0): super(MultiHeadSelfAttention, self).__init__() self.embed_dim embed_dim self.num_heads num_heads self.head_dim embed_dim // num_heads self.scale self.head_dim ** -0.5 self.qkv layers.Dense(embed_dim * 3, use_biasFalse) self.attn_drop layers.Dropout(dropout) self.proj layers.Dense(embed_dim) self.proj_drop layers.Dropout(dropout) def call(self, x): B, N, C tf.shape(x)[0], x.shape[1], x.shape[2] qkv self.qkv(x) qkv tf.reshape(qkv, [B, N, 3, self.num_heads, self.head_dim]) qkv tf.transpose(qkv, [2, 0, 3, 1, 4]) q, k, v qkv[0], qkv[1], qkv[2] attn tf.matmul(q, k, transpose_bTrue) * self.scale attn tf.nn.softmax(attn, axis-1) attn self.attn_drop(attn) x tf.matmul(attn, v) x tf.transpose(x, [0, 2, 1, 3]) x tf.reshape(x, [B, N, C]) x self.proj(x) return self.proj_drop(x)训练配置上我用的是AdamW优化器初始学习率1e-3warmup 5个epoch之后用cosine decay降到1e-5。权重衰减设为0.05。Transformer对weight decay比对CNN更敏感设置太高容易欠拟合太低又会过拟合。batch size 64总共训练100个epoch。3.3 Swin-T的Mask生成与窗口还原实现Swin-T实现中最大的坑是Shifted Window的attention mask。我用一个具体的例子来解释这个过程。假设输入特征图是8x8window_size7shift_size3。经过tf.roll位移之后原始特征图被切成4个不规则的窗口区域左上角的1x1块、右上角的1x7块、左下角的7x1块、右下角的7x7块。如果按7x7窗口直接切会得到多个不完整的窗口。Swin-T的原始代码通过给每个位置的token标记一个“窗口编号”region index然后比较两个token是否属于同一窗口如果不属于注意力分数设为-100相当于masked。我在TensorFlow里的实现如下def generate_mask(self, h, w, window_size, shift_size): img_mask tf.zeros([1, h, w, 1]) h_slices [ slice(0, -window_size), slice(-window_size, -shift_size), slice(-shift_size, None) ] w_slices h_slices cnt 0 for h_idx in h_slices: for w_idx in w_slices: img_mask[:, h_idx, w_idx, :] cnt cnt 1 # 把mask切成窗口 mask_windows window_partition(img_mask, window_size) mask_windows tf.reshape(mask_windows, [-1, window_size * window_size]) attn_mask tf.expand_dims(mask_windows, 1) - tf.expand_dims(mask_windows, 2) attn_mask tf.where(attn_mask ! 0, -100.0, 0.0) return attn_mask这段逻辑我调试了整整一天才跑通。最初遇到的问题是在位移后没有正确还原特征图导致训练时准确率在40%左右徘徊但loss却显示在下降。排查了半天发现是attention mask加错了位置——我在做完反向移位之后才加mask相当于所有mask都失效了。正确做法是位移 - 切窗口 - 加mask - 注意力计算 - 还原窗口 - 反向位移。3.4 训练完整代码与日志监控我把两个模型封装成统一的训练入口方便切换对比def train_model(model, train_ds, val_ds, config): lr_schedule tf.keras.optimizers.schedules.CosineDecay( initial_learning_rateconfig[lr], decay_stepsconfig[epochs] * config[steps_per_epoch], warmup_targetconfig[lr], warmup_stepsconfig[warmup_epochs] * config[steps_per_epoch], alpha1e-2 ) optimizer tfa.optimizers.AdamW( learning_ratelr_schedule, weight_decayconfig[weight_decay] ) model.compile( optimizeroptimizer, losstf.keras.losses.SparseCategoricalCrossentropy(from_logitsTrue), metrics[tf.keras.metrics.SparseCategoricalAccuracy(nameacc)] ) callbacks [ tf.keras.callbacks.ModelCheckpoint( filepathconfig[checkpoint_path], monitorval_acc, modemax, save_best_onlyTrue, save_weights_onlyTrue ), tf.keras.callbacks.ReduceLROnPlateau( monitorval_loss, factor0.5, patience10, min_lr1e-6 ), tf.keras.callbacks.CSVLogger(training_log.csv) ] history model.fit( train_ds, epochsconfig[epochs], validation_dataval_ds, callbackscallbacks ) return history用CSVLogger记录每一轮的loss和acc有很多好处我后面画loss曲线、分析是否过拟合全是靠这份日志。别只依赖TensorBoardCSV文件处理起来轻量得多而且可以离线分析。4. 常见问题与排查技巧实录4.1 模型不收敛或loss震荡这是Transformer类模型最常见的问题。我遇到过两种典型情况第一种是loss一开始就很高然后持续震荡完全不下降。这种情况十有八九是学习率太大。Transformer的优化空间比CNN更崎岖对学习率更敏感。CNN可以接受0.1的初始学习率Transformer一般只能从1e-3到3e-3开始。我排查的时候先固定warmup为5个epoch把峰值学习率降到3e-4问题立刻缓解。第二种是loss正常下降但到某个点后开始剧烈震荡。这通常是因为batch size太小导致梯度噪声过大或者是warmup阶段太短。ViT在CIFAR-100上batch size从32改到64之后loss曲线明显平滑很多。还有一个隐蔽的问题LayerNormalization的epsilon设置。默认是1e-6但在混合精度训练下有概率出现方差为0的极端情况导致NaN。我建议统一设置为1e-5或者1e-4几乎不影响精度但能大幅提高稳定性。4.2 显存占用过高与批量大小权衡Transformer比CNN吃显存是出了名的。我之前在3090上跑ViT-Base输入224x224、batch size 64显存占用直接冲到22G。这里有几个实用的省显存技巧使用混合精度训练tf.keras.mixed_precision.set_global_policy(mixed_float16)显存占用降低约40%而且RTX 30系显卡的Tensor Core能加速计算减少梯度检查点虽然能省显存但会显著拖慢训练速度不推荐先在小分辨率上训练再用大分辨率微调。比如ViT在CIFAR-100上patch_size4、32x32输入显存占用不到4G还有个很多人不知道的细节tf.data管线的prefetch会额外占用显存。如果prefetch的buffer开太大数据预处理的结果会缓存在GPU内存里。我通常用tf.data.AUTOTUNE自动调节而不是手动指定一个很大的数值。4.3 Swin-T训练速度比预期慢很多Swin-T在TensorFlow里的训练速度经常被人吐槽。在CIFAR-100上Swin-T一个epoch要比ViT慢将近50%。经过profiling我发现瓶颈不在注意力计算本身而是两个地方一是Window Partition操作。每次切窗口都要做tf.transpose和tf.reshape这两个操作在GPU上会触发数据拷贝开销很大。我优化了窗口切分的实现把多次小reshape合并成一次大reshape速度提升了30%。二是tf.roll在反向传播时效率很低。Swin-T每个Block都要做两次位移和两次还原GDDR的带宽都被这个操作吃掉了。如果你的GPU显存够大可以尝试把位移和mask融合到注意力计算里从数学上等价但能少几次显存读写。凭心而论如果只是做图像分类任务Swin-T相比ViT没有性能优势反而更慢。它的优势体现在检测、分割这种需要高分辨率特征图的任务上。选型时要先想清楚任务的本质需求。4.4 训练集小数据量下Transformer很容易过拟合Transformer的参数量巨大ViT-Base有8600万参数Swin-T有2800万参数在小数据集上过拟合几乎是必然的。CIFAR-100的5万张训练图对小Transformer来说勉强够用但必须配合强数据增强和正则化。我实测有效的防过拟合组合RandAugmentPython库tensorflow_addons.image里有现成实现或者用imgaugCutMix或MixUp这两种增强方式能显著提升Transformer的泛化能力在CIFAR-100上能带来3到5个点的提升。实现MixUp很简单把两张图按比例混合标签也按比例混合。CutMix稍微复杂一些要把一张图的某个区域裁剪后贴到另一张图上。我建议先用MixUp效果好且代码短DropPathStochastic Depth在训练时随机跳过某些残差分支。ViT和Swin-T的官方实现都带有DropPath概率一般设为0.1到0.3。TensorFlow里没有现成API需要自己写一个class DropPath(layers.Layer): def __init__(self, drop_probNone): super(DropPath, self).__init__() self.drop_prob drop_prob def call(self, x, trainingNone): if self.drop_prob 0.0 or not training: return x keep_prob 1.0 - self.drop_prob shape [tf.shape(x)[0]] [1] * (len(x.shape) - 1) random_tensor keep_prob tf.random.uniform(shape, dtypex.dtype) binary_tensor tf.floor(random_tensor) return x / keep_prob * binary_tensor4.5 常见问题速查表问题现象可能原因解决方案loss初始很高且持续震荡学习率过大初始学习率降到3e-4以下配合warmup训练后期loss突变为NaNLayerNorm的epsilon太小或混合精度溢出把epsilon设为1e-4检查是否用了float16但没有loss scaling验证集准确率远低于训练集过拟合增加DropPath率、使用MixUp/CutMix、提高weight decaySwin-T窗口注意力报维度错误mask和输入尺寸不匹配确认input size能被window_size整除或者正确设置shift_size训练速度忽快忽慢数据管线有瓶颈使用tf.data.AUTOTUNE调整prefetch buffer size加载预训练权重报错模型层命名不一致手动对比model.weights名称和权重文件中的名称用model.get_layer(name)定位4.6 从CIFAR-100迁移到自定义数据集的经验如果你的目标不是跑通Demo而是解决实际问题很可能需要在自己的数据集上训练。这里有几个操作层面的建议。自定义数据集建议用tf.keras.utils.image_dataset_from_directory加载代码量最少自动处理标签train_ds tf.keras.utils.image_dataset_from_directory( dataset/train, image_size(224, 224), batch_size64, shuffleTrue, seed42, validation_split0.1, subsettraining )目录结构要求是每个类一个文件夹文件夹名就是标签名这对中小规模项目足够用了。从CIFAR-100切到自己的数据集有三个必改项模型输入尺寸、patch大小或窗口大小、最后分类头的类别数。CIFAR-100是32x32小图很多自定义数据集是224x224甚至更大。如果直接换尺寸位置编码的维度就对不上了。ViT的位置编码是训练出来的换分辨率必须重新训练位置编码或者用插值的方式初始化。Swin-T就没有这个问题因为它的窗口大小固定分辨率变化后窗口数量自动变化。这也是Swin-T在实际项目中更受欢迎的原因之一。我遇到过好几次这种情况模型在低分辨率训练好了后来采集设备升级输入变成高分辨率图ViT需要额外处理位置编码的适配Swin-T几乎不用改代码。5. 一些关于模型部署和扩展的联想训练完模型不等于项目结束部署和推理优化在实际工作中同样关键。TensorFlow模型的导出很简单用model.save(vit_cifar100.keras)保存整个模型后续用tf.keras.models.load_model加载。如果要部署到服务端推荐转成TensorFlow SavedModel格式用model.export(saved_model_dir)然后在服务端用TensorFlow Serving加载。TensorFlow Serving支持动态batch可以合并并发请求吞吐量提升明显。如果是边缘设备部署需要转成TensorFlow Lite格式。Transformer类的模型转TFLite有一点要特别注意某些动态操作比如tf.roll、带动态shape的tf.reshape在TFLite里不支持需要提前把动态shape改成固定shape。Swin-T因为有窗口位移操作转换TFLite时最容易报错。推理优化方面一个非常有效的技术是模型蒸馏。把训练好的大模型当Teacher用一个小的CNN或者小Transformer当Student用Teacher的软标签指导Student学习。我试过用ViT-Base蒸馏出一个ResNet-18大小的学生模型准确率只掉了1.5个点但推理速度提升了8倍。这个方向在工业落地时非常实用。量化也是一个思路。TensorFlow Lite默认支持int8量化但Transformer的量化损失通常比CNN大。如果你要量化感知训练建议关注一下TensorFlow Model Optimization Toolkit的最新能力新版对Transformer的支持比老版本好很多。这个项目的代码我已经整理好放在GitHub仓库里包含完整的ViT和Swin-T实现、训练脚本、数据增强策略和推理示例。如果有人跑通了代码或者在实际应用中发现新的问题欢迎留言交流。我最初写这个项目完全是为了搞清楚Transformer在视觉上为什么能和CNN抗衡真正实践下来对两类架构的优缺点看得更透彻了。有时候读论文以为自己懂了动手实现一遍才会发现那些论文里不会写、但对结果至关重要的细节——比如Swin-T的mask要怎么加再比如ViT带着大初始化走不动路这些才是真正值得写进博客里的东西。
返回列表