What do you do about long sequences and sparse models?

Three axes in and the model fits, the batch is right and the pipeline is full. Then someone says the context window is going to 128k, or the model is a mixture of experts, and two more axes appear.

These are the ones people reach for last, and they are the ones an interviewer probes when a candidate says “we used 3D parallelism”. Both divide something the first three do not touch. Both are load-balance problems before they are bandwidth problems. And both have a specific detail that separates people who have run them from people who have read the paper.

They will ask How do you train on 128k-token sequences, and how do you train a mixture of experts?
The 30-second version Context parallelism splits the sequence itself. Activations are linear in sequence length, so 128k tokens is sixteen times the activation memory of 8k, and splitting the sequence sixteen ways puts it back. Attention is the only thing that couples the sequence, so the keys and values have to travel: either around a ring, overlapped with the block computation, or by all-gathering them, which is what Meta chose for Llama 3 because grouped-query attention makes K and V small and it composes with any mask. The detail is load balance: a causal mask makes the work triangular, so contiguous chunks mean the last rank does c times the work of the first, and the fix is to cut into 2c chunks and pair chunk i with chunk 2c minus 1 minus i. Expert parallelism splits a mixture of experts across GPUs, so tokens have to travel to their experts and back: two all-to-alls per layer in the forward pass and two more in the backward, sized by tokens times top-k times hidden. The detail is that routing is data-dependent, so the step waits for whichever GPU drew the most tokens, and balance has to be enforced with a capacity factor, an auxiliary loss, or a routing bias updated every step.

Why does the sequence need its own axis?

Because the sequence is a multiplier on every activation in the inventory, and none of the first three axes divides it.

Megatron’s per-layer activation count, $34\,s\,b\,h$, is linear in $s$. Push Llama 3 70B from 8192 tokens to 131,072 and the saved activations go from 22.8 GB per GPU to 365.1 GB, with tensor and sequence parallelism already applied.

Attention’s own arithmetic is worse than linear. It adds roughly $12\,L\,h\,s$ FLOPs per token, quadratic in the sequence, which at 128k tokens is 1031 GFLOP per token against the 423 GFLOP that the six-FLOPs-per-parameter rule accounts for. At long context, attention costs more than the rest of the model put together.

Context parallelism cuts the sequence into $c$ pieces and gives each GPU one of them. It divides activations and nothing else: no weight, no gradient, no optimizer state. That makes it the axis you reach for when the sequence, rather than the model, is what does not fit.

How does attention survive being split?

Everything except attention is pointwise in the sequence, so splitting it is free. Attention is not: every query needs every key up to it.

Ring attention’s answer is to move the keys and values instead of the queries. Each rank keeps its own block of queries and never moves it. Rank $i$ sends its key-value block to rank $i+1$ while computing attention against the block it already holds, and after $c-1$ hops every query block has seen every key block.

The partial results merge because of the online softmax. Attention over one block of keys, with its running maximum and running sum kept alongside, can be combined with attention over another block exactly, which is the same identity FlashAttention uses to avoid writing the score matrix. So ring attention is exact, not an approximation, and it is worth saying that out loud because people assume otherwise.

The memory on each rank is one query block, the key-value block it is using, and the block it is receiving. Nothing grows with the total sequence length.

When is the ring actually free?

When the arithmetic on a block outlasts the transfer of the next one. Liu, Zaharia and Abbeel write the condition down in one line: with $F$ FLOP/s per host and $B$ bytes per second between hosts, a block of $c$ tokens needs

\[\frac{4dc^2}{F} \ge \frac{4cd}{B} \quad \Longrightarrow \quad c \ge \frac{F}{B}\]

The block size has to exceed the machine’s FLOP-per-byte ratio. On an H100 over NVLink that is 989 TFLOP/s divided by 450 GB/s, so 2198 tokens per block, and the paper’s rule of thumb of six blocks per sequence puts the minimum sequence per rank at about 13,000.

Over a 400 Gb/s network port it is 19,780 tokens per block and about 119,000 tokens per rank. That second number is the practical story of ring attention: across a network it only pays at six-figure sequence lengths per device.

Why did Llama 3 not use a ring?

Meta say so directly, and the reasoning is a good example of choosing the simpler thing.

Their context parallelism all-gathers the key and value tensors first, then computes attention for the local query chunk. The all-gather latency is exposed on the critical path, which a ring would have hidden. They took that for two reasons.

It supports any attention mask, including the document mask they use to stop attention crossing document boundaries inside a packed sequence, which is awkward to express in a ring.

And with grouped-query attention the keys and values are much smaller than the queries. Llama 3 70B has 8 key-value heads against 64 query heads, so the tensors being gathered are an eighth of the size, and attention’s $O(s^2)$ arithmetic dwarfs an $O(s)$ transfer.

That is the pattern worth carrying: an architectural choice made for inference memory, grouped-query attention, turned out to decide which training-time context-parallel algorithm is worth implementing.

What breaks the load balance?

The causal mask. Split a sequence into $c$ contiguous chunks and give chunk $i$ to rank $i$. Rank 0’s queries attend to one chunk of keys. Rank $c-1$’s attend to all $c$.

So rank $i$ does $i+1$ units of work, the last rank does $c$ times the first, and every synchronisation in the layer waits for it. At $c = 8$ that is an eightfold spread and it happens in every layer of every step.

Meta’s fix in Llama 3 is one line of index bookkeeping: cut the sequence into $2c$ chunks and give rank $i$ chunks $i$ and $2c-1-i$. The work becomes $(i+1) + (2c-i) = 2c+1$ for every rank, exactly balanced, with no change to a single result.

Key idea Context parallelism and expert parallelism both look like bandwidth problems and are load-balance problems. One is triangular because of the causal mask and is fixed by pairing chunks; the other is uneven because the router is being trained, and has to be fixed continuously.

Why does a mixture of experts need its own axis?

Because most of the parameters are in places most tokens never visit. The Mixture of Experts series derives that bill from the model’s side. This part is the same layer seen from the cluster’s side: what the routing does to the wire.

DeepSeek-V3 is the reference to have in mind, and its shape is public: 61 layers, hidden size 7168, 256 routed experts plus one shared expert per MoE layer, eight routed experts activated per token, an expert intermediate width of 2048. That is 671 billion total parameters with 37 billion active for any given token.

The experts do not fit on one GPU and there is no reason to replicate them, so they are spread across $e$ GPUs. On DeepSeek’s layout, $e = 64$ across eight nodes, with four experts per GPU.

Now the tokens have to go where their experts are. That is not a reduction or a broadcast: every GPU has a different, data-dependent set of tokens for every other GPU. It is an all-to-all, which is the least friendly pattern a network has to serve.

What do the two all-to-alls move?

Two per MoE layer in the forward pass. The dispatch sends each token’s hidden vector to each of its chosen experts. The experts run. The combine brings the weighted outputs back to the GPU the token came from. The backward pass does both again.

Size them. For 4096 tokens on a GPU, top-8 routing and a hidden size of 7168, dispatching in FP8 and combining in BF16, and with routing limited to four nodes per token, the network sees 117.4 MB out and 234.9 MB back per layer, so 704.6 MB across the forward and backward passes.

Against that, the layer’s arithmetic on that GPU is about 4.97 TFLOP, which is 5.0 ms at peak, while 352 MB over a 50 GB/s InfiniBand port is 7.0 ms. Same order. DeepSeek describe the computation-to-communication ratio of cross-node expert parallelism as roughly one to one, and everything in their infrastructure section follows from that number.

The node limit is what keeps it there. Without it a token could need to reach eight different nodes; capping it at four halves the InfiniBand traffic. DeepSeek go further and exploit the two-tier structure: a token crosses InfiniBand once per target node and is then forwarded over NVLink to the specific experts inside it, so with their quoted 160 GB/s of NVLink against 50 GB/s of InfiniBand, each token reaches an average of 3.2 experts per node for the cost of one node hop.

What goes wrong when the routing is uneven?

Routing is data-dependent, so the token count per expert is a random variable that the model itself controls, and it drifts during training because the router is being trained too.

Two things break. Every GPU finishes its own experts and then waits, so the layer costs whatever the busiest rank costs. A modest 1.4-times imbalance throws away 30% of the expert FLOPs in that layer, in every layer, in every step. And if you cap the tokens per expert to keep the communication buffers a fixed size, the ones over the cap are dropped and their contribution to the layer is lost.

The fixes come in three flavours. A capacity factor above one reserves headroom and accepts some dropping. An auxiliary load-balancing loss adds a term that punishes uneven routing, which works and costs a little quality because it is a gradient pulling against the language-modelling objective.

DeepSeek-V3’s approach avoids that trade: they add a bias to each expert’s routing score and nudge it after every step, up for underloaded experts and down for overloaded ones, with no gradient term at all. The routing decision changes; the loss does not learn about balance.

How do the two axes compare?

Side by side, because they get confused.

Context parallelism divides the activations of the sequence and leaves weights, gradients and optimizer states alone. Its collective is the keys and values moving per layer, either around a ring or as an all-gather. It scales with sequence length. Its load-balance problem is the causal triangle, and it is fixed once, statically.

Expert parallelism divides the expert weights and their states and leaves attention, the norms and the router alone. Its collectives are two all-to-alls per layer. It scales with tokens times top-$k$ times hidden size. Its load-balance problem is the router, and it has to be managed continuously.

Neither is a first resort. Context parallelism exists because attention couples the whole sequence and the sequence got long. Expert parallelism exists because sparse models put most of their parameters somewhere most tokens never go.

They will ask Your MoE training run is at half the throughput you projected and the profiler shows the all-to-alls exposed. What is your first move?
The 30-second version Log the token counts per expert before touching anything, because an exposed all-to-all and an unbalanced router look identical in a summary and only one of them is a communication problem. If the counts are skewed, the fix is balance: a routing bias updated each step, or an auxiliary loss, and a capacity factor that is not silently dropping a large fraction of tokens. If the counts are even, then it really is bandwidth, and the levers are the number of nodes a token may reach, dispatching in FP8 rather than BF16, and whether the kernels exploit the fact that intra-node NVLink is several times the inter-node link so a token can cross the network once per node instead of once per expert. Then the scheduling question: an all-to-all this size will never be free, so it has to hide behind something, which is the whole reason DeepSeek built a bidirectional pipeline that overlaps one micro-batch's communication with another's computation.

Rapid fire: can you do these from memory?

  1. Say what context parallelism divides and what it leaves alone, and give the activation memory at 8k and 128k tokens.
  2. Explain why ring attention is exact, and name the identity that makes the merge work.
  3. State the overlap condition for ring attention and evaluate it for NVLink and for a 400 Gb/s port.
  4. Give two reasons an all-gather-based context parallelism can beat a ring.
  5. Describe the causal load imbalance and the pairing that fixes it.
  6. Name the two all-to-alls in a mixture-of-experts layer and say what each moves.
  7. What does node-limited routing bound, and why does it work?
  8. Name three ways to enforce expert balance and the cost of each.

Part 6 puts all five axes on one cluster, works out which one gets which link, and checks the answer against recipes that have actually been run.