Skip to main content

Hierarchical Attention in Vision Transformers

Trace how local windows, shifted cross-window exchange, and patch merging turn one high-resolution token grid into a multi-scale vision hierarchy.

A flat Vision Transformer keeps one spatial token count and width through its encoder. That makes every stage easy to describe, but dense prediction systems often need something else: high-resolution features for small details, lower-resolution features with larger receptive units, and a way to avoid one global N × N attention matrix at the finest scale.

A hierarchical vision transformer changes the tensor geometry between groups of blocks. Local attention limits early interaction, cross-window mechanisms reconnect those local groups, and patch merging trades spatial rows for channel width. The result is a feature pyramid rather than one flat sequence.

Follow the grid as locality becomes hierarchy

Click any image-derived token to follow its region through the stages, or play the complete sequence. Regular windows expose the local interaction boundary; the shifted stage shows where cross-window exchange and masking enter; merge stages preserve the selected image region while reducing the number of spatial rows.

Three mechanisms that must stay separate

1. Local window attention reduces the pair matrix

Suppose a stage has N = H × W spatial tokens and each window is M × M tokens. There are N/M² windows when the geometry divides evenly. Each window allocates (M²)² = M⁴ query-key pair slots, so the stage allocates

(N/M²) × M⁴ = N M²

pair slots instead of N² for global attention. With fixed M, the attention-pair term grows linearly with the number of image tokens. This does not make the entire block free or universally linear: Q/K/V projections, output projection, MLPs, normalization, layout changes, padding, and kernel behavior remain.

Local windows also restrict communication. Two tokens in different windows cannot exchange information inside that regular-window block, even when they are adjacent across a partition boundary.

2. Shifted windows reconnect prior boundaries

Swin alternates regular window attention with a half-window cyclic shift. The shift changes the partition, so tokens separated by an old boundary can land in the same allocated window. A pairwise mask blocks false neighbors introduced where cyclic wrap moves opposite image edges together, and the implementation reverses the shift after attention.

The shifted block allocates the same window-matrix shape as the regular block. It adds masking and data-movement concerns rather than restoring a full global matrix. Cross-window information spreads over multiple blocks; it is not globally available in one early local-attention step.

3. Patch merging creates the spatial hierarchy

Window attention alone changes who interacts. Patch merging changes how many spatial rows survive and therefore creates the hierarchy.

In a Swin-style merge, each 2 × 2 neighborhood contributes four C-wide rows. Concatenation produces 4C, then normalization and a learned linear projection typically produce one 2C-wide output row:

H × W × C → H/2 × W/2 × 2C

Spatial rows fall by 4×, channel width grows by 2×, and each output row represents a larger source region. Repeating this between stages yields fine, middle, and coarse maps for classification, detection, and segmentation heads.

A shape-faithful window partition

The layout operation should preserve values and make the window boundary explicit. For channel-last B × H × W × C tensors:

code
def window_partition(x, window_size): batch, height, width, channels = x.shape assert height % window_size == 0 assert width % window_size == 0 x = x.view( batch, height // window_size, window_size, width // window_size, window_size, channels, ) x = x.permute(0, 1, 3, 2, 4, 5).contiguous() return x.view(-1, window_size * window_size, channels)

Real implementations must also define padding, unpadding, attention-mask construction, relative-position bias, memory layout, and the reverse operation. A correct shape transformation does not by itself establish an efficient kernel.

Hierarchy is broader than shifted windows

Swin is one important hierarchical design, not the definition of every hierarchy.

  • Swin Transformer combines regular/shifted local windows with patch merging and typically exposes four stage resolutions.
  • Pyramid Vision Transformer (PVT) progressively shrinks spatial resolution while using spatial-reduction attention, where full-resolution queries read reduced K/V rows.
  • MViT pools queries, keys, and values across stages while increasing channel capacity.
  • Focal Transformer mixes fine local interactions with coarser contextual levels.

The common property is changing spatial scale across stages. Their attention operators, connectivity, pair counts, and implementation costs are not interchangeable.

Complexity and deployment boundaries

BoundaryGlobal flat attentionFixed local windowsHierarchical stage transition
Attention pair slotsN²N M² for M × M windowsDepends on the next stage's N and windowing
Immediate cross-window interactionAvailableNot in a regular-window blockAdded by shifts, bridges, pooling, or context
Spatial rowsUsually fixed across blocksFixed inside the windowed blockReduced by merge or pooling
Channel widthUsually fixed across blocksFixed inside the windowed blockOften increased after spatial reduction
Extra workLarge global matrixPartitioning and local kernelsReshape, concatenate/pool, project, move data

Analytical pair counts explain where the attention matrix shrinks; they do not guarantee wall-clock speed. Measure the same decode-to-output boundary and include padding, shifts, masks, merge projections, downstream heads, precision, batch size, compilation, and memory traffic.

Advantages and trade-offs

What the hierarchy buys:

  • fine and coarse features for dense tasks,
  • bounded early attention windows at high resolution,
  • larger effective source regions after merging,
  • reusable feature maps at several scales.

What it costs:

  • delayed global exchange in early stages,
  • geometry constraints, padding, and masks,
  • merge projections and activation traffic,
  • architecture-specific weights—windowing and merging are not drop-in inference switches for an arbitrary pretrained global-attention checkpoint.

If you found this explanation helpful, consider sharing it with others.

Mastodon