系列:Inference Systems Infrastructure Notes

Inference Systems Infrastructure Notes (2): Sequence Parallelism + Ring Attention — Making 1M+ Token Training Routine

0. One-Line Positioning

Sequence Parallelism shards a long sequence across GPUs along the token dimension; Ring Attention lets those GPUs compute a result mathematically identical to full attention by passing K/V blocks around a ring and accumulating with online softmax — with per-GPU activation memory independent of sequence length.

This is the key technique for breaking the single-GPU sequence-length ceiling, and it turns 1M+ token training and inference into routine work.

1. Landscape at a Glance

2. Core Method: Online Softmax + Ring K/V Communication

Assume an 8-GPU ring, with the sequence split into 8 token segments of length L/N. Each GPU holds its own query Q_i (which never moves) while K/V blocks flow around the ring:

m_i, l_i, o_i = -inf, 0, 0     # running max / running sum / output

for step in range(N):           # N steps around the ring, one K/V block per step
    K_block, V_block = recv_from_prev()   # from the previous GPU
    send_to_next(my_KV)                   # simultaneously send our own block onward

    s_ij = Q_i @ K_block^T / sqrt(d)               # (L/N, B)
    m_new = max(m_i, s_ij.max(-1))                 # update running max
    p_ij  = exp(s_ij - m_new[:, None])
    l_new = exp(m_i - m_new) * l_i + p_ij.sum(-1)
    o_i   = exp(m_i - m_new)[:, None] * o_i + p_ij @ V_block
    m_i, l_i = m_new, l_new

O_i = o_i / l_i[:, None]    # normalize; result identical to full attention

Notation: Q_i is the local query on GPU i; K_block, V_block are the K/V blocks traveling around the ring; m_i, l_i, o_i are the online softmax running max, running sum, and partial output.

3. Ring Communication Structure

Ring Attention: ring-passing K/V blocks plus online softmax accumulation GPU 0 Q0 (local, fixed) K/V block j GPU 1 Q1 (local, fixed) K/V block j GPU 2 Q2 (local, fixed) K/V block j GPU 3 Q3 (local, fixed) K/V block j K/V blocks flow around the ring (one block per step) GPU3 closes the ring back to GPU0 Online softmax accumulation (each step): s_ij = Q_i * K_block^T / sqrt(d); m_new = max(m, s.max); p = exp(s - m_new) l_new = exp(m - m_new) * l + p.sum; o_new = exp(m - m_new) * o + p * V_block; after N steps O_i = o_i / l_i Each GPU caches only its own Q plus one incoming K/V block, so activation memory is independent of total sequence length: 8 GPUs run 1M tokens, and in theory 100 GPUs reach 100M.

Figure: Ring Attention communication. The sequence is sharded across 8 GPUs by token; K/V blocks travel one step at a time around the ring, and each GPU accumulates the incoming blocks with online softmax. After N steps the result is mathematically equivalent to full attention.

4. Comparison with Naive Approaches

ApproachActivation memoryCan it run 1M tokens?Notes
Naive (full Q x K^T)O(L²)No, a single GPU always blows up1M tokens overflows one card outright
Megatron-LM SP (all-gather)Temporary full KVMarginal, tens of thousands of tokensEach GPU still needs the full KV temporarily inside attention
Ring Attention (this approach)Independent of LYes, 1M on 8 GPUs, 100M on 100Communication hides under matmul, wall clock barely grows

Communication overhead is hidden by the matmul (each hop’s transfer covers the previous matmul), so wall-clock latency barely increases.

5. Measured Results and Representative Work

6. Connection to Embodied Intelligence

  1. Video world models (V-JEPA 2 / Genie) need 1M+ token video clips as input; Ring Attention is standard in multi-node V-JEPA 2 training.
  2. Robot VLA scenarios with long operation logs plus video history planning: SP + Ring makes “compute the entire operation log at once” feasible.
  3. Synergy with HCA / CSA — HCA compresses KV one more notch, so what travels around the ring is a compressed “directory block”: 1M tokens become 8000 directory blocks, cutting bandwidth pressure by another 128x (see sys1 in this series on the NUMA path, and fa4 on HCA).

7. Study Tips and Common Pitfalls

  1. You must use online softmax with a running max, or accumulated error skews the result;
  2. Use Zig-Zag blocking under a causal mask, or later GPUs idle;
  3. Communication bandwidth is the real bottleneck — NVLink or InfiniBand is not optional (on plain Ethernet, once transfer time exceeds compute time, “compute while transferring” degrades into “waiting for data”);
  4. When mixing with tensor parallelism, keep the TP and SP communication directions orthogonal, or all-gather and ring traffic collide.

觉得有用?欢迎点赞、收藏,或请作者喝咖啡 ☕️

支付宝收款码

支付宝

微信收款码

微信

💬 留言

评论由 Giscus 驱动(基于 GitHub Discussions)。 当前仓库 NaphJohn/LLM-blog 尚未启用 Discussions:请在 GitHub 仓库 Settings → General → Features 勾选 Discussions 后刷新本页,评论区即自动显示。