Swin transformer
Swin runs attention inside small windows and shifts the windows every other layer, which makes attention affordable on large images while keeping the CNN's four-stage shape.
- 13 min read
- 3 reading levels
- Updated
Read these first
On this page 7
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
Swin is a transformer that only lets each patch talk to its near neighbours. Every other layer reshuffles the neighbourhoods.
Think about a classroom discussion of forty students. Letting everyone talk to everyone at once is noise. So the teacher makes groups of five.
Each group talks properly. Then the teacher reshuffles the groups, and now different people are together. After a few rounds, an idea from one corner has travelled to the other. Nobody ever had to hold one chaotic all-in-one conversation.
That reshuffling is the whole idea behind Swin. The window of attention stays small, and the windows move.
Why it had to exist
A plain vision transformer lets every patch of the image attend to every other patch. That is powerful and expensive.
The cost grows with the square of the number of patches. Double the width and the height of a photo and you get four times the patches. That is sixteen times the cost.
For a small classification image that is bearable. For the large images used in medical scans, satellite pictures or document photos, it stops being possible at all.
Windows fix that. Cost now grows in a straight line with picture size instead of squaring. The measurement further down shows a 64-fold saving on a standard input. On a large one it is thousands of times.
The problem windows create, and the fix
If every layer used the same fixed windows, information could never cross a window boundary. The picture would be processed as a set of sealed boxes.
So Swin alternates. One layer uses the plain grid of windows. The next layer slides the whole grid halfway across, so the boundaries land in new places.
layer 1 windows layer 2 windows (shifted by half a window)
+----+----+ --+--+------+--
| | | | | |
+----+----+ --+--+------+--
| | | | | |
+----+----+ --+--+------+--
a patch near a boundary in layer 1 sits in the middle of a window in layer 2Two layers together let information cross. Stack enough pairs and a patch in one corner can influence a patch in the other.
The part borrowed from convolutional networks
Swin also copies the four-stage shape that ResNet uses. It starts with a fine grid of small patches. Every stage merges groups of four neighbouring patches into one.
The grid halves at each stage while the description of each position gets richer. That is the same trade you saw in ResNet.
This is why detection and segmentation systems could adopt Swin without redesign. It hands them the same set of feature map sizes they already expected.
Where you have already seen this
- Swin backbones sit inside many document scanning and form reading systems.
- Satellite and aerial image analysis uses them, because the images are enormous.
- Several medical imaging models use Swin for the same reason.
The honest part
Swin is more complicated to implement than either a plain transformer or a plain convolutional network. Window partitioning, shifting, masking at the edges, and a position bias table all have to be right.
It is also slower in wall-clock time than a convolutional network of the same arithmetic cost. The reshaping work does not use hardware well. That gap is measured on the next page over.
Remember this
- Attention runs inside small windows, so cost grows in a straight line with image size.
- Alternate layers shift the windows, so information still crosses boundaries.
- The four-stage shrinking shape is borrowed from convolutional networks.
What to learn next
- Contrastive learning for images — training these backbones without labels.
- Vision transformers — the unwindowed original.
- Attention — the operation being bounded here.
Developer — Code and libraries.
Setup
pip install "torch==2.5.1" "torchvision==0.20.1"Run against PyTorch 2.5.1 and torchvision 0.20.1 on CPU. No downloads.
Do the window partition by hand and count the saving
import torch
M = 4 # window side, in tokens
def windows(x, shift=0):
"""Cut an HxW grid of token ids into MxM windows, optionally rolled first."""
if shift:
x = torch.roll(x, shifts=(-shift, -shift), dims=(0, 1))
H, W = x.shape
return x.view(H // M, M, W // M, M).permute(0, 2, 1, 3).reshape(-1, M, M)
grid = torch.arange(64).view(8, 8)
print("an 8x8 grid of tokens")
print(grid)
print(f"\nlayer 1: plain {M}x{M} windows. token 27 sits with:")
for win in windows(grid):
if 27 in win:
print(win)
print(f"\nlayer 2: the same grid rolled by {M // 2}, then cut the same way. token 27 sits with:")
for win in windows(grid, shift=M // 2):
if 27 in win:
print(win)
print("different neighbours, without any window ever growing")
print("\nattention pairs: every token against every token, versus inside windows only")
print(f"{'image':>7s} {'tokens':>8s} {'global pairs':>16s} {'windowed pairs':>16s} {'ratio':>8s}")
for side in (224, 448, 896, 1792):
n = (side // 4) ** 2 # Swin's first stage: one token per 4x4 patch
win = 7
n_windows = n // (win * win)
global_pairs = n * n
local_pairs = n_windows * (win * win) ** 2
print(f"{side:7d} {n:8,} {global_pairs:16,} {local_pairs:16,} {global_pairs / local_pairs:7.0f}x")
print("\nstage shapes of torchvision's swin_t on a 224x224 image")
from torchvision.models import swin_t
s = swin_t(weights=None).eval()
h = torch.zeros(1, 3, 224, 224)
names = ["patch embed", "stage 1", "merge", "stage 2", "merge", "stage 3", "merge", "stage 4"]
with torch.no_grad():
for name, layer in zip(names, s.features):
h = layer(h)
print(f" {name:12s} {tuple(h.shape)}")
print(" note the layout is (batch, height, width, channels), not the usual NCHW")an 8x8 grid of tokens
tensor([[ 0, 1, 2, 3, 4, 5, 6, 7],
[ 8, 9, 10, 11, 12, 13, 14, 15],
[16, 17, 18, 19, 20, 21, 22, 23],
[24, 25, 26, 27, 28, 29, 30, 31],
[32, 33, 34, 35, 36, 37, 38, 39],
[40, 41, 42, 43, 44, 45, 46, 47],
[48, 49, 50, 51, 52, 53, 54, 55],
[56, 57, 58, 59, 60, 61, 62, 63]])
layer 1: plain 4x4 windows. token 27 sits with:
tensor([[ 0, 1, 2, 3],
[ 8, 9, 10, 11],
[16, 17, 18, 19],
[24, 25, 26, 27]])
layer 2: the same grid rolled by 2, then cut the same way. token 27 sits with:
tensor([[18, 19, 20, 21],
[26, 27, 28, 29],
[34, 35, 36, 37],
[42, 43, 44, 45]])
different neighbours, without any window ever growing
attention pairs: every token against every token, versus inside windows only
image tokens global pairs windowed pairs ratio
224 3,136 9,834,496 153,664 64x
448 12,544 157,351,936 614,656 256x
896 50,176 2,517,630,976 2,458,624 1024x
1792 200,704 40,282,095,616 9,834,496 4096x
stage shapes of torchvision's swin_t on a 224x224 image
patch embed (1, 56, 56, 96)
stage 1 (1, 56, 56, 96)
merge (1, 28, 28, 192)
stage 2 (1, 28, 28, 192)
merge (1, 14, 14, 384)
stage 3 (1, 14, 14, 384)
merge (1, 7, 7, 768)
stage 4 (1, 7, 7, 768)
note the layout is (batch, height, width, channels), not the usual NCHWReading the output
Token 27 changes company between the two layers. In layer 1 it sits in the top-left window with tokens 0 to 27, right at the corner. After the shift it sits in the middle of a window holding 18 to 45, which spans what used to be a boundary. Two layers, and the boundary has stopped being a wall.
The torch.roll call is the entire shift. Swin implements it exactly this way: roll the feature map, partition normally, then roll back afterwards. The extra complication in a real implementation is masking, because rolling wraps tokens from one edge of the image to the other, and those must not attend to each other.
The saving grows with image size, which is the real claim. At 224 pixels windowed attention is 64 times cheaper. At 1792 pixels it is 4096 times cheaper. Global attention is quadratic in token count while windowed attention is linear, so the gap widens without limit.
The windowed count itself grows linearly. From 153,664 to 9,834,496 as the image side grows eightfold, which is a factor of 64, matching the 64-fold growth in token count. That linear growth is what makes high-resolution input possible at all.
The stage shapes are ResNet's. 56, 28, 14, 7 with 96, 192, 384, 768 channels. Compare with the ResNet-18 stage output in the ResNet lesson: the spatial sizes are identical. This is why a Swin backbone drops into an existing detection neck.
The tensor layout is channels-last. torchvision's Swin returns (N, H, W, C). Feed that into code expecting (N, C, H, W) and you get a confusing error, or worse, a silent wrong answer when the channel count happens to match the spatial size.
Common mistakes
Forgetting the attention mask after shifting. torch.roll moves tokens from the bottom edge to the top. Without a mask, a token from the top of the image attends to one from the bottom as if they were neighbours. This is the single hardest part of a from-scratch implementation.
An image size that is not divisible by the window size. Swin needs the feature map at each stage to divide by 7. A 225x225 input fails or requires padding. torchvision pads internally; many research implementations do not.
Assuming the relative position bias transfers across window sizes. The bias is a learned table indexed by relative offsets inside a window. Change the window size at fine-tuning time and the table must be interpolated, which is exactly what Swin V2's log-spaced continuous position bias was designed to fix.
Expecting speed from the FLOP saving. Measured on one desktop CPU at four threads, swin_t took about 160 ms per 224x224 image against about 80 ms for convnext_tiny, at nearly the same GMACs. Partitioning, rolling and masking are memory movement, not arithmetic. Timings vary by machine and between runs.
Treating Swin as the default vision model. For 224-pixel classification, a ConvNeXt or a well-trained ResNet of the same size is usually faster and about as accurate. Swin's advantage appears on large inputs and dense prediction.
Try it yourself
Change M to 2 and rerun the first section. With a smaller window, token 27's neighbours change more sharply and the information travels more slowly across the grid. Then set the shift to M instead of M // 2 and check that the windows come back to where they started, which is why the half-window shift is the one that works.
What to learn next
- Contrastive learning for images — training these backbones without labels.
- Vision transformers — the unwindowed original.
- Attention — the operation being bounded here.
Researcher — Mathematics and papers.
The construction
Liu, Lin, Cao, Hu, Wei, Zhang, Lin and Guo (2021), Swin Transformer: Hierarchical Vision Transformer using Shifted Windows, ICCV, arxiv.org/abs/2103.14030.
Global multi-head self-attention over $h \times w$ tokens of dimension $C$ costs
$$ \Omega(\text{MSA}) = 4 h w C^2 + 2 (hw)^2 C $$
while window-based attention with $M \times M$ windows costs
$$ \Omega(\text{W-MSA}) = 4 h w C^2 + 2 M^2 h w C . $$
The first term, the projections, is identical. The second term is the difference: quadratic in $hw$ against linear in $hw$ with a constant of $M^2$. With $M = 7$ fixed, cost is linear in image area.
Consecutive blocks alternate W-MSA and SW-MSA, the shifted variant, displacing the window grid by $\lfloor M/2 \rfloor$. The efficient implementation is cyclic shift by torch.roll plus a mask that blocks attention between tokens that were not spatially adjacent before the roll, which avoids the padding that a naive implementation would need.
Attention within a window uses a learned relative position bias $B$:
$$ \mathrm{Attention}(Q,K,V) = \mathrm{SoftMax}!\left(\frac{QK^{\top}}{\sqrt{d}} + B\right) V $$
$B$ is indexed from a table of size $(2M-1) \times (2M-1)$, covering every relative offset inside a window. The paper reports this bias as worth a substantial margin over absolute position embeddings, and over no position information at all.
Patch merging concatenates each $2 \times 2$ group of neighbouring tokens along the channel axis, giving $4C$, then applies a linear layer to $2C$. Four stages give strides 4, 8, 16 and 32, matching the ResNet hierarchy.
Reported results: 87.3% ImageNet-1K top-1, 58.7 box AP and 51.1 mask AP on COCO test-dev, and 53.5 mIoU on ADE20K val, with margins of +2.7 box AP, +2.6 mask AP and +3.2 mIoU over the previous state of the art.
What the design is trading
The shifted window scheme reintroduces two priors that a plain ViT discards: locality, since attention is bounded, and hierarchy, since resolution decreases through stages. Both are the priors that convolution has built in. Swin's position in the literature is therefore a hybrid, not a pure transformer, and the ConvNeXt paper reads it that way explicitly, treating Swin as the target to match with convolutions alone.
The receptive field grows by roughly $M/2$ tokens per shifted layer plus the doubling at each merge, so it is bounded and grows with depth exactly as in a CNN, rather than being global from layer one.
Swin V2
Liu et al. (2021), Swin Transformer V2: Scaling Up Capacity and Resolution, arxiv.org/abs/2111.09883, scale to 3 billion parameters and images up to 1536x1536 with three changes:
- Residual post-normalisation with cosine attention. Pre-norm transformers accumulate activation amplitude across depth, which destabilises very large models. Moving the norm after the residual branch, and replacing the dot product with a scaled cosine similarity, bounds the attention logits regardless of activation magnitude.
- Log-spaced continuous position bias. The bias table is replaced by a small network taking log-spaced relative coordinates, so a model trained at one window size transfers to another by evaluating the network rather than interpolating a table.
- SimMIM self-supervised pretraining, to reduce the labelled data required at that scale.
Point 1 is the transferable idea. Attention logit growth with scale is a general failure mode, not a Swin-specific one.
Where it stands now
Swin remains a strong dense-prediction backbone and the reference point for windowed attention. Three qualifications are worth stating.
- At classification scale with matched training, ConvNeXt matches or exceeds it at equal FLOPs, and is faster in wall-clock terms on most hardware.
- Plain ViT backbones became competitive for detection once Li et al. (2022), Exploring Plain Vision Transformer Backbones for Object Detection (ViTDet), showed that a simple feature pyramid built from a single-scale ViT works, removing the assumption that hierarchy must be built into the backbone.
- Attention implementations have improved enough that the practical cost argument has narrowed for moderate resolutions. The argument still holds at large resolutions, where the quadratic term dominates.
The durable lesson is structural: bounding an operator's context and then moving the bounds is a general way to make a quadratic mechanism linear. The same pattern appears in sliding-window and block-sparse attention for language models.
What to learn next
- Contrastive learning for images — training these backbones without labels.
- Vision transformers — the unwindowed original.
- Attention — the operation being bounded here.