Sliding Window Attention: Efficient Long-Context Attention¶
Overview¶
Sliding Window Attention (SWA) restricts attention to a fixed window of recent tokens instead of attending to all previous tokens. Each token attends to only W previous tokens, reducing KV cache and computation from O(NΒ²) to O(NΓW).
- Adopted by: Mistral 7B, Phi-2, Llama 2, modern efficient models
- Key Insight: Most attention focus is on recent context (local patterns)
- Computation: From O(NΒ²) to O(N) with moderate window size
- Limitation: Long-range dependencies require stacking multiple layers
- Trade-off: Linear complexity vs. slightly reduced quality
The Problem: Quadratic Attention is Expensive¶
Standard Attention Complexity¶
Attention Mechanism:
- Compute Q @ K^T
- K shape: (batch, seq_len, num_heads, head_dim)
- Output shape: (batch, num_heads, seq_len, seq_len)
- Complexity: O(seq_lenΒ²)
Concrete Example: seq_len = 4000 tokens
Step 1: Q @ K^T
- Q shape: (32, 4000, 128)
- K^T shape: (32, 128, 4000)
- Output shape: (32, 4000, 4000)
- Operations: 4000 Γ 4000 Γ 128 = 2B operations
Step 2: Softmax + V multiply
- Softmax: 4000 Γ 4000 = 16M
- Multiply with V: 4000 Γ 4000 Γ 128 = 2B operations
- Total: ~4B operations per layer!
For Llama 7B (32 layers):
- 32 layers Γ 4B ops = 128B operations
- At 100 TFLOPS GPU: ~1.3 seconds per forward pass
- Plus KV cache memory: 4000 Γ 32 Γ 128 Γ 4 bytes β 64MB
For longer sequences (8000 tokens):
- Operations: 8B per layer (2x)
- For 32 layers: 256B operations (2x)
- Time: ~2.6 seconds (linear time!)
- This is SLOW!
The Observation¶
Research finding: Most attention is LOCAL
In real language:
- Token attends to: Previous 128 tokens (90% of attention)
- Token attends to: Previous 512 tokens (99% of attention)
- Token attends to: All previous tokens (<1% attention on old tokens)
Question:
"Do we need to compute attention to ALL previous tokens?
When only recent tokens have significant attention?"
Answer:
NO! Only compute attention for recent tokens!
This is Sliding Window Attention.
How Sliding Window Attention Works¶
Architecture¶
Standard Attention:
- ββββββββββββββββββββββββββββββββββββββ
- Token at position 20 β
- Attends to: all tokens 1-20 β (quadratic!)
- Attention scores: 20 values β
- β
Sliding Window Attention (window=4):
- ββββββββββββββββββββββββββββββββββββββ
- Token at position 20 β
- Attends to: tokens 16-20 (only!) β (fixed window)
- Attention scores: 4 values β
- Ignores tokens 1-15 β (local attention)
- β
Sequence visualization:
Position: 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20
Token: a b c d e f g h i j k l m n o p q r s t
For "t" at position 20, with window_size=4:
- Attention to: p (16), q (17), r (18), s (19), t (20)
- Ignored: a-o (positions 1-15)
- Only 5 positions instead of 20!
Complexity Analysis¶
Standard Attention:
- For each of N tokens: compute N attention scores
- Total: N Γ N = NΒ²
- For 4000 tokens: 16M scores
Sliding Window Attention:
- For each of N tokens: compute W attention scores (W=window size)
- Total: N Γ W = O(N)
- For 4000 tokens with window=256: 1M scores
- 16x reduction!
Mathematical:
Standard: O(NΒ²) attention operations
SWA: O(NΓW) attention operations
= O(N) if W is fixed constant
Key insight: W is small (128-256), N is large (4000+)
So O(NΓW) β O(N), nearly linear!
Implementation¶
class SlidingWindowAttention(nn.Module):
def __init__(self, hidden_dim, num_heads, window_size=256):
super().__init__()
self.num_heads = num_heads
self.head_dim = hidden_dim // num_heads
self.window_size = window_size
self.q_proj = nn.Linear(hidden_dim, hidden_dim)
self.k_proj = nn.Linear(hidden_dim, hidden_dim)
self.v_proj = nn.Linear(hidden_dim, hidden_dim)
self.o_proj = nn.Linear(hidden_dim, hidden_dim)
def forward(self, x):
batch, seq_len, hidden_dim = x.shape
# Project Q, K, V
q = self.q_proj(x).view(batch, seq_len, self.num_heads, self.head_dim)
k = self.k_proj(x).view(batch, seq_len, self.num_heads, self.head_dim)
v = self.v_proj(x).view(batch, seq_len, self.num_heads, self.head_dim)
# Transpose to (batch, num_heads, seq_len, head_dim)
q = q.transpose(1, 2)
k = k.transpose(1, 2)
v = v.transpose(1, 2)
output = []
# Apply sliding window attention
for i in range(seq_len):
# Determine window range: max(0, i - window_size) to i
start = max(0, i - self.window_size)
end = i + 1
# Get K, V within window
k_window = k[:, :, start:end, :] # (batch, heads, window, head_dim)
v_window = v[:, :, start:end, :]
# Query for current position
q_i = q[:, :, i:i+1, :] # (batch, heads, 1, head_dim)
# Compute attention for this position
scores = q_i @ k_window.transpose(-2, -1) / sqrt(self.head_dim)
weights = softmax(scores, dim=-1)
out_i = weights @ v_window
output.append(out_i)
# Stack outputs
output = torch.cat(output, dim=2) # (batch, heads, seq_len, head_dim)
# Merge heads
output = output.transpose(1, 2).contiguous()
output = output.view(batch, seq_len, hidden_dim)
output = self.o_proj(output)
return output
# More efficient implementation (vectorized):
class VectorizedSlidingWindowAttention(nn.Module):
def __init__(self, hidden_dim, num_heads, window_size=256):
super().__init__()
self.num_heads = num_heads
self.head_dim = hidden_dim // num_heads
self.window_size = window_size
# Projections
self.q_proj = nn.Linear(hidden_dim, hidden_dim)
self.k_proj = nn.Linear(hidden_dim, hidden_dim)
self.v_proj = nn.Linear(hidden_dim, hidden_dim)
self.o_proj = nn.Linear(hidden_dim, hidden_dim)
def forward(self, x):
batch, seq_len, hidden_dim = x.shape
# Project and reshape
q = self.q_proj(x).view(batch, seq_len, self.num_heads, self.head_dim)
k = self.k_proj(x).view(batch, seq_len, self.num_heads, self.head_dim)
v = self.v_proj(x).view(batch, seq_len, self.num_heads, self.head_dim)
q = q.transpose(1, 2)
k = k.transpose(1, 2)
v = v.transpose(1, 2)
# Create attention mask for sliding window
# (seq_len, seq_len) matrix: True where position j is in window of i
mask = torch.zeros(seq_len, seq_len, dtype=torch.bool)
for i in range(seq_len):
start = max(0, i - self.window_size)
mask[i, start:i+1] = True
# Compute attention scores
scores = q @ k.transpose(-2, -1) / sqrt(self.head_dim) # (batch, heads, seq_len, seq_len)
# Apply mask: set out-of-window positions to -inf
scores = scores.masked_fill(~mask.unsqueeze(0).unsqueeze(0), float('-inf'))
# Softmax and apply to values
weights = softmax(scores, dim=-1)
output = weights @ v # (batch, heads, seq_len, head_dim)
# Merge heads
output = output.transpose(1, 2).contiguous()
output = output.view(batch, seq_len, hidden_dim)
output = self.o_proj(output)
return output
Addressing the Long-Range Problem¶
The Issue¶
Sliding window only attends to recent tokens.
What about dependencies between distant tokens?
Example:
"A conversation was happening in the park.
[1000 tokens of conversation]
The dog barked at a nearby tree."
"dog" needs to refer back to opening context (1000+ tokens back)
But with window=256, it can only see recent 256 tokens!
This seems like a big problem...
Solution: Multi-Layer Receptive Field¶
Intuition: Stack multiple attention layers
Receptive field grows with layer depth:
Layer 1:
- Each token sees 256 recent tokens
- Context: 256 tokens back
Layer 2:
- Each token sees 256 tokens from layer 1
- But layer 1 tokens represent 256-token context each
- Effective context: 256 Γ 256 = ~65K tokens! (transitive)
Layer 3:
- Effective context: 256 Γ 256 Γ 256 = ~16M tokens!
- Far beyond sequence length!
Key: Multi-layer stacking creates large receptive fields
Each layer adds one factor of window_size to context
For 32 layers with window=256:
- Theoretical receptive field: 256^32 (enormous!)
- Practical: Limited by sequence length
- Can attend to full history if needed
- But efficiency still maintained!
Quality Trade-off¶
Experiment: Llama 2 70B vs. Mistral 7B (uses SWA)
Task: Long-range reasoning (document retrieval in 4K context)
- Mistral: 256-token window attention
- Llama 2: Full dense attention
Results:
- Llama 2 (full): 92% accuracy
- Mistral (SWA): 88% accuracy
- Quality gap: 4% (acceptable trade-off!)
Why quality gap is small:
1. Multi-layer stacking: Receptive field grows
2. Most information needs local context
3. Layer 10+: Can still access full history if needed
4. Trade-off: 4% accuracy for 4x speedup is good deal
For most tasks:
- SWA has minimal impact on quality
When to Use Sliding Window Attention¶
Use SWA When¶
β
Inference speed is critical
- Get 4-8x faster inference
- Linear complexity instead of quadratic
β
Context length is moderate (4K-8K)
- Multi-layer stacking handles receptive field
β
Quality loss <1-5% is acceptable
- Minor degradation for major speedup
β
Mobile/edge deployment
- Reduced memory and computation
Avoid SWA When¶
β Quality must be pristine
- Dense attention can be needed for some tasks
β Very long-range dependencies critical
- Long-document QA at sequence ends
β Context length extremely long (>16K)
- Sliding window becomes limiting
- Use Sparse Attention or Recurrent models instead
Hybrid Approaches¶
Dilated / Strided Attention¶
Improvement: Mix local and sparse patterns
Layer 1: Sliding window (256 token window)
Layer 2: Sliding window (256 token window)
Layer 3: Dilated attention (every 4th token)
- Attends to: [0, 4, 8, 12, 16, ...] positions
(sparse coverage of history)
Layer 4: Sliding window (256 token window)
Benefit:
- Local layers: Keep quality high for local patterns
- Sparse layers: Maintain long-range connections
- Overall: Better quality than pure SWA
Adoption: Longformer, BigBird use this approach
Combined Sliding + Full Attention¶
Mistral approach:
- Most layers: Sliding window attention (window=256)
- Every 4th layer: Full dense attention
- Result: 4x speedup, minimal quality loss
Why it works:
- Most layers: Efficient computation
- Occasional full: Ensure long-range information flow
- Balanced: Speed and quality
Key Takeaways¶
π Sliding Window Attention: From O(NΒ²) to O(NΓW) complexity
β‘ 4-8x inference speedup with minimal quality loss
π Multi-layer stacking creates large receptive fields
π Local patterns dominate, distant tokens matter less
π― Hybrid approaches balance efficiency and quality
Related Notes¶
- 00 Attention Mechanisms - Foundation
- Kv Cache - What SWA reduces
- Length Extrapolation & Long Context - Related challenge
- Flash Attention - Can be combined with SWA
- Llm Inference Optimization - Complete stack