大概半年前,我在帮团队做一次视觉骨干网络选型时,遇到一个很有意思的僵局:方案评审会上,大家拿着 Swin-Transformer 的 ImageNet 精度、FLOPs 数据和 COCO 检测榜单投来期望的目光,但谁也说不清把它塞进我们这套高分辨率、多尺度输出的业务里,到底要改哪些源码、踩哪些坑、依赖是否可控。那段时间我翻了不少技术博客,大多是“如何快速上手”和“如何在 MMDetection 里调用”,真正从源码出发、站在工程治理视角去审计这个开源项目的资料,几乎没有。于是我自己把 Microsoft 官方仓库拆了一遍,从目录结构到核心算子,从配置系统到分布式训练,从模型导出到硬件部署,逐行做过核对。这篇就把当时的完整审计笔记整理出来,围绕 Swin-Transformer 的源码实现、工程质量和落地选型展开。它适合正在做检测/分割模型选型的算法工程师,也适合需要基于 Swin 做二次开发、推理优化或服务化改造的平台开发者。看完之后,你至少能回答三个问题:这个仓库到底能不能直接用、哪些地方隐藏了性能或精度陷阱、什么条件下选它才是划算的。
1. 为什么值得为 Swin-Transformer 做一次源码级审计
Swin-Transformer 来自微软研究院,2021 年发布后迅速成了视觉 Transformer 在密集预测任务里的标杆。它和 ViT 最大的区别,是把图像切块后的 token 序列重新组织成了金字塔式的层级结构,分辨率逐 stage 降低、通道数逐 stage 增加。这个设计和 CNN 的特征金字塔天然对齐,所以检测、分割这类需要多尺度特征的模型几乎是无痛接入。也正因为这一点,它一度成了各种检测框架里最常被当作默认 backbone 的 Transformer 结构之一。
但“论文精度高”和“工程可落地”是两回事。我见过不少团队拿着官方 README 里的 top-1 精度数字做决策,结果一跑源码,发现输入尺寸稍微一变就报 shape 错,或者把 window_size 从 7 改成 12 后精度反而严重下降,又或者不知道相对位置偏置表需要插值,直接加载官方预训练权重时报错。这些问题,只有把源码打开、把每个模块的前向逻辑摸透了才能解释清楚。
源码级审计这个动作,本质上是在做三重确认。第一重,确认代码结构和论文描述一致,确认官方实现里没有藏私货,比如某些关键操作被静默简化或加了额外处理。第二重,确认仓库的可维护性和可控性:配置和代码是否分离、依赖是否清晰、出问题时能不能快速定位到具体文件。第三重,确认选型时隐藏在 README 之外的算子和硬件约束。后面两重,恰恰是很多只看论文的人最容易忽略的。
我的审计方法是先把仓库所有核心文件过一遍,画出模块依赖关系,再挑出影响正确性和性能的关键实现逐段阅读,最后针对配置管理、预训练权重、混合精度、部署导出这些工程环节做实际验证。整条路径走下来,结论可以浓缩成一句话:Swin-Transformer 的工程质量在开源视觉模型里属于中上水平,但绝对不推荐当作黑盒使用。
2. 仓库结构审计:从根目录到模型定义,每份文件负责什么
先看官方仓库 microsoft/Swin-Transformer 的顶层布局。整个仓库大致是这样的:
Swin-Transformer/ ├── configs/ │ ├── swin_base_patch4_window7_224.yaml │ ├── swin_large_patch4_window7_224.yaml │ ├── swin_small_patch4_window7_224.yaml │ ├── swin_tiny_patch4_window7_224.yaml │ └── ... ├── datasets/ ├── models/ │ ├── build.py │ ├── swin_transformer.py │ ├── swin_transformer_v2.py │ └── swin_mlp.py ├── main.py ├── utils.py └── get_flops.py这个结构在工程上有一个很明显的优点:配置和代码分离。configs 目录下面放的都是 YAML 文件,模型结构、训练超参、数据路径几乎全部通过配置表达。models 目录只负责模型定义,main.py 负责训练和验证流程,utils.py 放工具函数,get_flops.py 负责统计计算量。模块边界是清楚的,对于想做二次开发的团队来说,这个组织方式比许多把超参散落在各处脚本里的学术项目要舒服得多。
2.1 配置系统:yacs 的用法与配置文件的“唯一真源”
官方用 yacs 管理配置。一个典型的 Swin-Tiny 配置里,关键字段包括 IMG_SIZE、PATCH_SIZE、EMBED_DIM、DEPTHS、NUM_HEADS、WINDOW_SIZE、MLP_RATIO、QKV_BIAS、APE、PATCH_NORM 这些模型结构参数,还包括训练相关参数。yacs 的特点是可以用配置文件里的嵌套结构映射到内存里的 CfgNode,代码中通过 cfg.MODEL.SWIN.EMBED_DIM 这样的路径直接访问,非常直观。
在 models/build.py 里,构建模型时会先调用 get_config 读入用户指定的 YAML 文件,再通过 build_model 根据配置文件里的 MODEL.TYPE 选择模型族。这个流程让实验管理变得很轻松:换一组参数只需要新增一个 YAML,不需要改动 Python 代码。但也有一个隐含约束,所有模型结构参数必须都在配置文件里出现过,否则代码访问时会报 KeyError。所以如果你从官方仓库开始改,最好保留原始配置文件的完整结构,只修改需要变更的字段。
2.2 训练入口与工具链:main.py、utils.py、get_flops.py
main.py 是整个仓库的训练和验证入口。流程很传统:解析命令行参数、加载配置、构建模型、构建数据加载器、进入训练循环或评估。官方支持分布式训练,对 Slurm 环境也有适配。utils.py 里封装了一些通用能力,例如学习率调度、指标统计和日志输出。由于官方实现主要面向 ImageNet 分类,数据加载里用的是 timm 的数据处理管线,这意味着对于自定义数据集,你大概率要替换数据读取部分。
get_flops.py 是审计时很容易被忽略的一个文件。它通过 thop 库统计模型的参数量和计算量。如果环境里没装 thop,直接运行会报错。如果你想快速验证不同配置下的 FLOPs 差异,可以用这个脚本,但要注意它的统计口径和论文有时不完全一致,尤其是 attention 里的相对位置偏置计算,不同 count 方式会有几个 GFLOPs 的差别。
2.3 多个模型族混杂在同一个仓库里的治理代价
models 目录下同时有 swin_transformer.py、swin_transformer_v2.py 和 swin_mlp.py。意味着 build.py 要维护多个模型族的构建逻辑。这种“多模型共存”的好处是方便横向对比,坏处是单文件体积会膨胀,swin_transformer.py 本身就有上千行,包含了从 Patch Embedding 到完整 SwinTransformer 的整条链路,可读性尚可,但对新加入的开发者来说,定位某个具体模块还是需要一点时间。
审计下来,这个仓库最明显的治理短板是缺少自动化测试和持续集成。官方没有提供跑通模型前向、验证 shape 的单元测试用例,也没有 CI 配置。对团队自用来说问题不大,但如果你想长期维护一个基于它的私有分支,我建议自己补上基本的回归测试,至少保证改完一个模块后,前向输出 shape 不会悄悄变化。
3. 核心实现拆解:六个必须看懂的机制
从源码角度讲,Swin-Transformer 的精华集中在六个机制里。把它们看透,后续所有问题都能推导出答案。
3.1 Patch Embedding:一个卷积完成了图像到 token 的转换
Patch Embedding 层的实现非常简短,本质就是一次卷积操作:用 kernel_size = patch_size、stride = patch_size 的 Conv2d,把输入图像从 (B, 3, H, W) 变成 (B, embed_dim, H/patch_size, W/patch_size)。官方默认 patch_size=4,所以输入 224x224 的图,会先变成 56x56 的 token 网格,总计 3136 个 token。注意这个操作是不重叠切块,没有像 ViT 那样额外加一个可学习的 position embedding(除非设置 APE=True),位置信息主要靠后续窗口注意力里的相对位置偏置来编码。
从这里可以推出第一个约束:输入图像的 H 和 W 必须能被 patch_size 整除。如果图片尺寸是 225x224,在第一层就会直接报错。实际工程中很多摄像头输入不是严格的 224 倍数,需要先做 Resize 或 Pad,这一点会在后面选型部分展开。
3.2 SwinTransformerBlock:标准 Transformer 块加上窗口分区
每个 SwinTransformerBlock 的前向逻辑可以概括为:先 LayerNorm,然后是窗口注意力(W-MSA 或 SW-MSA),残差连接,再 LayerNorm,接 MLP,再残差。这个结构和 ViT 的 Transformer Encoder 几乎一致,核心差异只在注意力计算的范围。
窗口注意力会把特征图按 window_size 切成不重叠的网格。默认 window_size=7,那么 56x56 的特征图会被分成 8x8=64 个窗口。每个窗口内部做标准多头自注意力。这样做最大的收益是计算复杂度大幅下降。全局注意力的复杂度随 HxW 的平方增长,而窗口注意力的复杂度只跟窗口大小相关,和输入分辨率近似线性关系。对于高分辨率输入来说,这是决定性的优势。
3.3 循环移位和遮罩掩码:窗口之间如何建立联系
只做窗口内注意力,窗口之间就没有信息交流,感受野会被限制在窗口大小里。Swin 的解法是交错使用 W-MSA 和 SW-MSA:偶数 block 用窗口注意力,奇数 block 把特征图沿左上方向平移 window_size//2 个像素后,再执行窗口注意力,最后把结果平移回来。通过这种循环移位,原本不相邻的窗口边界得以相互接触,实现跨窗口信息交换。
但循环移位有一个问题:移位后,一个窗口内可能包含原本不属于同一区域的切片,直接做注意力会把本不该相邻的 token 混在一起。为了解决这个问题,源码生成了一个 attention mask,在计算 softmax 之前把无效位置的分数设为一个非常大的负数。这个 mask 只在 SW-MSA 分支使用。审计源码时,这一段是公认最难读的部分,因为它涉及窗口数量变化和索引映射,没有足够的注释很容易看晕。我的建议是不要盯着代码硬啃,先打印出特征图 shape,手动推一遍 4x4 输入、window_size=2 的 mask 长什么样,再回来看代码就会豁然开朗。
3.4 相对位置偏置:模型里被低估的“位置记忆体”
Swin 没有用 ViT 那种绝对位置编码,而是给注意力 logits 加了一个可学习的相对位置偏置。偏置表 shape 是 (num_heads, (2window_size-1) * (2window_size-1))。为什么是这个形状?因为两个 token 在水平和垂直方向上的相对位置范围都是 -(window_size-1) 到 (window_size-1),一共 2*window_size-1 个取值,二维组合起来就是平方关系。源码通过一系列广播计算生成一个相对位置索引表,把每个 token pair 的偏移映射到表里的某一项。
这个偏置表是模型里少数会随着输入窗口尺寸变化而失效的参数。如果你在推理或迁移学习时把 window_size 从 7 改成 12,偏置表维度对不上,必须用插值把它从 13x13(即 27-1)放大到 23x23(即 212-1)。插值方式可以是双线性或双三次,但对精度的影响不可忽略,需要实验验证。
3.5 Patch Merging:跨阶段的降采样就这么简单
每个 stage 结束时,如果后面还有 stage,就执行 Patch Merging。它的作用是把分辨率减半、通道数翻倍。具体实现是先把特征图按 2x2 的邻域切分成四份,在通道维拼接,再通过一个 Linear 层把通道数调整到目标值。代码里通过 reshape 和 permute 完成重排。
这里有一个容易被忽略的点:Patch Merging 之后,特征图下的每个 token 拥有了 4 倍于之前的通道信息,但因为拼接顺序固定,后续 Linear 层必须学会融合这些分组信息。这也是为什么 Swin 的特征在检测任务中通常需要配合 FPN,才能把各阶段输出有效利用起来。
3.6 整体结构与参数量表:Swin-T/S/B 的选择依据
完整 SwinTransformer 由 4 个 stage 组成。每个 stage 包含若干 SwinTransformerBlock,前三个 stage 末尾接 Patch Merging。以 Swin-T 为例,depths=[2,2,6,2],embed_dim=96,heads=[3,6,12,24],window_size=7。Swin-B 则用 embed_dim=128、depths=[2,2,18,2],其余配置类似。下面是常见配置的粗略对比:
| 模型 | embed_dim | depths | 参数量 | ImageNet-1K Top-1(官方报告约) |
|---|---|---|---|---|
| Swin-T | 96 | [2,2,6,2] | 28M | 81.3% |
| Swin-S | 96 | [2,2,18,2] | 50M | 83.0% |
| Swin-B | 128 | [2,2,18,2] | 88M | 83.5% |
| Swin-L | 192 | [2,2,18,2] | 197M | 86.2%(ImageNet-22K 预训练后) |
这些数据来自官方报告,具体数值可能随训练配置浮动,但能看出一个大趋势:参数量翻倍带来的精度收益并不是很大,所以下游任务选型时不要盲目上大模型,Swin-T 和 Swin-S 往往是性价比最合适的区间。另外需要注意,Swin 的分类头用的是 LayerNorm + 全局平均池化 + Linear,和 ViT 的 CLS token 方案完全不同,迁移到检测框架时通常要丢掉分类头,只保留骨干部分。
4. 工程治理审计:配置管理、依赖控制与可维护性
这部分是我个人认为最“干货”但也最容易被忽略的一节。开源仓库能跑通是一回事,放进团队里长期维护是另一回事。我从配置管理、依赖、训练扩展、代码维护性、许可五个维度做了一遍审计。
4.1 配置管理的双刃剑:灵活与隐式覆盖
官方对 yacs 的使用是合理的。配置文件覆盖了几乎所有的模型结构参数,训练时通过命令行可以覆盖部分参数,例如--batch-size、--lr这些。这意味着实验记录里只要保存最终命令行和 YAML 文件,理论上就能还原模型。
不过这里有个治理风险:当你从命令行覆盖参数时,配置就不是唯一真源了。之前我见过团队里两个人用了同一个 YAML 文件,一个人手动调了 dropout,另一个人没有,最后对模型精度走势产生困惑。官方这个仓库没有“锁定配置”的机制,所以如果你要基于它搭平台,建议在保存实验记录时同时保存完整的解析后配置,而不是只保存 YAML。
4.2 依赖清单:torch、timm、apex 带来的复现摩擦
官方 requirements 里包含 torch、torchvision、timm、pyyaml、tensorboard 这些常见依赖。早期版本对 apex 有硬依赖,apex 的安装又是出了名的麻烦,很多人第一次跑官方训练脚本就在这一步卡住。后来仓库演进中加入了--enable_amp选项,可以走 PyTorch 原生的混合精度。如果你只是做推理或下游微调,完全不需要安装 apex;如果要完全复现官方训练流程,建议直接用官方 Docker 镜像或自己锁好版本。
timm 在这里的作用主要是数据增强和部分模型族支持。由于 timm 迭代很快,不同版本的接口有差异,我的经验是直接按照官方 README 指定的版本安装,不要装最新的,否则可能在数据处理接口上报兼容性问题。
4.3 训练扩展:分布式、恢复、混合精度的成熟度
官方代码在分布式训练上做得不错,支持 PyTorch DDP 和 Slurm。resume 机制是完整的,中断后能从 checkpoint 续训。混合精度方面,既有 apex 的 O1 路径,也有原生 AMP 路径。这些能力加在一起,意味着团队不需要在训练框架层面做太多额外开发,可以直接在一个多卡集群上把上游训练流程跑起来。
但要注意,官方并没有内置自动化评测指标统计、实验管理、超参搜索这类平台化能力。它是典型的“研究仓库”而非“训练平台”。你需要自己封装外层。
4.4 可维护性:单文件偏大、缺测试、维护活跃度中上
从长期维护角度看,这个仓库有两个明显的治理债。第一,swin_transformer.py 是个千行级的单体文件,Patch Embedding、WindowAttention、SwinTransformerBlock、Patch Merging、SwinTransformer 全部挤在一起。对阅读友好的部分是有序的,但如果你想替换其中一个模块,改动往往会在文件里多个位置牵扯。第二,仓库没有配套的测试用例,无法通过自动化手段快速验证修改是否破坏原有逻辑。
好在官方对 issue 和 PR 的响应还算及时,社区生态也足够大,很多坑在 GitHub issue 里都有人讨论过。这在一定程度上弥补了缺少测试的问题。但我的建议是,任何团队拿到这个仓库之后,第一步不是改功能,而是先补一个 smoke test,至少覆盖不同输入尺寸、不同 window_size 下的前向传播和梯度回传。
4.5 License 与使用边界
官方仓库使用的是 MIT License,这对商业使用很友好,不必担心传染性开源协议带来的合规负担。但 License 只覆盖代码本身,不覆盖论文、权重和训练过程中的技巧。如果要把预训练权重用于自己的模型,最好再确认一下权重发布页面的说明,通常这类权重可以商用,但这是可变的,实际使用时应该自行复核。
5. 落地选型指南:什么场景选 Swin,什么场景避开它
源码审计做完之后,回到最实际的选型问题。我总结了一个相对完整的决策路径,读者可以按自己的业务条件逐项对照。
5.1 适合选 Swin-Transformer 的场景
最典型的场景是目标检测、实例分割、语义分割这类需要多尺度特征的密集预测任务。Swin 的 4 个 stage 正好可以当作 FPN 的不同层输入,和 Detectron2、MMDetection 的架构天然契合。如果你既要高分辨率输入,又要 Transformer 的表达能力,Swin 几乎是当时唯一稳妥的选项。
第二个适合的场景是中小规模数据上的微调。相比于纯 ViT 需要动辄上亿数据的预训练,Swin 因为引入了局部性和金字塔结构,对归纳偏置的依赖更强,在几万到几十万张图片的自定义数据上表现通常比同尺寸的 ViT 更稳定。这一点在我们团队的实际评测中反复验证过。
第三个适合的场景是需要稳定复现和社区支持的场景。Swin 的权重、配置、下游框架集成都很齐全,如果你在 MMDetection 或 Detectron2 里用,甚至不需要自己写模型定义,直接加载官方转换好的权重,几行配置跑起来。这个生态成熟度对工程排期很重要。
5.2 不适合选 Swin-Transformer 的场景
如果输入尺寸在推理时高度动态,比如一张图里的目标区域大小变化极大,需要频繁处理任意分辨率,Swin 会有点麻烦。核心原因是 window_size 固定后,输入尺寸如果不是 window_size 的整数倍,就要做 padding 或调整,而 padding 会影响注意力计算范围,精度和速度都受影响。
其次,如果你的部署目标是边缘设备或对延迟极度敏感,需要仔细评估 ONNX 导出和 TensorRT 的效率。窗口划分、循环移位、mask 加法这类操作在导出后会产生不少 gather/scatter 类算子,部分加速卡支持得并不好。相比之下,ConvNeXt 这类纯卷积结构在部署工具链上要顺畅得多。我们实测下来,同一个 Swin-T 转 ONNX 后的推理延迟,在部分 GPU 上比原始 PyTorch 反而更差,就是因为算子融合不充分。
另一个需要避开的场景是超长序列或多模态输入。Swin 的设计是给图片用的,对任意长度序列没有特殊优化。如果你要处理视频多帧或文本加图像的联合序列,Swin 需要先切片再拼接,复杂度会上升,不如直接用标准 ViT 配合合适的 position embedding 灵活。
5.3 与 ViT、DeiT、ConvNeXt 的横向对比
| 模型 | 尺度结构 | 注意力范围 | 典型精度/计算量 | 部署友好度 | 适合任务 |
|---|---|---|---|---|---|
| ViT | 单尺度 | 全局 | 数据需求大,中小数据易过拟合 | 一般,序列长则贵 | 大模型预训练、分类 |
| DeiT | 单尺度 | 全局 | 数据增强技巧,适合中等数据 | 一般 | 分类、迁移学习 |
| Swin-T | 金字塔多尺度 | 窗口+移位窗口 | 高性价比,检测分割友好 | 中等,导出需优化 | 检测、分割、高分辨率任务 |
| ConvNeXt | 金字塔多尺度 | 卷积滑窗 | 精度同级别略优于 Swin | 好,部署生态成熟 | 通用视觉、边缘部署 |
这个表不是绝对的,但能反映大方向。我对选型的通常建议是:纯分类任务且数据充足,可以考虑 ViT/DeiT;密集预测任务,优先 Swin 或 ConvNeXt;需要端侧部署,ConvNeXt 往往更省心。
5.4 显存估算:一个粗糙但实用的公式
显存是决定训练方案能不能跑起来的关键。Swin-T 在 224x224 输入下的激活显存比同等规模的 ViT 小很多,因为注意力被限制在窗口内,K/V 张量不会指数扩张。一个粗糙的估算方法是按主干的特征图大小累加:
对 Swin-T,patch_size=4,经历四个 stage 后特征图分辨率从 56 降到 7(224/4/2/2/2),通道数从 96 升到 768。中间占显存最大的往往是第一个 stage:56x56x96 的特征图,一个 batch 64 下大约是 19MB,再算上注意力中间变量,整体还在可控范围。真正吃显存的是高分辨率输入和多卡同步 BN,以及后续检测头。所以如果你的任务需要 1024x1024 输入,我建议直接用 Swin-T 起步,batch 尽量小,配合梯度累积,而不是直接上 Swin-L。
5.5 版本选择:Swin V1 还是 Swin V2
官方仓库同时维护了 V1 和 V2。V2 主要做了几件事:把相对位置偏置换成了连续的 log-spaced 编码,把 attention 换成了 cosine attention,更深的模型改用 post-norm 结构,整体数值稳定性更好,也支持更大分辨率和更大窗口。
但 V2 的默认配置、输入分辨率和 V1 差异较大,很多下游框架默认接的还是 V1。我的建议是:如果你要快速落地,选 V1 官方权重,生态最成熟;如果你在做一个长期项目,愿意花时间适配下游框架,可以尝试 V2,它在高分辨率任务上的潜力更大。最忌讳的是把 V1 的预训练权重直接塞进 V2 的模型结构里,两代模型的权重定义不同,会直接报 key 不匹配。
5.6 与 MMDetection / Detectron2 的集成路径
实际业务里很少直接拿官方仓库做推理,通常是把 Swin 当作 backbone 接进检测框架。MMDetection 官方提供了从 Swin 权重转换到 mmdet 格式的脚本,配合 mmclassification 里的 backbone 定义,整个链路已经跑得很顺。如果你的团队已经基于 MMDetection 做检测,我的建议是不要自己从官方模型类复制代码,直接使用 mmdet 里的 Swin Transformer 实现,它能和 mmcv 的 checkpoint 转换机制自动对齐,省掉很多搬运功夫。
Detectron2 也有对应的第三方实现,但维护活跃度不如 MMDetection。从工程治理角度看,选一个社区支持度高的下游框架,比选一个“代码写得最优雅”的改造方案更重要。
6. 踩过的坑与复现要点
最后这部分是实操经验,每一条都是我或者团队在迁移、微调、部署中实际遇到过的问题。单独看都不难,但叠加在一起就很容易消耗掉一整天。
6.1 加载官方权重的隐藏条件:分类头与模型名的精确匹配
如果你用官方 main.py 做推理,直接指定 pretrained 路径就行。但如果你把 SwinTransformer 类搬到自己的代码里,加载官方权重时常常会遇到 key 不匹配。最常见的原因是你自定义了 num_classes,导致 head.weight 和 head.bias 的 shape 不一致。解决办法是加载时过滤掉 head 相关的 key,或者修改模型头后再加载。
另一个隐藏问题是官方权重文件里可能有model_ema这种带前缀的字段,取决于 checkpoint 的生成方式。加载时需要用torch.load(..., map_location='cpu')拿到 state_dict 后再做一层按键处理,不能直接load_state_dict到模型上,否则大概率报 missing key。
6.2 输入尺寸与 window_size 的整除关系
这是最容易踩的坑。官方 224 分辨率模型默认 window_size=7,224/7=32,整除没问题。官方 384 分辨率的 Swin-B 默认 window_size=12,384/12=32,也刚好整除。但如果你自己把输入设成 512,window_size=12,问题就来了:512/12=42.67,不是整数,代码会对输入做右下角 padding。padding 后的区域同样参与了 attention 计算,会引入额外的边界效应。
我在给一个遥感项目调参时,把输入从 384 改成 640 做多尺度,精度比预期低了 1.2 个点,查了半天才发现是整除问题。这里建议优先选择能被 32 和 window_size 同时整除的输入尺寸,比如 224、384、448、768 这类,或者自行实现新的 padding 策略。
6.3 相对位置偏置插值:改 window_size 前必须做的事
把官方权重从 window_size=7 微调到 window_size=12 时,由于偏置表本来就存在,直接替换权重文件会报 shape mismatch 的 RuntimeError。可以先加载原始权重、过滤掉相对位置偏置相关 key,然后初始化新尺寸的偏置,再对原偏置做一次双线性插值填充。插值是有效但不够精确的,实测在检测任务上精度会有小幅度下降,需要靠后续微调拉回来。
6.4 混合精度与 apex 的版本耦合
官方仓库里 AMP 和 apex 的路径是分开的。如果你选了--enable_amp,就用 PyTorch 原生 AMP,不需要 apex;如果你选了--use_apex,就需要安装和当前 CUDA 版本严格匹配的 apex。我见过太多人在ImportError: No module named 'apex'上浪费半小时。如果你不是要完全对齐官方训练策略,直接走原生 AMP 就好。原生 AMP 和 apex O1 在精度上几乎没有差别,但省掉一个编译依赖。
6.5 微调和迁移时的 batch size 与学习率关系
Swin 官方训练用的是大 batch(1024 左右)加较长 warmup。迁移到下游小数据集时,如果 batch 降到 8 或 16,学习率还按官方配置里的 1e-3 起步,很容易训练震荡。建议线性缩放:batch 缩小 64 倍,学习率也差不多缩小 64 倍开始,再小范围搜索。drop_path 在下游任务里建议适当降低,官方默认会根据模型尺寸在 0.1 到 0.3 之间浮动,但小数据集上 drop_path 太大容易欠拟合。
6.6 导出 ONNX 时的算子兼容与验证方法
把 Swin 导出到 ONNX 时,第一个建议是最小化动态轴。由于窗口划分和 mask 生成在实现时都假设了固定分辨率,动态尺寸导出会让 ONNX 图里出现大量非标准 shape 推导,很多引擎直接不支持。输出固定尺寸后,用 onnxsim 做一轮化简,能去掉不少冗余的 reshape 和 gather 操作。
导出完成后不要只看能跑通,还要做数值对比。用同一张输入图分别跑 PyTorch 模型和 ONNX Runtime,对比到最后十层输出的余弦相似度。相似度低于 0.99 就要检查哪些算子在转换中被改写。相对位置偏置因为是 buffer,通常会变成常量,这没问题;但循环移位和 reverse 操作在不同导出工具链下会生成不同算子,数值差异往往就藏在这几个点。
6.7 检测框架中多尺度特征的使用习惯
在 MMDetection 里用 Swin 做 backbone,通常取后三个 stage 的输出喂给 FPN。这是因为第一个 stage 分辨率太高、通道数太少,语义信息不足。如果你把四个 stage 全塞进 FPN,并不会让精度明显提升,反而增加显存和延迟。这一点在选型评审时值得先对齐,避免团队里有人误以为“用满全部 stage 效果最好”。
如果要用 Swin 做实例分割,mask head 的输入特征来自 FPN 的不同层,分辨率差异较大,建议给高分辨率分支适当降低通道数,因为 Swin 高分辨率分支的通道数往往比 ResNet 对应层更少,直接套用 ResNet 的 neck 结构不一定最优。
其实做完这次源码级审计,我心里最大的感受是:Swin-Transformer 在论文和工程之间取得了很好的平衡,但它的“好用”是建立在完整理解其窗口机制和配置体系之上的。任何跳过源码审计、直接把它当作模块塞进业务系统的做法,都会在后续某个环节付出代价。如果你正准备选型,我建议先花半天时间把我上面提到的几个关键源码位置读一遍,再决定要不要接进来。这一步省不掉。