Reference architecture

Long-Context Attention: Ring Attention, RoPE Frequency Scaling & Quantized KV-Caches

Engineering million-token context windows: circular peer-to-peer Ring Attention, 2D rotary embedding extrapolation (YaRN & NTK-aware), and FP8/FP4 KV-cache memory budgeting.

22 minVerified 2026-09-292 primary sources
A governed production AI reference architecture with observable, secured service boundaries.

Architecture

Extending large language model context windows beyond 128k tokens toward 1M+ tokens exposes fundamental limits in transformer memory scaling and hardware communications. Standard multi-head attention scales quadratically with sequence length N in activation memory (O(N^2) without memory-efficient attention) and linearly in Key-Value (KV) cache storage (O(N)). When a single forward pass processes 1,000,000 tokens on a 70B parameter model, the KV cache alone demands hundreds of gigabytes of high-bandwidth memory (HBM), far exceeding the capacity of any single GPU (such as an NVIDIA H100 with 80 GB or H200 with 141 GB).

text(6 lines)
1Sequence Scaling Challenge:
2- Single Token (Llama-3-, GQA 8 KV heads, d_head=128, 80 layers):
3 KV Cache per token = 2 * 80 * 8 * 128 * 2 bytes (FP16) = 327,680 bytes (~320 KB/token)
4- 128,000 Tokens = 40.96 GB VRAM (KV Cache only)
5- 1,000,000 Tokens = 320.0 GB VRAM (Exceeds H100 GPUs solely for KV cache)

To break this memory wall without sacrificing exact attention accuracy, production inference architectures combine three synergistic systems:

  1. Ring Attention (Blockwise Sequence Parallelism): Partitions context across a circular ring of accelerator devices, overlapping asynchronous peer-to-peer (P2P) KV transfers with local tensor core matrix multiplications.
  2. RoPE Frequency Scaling & Extrapolation: Extends 2D rotary positional embeddings to long sequences via Neural Tangent Kernel (NTK)-aware and YaRN interpolation without catastrophic perplexity degradation.
  3. Low-Precision Quantized KV-Caches: Compresses KV tensors to FP8 (E4M3/E5M2) or FP4 (NVFP4) block representations, slashing per-token memory footprint by 50% to 75% while maintaining generation fidelity.
Sequence sharding and local Query Key Value block allocation across GPU ring
Asynchronous non-blocking P2P send and receive of Key Value blocks via NCCL
Overlapped local FlashAttention tile GEMM execution on current block
Milakov-Gimelshein online softmax rescaling with running statistics
Ring rotation barrier synchronization and causal triangular mask pruning
Final attention output reduction and normalization across full context
Conceptual teaching model synthesized from:Ring Attention with Blockwise Transformers for Near-Infinite ContextRoFormer: Enhanced Transformer with Rotary Position Embedding

The diagram above illustrates the asynchronous Ring Attention execution cycle: while GPU r computes FlashAttention between its local Query block Q_r and received Key/Value block K_i, V_i, it concurrently transmits K_i, V_i to rank (r + 1) mod P and receives K_(i-1), V_(i-1) from rank (r - 1) mod P via NVLink.

2D orthogonal rotation of query and key coordinate pairs in complex plane
Base frequency decomposition across head dimensions with geometric progression
High-frequency token dimension preservation for local token discrimination
Low-frequency wavelength stretching and YaRN ramped interpolation
NTK-aware base frequency expansion to prevent out-of-distribution angles
Attention temperature scaling to prevent softmax entropy collapse at scale
Conceptual teaching model synthesized from:RoFormer: Enhanced Transformer with Rotary Position EmbeddingRing Attention with Blockwise Transformers for Near-Infinite Context

Complementing distributed communication, rotary position embedding scaling (shown above) decomposes embedding dimensions into frequency bands, preserving rapid local phase changes while interpolating long-wavelength dimensions to avoid out-of-distribution phase saturation.


Sequence Parallelism Comparison: Megatron vs. Ulysses vs. Ring Attention

Distributed training and inference employ different sequence parallelism paradigms to split sequence length N across P GPUs. Understanding their communication primitives reveals why Ring Attention is uniquely suited for million-token contexts.

| Parallelism Paradigm | Sharding Strategy | Communication Primitives | Communication Volume per Layer | Interconnect Dependency | Max Scalable Context | |---|---|---|---|---|---| | Megatron-LM Sequence Parallelism | Splits activations across sequence dimension along Tensor Parallel (TP) group | All-Gather (before Self-Attention/MLP) & Reduce-Scatter (after) | 4 * b * s * h bytes | High-bandwidth NVLink only (typically intra-node P <= 8) | Under 32k tokens | | DeepSpeed Ulysses | Splits tokens across sequence dimension; gathers all heads for attention | All-to-All (Q, K, V transpose) before and after attention | 4 * b * s * h bytes | High bisection bandwidth (All-to-All degrades over InfiniBand/RoCE) | Under 128k tokens | | Ring Attention | Blockwise circular sharding; each GPU holds N/P tokens of Q, K, V | P2P ncclSend and ncclRecv in a circular ring | 2 * ((P - 1) / P) * b * s * h_kv bytes | Point-to-Point bidirectional bandwidth (NVLink or multi-node RDMA) | 1M to 10M+ tokens | | Context Parallelism (Megatron-Core) | Hybrid Ulysses / Ring blockwise formulation with CUDA Graph optimization | Combined All-to-All or Ring P2P with causal masking | Variable (2 to 4 * b * s * h) | Dual-tier NVLink + RoCE v2 cluster networking | Up to 1M tokens |

The Ring Attention Breakthrough

Standard sequence parallelism requires collective communications (All-Gather or All-to-All) that create global synchronization barriers across all P devices. If one device encounters a tail latency spike, the entire cluster stalls.

Ring Attention replaces collective communication with asynchronous peer-to-peer ring communication. For an input sequence of length N distributed over P processors, each processor holds a local slice of length B = N / P:

  • Device r stores fixed local query block Q_r.
  • Device r initially stores local key block K_r and value block V_r.
  • Over P circular ring steps, K and V blocks rotate along the ring:
text(4 lines)
1Ring Rotation Primitives:
2Send: Rank r -> (r + 1) mod P
3Receive: Rank (r - 1) mod P -> r
  • In step k (where k = 0, 1, ..., P - 1), device r computes attention between its static query Q_r and the currently resident key-value block K_((r - k) mod P), V_((r - k) mod P).

Because the computational cost of the local attention block GEMM scales as O(B^2 * d) = O((N/P)^2 * d), while the P2P transfer cost of sending K and V scales linearly as O(B * d) = O((N/P) * d), computation completely hides communication latency whenever the block size B exceeds a critical threshold:

text(8 lines)
1Computation-Communication Overlap Condition:
2Time_GEMM(B) >= Time_P2P(B)
3(2 * B^2 * d) / FLOPS_TensorCore >= (2 * B * d * sizeof(FP16)) / Bandwidth_Interconnect
4
5Minimum Block Size Floor:
6B_min >= (FLOPS_TensorCore * sizeof(FP16)) / Bandwidth_Interconnect
7B_min >= ( * 2) / () ≈ 2,198 tokens

On an NVIDIA H100 SXM5 GPU (989 TFLOPS BF16 Tensor compute, 900 GB/s bidirectional NVLink bandwidth), as long as each GPU holds at least 2,048 to 4,096 tokens, communication is completely hidden behind computation, achieving near-perfect linear scaling up to millions of tokens.


Mathematical Foundation: Online Softmax Rescaling Across Ring Steps

The mathematical hurdle in computing attention over segmented blocks without global communication is the non-linear softmax operation:

text(3 lines)
1Standard Attention:
2Attention(Q, K, V) = softmax((Q * K^T) / sqrt(d_k)) * V

In standard attention, calculating softmax(S)_ij requires knowing the global maximum max_k S_ik and global normalizer across the entire context length N. Ring Attention solves this by adapting the Milakov-Gimelshein and FlashAttention online softmax algorithm across distributed ring steps.

Step-by-Step Online Softmax Formulation

Let device r hold query block Q_r. Across ring steps k = 0, 1, ..., P - 1, device r receives key block K^(k) and value block V^(k). Device r maintains three running variables in SRAM:

  1. m^(k): The running row-wise maximum of attention logits.
  2. l^(k): The running row-wise unnormalized softmax denominator.
  3. O^(k): The running unnormalized attention output accumulator.
text(19 lines)
1Step 1: Base Case Initialization (k = 0)
2S^(0) = (Q_r * (K^(0))^T) / sqrt(d_k)
3m^(0) = max_row(S^(0))
4P^(0) = exp(S^(0) - m^(0))
5l^(0) = sum_cols(P^(0))
6O^(0) = P^(0) * V^(0)
7
8Step 2: Inductive Ring Step Update (k >= 1)
9S^(k) = (Q_r * (K^(k))^T) / sqrt(d_k)
10m_tilde^(k) = max_row(S^(k))
11m^(k) = max(m^(k-1), m_tilde^(k))
12alpha = exp(m^(k-1) - m^(k))
13beta = exp(m_tilde^(k) - m^(k))
14l^(k) = l^(k-1) * alpha + sum_cols(exp(S^(k) - m^(k)))
15O^(k) = diag(alpha) * O^(k-1) + exp(S^(k) - m^(k)) * V^(k)
16
17Step 3: Terminal Normalization (k = P - 1)
18O_final = diag(1 / l^(P-1)) * O^(P-1)

This distributed formulation is numerically identical to computing full attention over all N tokens simultaneously, while requiring only O(B) memory per GPU.

python(39 lines)
1# Pure PyTorch Reference: Ring Attention Online Softmax Accumulation
2import torch
3import torch.nn.functional as F
4
5def ring_attention_step(
6 Q: torch.Tensor, # [B, d]
7 K_block: torch.Tensor, # [B, d]
8 V_block: torch.Tensor, # [B, d]
9 prev_m: torch.Tensor, # [B, 1] running max
10 prev_l: torch.Tensor, # [B, 1] running sum
11 prev_O: torch.Tensor, # [B, d] running output
12 scale: float,
13 causal_mask: torch.Tensor = None
14) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
15 # 1. Compute local tile scores
16 S = torch.matmul(Q, K_block.transpose(-1, -2)) * scale # [B, B]
17 if causal_mask is not None:
18 S = S + causal_mask
19
20 # 2. Local block statistics
21 curr_m = torch.max(S, dim=-1, keepdim=True).values # [B, 1]
22
23 # 3. New combined maximum
24 new_m = torch.maximum(prev_m, curr_m)
25
26 # 4. Rescaling factors
27 alpha = torch.exp(prev_m - new_m)
28 beta = torch.exp(curr_m - new_m)
29
30 # 5. Local exponentiated probabilities with updated max
31 P = torch.exp(S - new_m)
32 curr_l = torch.sum(P, dim=-1, keepdim=True)
33
34 # 6. Online update of denominator and output
35 new_l = prev_l * alpha + curr_l
36 new_O = prev_O * alpha + torch.matmul(P, V_block)
37
38 return new_m, new_l, new_O
19 lines hidden

Causal Masking Optimization in Ring Attention

In autoregressive decoder models, token i can only attend to tokens j <= i. When sequence N is divided into P blocks along the ring:

  • If block K_j contains indices strictly greater than Query block Q_i (j > i), the entire block computation is skipped (zero communication and zero FLOPs).
  • If j < i, all tokens in Q_i attend to all tokens in K_j (full dense GEMM without masking).
  • If j = i (diagonal block), a lower-triangular causal mask is applied.

By organizing the ring communication to prioritize non-masked blocks or employing a striped sequence layout (where token indices are interleaved across GPUs modulo P), compute loads remain balanced across all ranks throughout causal decoding.


Rotary Position Embedding (RoPE) Extrapolation & Frequency Scaling

Even when Ring Attention overcomes physical memory and communication limits, attention mechanisms fail mathematically if positional representations degrade. Transformers trained with Rotary Position Embeddings (RoPE) exhibit catastrophic perplexity spikes when evaluated on sequences beyond their pre-training length L_train (e.g., evaluating an 8k-trained model on 128k tokens).

Mathematical Definition of RoPE

Given a 2D component of a query or key vector at sequence position m, RoPE applies an orthogonal rotation in the complex plane:

text(13 lines)
1RoPE Orthogonal Rotation:
2R_(Theta, m)^d * x = [ cos(m * theta_i) -sin(m * theta_i) ] [ x_0 ]
3 [ sin(m * theta_i) cos(m * theta_i) ] [ x_1 ]
4
5Rotary Base Frequency:
6theta_i = b^(- / d), where base b = 10,000
7
8Relative Position Invariance:
9<R_(Theta, m)^d * q, R_(Theta, n)^d * k> = q^T * R_(Theta, n - m)^d * k
10
11Wavelength per Dimension:
12lambda_i = 2 * pi / theta_i = 2 * pi * b^( / d)
  • High-frequency dimensions (low index i): Wavelength lambda_i < L_train. A single rotation completes within a few tokens. These dimensions capture local syntactic structure.
  • Low-frequency dimensions (high index i): Wavelength lambda_i >> L_train. The phase angle m * theta_i rotates only a fraction of a circle during pre-training.

The Breakdown of Naive Extrapolation

When evaluating at position m > L_train:

  1. Low-frequency dimensions encounter phase angles m * theta_i never observed during training, forcing the attention layer into out-of-distribution rotational states.
  2. The attention logits shift distributionally, causing softmax entropy to collapse or diverge, leading to generation loops or incoherent outputs.
text(5 lines)
1RoPE Frequency Spectrum across Head Dimensions (d=128, b=10,000):
2Dim 0-1: theta = 1.0000, wavelength = 6.28 tokens --> Rapid local oscillation
3Dim 32-33: theta = 0.0100, wavelength = 628 tokens --> Intermediate phrase scope
4Dim 62-63: theta = 0.0001, wavelength = 62,831 tokens --> Global document scope

Extrapolation Strategies: Linear, NTK-Aware, and YaRN

text(19 lines)
1Comparison of RoPE Extension Methods:
2
31. Linear Position Interpolation (PI):
4 m' = m / s
5 Uniformly compresses all frequencies by context ratio s = L_test / L_train.
6 Problem: Destroys high-frequency local resolution, degrading fine-grained syntax.
7
82. NTK-Aware Interpolation:
9 b' = b * s^(d / (d - 2))
10 Applies Neural Tangent Kernel scaling to base b. High frequencies are
11 barely scaled (preserving local structure), while low frequencies are scaled aggressively.
12
133. YaRN (Yet another RoPE extensioN):
14 Partitions dimensions into three bands based on ratio r = L_train / lambda_i:
15 - r > beta: Pure extrapolation (unscaled)
16 - r < alpha: Pure linear interpolation (scaled by s)
17 - alpha <= r <= beta: Smooth transition via ramp function gamma(r)
18 Includes attention temperature factor t = 1 + 0.1 * ln(s)

The YaRN Mathematical Formulation

YaRN computes a dimension-dependent interpolation factor gamma_i in [0, 1] based on the ratio of training context length L_train to wavelength lambda_i:

text(15 lines)
1YaRN Wavelength Ratio:
2r_i = L_train / lambda_i = (L_train * theta_i) / (2 * pi)
3
4Dimension Interpolation Weight gamma_i:
5gamma_i = 0, if r_i < alpha (interpolate)
6gamma_i = 1, if r_i > beta (extrapolate)
7gamma_i = (r_i - alpha) / (beta - alpha), if alpha <= r_i <= beta (smooth ramp)
8
9Blended Frequency:
10theta_i' = (1 - gamma_i) * (theta_i / s) + gamma_i * theta_i
11
12Attention Softmax Temperature Scaling:
13Attention(Q, K, V) = softmax((Q * K^T) / (t * sqrt(d_k))) * V
14where sqrt(t) ≈ sqrt(1 + 0.07 * ln(s))

For a 16x context extension (s = 16), sqrt(t) ≈ sqrt(1 + 0.07 * 2.77) ≈ 1.093, perfectly stabilizing attention distribution entropy.


Quantized KV-Cache Memory Architecture: FP16 vs. FP8 vs. FP4

Extending context to 1,000,000 tokens on production inference engines (vLLM, TensorRT-LLM) requires rigorous memory budgeting. The aggregate memory required for storing Key-Value activations during autoregressive generation is calculated deterministically.

text(3 lines)
1KV Cache Memory Formula:
2Memory_KV = 2 * L * n_kv * d_head * N_tokens * bytes_per_element

Where:

  • 2: Two separate tensors (Key and Value).
  • L: Number of transformer layers.
  • n_kv: Number of Key-Value attention heads (under Grouped Query Attention / GQA).
  • d_head: Dimension per attention head (typically 128).
  • N_tokens: Active context length (prompt tokens + generated tokens).
  • bytes_per_element: Precision format size (2 for FP16/BF16, 1 for FP8, 0.5 for FP4).

Model Memory Footprint Matrix (KV Cache Only)

| Model Architecture | Layers (L) | KV Heads (n_kv) | Head Dim (d_h) | Bytes/Token (FP16) | Bytes/Token (FP8 E4M3) | Bytes/Token (FP4 NVFP4) | 128k Tokens (FP16 / FP8 / FP4) | 1M Tokens (FP16 / FP8 / FP4) | |---|---|---|---|---|---|---|---|---| | Llama-3-8B | 32 | 8 | 128 | 131,072 B (128 KB) | 65,536 B (64 KB) | 32,768 B (32 KB) | 16.0 GB / 8.0 GB / 4.0 GB | 128 GB / 64 GB / 32 GB | | Llama-3-70B | 80 | 8 | 128 | 327,680 B (320 KB) | 163,840 B (160 KB) | 81,920 B (80 KB) | 40.0 GB / 20.0 GB / 10.0 GB | 320 GB / 160 GB / 80 GB | | Llama-3-405B | 126 | 16 | 128 | 1,032,192 B (1,008 KB) | 516,096 B (504 KB) | 258,048 B (252 KB) | 126.0 GB / 63.0 GB / 31.5 GB | 1,008 GB / 504 GB / 252 GB |

FP8 Precision Formats: E4M3 vs. E5M2

When quantizing KV caches to 8 bits, inference engines select between two IEEE FP8 formats:

text(14 lines)
1FP8 Format Comparison:
21. FP8 E4M3 (1 sign bit, 4 exponent bits, 3 mantissa bits):
3 - Dynamic Range: ~[-448, 448]
4 - Precision: Higher mantissa resolution (3 bits vs 2 bits)
5 - Primary Use: GEMM activations and Key tensors (sensitive to fine-grained phase differences)
6
72. FP8 E5M2 (1 sign bit, 5 exponent bits, 2 mantissa bits):
8 - Dynamic Range: ~[-57344, 57344] (Identical dynamic range to FP16)
9 - Precision: Lower mantissa resolution (coarser quantization steps)
10 - Primary Use: Value tensors or layers with extreme logit variance
11
12Hardware GEMM Execution:
13S = ((Q * s_Q) * (K_fp8 * s_K)^T) / sqrt(d_k) = (s_Q * s_K) * (Q * K_fp8^T) / sqrt(d_k)

In modern inference runtimes (e.g. vLLM with FlashAttention-3 or CUTLASS FP8 kernels), FP8 E4M3 with per-tensor or per-channel scaling is standard. By factoring the scalar dequantization coefficients (s_Q * s_K) out of the matrix multiplication, Tensor Cores execute FP8 GEMMs at up to 1,979 TFLOPS on H100 (double the speed of FP16), simultaneously halving HBM bandwidth pressure.


Decisions

| Decision | Required evidence | Review trigger | |---|---|---| | Deploy Ring Attention combined with TP=8 for million-token inference. | Benchmark showing computation completely overlaps P2P communication when B >= 2,048 | Scaling single-session context past 128k tokens | | Use YaRN dynamic frequency scaling over linear position interpolation. | Zero perplexity degradation on needle-in-a-haystack retrieval evaluations | Context perplexity degradation on long-context benchmarks | | Quantize KV cache to FP8 E4M3 with per-channel scaling. | VRAM savings exceeding 50% with under 0.05 perplexity shift | GPU VRAM utilization exceeding 85% during decode | | Double-buffer asynchronous CUDA streams for P2P KV transfers. | Profiling showing under 3% interconnect stall during decode steps | Interconnect stall telemetry exceeding 3% of step time |


Alternatives and trade-offs

Megatron-LM Sequence Parallelism and DeepSpeed Ulysses rely on collective All-Gather or All-to-All communication, which creates global synchronization bottlenecks across large GPU rings. Ring Attention replaces collective communication with non-blocking circular P2P transfers, enabling near-linear context scaling up to millions of tokens over InfiniBand or RoCE v2 interconnects at the cost of slightly higher ring scheduling complexity.


Failure modes

  • NCCL Circular Ring Deadlocks: If a single GPU in the ring encounters a transient CUDA exception during non-blocking ncclSend/ncclRecv, upstream and downstream ranks block indefinitely on ring barriers. Mitigate with strict watchdog timeouts (NCCL_COMM_BLOCKING=0, NCCL_ASYNC_ERROR_HANDLING=1).
  • Numerical Underflow in Online Softmax Exponentiation: When running max differences (m^(k-1) - m^(k)) exceed -88.7 in single precision, exponentiation underflows to zero, producing NaN activations. Mitigate by accumulating all statistics in FP32 registers with epsilon floor clamping.
  • RoPE High-Frequency Phase Wrapping (Aliasing): Applying naive linear position interpolation across large scaling factors compresses high-frequency dimensions so severely that adjacent tokens lose positional discrimination. Mitigate by gating with YaRN frequency boundaries (alpha = 1, beta = 32).
  • FP8 KV-Cache Outlier Magnitude Saturation: Deep transformer layers develop persistent outlier activation channels. Coarse per-tensor quantization destroys dynamic range. Mitigate with per-head or per-channel block-scaled quantization (tile size 128).

Operational checklist

  • [ ] Cluster Interconnect Verification: Confirm bidirectional peer-to-peer bandwidth between all adjacent nodes in the ring meets minimum line rate (at least 400 Gbps over RoCE v2 / InfiniBand).
  • [ ] Block Size Floor Verification: Ensure Ring Attention block size B = N / P satisfies B >= 2,048 tokens to maintain compute-communication overlap.
  • [ ] Precision Alignment: Set Key-Value cache dtype to fp8_e4m3 with per-channel scaling enabled in the model runner configuration.
  • [ ] YaRN Temperature Calibration: Validate attention logit scale factor matching target context length.
  • [ ] Causal Ring Pruning: Confirm that causal mask triangular skipping is enabled in the attention engine, eliminating redundant upper-diagonal block calculations.
  • [ ] Telemetry Instrumentation: Configure Prometheus alerts for nccl_p2p_stall_ratio (alert if above 0.05) and kv_cache_fragmentation_ratio (alert if above 0.15).

Connected practice


Sources

  • ring-attention-liu-2023
  • roformer-su-2021