学习 VLA 第6天:SWIN VIT算法原理以及复现

1. 引言

在视觉语言动作模型(VLA)的学习路线中,视觉骨干网络(Vision Backbone)是理解环境、提取特征的第一环。第 6 天聚焦 Swin Transformer(Swin ViT)------一种基于移位窗口的层次化视觉 Transformer。它不仅是众多 VLA 模型(如 RT-2、OpenVLA 等)中常见的视觉编码器选项,也是从 CNN 过渡到纯 Transformer 架构的关键一步。## 2. 为什么 VLA 需要 Swin Transformer

在 VLA 模型中,视觉编码器负责将高维图像转换为紧凑、富有语义的特征序列,供语言模型与动作解码器使用。传统 CNN(如 ResNet)虽然计算高效,但在捕捉长距离依赖和全局上下文方面存在局限。而标准 ViT 虽然具备全局感受野,却面临两个问题:

  • 计算复杂度随图像分辨率呈平方级增长;
  • 缺乏多尺度特征,不利于处理不同尺寸的物体。

Swin Transformer 通过层次化设计移位窗口自注意力,在保持线性计算复杂度的同时,兼顾了全局与局部信息,非常适合作为 VLA 的视觉编码器。

2.1 Swin Transformer 与 ViT 的核心对比

为了更直观地理解 Swin Transformer 的改进,下面从多个维度对比它与标准 ViT 的差异:

对比维度 ViT Swin Transformer
Patch 大小 16×16,粒度较粗 4×4,粒度更细,Token 数量更多
特征层级 单一尺度,所有层输出分辨率一致 层次化金字塔结构,多尺度特征
自注意力范围 全局自注意力,所有 Token 两两交互 窗口内自注意力,局部计算
计算复杂度 随图像分辨率呈平方级增长 随图像分辨率呈线性增长
跨区域信息交互 天然全局,无需额外机制 通过移位窗口(SW-MSA)实现窗口间交互
位置编码 绝对位置编码(可学习) 相对位置偏置(可学习参数表)
下采样方式 无,直接输出分类 Token Patch Merging 逐阶段下采样
适用场景 图像分类等单尺度任务 检测、分割、VLA 等多尺度密集任务

核心差异总结

  • 计算效率:ViT 的全局注意力在输入分辨率增大时开销急剧上升;Swin 将注意力限制在局部窗口内,复杂度从平方级降为线性,更适合高分辨率输入。
  • 多尺度能力:ViT 输出单一分辨率的特征序列,难以直接用于需要多尺度信息的任务;Swin 通过 Patch Merging 形成金字塔结构,天然适配 FPN 等特征融合范式。
  • 归纳偏置:Swin 的窗口机制引入了更强的局部性先验,在中小规模数据集上更容易收敛,对训练数据的依赖低于 ViT。
  • 信息交互:ViT 靠全局注意力直接建模长距离依赖;Swin 则通过 W-MSA 与 SW-MSA 交替,以更低的代价逐步实现跨窗口信息流通。

在 VLA 场景中,视觉编码器需要同时兼顾高分辨率细节与多尺度语义,Swin Transformer 的上述特性使其往往优于直接使用 ViT 作为骨干网络。

3. Swin Transformer 核心原理

3.1 整体架构

Swin Transformer 的整体结构分为四个阶段(Stage),每个阶段由若干 Swin Transformer Block 组成,并在阶段之间进行 Patch Merging 下采样,从而形成类似 CNN 的金字塔结构:

text 复制代码
输入图像 (H x W x 3)
   → Patch Partition (4x4)
   → Stage 1: Linear Embedding + Swin Block x2
   → Stage 2: Patch Merging + Swin Block x2
   → Stage 3: Patch Merging + Swin Block x6
   → Stage 4: Patch Merging + Swin Block x2
   → 输出特征图

3.2 Patch Partition 与 Linear Embedding

与 ViT 使用 16x16 的 Patch 不同,Swin Transformer 采用 4x4 的 Patch 大小,将图像划分为更细粒度的 Token。对于输入图像 x ∈ R^(H×W×3),经过 Patch Partition 后得到 (H/4) × (W/4) 个 Patch,每个 Patch 展平后维度为 4×4×3 = 48,再通过 Linear Embedding 映射到任意维度 C

3.3 移位窗口自注意力(Shifted Window Attention)

移位窗口自注意力是 Swin Transformer 最核心的创新点。它将特征图划分为不重叠的窗口(Window),在每个窗口内部计算自注意力,从而将计算复杂度从平方级降为线性。

窗口划分 :对于 (H/4) × (W/4) 的特征图,按 M × M(默认 M=7)划分窗口。

循环移位:为了让不同窗口之间能够进行信息交互,Swin Transformer 在相邻层之间交替使用两种窗口划分方式:

  • W-MSA(Window Multi-head Self-Attention):常规窗口划分;
  • SW-MSA (Shifted Window Multi-head Self-Attention):将特征图向左上角循环移位 (⌊M/2⌋, ⌊M/2⌋) 后再划分窗口。

掩码机制:由于循环移位后,部分窗口会包含来自不同原始区域的 Patch,因此需要引入 Attention Mask,屏蔽掉不属于同一原始窗口的 Patch 对。

下面是移位窗口自注意力的工作流程:
#mermaid-svg-UBFf1CWzXwKhlFzi{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-UBFf1CWzXwKhlFzi .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-UBFf1CWzXwKhlFzi .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-UBFf1CWzXwKhlFzi .error-icon{fill:#552222;}#mermaid-svg-UBFf1CWzXwKhlFzi .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-UBFf1CWzXwKhlFzi .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-UBFf1CWzXwKhlFzi .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-UBFf1CWzXwKhlFzi .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-UBFf1CWzXwKhlFzi .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-UBFf1CWzXwKhlFzi .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-UBFf1CWzXwKhlFzi .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-UBFf1CWzXwKhlFzi .marker{fill:#333333;stroke:#333333;}#mermaid-svg-UBFf1CWzXwKhlFzi .marker.cross{stroke:#333333;}#mermaid-svg-UBFf1CWzXwKhlFzi svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-UBFf1CWzXwKhlFzi p{margin:0;}#mermaid-svg-UBFf1CWzXwKhlFzi .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-UBFf1CWzXwKhlFzi .cluster-label text{fill:#333;}#mermaid-svg-UBFf1CWzXwKhlFzi .cluster-label span{color:#333;}#mermaid-svg-UBFf1CWzXwKhlFzi .cluster-label span p{background-color:transparent;}#mermaid-svg-UBFf1CWzXwKhlFzi .label text,#mermaid-svg-UBFf1CWzXwKhlFzi span{fill:#333;color:#333;}#mermaid-svg-UBFf1CWzXwKhlFzi .node rect,#mermaid-svg-UBFf1CWzXwKhlFzi .node circle,#mermaid-svg-UBFf1CWzXwKhlFzi .node ellipse,#mermaid-svg-UBFf1CWzXwKhlFzi .node polygon,#mermaid-svg-UBFf1CWzXwKhlFzi .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-UBFf1CWzXwKhlFzi .rough-node .label text,#mermaid-svg-UBFf1CWzXwKhlFzi .node .label text,#mermaid-svg-UBFf1CWzXwKhlFzi .image-shape .label,#mermaid-svg-UBFf1CWzXwKhlFzi .icon-shape .label{text-anchor:middle;}#mermaid-svg-UBFf1CWzXwKhlFzi .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-UBFf1CWzXwKhlFzi .rough-node .label,#mermaid-svg-UBFf1CWzXwKhlFzi .node .label,#mermaid-svg-UBFf1CWzXwKhlFzi .image-shape .label,#mermaid-svg-UBFf1CWzXwKhlFzi .icon-shape .label{text-align:center;}#mermaid-svg-UBFf1CWzXwKhlFzi .node.clickable{cursor:pointer;}#mermaid-svg-UBFf1CWzXwKhlFzi .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-UBFf1CWzXwKhlFzi .arrowheadPath{fill:#333333;}#mermaid-svg-UBFf1CWzXwKhlFzi .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-UBFf1CWzXwKhlFzi .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-UBFf1CWzXwKhlFzi .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-UBFf1CWzXwKhlFzi .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-UBFf1CWzXwKhlFzi .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-UBFf1CWzXwKhlFzi .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-UBFf1CWzXwKhlFzi .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-UBFf1CWzXwKhlFzi .cluster text{fill:#333;}#mermaid-svg-UBFf1CWzXwKhlFzi .cluster span{color:#333;}#mermaid-svg-UBFf1CWzXwKhlFzi div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-UBFf1CWzXwKhlFzi .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-UBFf1CWzXwKhlFzi rect.text{fill:none;stroke-width:0;}#mermaid-svg-UBFf1CWzXwKhlFzi .icon-shape,#mermaid-svg-UBFf1CWzXwKhlFzi .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-UBFf1CWzXwKhlFzi .icon-shape p,#mermaid-svg-UBFf1CWzXwKhlFzi .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-UBFf1CWzXwKhlFzi .icon-shape .label rect,#mermaid-svg-UBFf1CWzXwKhlFzi .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-UBFf1CWzXwKhlFzi .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-UBFf1CWzXwKhlFzi .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-UBFf1CWzXwKhlFzi :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} 否

输入特征图 (H/4 x W/4)
是否为 SW-MSA?
W-MSA: 常规窗口划分
循环移位 (M/2, M/2)
窗口划分
计算 Attention Mask
SW-MSA: 窗口内自注意力
窗口内自注意力
还原窗口并反向移位
输出特征图

3.4 相对位置偏置(Relative Position Bias)

Swin Transformer 在自注意力计算中加入了相对位置偏置 B,公式如下:

text 复制代码
Attention(Q, K, V) = SoftMax(QK^T / √d + B) V

其中 B 是从可学习的参数表 B_hat ∈ R^((2M-1)×(2M-1)) 中索引得到的。相对位置偏置显著提升了模型性能,是 Swin Transformer 超越 ViT 的关键因素之一。

3.5 层次化特征与 Patch Merging

在每个 Stage 之间,Swin Transformer 通过 Patch Merging 进行下采样:将 2×2 邻域的 4 个 Patch 拼接,经过 LayerNorm 和线性层,将通道数翻倍、空间尺寸减半。这使模型天然具备多尺度特征表达能力,可直接输出类似 FPN 的特征金字塔,便于 VLA 模型融合不同粒度的视觉信息。

下面是 Swin Transformer 的整体架构流程图:
#mermaid-svg-nz0bLMpjFmHkFBry{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-nz0bLMpjFmHkFBry .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-nz0bLMpjFmHkFBry .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-nz0bLMpjFmHkFBry .error-icon{fill:#552222;}#mermaid-svg-nz0bLMpjFmHkFBry .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-nz0bLMpjFmHkFBry .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-nz0bLMpjFmHkFBry .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-nz0bLMpjFmHkFBry .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-nz0bLMpjFmHkFBry .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-nz0bLMpjFmHkFBry .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-nz0bLMpjFmHkFBry .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-nz0bLMpjFmHkFBry .marker{fill:#333333;stroke:#333333;}#mermaid-svg-nz0bLMpjFmHkFBry .marker.cross{stroke:#333333;}#mermaid-svg-nz0bLMpjFmHkFBry svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-nz0bLMpjFmHkFBry p{margin:0;}#mermaid-svg-nz0bLMpjFmHkFBry .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-nz0bLMpjFmHkFBry .cluster-label text{fill:#333;}#mermaid-svg-nz0bLMpjFmHkFBry .cluster-label span{color:#333;}#mermaid-svg-nz0bLMpjFmHkFBry .cluster-label span p{background-color:transparent;}#mermaid-svg-nz0bLMpjFmHkFBry .label text,#mermaid-svg-nz0bLMpjFmHkFBry span{fill:#333;color:#333;}#mermaid-svg-nz0bLMpjFmHkFBry .node rect,#mermaid-svg-nz0bLMpjFmHkFBry .node circle,#mermaid-svg-nz0bLMpjFmHkFBry .node ellipse,#mermaid-svg-nz0bLMpjFmHkFBry .node polygon,#mermaid-svg-nz0bLMpjFmHkFBry .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-nz0bLMpjFmHkFBry .rough-node .label text,#mermaid-svg-nz0bLMpjFmHkFBry .node .label text,#mermaid-svg-nz0bLMpjFmHkFBry .image-shape .label,#mermaid-svg-nz0bLMpjFmHkFBry .icon-shape .label{text-anchor:middle;}#mermaid-svg-nz0bLMpjFmHkFBry .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-nz0bLMpjFmHkFBry .rough-node .label,#mermaid-svg-nz0bLMpjFmHkFBry .node .label,#mermaid-svg-nz0bLMpjFmHkFBry .image-shape .label,#mermaid-svg-nz0bLMpjFmHkFBry .icon-shape .label{text-align:center;}#mermaid-svg-nz0bLMpjFmHkFBry .node.clickable{cursor:pointer;}#mermaid-svg-nz0bLMpjFmHkFBry .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-nz0bLMpjFmHkFBry .arrowheadPath{fill:#333333;}#mermaid-svg-nz0bLMpjFmHkFBry .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-nz0bLMpjFmHkFBry .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-nz0bLMpjFmHkFBry .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-nz0bLMpjFmHkFBry .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-nz0bLMpjFmHkFBry .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-nz0bLMpjFmHkFBry .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-nz0bLMpjFmHkFBry .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-nz0bLMpjFmHkFBry .cluster text{fill:#333;}#mermaid-svg-nz0bLMpjFmHkFBry .cluster span{color:#333;}#mermaid-svg-nz0bLMpjFmHkFBry div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-nz0bLMpjFmHkFBry .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-nz0bLMpjFmHkFBry rect.text{fill:none;stroke-width:0;}#mermaid-svg-nz0bLMpjFmHkFBry .icon-shape,#mermaid-svg-nz0bLMpjFmHkFBry .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-nz0bLMpjFmHkFBry .icon-shape p,#mermaid-svg-nz0bLMpjFmHkFBry .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-nz0bLMpjFmHkFBry .icon-shape .label rect,#mermaid-svg-nz0bLMpjFmHkFBry .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-nz0bLMpjFmHkFBry .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-nz0bLMpjFmHkFBry .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-nz0bLMpjFmHkFBry :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} 输入图像 (H x W x 3)
Patch Partition (4x4)
Linear Embedding (C=96)
Stage 1: Swin Block x2
Patch Merging
Stage 2: Swin Block x2
Patch Merging
Stage 3: Swin Block x6
Patch Merging
Stage 4: Swin Block x2
输出多尺度特征

4. Swin Transformer Block 内部结构

每个 Swin Transformer Block 由以下模块组成:

text 复制代码
输入 x
   → LayerNorm
   → (W-MSA 或 SW-MSA)
   → 残差连接
   → LayerNorm
   → MLP (GELU, 4x 扩展)
   → 残差连接
   → 输出

其中 W-MSA 与 SW-MSA 在相邻 Block 间交替使用,保证窗口间信息流通。

下面是 Swin Transformer Block 的内部结构图:
#mermaid-svg-Uj7l3mQFNZ0MW4Uj{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .error-icon{fill:#552222;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .marker{fill:#333333;stroke:#333333;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .marker.cross{stroke:#333333;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj p{margin:0;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .cluster-label text{fill:#333;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .cluster-label span{color:#333;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .cluster-label span p{background-color:transparent;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .label text,#mermaid-svg-Uj7l3mQFNZ0MW4Uj span{fill:#333;color:#333;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .node rect,#mermaid-svg-Uj7l3mQFNZ0MW4Uj .node circle,#mermaid-svg-Uj7l3mQFNZ0MW4Uj .node ellipse,#mermaid-svg-Uj7l3mQFNZ0MW4Uj .node polygon,#mermaid-svg-Uj7l3mQFNZ0MW4Uj .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .rough-node .label text,#mermaid-svg-Uj7l3mQFNZ0MW4Uj .node .label text,#mermaid-svg-Uj7l3mQFNZ0MW4Uj .image-shape .label,#mermaid-svg-Uj7l3mQFNZ0MW4Uj .icon-shape .label{text-anchor:middle;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .rough-node .label,#mermaid-svg-Uj7l3mQFNZ0MW4Uj .node .label,#mermaid-svg-Uj7l3mQFNZ0MW4Uj .image-shape .label,#mermaid-svg-Uj7l3mQFNZ0MW4Uj .icon-shape .label{text-align:center;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .node.clickable{cursor:pointer;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .arrowheadPath{fill:#333333;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .cluster text{fill:#333;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .cluster span{color:#333;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj rect.text{fill:none;stroke-width:0;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .icon-shape,#mermaid-svg-Uj7l3mQFNZ0MW4Uj .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .icon-shape p,#mermaid-svg-Uj7l3mQFNZ0MW4Uj .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .icon-shape .label rect,#mermaid-svg-Uj7l3mQFNZ0MW4Uj .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-Uj7l3mQFNZ0MW4Uj :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} 输入 x
LayerNorm
W-MSA 或 SW-MSA
残差连接
LayerNorm
MLP (GELU, 4x 扩展)
残差连接
输出

5. 基于 PyTorch 的 Swin Transformer 复现

下面给出一个精简但完整的 Swin Transformer 实现,包含核心模块:Patch Embedding、窗口注意力、移位窗口、相对位置偏置和 Patch Merging。

python 复制代码
import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import Optional


class PatchEmbed(nn.Module):
    """4x4 Patch Embedding"""
    def __init__(self, in_chans=3, embed_dim=96):
        super().__init__()
        self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=4, stride=4)

    def forward(self, x):
        # x: (B, 3, H, W) -> (B, embed_dim, H/4, W/4)
        x = self.proj(x)
        # -> (B, H/4 * W/4, embed_dim)
        x = x.flatten(2).transpose(1, 2)
        return x


class PatchMerging(nn.Module):
    """2x2 Patch Merging 下采样"""
    def __init__(self, dim):
        super().__init__()
        self.norm = nn.LayerNorm(4 * dim)
        self.reduction = nn.Linear(4 * dim, 2 * dim, bias=False)

    def forward(self, x, H, W):
        B, L, C = x.shape
        x = x.view(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 = torch.cat([x0, x1, x2, x3], dim=-1)  # (B, H/2, W/2, 4C)
        x = x.view(B, -1, 4 * C)

        x = self.norm(x)
        x = self.reduction(x)
        return x, H // 2, W // 2


def window_partition(x, window_size):
    """将特征图划分为窗口"""
    B, H, W, C = x.shape
    x = x.view(B, H // window_size, window_size, W // window_size, window_size, C)
    windows = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_size, window_size, C)
    return windows


def window_reverse(windows, window_size, H, W):
    """将窗口还原为特征图"""
    B = int(windows.shape[0] / (H * W / window_size / window_size))
    x = windows.view(B, H // window_size, W // window_size, window_size, window_size, -1)
    x = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, H, W, -1)
    return x


class WindowAttention(nn.Module):
    """带相对位置偏置的窗口多头自注意力"""
    def __init__(self, dim, window_size, num_heads):
        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 = nn.Linear(dim, dim * 3, bias=True)
        self.proj = nn.Linear(dim, dim)

        # 相对位置偏置表
        self.relative_position_bias_table = nn.Parameter(
            torch.zeros((2 * window_size[0] - 1) * (2 * window_size[1] - 1), num_heads)
        )
        nn.init.trunc_normal_(self.relative_position_bias_table, std=0.02)

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

    def forward(self, x, mask: Optional[torch.Tensor] = None):
        B_, N, C = x.shape
        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]

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

        relative_position_bias = self.relative_position_bias_table[self.relative_position_index.view(-1)]
        relative_position_bias = relative_position_bias.view(
            self.window_size[0] * self.window_size[1],
            self.window_size[0] * self.window_size[1],
            -1,
        ).permute(2, 0, 1).contiguous()
        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)
        else:
            attn = F.softmax(attn, dim=-1)

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


class SwinBlock(nn.Module):
    """Swin Transformer Block(W-MSA 或 SW-MSA)"""
    def __init__(self, dim, num_heads, window_size=7, shift=False):
        super().__init__()
        self.norm1 = nn.LayerNorm(dim)
        self.attn = WindowAttention(dim, (window_size, window_size), num_heads)
        self.shift = shift
        self.window_size = window_size

        self.norm2 = nn.LayerNorm(dim)
        self.mlp = nn.Sequential(
            nn.Linear(dim, 4 * dim),
            nn.GELU(),
            nn.Linear(4 * dim, dim),
        )

    def forward(self, x, H, W):
        B, L, C = x.shape
        shortcut = x
        x = self.norm1(x)
        x = x.view(B, H, W, C)

        # 循环移位
        if self.shift:
            shift_size = self.window_size // 2
            x = torch.roll(x, shifts=(-shift_size, -shift_size), dims=(1, 2))

        # 划分窗口
        x_windows = window_partition(x, self.window_size)
        x_windows = x_windows.view(-1, self.window_size * self.window_size, C)

        # 计算注意力掩码
        if self.shift:
            H_windows = H // self.window_size
            W_windows = W // self.window_size
            mask = torch.zeros((1, H, W, 1), device=x.device)
            cnt = 0
            for h in range(H_windows):
                for w in range(W_windows):
                    mask[:, h * self.window_size:(h + 1) * self.window_size,
                         w * self.window_size:(w + 1) * self.window_size, :] = cnt
                    cnt += 1
            mask_windows = window_partition(mask, self.window_size)
            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

        attn_windows = self.attn(x_windows, mask=attn_mask)
        attn_windows = attn_windows.view(-1, self.window_size, self.window_size, C)

        # 还原窗口
        x = window_reverse(attn_windows, self.window_size, H, W)
        x = x.view(B, H * W, C)

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

        x = shortcut + x
        x = x + self.mlp(self.norm2(x))
        return x


class SwinTransformer(nn.Module):
    """Swin Transformer 精简实现"""
    def __init__(self, in_chans=3, embed_dim=96, depths=(2, 2, 6, 2), num_heads=(3, 6, 12, 24), window_size=7):
        super().__init__()
        self.patch_embed = PatchEmbed(in_chans, embed_dim)
        self.layers = nn.ModuleList()
        self.num_layers = len(depths)

        for i in range(self.num_layers):
            layer = nn.ModuleList()
            for j in range(depths[i]):
                layer.append(SwinBlock(embed_dim * (2 ** i), num_heads[i], window_size, shift=(j % 2 == 1)))
            self.layers.append(layer)

            if i < self.num_layers - 1:
                self.layers.append(PatchMerging(embed_dim * (2 ** i)))

    def forward(self, x):
        x = self.patch_embed(x)  # (B, L, C)
        B, L, C = x.shape
        H = W = int(L ** 0.5)

        features = []
        idx = 0
        for i in range(self.num_layers):
            for block in self.layers[idx]:
                x = block(x, H, W)
            idx += 1
            features.append(x)

            if i < self.num_layers - 1:
                x, H, W = self.layers[idx](x, H, W)
                idx += 1

        return features


if __name__ == "__main__":
    model = SwinTransformer()
    x = torch.randn(2, 3, 224, 224)
    features = model(x)
    for i, f in enumerate(features):
        print(f"Stage {i + 1}: {f.shape}")

运行上述代码,输出如下:

text 复制代码
Stage 1: torch.Size([2, 3136, 96])
Stage 2: torch.Size([2, 784, 192])
Stage 3: torch.Size([2, 196, 384])
Stage 4: torch.Size([2, 49, 768])

可以看到,特征图的空间尺寸逐阶段减半,通道数逐阶段翻倍,形成了类似 CNN 的金字塔结构。

6. 在 VLA 中使用 Swin Transformer

在 VLA 模型中,Swin Transformer 通常作为视觉编码器,输出多尺度特征。使用时需要注意以下几点:

  • 输入分辨率:Swin Transformer 要求输入尺寸能被 32 整除(4 次下采样),建议使用 224×224 或 256×256;
  • 特征融合:可将 Stage 3 和 Stage 4 的特征拼接或加权融合,作为语言模型的视觉 Token 输入;
  • 预训练权重:推荐加载 ImageNet-22K 预训练权重,可显著加速收敛并提升最终性能;
  • 与语言模型对齐:通过一个线性投影层将视觉特征映射到语言模型的嵌入空间,实现跨模态对齐。

下面是 Swin Transformer 在 VLA 模型中的接入流程:
#mermaid-svg-b44AjsxBm7WWmGo1{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-b44AjsxBm7WWmGo1 .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-b44AjsxBm7WWmGo1 .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-b44AjsxBm7WWmGo1 .error-icon{fill:#552222;}#mermaid-svg-b44AjsxBm7WWmGo1 .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-b44AjsxBm7WWmGo1 .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-b44AjsxBm7WWmGo1 .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-b44AjsxBm7WWmGo1 .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-b44AjsxBm7WWmGo1 .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-b44AjsxBm7WWmGo1 .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-b44AjsxBm7WWmGo1 .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-b44AjsxBm7WWmGo1 .marker{fill:#333333;stroke:#333333;}#mermaid-svg-b44AjsxBm7WWmGo1 .marker.cross{stroke:#333333;}#mermaid-svg-b44AjsxBm7WWmGo1 svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-b44AjsxBm7WWmGo1 p{margin:0;}#mermaid-svg-b44AjsxBm7WWmGo1 .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-b44AjsxBm7WWmGo1 .cluster-label text{fill:#333;}#mermaid-svg-b44AjsxBm7WWmGo1 .cluster-label span{color:#333;}#mermaid-svg-b44AjsxBm7WWmGo1 .cluster-label span p{background-color:transparent;}#mermaid-svg-b44AjsxBm7WWmGo1 .label text,#mermaid-svg-b44AjsxBm7WWmGo1 span{fill:#333;color:#333;}#mermaid-svg-b44AjsxBm7WWmGo1 .node rect,#mermaid-svg-b44AjsxBm7WWmGo1 .node circle,#mermaid-svg-b44AjsxBm7WWmGo1 .node ellipse,#mermaid-svg-b44AjsxBm7WWmGo1 .node polygon,#mermaid-svg-b44AjsxBm7WWmGo1 .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-b44AjsxBm7WWmGo1 .rough-node .label text,#mermaid-svg-b44AjsxBm7WWmGo1 .node .label text,#mermaid-svg-b44AjsxBm7WWmGo1 .image-shape .label,#mermaid-svg-b44AjsxBm7WWmGo1 .icon-shape .label{text-anchor:middle;}#mermaid-svg-b44AjsxBm7WWmGo1 .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-b44AjsxBm7WWmGo1 .rough-node .label,#mermaid-svg-b44AjsxBm7WWmGo1 .node .label,#mermaid-svg-b44AjsxBm7WWmGo1 .image-shape .label,#mermaid-svg-b44AjsxBm7WWmGo1 .icon-shape .label{text-align:center;}#mermaid-svg-b44AjsxBm7WWmGo1 .node.clickable{cursor:pointer;}#mermaid-svg-b44AjsxBm7WWmGo1 .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-b44AjsxBm7WWmGo1 .arrowheadPath{fill:#333333;}#mermaid-svg-b44AjsxBm7WWmGo1 .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-b44AjsxBm7WWmGo1 .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-b44AjsxBm7WWmGo1 .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-b44AjsxBm7WWmGo1 .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-b44AjsxBm7WWmGo1 .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-b44AjsxBm7WWmGo1 .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-b44AjsxBm7WWmGo1 .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-b44AjsxBm7WWmGo1 .cluster text{fill:#333;}#mermaid-svg-b44AjsxBm7WWmGo1 .cluster span{color:#333;}#mermaid-svg-b44AjsxBm7WWmGo1 div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-b44AjsxBm7WWmGo1 .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-b44AjsxBm7WWmGo1 rect.text{fill:none;stroke-width:0;}#mermaid-svg-b44AjsxBm7WWmGo1 .icon-shape,#mermaid-svg-b44AjsxBm7WWmGo1 .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-b44AjsxBm7WWmGo1 .icon-shape p,#mermaid-svg-b44AjsxBm7WWmGo1 .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-b44AjsxBm7WWmGo1 .icon-shape .label rect,#mermaid-svg-b44AjsxBm7WWmGo1 .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-b44AjsxBm7WWmGo1 .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-b44AjsxBm7WWmGo1 .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-b44AjsxBm7WWmGo1 :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} 图像输入
Swin Transformer 视觉编码器
多尺度特征 (Stage 3 & 4)
特征融合
线性投影层
语言模型嵌入空间
动作解码器
动作输出

8. 总结问题与调试建议

在 ResNet 或 ViT 上能跑通的代码,直接替换成 Swin Transformer 后经常会遇到几类问题。下面按"输入尺寸 → 显存 → 训练收敛 → 窗口划分"的顺序,给出排查思路与可直接复用的代码示例。

7.1 输入尺寸不匹配

Swin Transformer 经过 4 次 Patch Merging 下采样,通常要求输入宽高都能被 32 整除。如果不满足,PatchMerging.forward 中的 x.view(B, H, W, C)SwinBlock 里的窗口划分会因形状不符而报错,常见错误形如:

text 复制代码
RuntimeError: shape '[B, ...]' is invalid for input of size ...

解决方案:在进入模型前对图像进行 padding,使宽高对齐到 32 的倍数;或在数据预处理阶段直接统一 resize 到 224×224 或 256×256。

python 复制代码
import torch
import torch.nn.functional as F

def pad_to_multiple(x, divisor=32):
    """将输入 H、W padding 到 divisor 的整数倍。"""
    B, C, H, W = x.shape
    pad_h = (divisor - H % divisor) % divisor
    pad_w = (divisor - W % divisor) % divisor
    if pad_h or pad_w:
        # F.pad 顺序为 (left, right, top, bottom)
        x = F.pad(x, (0, pad_w, 0, pad_h))
    return x

x = torch.randn(2, 3, 240, 300)
x = pad_to_multiple(x, divisor=32)  # (2, 3, 256, 320)

7.2 显存不足(OOM)

Swin 使用窗口注意力将计算复杂度降到线性,但高分辨率输入、较大的 window_size 或过大的 batch size 仍会带来较高显存占用,尤其是 Stage 1 的 Token 数量很多。

解决方案,按优先级依次尝试:

  1. 降低 batch_size 或输入分辨率;
  2. 训练时开启 AMP 混合精度,减少激活显存占用;
  3. 对深层 Block 使用 torch.utils.checkpoint 进行梯度检查点;
  4. 推理阶段使用 torch.no_grad(),并逐张处理大尺寸图像。
python 复制代码
from torch.cuda.amp import autocast, GradScaler

scaler = GradScaler()

for x, target in dataloader:
    optimizer.zero_grad()
    with autocast():
        features = model(x)
        # 示例损失,实际任务请替换为真实 loss
        loss = (features[-1].sum() - target).abs().mean()

    scaler.scale(loss).backward()
    scaler.step(optimizer)
    scaler.update()

7.3 训练不收敛

如果 loss 长时间不下降或震荡明显,通常不是"训练步数不够",而是优化设置或初始化不匹配。

排查与解决建议

  • 使用较小的峰值学习率,并配合 warmup,让学习率在训练初期线性上升;
  • 加载 ImageNet-22K 预训练权重,能明显改善收敛速度和最终效果;
  • 检查 embed_dim 能否被 num_heads 整除,避免出现非整数 head_dim
  • 保持相对位置偏置表的初始化方式为 trunc_normal_,不要把标准差设置得过大。
python 复制代码
def linear_warmup(step, warmup_steps, peak_lr):
    """前 warmup_steps 步内学习率线性上升到 peak_lr。"""
    if step < warmup_steps:
        return peak_lr * step / max(1, warmup_steps)
    return peak_lr

lr = linear_warmup(global_step, warmup_steps=1000, peak_lr=1e-4)

7.4 窗口划分边界问题

window_partition 默认要求 HW 能被 window_size 整除;而 SW-MSA 在循环移位后,窗口内会混入不同原始区域的 patch,如果 attention mask 计算错误,模型会错误地在这些区域之间做注意力,导致特征质量下降。

解决方案

  • 确保经过 padding 后,H/4W/4 都能被 window_size 整除;
  • 确认 mask 对 attn_mask != 0 的位置使用足够小的负值(如 -100.0)进行屏蔽;
  • 如果自定义窗口函数,保持 window_partitionwindow_reverse 的维度排列一致。
python 复制代码
def check_window_size(H, W, window_size=7):
    assert H % window_size == 0 and W % window_size == 0, (
        f"特征图尺寸 ({H}, {W}) 必须是 window_size={window_size} 的整数倍,"
        "需先 padding 或调整 window_size。"
    )
    return True

check_window_size(56, 56, window_size=7)

7.5 排查路线图

遇到问题时,可按下面的流程快速定位:
#mermaid-svg-Fy19JefcKfhVfZD6{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-Fy19JefcKfhVfZD6 .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-Fy19JefcKfhVfZD6 .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-Fy19JefcKfhVfZD6 .error-icon{fill:#552222;}#mermaid-svg-Fy19JefcKfhVfZD6 .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-Fy19JefcKfhVfZD6 .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-Fy19JefcKfhVfZD6 .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-Fy19JefcKfhVfZD6 .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-Fy19JefcKfhVfZD6 .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-Fy19JefcKfhVfZD6 .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-Fy19JefcKfhVfZD6 .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-Fy19JefcKfhVfZD6 .marker{fill:#333333;stroke:#333333;}#mermaid-svg-Fy19JefcKfhVfZD6 .marker.cross{stroke:#333333;}#mermaid-svg-Fy19JefcKfhVfZD6 svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-Fy19JefcKfhVfZD6 p{margin:0;}#mermaid-svg-Fy19JefcKfhVfZD6 .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-Fy19JefcKfhVfZD6 .cluster-label text{fill:#333;}#mermaid-svg-Fy19JefcKfhVfZD6 .cluster-label span{color:#333;}#mermaid-svg-Fy19JefcKfhVfZD6 .cluster-label span p{background-color:transparent;}#mermaid-svg-Fy19JefcKfhVfZD6 .label text,#mermaid-svg-Fy19JefcKfhVfZD6 span{fill:#333;color:#333;}#mermaid-svg-Fy19JefcKfhVfZD6 .node rect,#mermaid-svg-Fy19JefcKfhVfZD6 .node circle,#mermaid-svg-Fy19JefcKfhVfZD6 .node ellipse,#mermaid-svg-Fy19JefcKfhVfZD6 .node polygon,#mermaid-svg-Fy19JefcKfhVfZD6 .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-Fy19JefcKfhVfZD6 .rough-node .label text,#mermaid-svg-Fy19JefcKfhVfZD6 .node .label text,#mermaid-svg-Fy19JefcKfhVfZD6 .image-shape .label,#mermaid-svg-Fy19JefcKfhVfZD6 .icon-shape .label{text-anchor:middle;}#mermaid-svg-Fy19JefcKfhVfZD6 .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-Fy19JefcKfhVfZD6 .rough-node .label,#mermaid-svg-Fy19JefcKfhVfZD6 .node .label,#mermaid-svg-Fy19JefcKfhVfZD6 .image-shape .label,#mermaid-svg-Fy19JefcKfhVfZD6 .icon-shape .label{text-align:center;}#mermaid-svg-Fy19JefcKfhVfZD6 .node.clickable{cursor:pointer;}#mermaid-svg-Fy19JefcKfhVfZD6 .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-Fy19JefcKfhVfZD6 .arrowheadPath{fill:#333333;}#mermaid-svg-Fy19JefcKfhVfZD6 .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-Fy19JefcKfhVfZD6 .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-Fy19JefcKfhVfZD6 .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-Fy19JefcKfhVfZD6 .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-Fy19JefcKfhVfZD6 .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-Fy19JefcKfhVfZD6 .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-Fy19JefcKfhVfZD6 .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-Fy19JefcKfhVfZD6 .cluster text{fill:#333;}#mermaid-svg-Fy19JefcKfhVfZD6 .cluster span{color:#333;}#mermaid-svg-Fy19JefcKfhVfZD6 div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-Fy19JefcKfhVfZD6 .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-Fy19JefcKfhVfZD6 rect.text{fill:none;stroke-width:0;}#mermaid-svg-Fy19JefcKfhVfZD6 .icon-shape,#mermaid-svg-Fy19JefcKfhVfZD6 .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-Fy19JefcKfhVfZD6 .icon-shape p,#mermaid-svg-Fy19JefcKfhVfZD6 .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-Fy19JefcKfhVfZD6 .icon-shape .label rect,#mermaid-svg-Fy19JefcKfhVfZD6 .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-Fy19JefcKfhVfZD6 .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-Fy19JefcKfhVfZD6 .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-Fy19JefcKfhVfZD6 :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} 是



报错或效果异常
是否为 shape 相关报错?
检查 H/W 能否被 32 整除
pad_to_multiple 或统一 resize
是否显存不足 OOM?
降 batch size / AMP / 梯度检查点
检查 loss 是否收敛
warmup / 预训练权重 / 调整学习率

7. 总结

本文系统梳理了 Swin Transformer 的核心原理,包括 Patch Embedding、移位窗口自注意力、相对位置偏置和层次化特征提取,并给出了完整的 PyTorch 复现。Swin Transformer 凭借其线性计算复杂度和多尺度特征表达能力,已成为 VLA 模型中视觉编码器的主流选择之一。

相关推荐
平原201813 分钟前
AI岗位需求增长244%背后:从任务重组到可验证学习闭环
大数据·人工智能·学习
ouynagda17 分钟前
Linux信号与进程间通信学习笔记
linux·笔记·学习
️学习的小王23 分钟前
AI Agent Skills进阶:调试、排错、实战开发与工程化落地
人工智能·经验分享·笔记·学习
kkkkkkkkkk_Z26 分钟前
学嵌入式和Linux应用编程|学习日记Day26:Linux进程完整学习笔记
linux·笔记·学习
kmomo..37 分钟前
Linux 进程间通信(IPC)学习笔记
linux·笔记·学习
2601_949950631 小时前
在线刷题用什么小程序?认识一下能导资料、AI出题的练题簿
人工智能·学习·小程序·刷题·小程序推荐
猎头南楼1 小时前
杭州算法工程师,偏大模型应用落地
算法
HugoStudio_SWAN1 小时前
【按键反馈】C++ 实现键盘烟花秀
开发语言·c++·学习·程序人生