Skip to main content

Hilbert-Guided Sparse Local Attention

Reordering image tokens along a Hilbert curve makes 2D local attention block-sparse, for about 4x faster window attention and 18x faster slide attention.

TL;DR

  • Local attention for images (Swin windows, slide attention, neighborhood attention) is sparse on paper. A block-sparse kernel such as FlexAttention only saves time on blocks of the attention matrix that are completely empty.
  • Flattening an image in row-major order scatters each 2D window across many rows of the sequence, so most non-empty blocks are partial and need element-wise masking. FlexAttention running plain window attention is slower than a dense kernel in every setting the paper tests.
  • The fix is a reordering, not a new kernel: flatten the image along a Hilbert curve, then build windows and neighborhoods on the 1D sequence. Nearby tokens stay nearby, the attention pattern collapses toward the diagonal, and far more blocks are empty.
  • On an RTX 3080, Hilbert window attention runs 4.0× faster than dense window attention at 128×128 tokens, and Hilbert slide attention runs 18× faster than naive slide attention at 56×56. Swin-T and NAT-mini variants built on these patterns lose at most 0.2 points of ImageNet top-1 accuracy.

Local attention is sparse, but the kernel cannot see it

Global self-attention over N image tokens costs O(N2). Local attention restricts each query to a neighborhood: a fixed window in Swin Transformer, a sliding window in Slide Transformer, a neighborhood that shifts inward at the image border in NAT. Most of the N × N attention matrix is then masked out.

Block-sparse kernels turn that mask into saved work. FlexAttention tiles the matrix into blocks of bq × bk entries and classifies each one (sparse attention patterns covers the general idea):

  • Empty block: every entry is masked. The kernel skips it and never loads its keys and values.
  • Full block: every entry is kept. The kernel runs a plain dense tile.
  • Partial block: some entries are masked. The kernel computes the whole tile and applies the mask element by element.

Empty blocks are cheapest, then full blocks, then partial blocks. The paper's Figure 1 shows FlexAttention's speedup rising with the share of empty blocks, most steeply above about 80%.

The catch is the order of the tokens. Vision models flatten the feature map row by row, so a 2D window such as tokens (1, 2, 5, 6) in a 4×4 map sits in two separate stretches of the sequence. Its attention lands in several blocks, and each of those blocks is only partly used. At 128×128 tokens with 16×16 windows, 87.5% of the blocks are already empty, yet window attention on FlexAttention takes 5.68 ms against 2.74 ms for a dense per-window kernel. Every remaining block is partial, and masking them costs more than the skipped blocks save.

Reordering tokens along a Hilbert curve

A Hilbert curve visits every cell of a grid, moving one step at a time, and fills each aligned square completely before moving to the next. Consecutive positions on the curve are neighbors in 2D, and tokens close in 2D tend to stay close in the sequence. The paper flattens the feature map along this curve and then defines local attention directly on the 1D sequence:

  • Hilbert Window Attention (HWA): a window is a run of W2 consecutive tokens. When the map and window sides are powers of two, this run is exactly an aligned W × W square.
  • Hilbert Slide Attention (HSA): each token attends to the K2 tokens around it in the sequence, a band along the diagonal.
  • Hilbert Neighborhood Attention (HNA): the same band, shifted inward at the ends of the sequence, which is 1D neighborhood attention.

The paper's toy example (Figure 3) uses 16 tokens, 2×2 windows and 4-token blocks. In row-major order the first eight tokens form four half-used partial blocks. In Hilbert order the first window is tokens 1 to 4 and the second is tokens 5 to 8, so the same attention forms two full blocks and two empty ones. Slide attention keeps its empty-block count but halves its partial blocks.

The authors' code uses a generalized Hilbert curve that works for any height and width. When the map or window size is not a power of two, a Hilbert window is an irregular region rather than a square, but its tokens are still adjacent in the image.

Why empty blocks translate into time

The paper models the runtime of a block-sparse kernel as the work of all its thread blocks (CTAs) divided by how many the GPU can run at once:

T ≈ Σi=1M (α + β · ri)P\text{eff}

Here α is the fixed cost of one CTA (launch, loading its query block, setup), β is the cost of one non-empty block (loading keys and values, the QK^\top tile, score modification, the online softmax and the product with V), ri is the number of non-empty blocks the i-th CTA processes, and P\text{eff} is the effective parallelism across the streaming multiprocessors. Hilbert order lowers the ri, and for the window pattern it also turns partial blocks into full ones.

Block size and window size interact. At 128×128 tokens with 16×16 windows, 128- and 256-token blocks leave 98.44% of blocks empty and run 4.0× faster than dense window attention. With 1024-token blocks each block spans several windows, sparsity falls to 93.75%, and the speedup disappears (2.78 ms against 2.74 ms).

Two backbones: HWT and HNT

Hilbert Window Transformer (HWT) follows Swin Transformer. The patch tokens are reordered along the Hilbert path, which depends only on the feature map size, so it is computed once and cached. Blocks come in pairs as in Swin: the first applies HWA, and the second applies Hilbert shifted window attention, which moves the windows forward along the 1D sequence by a fixed offset. Tokens from the start and the end of the sequence that land in the same shifted window are not neighbors in 2D, so their attention is masked out. Because Hilbert windows can be irregular, Swin's per-window relative position bias no longer fits, and HWT uses a single relative position bias over the whole feature map. Both attention types are expressed as a FlexAttention mask_mod and score_mod, with no hand-written kernel.

Hilbert Neighborhood Transformer (HNT) follows NAT. After the Hilbert reordering, 2D neighborhood attention becomes 1D neighborhood attention, which can run on NATTEN's na1d or on FlexAttention. The attention pattern is the same either way, so the backend changes speed but not accuracy.

Results

The kernel benchmarks use an RTX 3080 with CUDA 12.6 and PyTorch 2.7.0, batch size 16, 2 heads of dimension 64 and 128-token blocks; the appendix repeats them on an A100. Times below are forward passes.

Window attention (Table 1). Dense WSA computes each window separately. WSA (Flex) and HWA (Flex) run on FlexAttention.

Input, windowWSAWSA (Flex)HWA (Flex)Empty blocks, WSA to HWA
64×64, 80.28 ms0.40 ms0.12 ms (2.3×)87.50% to 96.88%
96×96, 161.56 ms2.63 ms0.40 ms (3.9×)83.33% to 97.22%
128×128, 162.74 ms5.68 ms0.68 ms (4.0×)87.50% to 98.44%

The FlexAttention variants also use far less memory than dense WSA, which materializes each window's attention matrix: 66 MB against 520 MB in the forward pass at 128×128. HWA also runs on xFormers' block-diagonal dense kernel, which plain window attention cannot use.

Slide and neighborhood attention (Table 3). SA gathers each neighborhood explicitly. Because HSA is a plain band, it also runs on FlashAttention-2's built-in sliding-window kernel (FA2).

Input, kernelSASA (Flex)HSA (Flex)HSA (FA2)NA2D (NATTEN)HNA (NATTEN)
56×56, 75.85 ms0.56 ms0.32 ms0.25 ms0.21 ms0.12 ms
96×96, 11out of memory1.83 ms0.61 ms0.69 ms1.07 ms0.51 ms

The headline 18× is HSA (Flex) against naive SA at 56×56 (5.85 ms to 0.32 ms). Against the optimized NATTEN kernel the gain is smaller but still there: HNA runs 1.8× faster than NA2D at 56×56 and 2.1× faster at 96×96.

Full layers. Counting the reshape (window partition or Hilbert reorder) and the QKV projection, HWA (Flex) takes 1.63 ms at 128×128 with 16×16 windows, against 4.22 ms for WSA and 6.39 ms for WSA (Flex). At 56×56 with 7×7 kernels, NA2D and HNA on NATTEN tie at 0.45 ms; HNA pulls ahead only at larger inputs and kernels.

Models (Table 4 and Figure 7). HWT-T and HNT-mini were trained on 8 V100 GPUs with the Swin-T and NAT-mini recipes. The V100 does not support FlexAttention well, so HWT trained with a dense kernel; the authors checked that the dense and FlexAttention versions give identical outputs. The speedups are therefore inference measurements.

ModelImageNet-1K top-1 (224×224)Images/s at 224×224Images/s at 512×512
Swin-T81.2%69696
HWT-T81.0%730162
NAT-mini81.8%818134
HNT-mini81.6%921160

The throughput runs use 7×7 windows or kernels at 224×224 and 16×16 at 512×512; the paper does not name the GPU for this figure. At 256×256, HWT-T reaches 81.5% against 81.6% for Swin-T. On CIFAR-10 and CIFAR-100, both models stay within 0.2 points of Swin and NAT.

Critical analysis

Strengths:

  • A reordering instead of a kernel. The whole method is a token permutation plus a mask definition. It runs on FlexAttention, FlashAttention-2, xFormers and NATTEN without any custom CUDA.
  • It fixes the right bottleneck. The paper identifies partial blocks, not attention FLOPs, as what makes local attention slow on block-sparse kernels, and the measured speedups track its sparsity numbers.
  • Near-free accuracy. Changing the window shape from squares to Hilbert regions costs at most 0.2 points on ImageNet at the Swin-T and NAT-mini scale.

Limitations:

  • It needs enough sparsity to pay off. At 56×56 with 7×7 windows, HWA (Flex) takes 0.31 ms against 0.22 ms for dense WSA, 0.7× the speed. The gain appears only at higher resolutions and with a good block-window pairing.
  • Speedups depend on hardware and batch size. The main results are on a consumer RTX 3080. At batch size 1 on that GPU the gain nearly vanishes, and it stops growing beyond a batch of about 32. At high resolution with large windows, PyTorch's SDPA running dense window attention on FlashAttention can match HWA (Flex).
  • Irregular windows complicate the model. Swin's per-window position bias has to become a global one, and shifted Hilbert windows need extra masks for tokens that wrap around the sequence.
  • Narrow evaluation. Accuracy is reported only for image classification at small model sizes. Detection, segmentation and generation, where high resolution matters most, are left to future work.
  • Swin Transformer: the window attention and shifted-window backbone that HWT reorders
  • FlashAttention: the tiled attention kernel that FlexAttention and FlashAttention-2 build on
  • Vision Transformer (ViT): global attention over image patches, the quadratic cost local attention avoids
  • Making Deep Learning Go Brrrr: why overhead and memory traffic, not FLOPs, often decide GPU runtime
  • VIOLIN: the same generalized Hilbert curve, used as a spatial prior inside global attention instead of for speed

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

Mastodon