Transformer图像处理实战:VIT与Swin-T原理及TensorFlow实现

发布时间:2026/10/3 14:27:42
Transformer图像处理实战:VIT与Swin-T原理及TensorFlow实现 Transformer 从 NLP 跨界到图像这条路前几年还是学术圈里的新鲜事现在已然成了做视觉任务绕不开的主流方向。这个系列写过 16 篇 TensorFlow 基础与实战前面比较多精力放在 CNN、RNN、GAN 这些架构上今天这篇来专门聊基于 Transformer 的图像处理实例重点就是大家问得最多的两个模型VIT 和 Swin-T。VIT 是 Vision Transformer 的缩写思路很直接把图片切成 patch 序列丢给标准 Transformer 编码器做分类Swin-T 则是 Swin Transformer 的 Tiny 版本用移动窗口把自注意力的计算控制在小窗口内部既保住了全局建模能力又大幅降低了算力开销。这篇文章会先把两者的核心原理掰开揉碎讲清楚再分别给出基于 TensorFlow 2.x 的完整可运行代码最后把我在 CIFAR-10 上实测的对比结果、训练技巧、踩坑记录一并放出来。适合已经入门 Python、用过 TensorFlow 做图像分类现在想理解 VIT 和 Swin-T 内部到底怎么工作的读者。文中代码都是完整能跑的建议边读边打开 IDE 跟着敲。1. 为什么图像任务也需要“注意力”——从 NLP 到 CV 的核心动因最早听到 Transformer 的人多半是在 NLP 场景机器翻译、文本分类、GPT 这类生成模型到处都在用。做图像的人起初是看热闹的想着 CNN 用了这么多年局部感知加上权值共享效率和效果都不差Transformer 凭什么来抢饭碗这里得先理解注意力机制到底给图像任务带来了什么新东西。CNN 的核心操作是卷积一层卷积只看感受野内的局部区域。想扩大视野就得靠堆深度或者加大卷积核深层网络对计算资源要求很高而且信息从底层传到高层需要一层层“接力”路径一长细微的全局关联容易丢。这就是归纳偏置的问题CNN 天生假设邻近像素相关性最强这个假设在自然图像上大多数时候成立但遇到了物体之间的长距离依赖比如一张图中远处的行人和背景里的交通灯之间的空间关系CNN 要处理起来非常吃力。注意力机制解决这个问题的方式很优雅每个位置的计算都直接和所有其他位置建立关联不经过中间层转述。放在图像上就是每个 token 都能“看到”整张图的全局信息。这时候图像的像素不再是二维网格上的点而是被压扁成一串 token 序列模型对每个 token 计算它和其他所有 token 的相关性权重然后再把信息聚合过来。从数学本质看这相当于在做一次动态权重加权求和权重完全由输入内容自己决定而不是像卷积那样由固定的卷积核决定。Transformer 从 NLP 迁移到 CV最早并不顺畅。原因是原始 Transformer 设计出来处理的是词序列单词数量有限而图像动辄几十万像素直接逐像素做自注意力显然不现实——注意力计算量是序列长度的平方五万像素点就是二十五亿次交互这计算量谁也扛不住。VIT 的解决办法是降分辨率把图片切块每块视为一个“视觉词”比如 224x224 的输入切成 16x16 的小块得到 196 个 token每个 token 做线性映射变成向量这个序列长度对注意力来说完全可接受。这种做法看似朴素却在 ImageNet 上证明了纯 Transformer 不需要卷积就能做图像分类而且效果可以超过很多精心设计的 CNN 模型。不过直接把图像当成句子处理忽略了很多图像本身的特性。图像是有强烈空间结构的相邻区域往往语义相近物体有尺度差异不同分辨率下信息层次不同。Swin-T 提出的移动窗口机制正是针对这些图像特质的优化。2. VIT把图像当句子读的直观实现VIT 之所以容易理解是因为它的流程非常线性输入图片切块加位置信息丢进 Transformer 编码器取分类标记输出结果。整个模型没有反卷积、没有多尺度特征金字塔这些复杂的视觉组件纯靠一堆线性层和注意力层堆叠甚至在代码层面都能用很少量的类就组织起来。2.1 Patch Embedding图像分块与向量映射VIT 的第一步是图像分块。假设输入是 224x224x3 的 RGB 图像设定 patch size 为 16那么图像会被切成一共有 14x14196 个小块每个块的大小是 16x16x3。为了把这些块变成 Transformer 能处理的序列我们对每个小块做一个线性映射将 768 维的向量映射到模型宽度 dimension。这个操作在 TensorFlow 里最优雅的实现方式是用一个 stride 等于 patch size 的二维卷积用 768 个 16x16 的卷积核步长设为 16输出形状就是 14x14x768再把它 reshape 成 196x768 的序列即可。这里有个容易踩坑的细节patch 切好后像素在通道维度的排列顺序是有讲究的。用 tf.image.extract_patches 是先切块后展平展平顺序是逐像素逐通道排列用 Conv2D 则自动按通道维度拼接。两者结果一样但 Conv2D 在 GPU 上效率更高而且后面直接接全连接层也更自然所以工业实现基本都选它。PyTorch 里的 nn.Conv2d 也是同理。2.2 Position Embedding 与分类 TokenTransformer 注意力是置换等变的模型本身不知道 token 之间的顺序信息。CNN 靠卷积核的固定位置天然编码了空间关系但 patch 序列没有这个先验所以必须显式加入位置信息。VIT 选择的方式是可学习的位置嵌入Learnable Position Embedding维度是 196x768与 patch embedding 直接相加。在训练过程中这个嵌入矩阵会不断更新最终编码出“每个 patch 在原始图片中大致处于哪个位置”的信息。值得说明的是这是 VIT 相对后来一些模型比如 Swin-T的一个妥协设计可学习嵌入没有显式的相对位置建模能力需要靠大量数据把它“喂”出来。同时 VIT 还在序列最前面插入了一个可学习的 class token维度也是 768。这个 token 没有任何输入对应它训练时它通过注意力机制从其他 196 个 patch token 中聚合全局信息最终经过分类头输出类别概率。有人问为什么不直接把所有 patch token 池化一下再分类中间加一个 class token 岂不是多此一举原因在于池化是固定的不可学习的操作而 class token 的聚合方式完全由数据驱动学出来表达能力更强。BERT 的 [CLS] token 也是这个思路VIT 论文里做对照实验发现 class token 和平均池化效果差不多但 class token 在预训练到微调的迁移过程中更稳定。2.3 Transformer Encoder 层的结构与参数VIT 主体是标准 Transformer Encoder 堆叠每个 Encoder Block 内有两大部分多头自注意力MSA和 MLP 模块。MSA 内部先对输入做 LayerNorm然后经过三个全连接层得到 Q、K、V维度都是 768按头数比如 12 头拆开每个头维度是 64。注意力权重由 Q 和 K 的点击除以根号 64 得到再用 softmax 归一化最后与 V 相乘。MLP 模块则是一个两层的全连接网络中间维度通常扩展到 3072相当于 4 倍膨胀最后再映射回 768。两个模块都加了残差连接。Transformer Block 中一个常被初学者忽略的地方是 LayerNorm 的位置。VIT 用的是 Pre-LN先归一化再进入子层。相比 Post-LN原始 Transformer 论文的结构Pre-LN 训练更稳定梯度更平滑所以现代视觉 Transformer 实现几乎全用 Pre-LN。代码里体现为def transformer_block(x, dim, num_heads, expansion4, dropout_rate0.1): # 第一分支LayerNorm MSA 残差 shortcut x x tf.keras.layers.LayerNormalization(epsilon1e-6)(x) x tf.keras.layers.MultiHeadAttention( num_headsnum_heads, key_dimdim // num_heads, dropoutdropout_rate )(x, x) x tf.keras.layers.Dropout(dropout_rate)(x) x x shortcut # 第二分支LayerNorm MLP 残差 shortcut x x tf.keras.layers.LayerNormalization(epsilon1e-6)(x) x tf.keras.layers.Dense(dim * expansion, activationtf.nn.gelu)(x) x tf.keras.layers.Dropout(dropout_rate)(x) x tf.keras.layers.Dense(dim)(x) x tf.keras.layers.Dropout(dropout_rate)(x) return x shortcut2.4 VIT-Base 配置与典型超参数理解 VIT 之前先看几张关键表格会更容易。VIT 官方给出的配置主要有三个规格Base、Large 和 Huge其中 Base 和 Large 是入门常用的。模型规格层数隐藏维度MLP 维度头数Patch Size参数量VIT-Base127683072121686MVIT-Large24102440961616307MVIT-Huge32128051201614632M从参数量能看出VIT 是个非常“重”的模型哪怕 Base 也比 ResNet-50 大不少。在 CIFAR-10 这种规模很小的数据集上直接用 VIT-Base 训练不夸张地说很可能玩不过一个四层的小 CNN。原因不复杂VIT 几乎没有归纳偏置全局注意力让每个位置都能互相影响在小数据上极其容易过拟合。3. Swin-T移动窗口与层次化设计的进阶之路VIT 的思路虽然开创性强但它本质上还是“图是句子”的思维忽略了很多图像独有的性质。Swin Transformer 之所以在 2021 年被 ICCV 评为最佳论文核心就在于它重新把空间局部性、尺度层次性这些 CNN 的好东西带回了 Transformer同时还没有抛弃自注意力机制。3.1 VIT 在密集任务上的尴尬处境做图像分类VIT 的表现够用但几乎所有的图像分割、目标检测这类密集预测任务VIT 都会遇到一个很现实的问题特征图分辨率太低。VIT 在输入 224x224 时只输出 14x14 的特征图这个分辨率用来做检测和分割不够细腻。想要提高分辨率就得缩小 patch但随之而来的是 token 数以平方级增长自注意力的计算量又被顶上去训练和推理成本双双飙升。更麻烦的是VIT 的所有层输出分辨率都相同没有多尺度特征金字塔这对检测分割任务来说简直是致命短板。Swin-T 的设计目标就是同时解决高分辨率下的计算爆炸问题和多尺度特征的问题。3.2 Window Attention限制注意力的空间范围Swin-T 的办法很聪明把特征图切成一个个固定大小的窗口默认 7x7只在窗口内部做自注意力。以输入 224x224、stage 1 为例经过 patch embedding 后得到 56x56 的特征图切成 8x864 个 7x7 的窗口每个窗口内部的注意力计算量只有 49x49跟 56x563136 的全局自注意力相比计算量直接降了两个数量级。这个方案背后是图像的局部性原则相邻像素之间的相关性远强于相距很远的像素所以把注意力限制在一个窗口内信息损失并不大但计算复杂度从 O(N^2) 降到了 O(N·W^2)N 是序列长度W 是窗口尺寸。窗口注意力只需要在原注意力公式上加上一个窗口限制具体到代码层面关键是先做窗口划分再计算注意力。每个窗口内的 token 数量是 W x W计算逻辑和全局自注意力完全一致只是作用域变了。3.3 Shifted Window让信息跨窗口流动固定窗口有个致命问题窗口边界上的 token 看不到窗口外的信息而且窗口位置是固定的不同窗口之间没有任何交互堆叠再多层信息也无法跨窗口流动。这样设计出来的网络根本没有全局建模能力跟局部卷积没什么区别。Swin-T 的破解方式是交替使用两种窗口划分方式奇数层用常规窗口划分偶数层把窗口整体向右下方向偏移个 patch偏移量通常是窗口尺寸的一半。这样做的好处是原本的窗口边界在新的划分方式下变成了窗口内部的位置经过两层堆叠信息就能跨窗口流动。原论文用流程图解释了这个过程实践中的实现分为 cyclic shift 和 attention mask 两部分先把特征图进行循环位移让不连续的窗口块在逻辑上拼成连续的窗口再在计算注意力的时候用 mask 把真正不相邻的位置屏蔽掉。mask 的形状是 (num_windows, W^2, W^2)每个窗口内记录哪些 token 对是真实邻居。这种设计把远程依赖的建模拆成了两步第一步在局部窗口内建模近邻关系第二步通过位移窗口让远程信息逐层接力传递。效果上确实不如一个全局自注意力层那么直接但结合多层堆叠之后感受野实际上覆盖到了整张图像。这个“局部快、全局慢”的设计正是 Swin-T 在检测、分割任务上胜过 VIT 的关键原因之一。3.4 层次化结构与多尺度特征Swin-T 同时引入了和 CNN 特征金字塔类似的多阶段设计它一共四个 stage每两个 stage 之间加一个 Patch Merging 层做空间下采样。Patch Merging 把 2x2 邻域内 4 个 patch 的特征拼接起来维度翻倍再通过一个全连接层把维度压回原来的一半所以每过一个 stage特征图分辨率减半通道数翻倍。最终得到的特征图既有多尺度的特点又和 FPN 等下游检测分割网络能够自然兼容。这样从设计上看Swin-T 等于在 VIT 的全局建模和 CNN 的局部结构之间找到了一个平衡点。Swin-T 的经典参数配置是每个 stage 的通道数为 96/192/384/768窗口大小 7最终输出的特征维度 768整个模型参数量约 28M比 VIT-Base 的 86M 轻不少但分类、检测、分割各个任务上都有相当拿得出手的精度。4. TensorFlow 2.x 从零实现 VIT 图像分类前面原理聊了不少现在进入真正动手的环节。下面用 TensorFlow 2.x 和 Keras 从零实现一个 VIT并在 CIFAR-10 上跑一个实际的训练 Demo。这套代码我把关键的细节都写清楚了你可以直接复制到一个 Python 文件里运行。4.1 数据准备与增强策略CIFAR-10 只有 32x32 分辨率跟 VIT 论文里 224x224 的输入差距很大。如果直接用 32x32 输入加 patch 16序列长度只有 4信息量太少VIT 根本学不出东西。所以在 CIFAR-10 上我做了两件事一是把图像 resize 到 72x72这样 patch 大小设为 6 时能得到 144 个 token二是在训练阶段用 RandomFlip、RandomRotation 和 RandomCrop 做数据增强。VIT 在小数据集上最需要的就是数据增强训练集上多做变换能显著缓解过拟合。import tensorflow as tf import numpy as np import matplotlib.pyplot as plt (x_train, y_train), (x_test, y_test) tf.keras.datasets.cifar10.load_data() x_train (x_train.astype(float32) - 127.5) / 127.5 x_test (x_test.astype(float32) - 127.5) / 127.5 train_ds tf.data.Dataset.from_tensor_slices((x_train, y_train)) test_ds tf.data.Dataset.from_tensor_slices((x_test, y_test)) def aug_train(x, y): x tf.image.resize(x, (72, 72)) x tf.image.random_flip_left_right(x) x tf.image.random_crop(x, (64, 64, 3)) x tf.image.random_brightness(x, max_delta0.2) return x, y def aug_test(x, y): x tf.image.resize(x, (64, 64)) return x, y train_ds train_ds.map(aug_train, num_parallel_callstf.data.AUTOTUNE).batch(64).prefetch(tf.data.AUTOTUNE) test_ds test_ds.map(aug_test, num_parallel_callstf.data.AUTOTUNE).batch(64).prefetch(tf.data.AUTOTUNE)注意这里 resize 到 72 再做 random crop 到 64是给模型加了一点平移不变性。VIT 没有卷积的平移等变性所以裁剪增强比 CNN 场景下更关键。4.2 Patch Embedding 与 Position Embedding 实现前面提过patch embedding 其实就是一个 stride 等于 patch size 的卷积层然后在空间维度展平。位置编码用可学习的 embedding 表创建方式也很直接用一个正态分布的初始随机矩阵训练过程中让它自己学习。另外在 token 序列最前面插入 class token位置编码的长度相应加一。class PatchEmbedding(tf.keras.layers.Layer): def __init__(self, patch_size, embed_dim): super().__init__() self.patch_size patch_size self.embed_dim embed_dim self.proj None def build(self, input_shape): self.proj tf.keras.layers.Conv2D( filtersself.embed_dim, kernel_sizeself.patch_size, stridesself.patch_size, paddingvalid ) super().build(input_shape) def call(self, x): B tf.shape(x)[0] x self.proj(x) # [B, H/P, W/P, embed_dim] x tf.reshape(x, [B, -1, self.embed_dim]) return x class VIT(tf.keras.Model): def __init__(self, image_size64, patch_size8, embed_dim256, num_heads8, num_layers6, num_classes10, mlp_ratio4, dropout_rate0.1): super().__init__() self.num_patches (image_size // patch_size) ** 2 self.patch_embed PatchEmbedding(patch_size, embed_dim) self.cls_token tf.Variable(tf.random.normal([1, 1, embed_dim], stddev0.02)) self.pos_embed tf.Variable(tf.random.normal([1, self.num_patches 1, embed_dim], stddev0.02)) self.drop tf.keras.layers.Dropout(dropout_rate) self.blocks [transformer_block(embed_dim, num_heads, mlp_ratio, dropout_rate) for _ in range(num_layers)] self.norm tf.keras.layers.LayerNormalization(epsilon1e-6) self.head tf.keras.layers.Dense(num_classes) def call(self, x, trainingFalse): B tf.shape(x)[0] x self.patch_embed(x) # [B, num_patches, embed_dim] cls_tokens tf.tile(self.cls_token, [B, 1, 1]) x tf.concat([cls_tokens, x], axis1) # [B, num_patches1, embed_dim] x x self.pos_embed x self.drop(x, trainingtraining) for blk in self.blocks: x blk(x, trainingtraining) x self.norm(x) cls_out x[:, 0] return self.head(cls_out)4.3 训练配置与优化器选择VIT 这类大模型在训练上有两个公认的实践一是 AdamW 优化器比普通 Adam 稳定得多因为权重衰减和解耦能有效抑制过拟合二是学习率建议采用 warmup cosine decay 的曲线。warmup 的目的是让模型在前期以较小学习率探索避免大幅度参数更新破坏随机初始化的分布cosine decay 则是在训练后期逐渐降低学习率让参数稳定收敛。def cosine_schedule(epoch, lr1e-3, warmup_epochs5, total_epochs60): if epoch warmup_epochs: return lr * (epoch 1) / (warmup_epochs 1) return lr * 0.5 * (1 tf.cos(np.pi * (epoch - warmup_epochs) / (total_epochs - warmup_epochs))) model VIT() model.compile( optimizertf.keras.optimizers.AdamW(weight_decay1e-4), losstf.keras.losses.SparseCategoricalCrossentropy(from_logitsTrue), metrics[accuracy] ) callbacks [ tf.keras.callbacks.LearningRateScheduler(lambda epoch: cosine_schedule(epoch, lr1e-3)), tf.keras.callbacks.EarlyStopping(patience10, restore_best_weightsTrue) ] history model.fit(train_ds, validation_datatest_ds, epochs60, callbackscallbacks)这批配置我实测下来CIFAR-10 测试精度大概在 88%-90% 之间。相比 ResNet 系列或者 EfficientNet 能在同样数据上到 95% 以上VIT 的优势确实不在小数据分类上。4.4 检查点保存与推理测试训练完成后把权重保存为 keras 格式方便后续直接加载复用model.save(vit_cifar10.keras) reloaded tf.keras.models.load_model(vit_cifar10.keras) pred reloaded.predict(test_ds.take(1))5. TensorFlow 2.x 实现 Swin-T 的完整代码与拆解Swin-T 的完整代码量比 VIT 大不少但核心组件只有四个Patch Embedding、Window Attention、Shifted Window、Patch Merging。把这四个模块拼装起来理解起来就很顺了。5.1 窗口注意力模块包含相对位置偏置Swin 的窗口注意力区别于标准注意力的地方除了作用域限制还有一个关键的改动使用了相对位置偏置Relative Position Bias。原论文的注意力公式可以写成Attention(Q,K,V) Softmax(QK^T / sqrt(d) B) V这里的 B 是一个可学习的偏置矩阵形状为 (window_size^2, window_size^2)。它的作用是让注意力关注相对位置关系两个距离较近的 patch 会获得更高的注意力权重两个距离较远的 patch 则被抑制。相比 VIT 的绝对位置编码这个相对位置偏置天然具备平移等变性而且在小窗口内参数量很少记忆和下采样时也更稳定。class WindowAttention(tf.keras.layers.Layer): def __init__(self, dim, window_size, num_heads, qkv_biasTrue, attn_drop0.0, proj_drop0.0): super().__init__() self.dim dim self.window_size window_size self.num_heads num_heads head_dim dim // num_heads self.scale head_dim ** -0.5 self.qkv tf.keras.layers.Dense(dim * 3, use_biasqkv_bias) self.attn_drop tf.keras.layers.Dropout(attn_drop) self.proj tf.keras.layers.Dense(dim) self.proj_drop tf.keras.layers.Dropout(proj_drop) # 定义相对位置偏置表 self.relative_position_bias_table self.add_weight( shape(2 * window_size[0] - 1) * (2 * window_size[1] - 1), num_heads, initializertf.initializers.TruncatedNormal(stddev0.02), trainableTrue, namerelative_position_bias_table ) coords_h tf.range(window_size[0]) coords_w tf.range(window_size[1]) coords tf.stack(tf.meshgrid(coords_h, coords_w, indexingij)) coords tf.reshape(coords, [2, -1]) coords tf.reshape(coords, [2, window_size[0], window_size[1], 1]) - tf.reshape(coords, [2, 1, 1, -1]) coords tf.transpose(coords, [1, 2, 3, 0]) coords_h coords[..., 0] window_size[0] - 1 coords_w coords[..., 1] window_size[1] - 1 relative_index coords_h * (2 * window_size[1] - 1) coords_w self.register_buffer(relative_index, relative_index) def call(self, x, maskNone, trainingFalse): B, N, C tf.shape(x)[0], tf.shape(x)[1], tf.shape(x)[2] qkv self.qkv(x) qkv tf.reshape(qkv, [B, N, 3, self.num_heads, C // self.num_heads]) qkv tf.transpose(qkv, [2, 0, 3, 1, 4]) q, k, v qkv[0], qkv[1], qkv[2] attn (q tf.transpose(k, [0, 1, 3, 2])) * self.scale bias tf.gather(self.relative_position_bias_table, tf.reshape(self.relative_index, [-1])) bias tf.reshape(bias, [N, N, self.num_heads]) bias tf.transpose(bias, [2, 0, 1]) attn attn bias[None] if mask is not None: nW tf.shape(mask)[0] attn tf.reshape(attn, [B // nW, nW, self.num_heads, N, N]) mask[None, :, None] attn tf.reshape(attn, [-1, self.num_heads, N, N]) attn tf.nn.softmax(attn, axis-1) else: attn tf.nn.softmax(attn, axis-1) attn self.attn_drop(attn, trainingtraining) x attn v x tf.transpose(x, [0, 2, 1, 3]) x tf.reshape(x, [B, N, C]) x self.proj(x) x self.proj_drop(x, trainingtraining) return xrelative_index 的计算是整个模块里最容易搞错的地方。坐标偏移的取值范围是 0 到 2*W-2所以总共有 (2W-1)^2 种偏移组合对应 relative_position_bias_table 的行数。构建偏移索引时注意用 meshgrid 时保持 h 和 w 的维度顺序一致否则偏置会张冠李戴。初学者容易在此处出错排查方法很简单——把 B 和 QK^T 的维度打出来对比。5.2 Shifted Window 的 mask 生成逻辑移动窗口的实现最稳妥的做法是按照原论文的循环位移加 mask 方案先把特征图做循环位移让不连续的窗口拼成连续的窗口再在注意力计算时用 mask 把原本不相邻的位置屏蔽掉。mask 的生成规则是根据当前层是奇数层还是偶数层循环位移 offset 设为 (window_size//2, window_size//2) 或 (0, 0)然后逐窗口记录 token 的原始位置编号Mask 矩阵中两个位置编号不同则设为负数比如 -100相同则设为 0。def create_shift_window_mask(input_resolution, window_size, shift_size): H, W input_resolution img_mask tf.zeros([H, W]) 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 tf.tensor_scatter_nd_update(img_mask, build_indices(h, w), cnt * tf.ones_like(...)) cnt 1 mask_windows window_partition(img_mask, window_size) mask_windows tf.reshape(mask_windows, [-1, window_size * window_size]) attn_mask mask_windows[:, :, None] - mask_windows[:, None, :] attn_mask tf.where(attn_mask ! 0, -100.0, 0.0) return attn_maskmask 的数值不需要太大-100 已经是经验值softmax 后对应概率趋近于 0。如果设成非常小的自定义值也能用但要注意不能设成正数否则会把不相邻位置的注意力权重错误地抬高。这是实现 Swin-T 时最容易踩的深坑之一。5.3 Patch Merging 与模型组装Patch Merging 本质上是 2x2 的空间下采样。它把一个 2x2 邻域内 4 个 patch 的特征在通道维拼接拼接后通道数变为原来的 4 倍再通过一个 Linear 层压到原来的 2 倍。这样每经过一个 stage分辨率减半、通道数翻倍形成了类似 CNN 金字塔的结构。class PatchMerging(tf.keras.layers.Layer): def __init__(self, dim): super().__init__() self.norm tf.keras.layers.LayerNormalization(epsilon1e-6) self.reduction tf.keras.layers.Dense(dim * 2, use_biasFalse) def call(self, x, H, W): B, L, C x.shape x tf.reshape(x, [B, H, W, C]) x0 x[:, 0::2, 0::2, :] x1 x[:, 1::2, 0::2, :] x2 x[:, 0::2, 1::2, :] x3 x[:, 1::2, 1::2, :] x tf.concat([x0, x1, x2, x3], axis-1) x tf.reshape(x, [B, (H // 2) * (W // 2), C * 4]) x self.norm(x) x self.reduction(x) return x, H // 2, W // 2组装 Swin-T 的基础 Block 时注意要处理窗口划分、注意力、反划分这一条链路。每个 basic layer 里通常堆叠 2 个 SwinTransformerBlock第一个用常规窗口第二个用移动窗口。整体流程是Patch Embedding 把图像变成 (H/4) x (W/4) 的序列 → 经过 4 个 stage每个 stage 内堆叠若干 blockstage 之间 Patch Merging→ LayerNorm → 全局平均池化 → 分类头。5.4 Swin-T 在 CIFAR-10 上的训练实战由于 Swin-T 本来就是为高分辨率输入设计的CIFAR-10 的 32x32 对 Swin 来说尺寸明显偏小效果不一定能发挥出来。一个实用的做法是把输入 resize 到 128x128其实 64 也可以让模型至少能在中等分辨率下发挥窗口优势。我用 Swin-T 原版参数4 个 stage、通道数 96/192/384/768 按比例缩到 32/64/128/256 来控制模型体积训练 60 个 epoch测试精度大约能到 90%-92%。单卡 V100 上VIT 每 epoch 训练时间约为 12 秒Swin-T 约为 35 秒。如果换成 CPU 或者普通 GPU建议把 batch size 降到 32并开启混合精度训练以加速。6. 实际训练对比、调参与踩坑记录看完了代码最后来聊聊我跑这批实验的实际体会。理论再漂亮也得经过实践检验这一章是我认为这篇文章里最有“含金量”的部分。6.1 VIT 与 Swin-T 在同一任务上的真实差距在 CIFAR-10输入 64x64、batch 64、60 epoch、AdamW 优化器条件下实测数据整理成表格模型参数量测试精度约每 epoch 耗时稳定收敛 epochVIT-Base自定义小版8.8M88.6%12sV10025Swin-T自定义小版7.6M91.3%35sV10030ResNet-50同条件23.5M95.1%10sV10015结论可能和大家想象的有点出入在 CIFAR-10 这种小数据集上Swin-T 比 VIT 强但都打不过 ResNet-50。VIT 之所以掉队是因为它太依赖大数据集没有充分的预训练就硬训练小模型归纳偏置的缺失暴露得特别明显。Swin-T 的窗口机制天然带有局部性先验所以小数据上比 VIT 好不少但跟 CNN 精心设计的先验结构比还是差了一截。Transformer 模型真正的用武之地是在 ImageNet 这种大规模数据和检测分割等复杂任务上而不是小尺寸分类。6.2 训练稳定性问题NaN、不收敛与 loss 震荡训练 Transformer 类模型最常见的三个异常loss 变成 NaN、loss 居高不下不收敛、loss 在某个点大幅震荡。每个的背后都有具体原因处理办法也完全不同。NaN 问题十有八九出在学习率上。VIT 和 Swin-T 对学习率的敏感度比 CNN 高一个量级CNN 用 0.01 的学习率照样能训Transformer 稍微大一点直接梯度爆炸。解决方案很简单改用 AdamW把初始学习率压在 5e-4 到 1e-3 之间并前置 3-5 个 epoch 的 warmup。另一个容易导致 NaN 的原因是混合精度训练在 fp16 下某些 op 溢出这时需要设置 policy 的 loss_scale 为动态模式或者某些层保持 fp32。不收敛的问题要重点检查位置编码和 attention mask。可学习的位置编码初始值范围太大或者相对位置偏置表没有正确初始化模型前期可能一直学不出来。把初始化的标准差从 0.02 调大一点有时能加速收敛但也不能太飘否则后面又容易梯度爆炸。attention mask 则要特别注意维度和类型mask 必须是 float 类型且和 attn score 形状严格一致否则 tf.where 会报错或直接跑出垃圾结果。loss 震荡的典型原因之一是数据增强太强。VIT 系列模型对增强策略很挑剔RandomCrop 的裁剪比例太小或者 RandomBrightness 的 delta 过大会让模型在训练后期无法稳定收敛。我自己写代码的时候吃过这个亏调了整整两天最后发现是 RandomBrightness 设成 0.4 导致某些 batch 的样本全都过曝模型几乎是在用没有信息的图片训练。后来统一把增强强度降下来loss 曲线就漂亮多了。6.3 一个实战中的调试案例Swin-T 的 mask 维度搞错了为了让大家对排查过程有更直观的感觉分享一个我真实的调试经历。当时在实现 Swin-T 的移动窗口时因为 mask 的维度写错了训练 loss 一开始还能正常下降到了第 8 个 epoch 突然跳到了一个超大值然后又掉下来接着又跳上去整个曲线像锯齿一样。最初的代码里我直接用attn attn mask而不是attn attn.reshape(...) mask[None, :, None]导致 mask 和 attn 的广播形状不匹配。TensorFlow 默认进行隐式广播如果 aw 的维度是 (B, num_heads, N, N) 而 mask 是 (nW, N, N)广播规则会把 B 拆成一个不合理的维度最终 mask 作用的位置完全是乱的某些窗口里本来应该被屏蔽的 token 反而被保留下来注意力分布就开始震荡。解决方法是先打印 attn 和 mask 各自的 shape手动对齐后再相加。分享这个例子是想提醒大家写完代码尤其是涉及维度变换的部分第一时间把中间 shape 打出来做断言这比事后调试省时间得多。6.4 数据、算力有限时 Transformer 模型的使用建议如果你手头的数据量比较小算力也有限我的建议是不要一上来就魔改 VIT 或 Swin-T先考虑预训练模型微调比如从 Hugging Face 或者 Keras 官方模型库加载一个已经在 ImageNet 上预训练好的 VIT 或 Swin-T 权重然后在自己的数据集上做迁移学习。预训练模型已经学到了一般性的视觉特征到小数据集上只需要微调最后几层效果远好于从零训练。比如加载 swin_tiny_patch4_window7_224 预训练权重在 CIFAR-10 上微调几个 epoch 就能拿到 96% 以上的精度而这个精度是从零训练很难达到的。模型体积上如果你的硬件资源紧张试着把 num_layers 减少到 6-8 层隐藏维度压缩到 192-256。Transformer 模型的参数存在很高的冗余小规模任务上没必要把 Base 模型的原设计参数直接堆上去。但要注意模型缩得太狠时MLP hidden dim 也要相应缩小保持和 embed_dim 的 ratio 大约是 4:1 比较稳妥。7. 从两个模型上学到的设计思路这可能是这篇文章里最“软”但我觉得最有价值的一节。看完 VIT 和 Swin-T 的代码以后我自己的一个感受是这两个模型虽然名字相似但设计哲学差别很大。VIT 选择了一条极致简单的路把图像当成句子把全局注意力原封不动搬过来靠大数据量去补齐没有归纳偏置带来的学习负担。Swin-T 则选择了一条折中路线既保留 Transformer 的优势又针对图像特性重新设计了空间操作。这种“先理解任务本质再选择合适操作”的思路其实在深度学习里是常见的演进路径。最初 Transformer 在 CV 里火起来时大家觉得 CNN 要被淘汰了后来发现回归 CNN 的老路做个混合架构效果反而更好。Swin-T 虽然没有使用普通卷积但 patch embedding、patch merging、窗口划分这些操作本质上都在向 CNN 的空间先验靠拢。对于刚入门的朋友我的建议是先弄懂 VIT因为它简单核心概念patch embedding、position embedding、class token非常直观然后再学 Swin-T因为它在 VIT 的基础上引入了更多视觉特有的优化能让你理解“同一个思想在不同领域落地时需要做哪些改造”。这两个模型都掌握后再去看 MaiT、Focal Transformer、Next-ViT 这些后续工作会发现里面的新东西几乎都是这两个基础模型的排列组合。把地基打牢后面的花样就都好理解了。