← All Posts
From Scratch · C++ · Transformers · Attention Variants · Part 5 of 13

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$:

For token i at row i and key j at column j: Keep (mask=1) if: abs(i - j) <= W Mask (mask=0) if: abs(i - j) > W Example W=1 (4x4 matrix): [ 1 1 0 0 ] token 0 can see 0, 1 [ 1 1 1 0 ] token 1 can see 0, 1, 2 [ 0 1 1 1 ] token 2 can see 1, 2, 3 [ 0 0 1 1 ] token 3 can see 2, 3 Total kept cells: ~3*T (linear in T)

LibTorch Implementation

struct SlidingWindowAttentionImpl : public torch::nn::Module { torch::nn::Linear wq{nullptr}, wk{nullptr}, wv{nullptr}; int64_t d_model, d_k, window_size; SlidingWindowAttentionImpl(int64_t d_model_, int64_t window_size_) : d_model(d_model_), d_k(d_model_), window_size(window_size_) { wq = register_module("wq", torch::nn::Linear( torch::nn::LinearOptions(d_model, d_k).bias(false) ) ); wk = register_module("wk", torch::nn::Linear( torch::nn::LinearOptions(d_model, d_k).bias(false) ) ); wv = register_module("wv", torch::nn::Linear( torch::nn::LinearOptions(d_model, d_k).bias(false) ) ); } torch::Tensor forward(torch::Tensor x) { int64_t T = x.size(1); auto Q = wq(x); auto K = wk(x); auto V = wv(x); auto scores = torch::matmul(Q, K.transpose(-2, -1)); scores = scores / sqrt((float)d_k); // Create band mask: keep where abs(i - j) <= window_size auto i_idx = torch::arange(T, x.options()); auto j_idx = torch::arange(T, x.options()); auto i_exp = i_idx.unsqueeze(1); // (T, 1) auto j_exp = j_idx.unsqueeze(0); // (1, T) auto band_mask = torch::abs(i_exp - j_exp) <= window_size; // Apply mask: set outside band to -inf scores.masked_fill_(!band_mask, std::numeric_limits<float>::lowest()); auto attn_weights = torch::softmax(scores, -1); auto output = torch::matmul(attn_weights, V); return output; } }; TORCH_MODULE(SlidingWindowAttention);

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 TypeMemoryComputeBehavior
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

StepVariableShapeNotes
Inputx(batch, T, d_model)Tokens
Project Q, K, VQ, K, V(batch, T, d_k)Linear projections
Scoresscores(batch, T, T)Q @ K^T / sqrt(d_k)
Index gridsi_exp, j_exp(T, 1) and (1, T)For broadcasting
Distance matrixi_exp - j_exp(T, T)$|i - j|$ for each (i,j)
Band maskband_mask(T, T) booleanTrue where $|i-j| \leq W$
Masked scoresscores(batch, T, T)Outside band set to -∞
Weightsattn_weights(batch, T, T)softmax; sparse row sums to 1
Outputoutput(batch, T, d_k)Weighted sum of values

Key Insights

Sliding-window is linear in sequence length. For constant $W$, the score matrix has $O(T \cdot W)$ non-masked cells, reducing both memory and compute from $O(T^2)$ to $O(T \cdot W)$.
Trade-off between locality and context. Smaller $W$ saves more computation but limits context. Larger $W$ approaches full attention. Choose $W$ based on your domain's typical dependency distance.
Broadcasting constructs the mask efficiently. Using 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.

Previous Cross-Attention