Swin Transformer
Self-attention is computed inside small local windows that shift across consecutive layers to achieve linear computational complexity with image size.
Why Does This Exist?
The original Vision Transformer (ViT) calculates global self-attention across every patch in the image. If an image produces patch tokens, computing the pairwise attention matrix requires floating-point operations and memory. For a standard classification image of pixels split into patches, , which is modest for modern accelerators.
However, dense computer vision tasks such as semantic segmentation and object detection require high-resolution inputs like or to delineate precise object boundaries. At with patches, the sequence length explodes to . A single global attention matrix would require elements per head, consuming over 17 gigabytes of GPU memory for one attention layer alone and triggering immediate out-of-memory crashes.
Furthermore, ViT produces single-scale feature maps at a constant downsampling factor (), making it incompatible with Feature Pyramid Networks (FPN) and UNet architectures that rely on multi-scale spatial representations.
The Swin Transformer (Shifted Window Transformer) solves both problems. It restricts self-attention to small, fixed-size local windows of size (typically ), reducing computational complexity from quadratic to strictly linear . To allow information to cross window boundaries without reintroducing global computation, Swin shifts the window partitioning grid by half a window width in alternating layers. Optimization relies on gradient-descent across multi-scale backbones.
Think of It Like This
Staggered bricklaying across building floors
Imagine masons constructing a brick wall. On the first floor, bricks are laid side by side in neat pairs. If each pair of bricks is glued together internally (local window attention), the two bricks know each other perfectly, but no glue connects pair A to pair B. The wall would crack along the seams between pairs under slight pressure.
On the next floor, the masons do not stack bricks in identical columns. They shift every brick by half its width. The joint between the two bricks below now sits directly under the center of a solid brick above.
In the Swin Transformer, layer computes attention within static local windows (the pairs on the lower floor). Layer shifts the window boundaries by half a window. Patches that were separated at the borders in layer now find themselves inside the same window in layer . By alternating regular and shifted windows across successive layers, information flows seamlessly across the entire image canvas while keeping each mason's view strictly local.
How It Actually Works
Window and Shifted Window Multi-Head Attention
A feature map with patches is partitioned into non-overlapping windows of size patches. The computational complexities of standard global Multi-Head Self-Attention () and Window-based Multi-Head Self-Attention () for an image of patches and hidden dimension are:
When is fixed (commonly ), scales linearly with image area , whereas scales quadratically.
To enable cross-window communication, consecutive Swin Transformer blocks are deployed in pairs. The first block computes regular window attention ():
The second block computes Shifted Window attention () by translating the window partitioning grid by pixels from the top-left corner:
Shifting the window grid creates up to 9 smaller sub-windows along the boundary edges. A naive approach would pad these sub-windows to size , substantially increasing computation. Instead, Swin uses a cyclic shift: the boundary slices on the top and left are rolled to the bottom and right edges of the feature map, forming exactly regular windows.
Because cyclic rolling groups non-adjacent spatial regions together into the same window, a specialized attention mask is added during the dot-product computation:
where is a learnable relative position bias matrix, and contains zeros for positions originating from the same spatial area and (effectively ) for mismatched regions, completely blocking spurious cross-boundary attention. After self-attention, a reverse cyclic shift restores the feature map to its original spatial layout.
Hierarchical Representation via Patch Merging
As the network goes deeper, resolution must decrease while channel capacity increases, matching standard vision pyramids. Swin achieves this through Patch Merging:
- Groups of neighboring patches are concatenated along the channel dimension, transforming a spatial patch tensor of shape into .
- A linear projection layer maps the channels to .
- The spatial resolution is halved while channel capacity doubles, creating feature maps at and scale.
Worked Example
Compare computational FLOPs on an input feature map of size patches with window size , channel dimension , and total patch count :
-
Global Attention (ViT style):
- Self-attention matrix: attention scores.
- Matrix multiplication FLOPs: FLOPs ( GFLOPs).
-
Swin Local Window Attention:
- Number of windows: windows.
- Number of patches per window: patches.
- Attention scores per window: .
- Total attention scores across all 64 windows: .
- Matrix multiplication FLOPs: FLOPs ( GFLOPs).
The local window mechanism achieves a 64-fold reduction in attention matrix computation while shifted windows guarantee cross-window receptive field propagation.
Code
import torchimport torch.nn as nn
def window_partition(x: torch.Tensor, window_size: int) -> torch.Tensor: """Partitions a feature map into non-overlapping local windows. Args: x: Tensor of shape (B, H, W, C) window_size: Window dimension M Returns: windows: Tensor of shape (num_windows * B, window_size, window_size, C) """ b, h, w, c = x.shape x = x.view(b, h // window_size, window_size, w // window_size, window_size, c) # Permute to (B, H/M, W/M, M, M, C) and flatten window batch windows = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_size, window_size, c) return windows
def window_reverse(windows: torch.Tensor, window_size: int, h: int, w: int) -> torch.Tensor: """Reverses window partition back to standard feature map layout.""" 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
# Verify partition and reverse cyclic shiftbatch_size, height, width, channels = 2, 14, 14, 96window_size = 7feature_map = torch.randn(batch_size, height, width, channels)
# Partition into 4 windows per image -> total 8 windowswindows = window_partition(feature_map, window_size=window_size)print("Windows shape:", windows.shape)# -> Windows shape: torch.Size([8, 7, 7, 96])
# Reverse back to original spatial layoutreconstructed = window_reverse(windows, window_size=window_size, h=height, w=width)print("Reconstruction matches input:", torch.allclose(feature_map, reconstructed))# -> Reconstruction matches input: True
# Cyclic shift translation demonstration (shift_size = M // 2 = 3)shift_size = window_size // 2shifted_x = torch.roll(feature_map, shifts=(-shift_size, -shift_size), dims=(1, 2))print("Shifted tensor shape:", shifted_x.shape)# -> Shifted tensor shape: torch.Size([2, 14, 14, 96])Watch Out For
Information leakage in cyclic shifts without relative attention masking
When implementing Shifted Window Attention with cyclic shifts, boundary slices are rolled around the tensor to maintain a uniform rectangular grid. Consequently, a single window in the bottom-right corner contains patches originating from four completely disconnected quadrants of the original image.
If self-attention is computed without an attention mask, tokens from the top-left corner attend directly to tokens from the bottom-right corner purely because they were rolled into the same computational buffer. This introduces artificial spatial shortcuts that corrupt position encoding and degrade segmentation boundaries. You must construct an attention mask containing large negative values (e.g. ) between sub-regions that do not share true spatial proximity, ensuring the softmax sets their attention weight strictly to zero.
The Quick Version
- Swin Transformer limits self-attention to non-overlapping local windows (), reducing computational complexity from to strictly linear .
- Alternating layers shift the window boundaries by half a window width, connecting neighboring regions and expanding the receptive field across the entire canvas.
- Cyclic shifts and attention masking compute shifted-window attention with zero padding overhead, while Patch Merging builds a hierarchical multi-scale feature pyramid.