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
structGlobalAttentionImpl : publictorch::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::Tensorforward(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 maskauto 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 maskauto band = torch::abs(i_exp - j_exp) <= window_size;
// Part 2 & 3: Global rows and columnsauto 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 globalauto 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 Type
Score Matrix Cells
For 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
Step
Variable
Shape
Notes
Input
x
(batch, T, d_model)
Tokens
Projects
Q, K, V
(batch, T, d_k)
Linear projections
Scores
scores
(batch, T, T)
Q @ K^T / sqrt(d_k)
Global indices
global_tensor
(|g|,)
List of global positions
Is global row
is_global_row
(T,) boolean
Which rows are global
Is global col
is_global_col
(T,) boolean
Which cols are global
Hybrid mask
hybrid_mask
(T, T) boolean
Band | global rows | global cols
Masked scores
scores
(batch, T, T)
Outside hybrid set to -∞
Weights
attn_weights
(batch, T, T)
softmax; sparse rows sum to 1
Output
output
(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.