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

Global Attention: Hybrid Sparse Pattern

Introduction

Sliding-window attention is efficient but myopic: tokens only see nearby neighbors. Some information is global: a word's meaning or a document's theme spans the entire sequence. Global attention combines sliding-window with a small set of "global" tokens that can attend to all positions and are attended to by all. The result is a cross-shaped sparse pattern: a band around the diagonal plus full rows/columns for global indices.

This post implements the hybrid mask and visualizes the cross-shaped pattern.

The Hybrid Pattern

The global mask is the union of three sets:

Mask(i, j) = 1 if: 1. abs(i - j) <= W [sliding-window band] 2. i in global_indices [full row i] 3. j in global_indices [full column j] Example: T=8, W=1, global=[0, 7] Column indices → Row 0 (global): [ 1 1 1 1 1 1 1 1 ] full row Row 1: [ 1 1 1 0 0 0 0 1 ] band + global column Row 2: [ 0 1 1 1 0 0 0 1 ] band + global column ... Row 7 (global): [ 1 1 1 1 1 1 1 1 ] full row Visual: cross (band + full rows for global + full columns for global)

LibTorch Implementation

struct GlobalAttentionImpl : public torch::nn::Module { torch::nn::Linear wq{nullptr}, wk{nullptr}, wv{nullptr}; int64_t d_model, d_k, window_size; std::vector<int64_t> global_indices; GlobalAttentionImpl(int64_t d_model_, int64_t window_size_, std::vector<int64_t> global_) : d_model(d_model_), d_k(d_model_), window_size(window_size_), global_indices(global_) { 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); // Construct hybrid mask 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) // Part 1: Band mask auto band = torch::abs(i_exp - j_exp) <= window_size; // Part 2 & 3: Global rows and columns auto global_tensor = torch::tensor(global_indices); auto is_global_row = torch::isin(i_idx, global_tensor); auto is_global_col = torch::isin(j_idx, global_tensor); auto global_mask = (is_global_row.unsqueeze(1)) | (is_global_col.unsqueeze(0)); // Union of band and global auto hybrid_mask = band | global_mask; // Apply mask scores.masked_fill_(!hybrid_mask, std::numeric_limits<float>::lowest()); auto attn_weights = torch::softmax(scores, -1); auto output = torch::matmul(attn_weights, V); return output; } }; TORCH_MODULE(GlobalAttention);

Mask Construction Details

The band mask identifies local neighbors. The global rows mask uses torch::isin to check if index i is in the global list. If true, that entire row is kept. Similarly for columns.

unsqueeze(1) and unsqueeze(0) reshape the 1D boolean tensors for broadcasting. The OR operation | combines masks element-wise: a cell is kept if band OR global row OR global column.

Common choices: [0, T-1] (first and last token, for document boundaries) or [0] (just CLS token). In practice, you might learn which tokens should be global. For T=4096 and |global|=2, we save ~96% of computation vs. full attention.

Interactive Hybrid Pattern

▶ Global Attention: Hybrid Sparse Pattern

Watch the cross-shaped pattern emerge as global tokens are added.

Complexity Analysis

Attention TypeScore Matrix CellsFor T=4096, W=64, |g|=2
Full$T^2$16.8M
Sliding window only$2WT$524K
Global hybrid$2WT + 2|g|T$536K

The global tokens add minimal overhead ($O(|g| \cdot T)$) but provide crucial long-range context. With just 2 global tokens, we capture 97% of speedup compared to full attention.

Complete Shape Trace

StepVariableShapeNotes
Inputx(batch, T, d_model)Tokens
ProjectsQ, K, V(batch, T, d_k)Linear projections
Scoresscores(batch, T, T)Q @ K^T / sqrt(d_k)
Global indicesglobal_tensor(|g|,)List of global positions
Is global rowis_global_row(T,) booleanWhich rows are global
Is global colis_global_col(T,) booleanWhich cols are global
Hybrid maskhybrid_mask(T, T) booleanBand | global rows | global cols
Masked scoresscores(batch, T, T)Outside hybrid set to -∞
Weightsattn_weights(batch, T, T)softmax; sparse rows sum to 1
Outputoutput(batch, T, d_k)Weighted sum of values

Key Insights

Hybrid attention combines local and global information. Sliding-window handles local interactions efficiently; global tokens provide long-range context. Together, they achieve both speedup and expressiveness.
Global tokens create a cross pattern. Full rows and columns for global positions form a cross visually, enabling them to attend to and be attended by any position in the sequence.
Minimal overhead for significant benefit. Even 1-2 global tokens (e.g., CLS tokens) add negligible cost but unlock document-level reasoning. This makes global attention practical for language models.

What Comes Next

Sliding-window and global attention both reduce complexity through sparsity, but still require $O(T \cdot W)$ or $O(T \cdot d^2)$ computation. Linear Attention takes a different approach: replace softmax with a positive feature map and reorder operations to achieve $O(T \cdot d^2)$ or even $O(T)$ complexity.

Previous Sliding-Window Attention