Exploring Attention Variants: Flash Attention, Sliding-Window, and Sparse Attention Implementation

Flash Attention eliminates memory bottlenecks through kernel fusion and tiling, while Sliding-Window and Sparse attention reduce KV-cache usage via structured sparsity patterns, enabling 128K+ context windows on consumer GPUs.

Modern transformer models face a fundamental scaling challenge: standard self-attention incurs quadratic memory and compute costs with respect to sequence length. This article examines three production-ready attention variants—Flash Attention, Sliding-Window Attention (SWA), and Sparse attention—as implemented in the rohitg00/ai-engineering-from-scratch curriculum, demonstrating how each optimizes the attention mechanism for long-context inference.

The Quadratic Bottleneck in Standard Attention

Standard causal self-attention computes an N×N attention matrix for a sequence of length N, resulting in O(N²) time and memory complexity. For long contexts (e.g., 128K tokens), materializing this matrix exhausts GPU memory even before storing the KV-cache—the persistent key and value tensors required for autoregressive generation. The repository addresses this limitation through three distinct architectural approaches that modify how tokens attend to one another.

Flash Attention: Memory-Efficient Attention Through Tiling

Flash Attention reorders the softmax and matrix-multiplication operations to ensure the full N×N attention matrix never materializes in high-bandwidth memory (HBM). Instead, it employs cache-friendly tiling and fused GPU kernels to compute attention in blocks.

Unlike other variants, Flash Attention uses an implicit causal mask without constructing an explicit mask matrix. The KV-cache remains unchanged (storing full-length keys and values), but memory usage during the attention step drops dramatically, allowing models to process up to 128K-token contexts on a single GPU. While the repository explains the algorithmic rearrangement in phases/07_transformers_deep_dive/15-attention-variants/docs/en.md, the actual implementation relies on custom CUDA kernels rather than pure Python code.

Sliding-Window Attention (SWA): Constraining the Receptive Field

Sliding-Window Attention (SWA) restricts each token’s receptive field to a fixed window size W of recent tokens, preserving causality while reducing the attention map from N×N to N×W. This is implemented via a mask where M[i][j] = 0 for j ∈ [i-W+1, i] and -inf otherwise.

The critical advantage appears in KV-cache scaling. Rather than growing linearly with sequence length N, the cache scales with window size W:

KV_bytes ∝ window_size  # Instead of sequence_length

According to the kv_cache_bytes calculations in phases/07_transformers_deep_dive/15_attention_variants/code/main.py, this yields up to a 6× memory reduction for 128K context windows when using a modest window size.

Sparse Attention: Hybrid Local and Global Context

Sparse (strided) attention augments the sliding window with periodic "strided" connections that attend to older tokens at fixed intervals. This creates a hybrid pattern that captures coarse-grained long-range dependencies without full attention costs.

The implementation combines masks through a logical OR operation:

  1. Sliding-window mask: Attends to recent local context
  2. Stride mask: Attends to tokens at regular intervals (e.g., every 3rd token)

The resulting sparse pattern maintains most of the memory savings from SWA while adding only a few extra KV-cache entries per token. This variant excels in retrieval-augmented generation scenarios requiring occasional distant context.

Quantifying Memory Impact with KV-Cache Calculations

The repository provides the kv_cache_bytes function in phases/07_transformers_deep_dive/15_attention_variants/code/main.py to quantify memory requirements precisely. The function calculates storage needs based on layer count, KV heads, head dimension, and sequence length:

from phases.07_transformers_deep_dive.15_attention_variants.code.main import kv_cache_bytes

layers, kv_heads, d_head = 80, 8, 128
seq_len = 131_072  # 128K tokens

window = 4096      # 4K sliding window

full_kv = kv_cache_bytes(layers, kv_heads, d_head, seq_len)
swa_kv = full_kv * (window / seq_len)  # Linear scaling with window

For a typical 80-layer model, full attention requires approximately 17 GB of KV-cache memory, while sliding-window attention with a 4K window reduces this to roughly 2.8 GB.

Implementing Attention Variants in Pure Python

The repository implements mask construction and attention computation without external dependencies in phases/07_transformers_deep_dive/15_attention_variants/code/main.py. Below are minimal, runnable snippets demonstrating each variant:

import numpy as np
from phases.07_transformers_deep_dive.15_attention_variants.code.main import (
    causal_mask, swa_mask, strided_mask, kv_cache_bytes, attention_row
)

# 1️⃣ Full-causal baseline mask

n = 8
full_mask = causal_mask(n)
print("Full-causal mask (first row):", full_mask[0])

# 2️⃣ Sliding-window mask (W=4)

window = 4
swa = swa_mask(n, window)
print("\nSliding-window mask (first row):", swa[0])

# 3️⃣ Strided sparse mask (window=2, stride=3)

sparse = strided_mask(n, window=2, stride=3)
print("\nStrided mask (first row):", sparse[0])

# 4️⃣ Memory estimation for 128K context

layers, kv_heads, d_head = 80, 8, 128
seq_len = 131_072
full_kv = kv_cache_bytes(layers, kv_heads, d_head, seq_len)
swa_kv = full_kv * (window / seq_len)
print(f"\nFull KV: {full_kv / 1e9:.2f} GB")
print(f"SWA KV: {swa_kv / 1e9:.2f} GB")

# 5️⃣ Execute single attention row with SWA

rng = np.random.default_rng(0)
d = 8
q = rng.normal(size=d)
K = rng.normal(size=(n, d))
V = rng.normal(size=(n, d))

mask_row = swa[7]  # Last token's view

out, weights = attention_row(q, K, V, mask_row)
print("\nAttention output:", out)
print("Attention weights:", weights)

These primitives compose into full multi-head modules via the SelfAttention and MultiHeadSelfAttention classes defined in phases/07_transformers_deep_dive/02-self-attention-from-scratch/code/self_attention.py.

Selecting the Right Attention Topology

The attention-variant picker documented in phases/07_transformers_deep_dive/15-attention-variants/outputs/skill-attention-variant-picker.md provides concrete decision rules:

  • Full attention + Flash Attention: Default choice when context ≤ 16K and retrieval demand is moderate
  • Sliding-window attention: Preferred when context exceeds 16K and local coherence dominates
  • Hybrid SWA + global: Use when occasional global context is required (e.g., Gemma-3 architecture)
  • Sparse or differential patterns: Reserve for specialized retrieval-heavy workloads with dedicated kernel support

Summary

  • Flash Attention eliminates the N×N memory materialization through tiling and fused kernels, maintaining full context capability without changing KV-cache structure.
  • Sliding-Window Attention reduces KV-cache memory linearly with window size W, enabling 6× memory savings on 128K contexts compared to full attention.
  • Sparse attention combines sliding windows with strided connections to capture long-range dependencies while preserving most memory benefits.
  • The kv_cache_bytes function in main.py quantifies exact memory trade-offs for any configuration.
  • Selection depends on context length thresholds (16K), retrieval requirements, and available kernel support.

Frequently Asked Questions

What is the primary advantage of Flash Attention over standard attention?

Flash Attention reduces memory usage during the attention computation from O(N²) to O(N) by reordering the softmax and matrix-multiplication operations and using tiling, allowing processing of longer sequences without materializing the full attention matrix in GPU memory.

How does Sliding-Window Attention reduce memory usage?

Sliding-Window Attention restricts each token to attend only to the previous W tokens rather than the entire sequence, reducing the KV-cache storage requirements from proportional to sequence length N to proportional to window size W, as implemented in the swa_mask function.

When should I use Sparse attention instead of pure Sliding-Window?

Use Sparse attention when your application requires occasional access to distant context (such as retrieval-augmented generation) but cannot afford the memory cost of full attention; the strided pattern adds periodic global connections to the local sliding window.

Does Flash Attention change the KV-cache requirements?

No, Flash Attention maintains the same KV-cache size as standard full attention because it optimizes the computation of attention scores without altering which keys and values must be stored; it reduces only the temporary memory used during the attention operation itself.

Have a question about this repo?

These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →