Sliding-Window Attention: Efficient Sparse Pattern
Introduction
Full attention on a sequence of length $T$ requires $O(T^2)$ memory and computation for the score matrix. On long documents (thousands of tokens), this is expensive. Sliding-window attention restricts each token to attend only to a band of width $W$ around itself: token $i$ attends to tokens $[i-W, i+W]$. This reduces complexity to $O(T \cdot W)$, making it linear in sequence length when $W$ is constant.
This post implements the band mask in C++ and visualizes which cells are kept vs. masked as the window width varies.
The Band Mask
For window size $W$, the mask includes positions where $|i - j| \leq W$:
LibTorch Implementation
Mask Construction Details
To construct the band mask, we create two index grids using torch::arange and broadcasting. i_idx.unsqueeze(1) gives shape (T, 1); j_idx.unsqueeze(0) gives shape (1, T). Subtracting them broadcasts to (T, T) with each entry $(i, j) = i - j$.
torch::abs(i_exp - j_exp) computes $|i - j|$ for every cell. We then check if this distance is $\leq \text{window\_size}$, producing a boolean mask.
masked_fill_ fills where the condition is true. We want to fill where the mask is false (outside the band), so we use !band_mask.
Interactive Band Visualization
▶ Sliding-Window Band Mask as W Increases
Adjust the window size and watch how the band grows.
Complexity Analysis
| Attention Type | Memory | Compute | Behavior |
|---|---|---|---|
| Full | $O(T^2)$ | $O(T^2 \cdot d)$ | All-to-all; quadratic cost |
| Sliding (W=const) | $O(T \cdot W)$ | $O(T \cdot W \cdot d)$ | Linear in T; practical for long docs |
| Sliding W=T/2 | $O(T^2)$ | $O(T^2 \cdot d)$ | Degrades to full attention |
The key is to choose $W$ such that it's sufficient for your task but small enough to gain speedup. Common choices: $W = 64, 128, 256$ for language modeling and document processing.
Complete Shape Trace
| Step | Variable | Shape | Notes |
|---|---|---|---|
| Input | x | (batch, T, d_model) | Tokens |
| Project Q, K, V | Q, K, V | (batch, T, d_k) | Linear projections |
| Scores | scores | (batch, T, T) | Q @ K^T / sqrt(d_k) |
| Index grids | i_exp, j_exp | (T, 1) and (1, T) | For broadcasting |
| Distance matrix | i_exp - j_exp | (T, T) | $|i - j|$ for each (i,j) |
| Band mask | band_mask | (T, T) boolean | True where $|i-j| \leq W$ |
| Masked scores | scores | (batch, T, T) | Outside band set to -∞ |
| Weights | attn_weights | (batch, T, T) | softmax; sparse row sums to 1 |
| Output | output | (batch, T, d_k) | Weighted sum of values |
Key Insights
unsqueeze and broadcasting to compute the distance matrix is both elegant and efficient in LibTorch.
What Comes Next
Sliding-window attention is local; it never attends to distant tokens. Some information (e.g., document-level semantic coherence) requires long-range dependencies. Global Attention combines sliding-window with a small set of "global" tokens that can see the entire sequence.