Trust is earned, not given

A different perspective

2024-04-04 · Projects

AI Frontiers, part 38: Training infrastructure — FlashAttention, FSDP, and the plumbing of scale

Part 38from the AI Frontiers series · 65 parts in all

The interesting thing about a frontier training run is that almost none of it is machine learning. The model definition is a few hundred lines. The optimizer is a known algorithm. What takes months of engineering is the question of how to keep ten thousand accelerators busy for weeks without losing a run, and the answers are borrowed wholesale from distributed systems: sharding, replication, checkpointing, scheduling, and failure recovery.

This entry is about that layer, because it is where the gap between a good idea and a trained model actually lives. It is also the layer that explains why the efficiency results in part 13 were an engineering achievement rather than an algorithmic one — a point that is easy to miss when the headline is a dollar figure.

Four kinds of parallelism, each solving a different problem

There are four standard ways to split a training job, and they compose because they address different bottlenecks.

Data parallelism replicates the model and splits the batch. Each worker computes gradients on different data, and the gradients are averaged. It is the simplest form and it works until the model no longer fits in one device's memory, which for modern models happens at a scale far below the frontier. Its refined versions are the ones actually used: ZeRO (Rajbhandari et al.) partitions optimizer state, gradients and parameters across workers so that no single device holds a full copy, and PyTorch's fully sharded data parallel (Zhao et al.) implements the same idea inside a mainstream framework. The cost is communication — sharded parameters must be gathered before each layer's forward pass and released afterwards — and the benefit is that memory per device scales inversely with the number of devices.

Tensor parallelism splits individual matrix multiplications across devices, so that a single layer's weights are distributed and the partial results combined. It is the most communication-intensive form, which is why it is confined within a node where the interconnect is fastest, usually over NVLink rather than the network. Megatron-LM (Shoeybi et al.) established the standard approach and the follow-up work (Narayanan et al.) established how to combine it with the rest.

Pipeline parallelism splits the model by layer into stages on different devices, with micro-batches flowing through the stages like an assembly line. GPipe (Huang et al.) and PipeDream (Narayanan et al., 2018) are the canonical formulations, and the engineering problem is bubble minimization: keeping every stage busy despite the sequential dependency. The reason it matters is that it is the only one of the four that reduces memory without proportionally increasing communication.

Expert parallelism is the newest and the one that sparse models require. When a mixture-of-experts layer routes each token to a subset of experts, the natural placement is to put different experts on different devices, which means every token has to be sent to wherever its experts live and the result sent back. GShard (Lepikhin et al.) and the Switch Transformer work established the pattern; MegaBlocks (Gale et al.) and DeepSpeed-MoE (Rajbhandari et al.) solved the systems problem of batching variable numbers of tokens per expert efficiently. This is the mechanism that made the sparse architectures of part 8 trainable at all.

The practical insight is that these are not alternatives. A real run uses all four, arranged so that the highest-bandwidth communication stays within a node and the lowest-bandwidth crosses the network, because the topology of the cluster determines which splits are affordable.

Attention, and the memory that started the whole discipline

The single largest software win in this period came from noticing that the attention kernel was wasting memory. Computing attention naively materializes an n-by-n score matrix, which for a long sequence is enormous and is written to and read from slow memory repeatedly. FlashAttention (Dao et al.) restructured the computation into tiles that fit in fast on-chip memory, computed the softmax incrementally with a running normalization, and avoided ever storing the full matrix.

The consequence was not just a constant-factor speedup. It changed what was feasible: longer sequences became trainable within a fixed memory budget, which is the precondition for the context-length progression in this series. The follow-up version (Dao) refined the work partitioning and parallelism, and the kernel is now standard in every serious framework.

The same trade recurs across the stack: when memory is the binding constraint, spend compute to save memory. Activation recomputation (Chen et al.) is the canonical instance — discard intermediate activations during the forward pass and recompute them during the backward pass, trading roughly a third more compute for a large reduction in memory. Sequence parallelism (Korthikanti et al.) partitions the activations along the sequence dimension rather than replicating them across tensor-parallel workers, which removes a redundancy that had been hiding in plain sight. Every one of these is the same decision applied at a different granularity.

Numerics, which decide whether a run succeeds

Precision choices look like a detail and behave like a constraint. Training in 16-bit precision requires loss scaling to keep gradients from underflowing, which is why the mixed-precision recipe (Micikevicius et al.) became standard practice. BF16 replaced FP16 for most purposes because its wider exponent range removes the need for loss scaling in the common case, at the cost of mantissa precision that rarely matters for training.

By 2024 the frontier had moved to FP8 for parts of the computation, which is a substantially harder engineering problem than FP16 because the dynamic range is narrow enough that scaling has to be managed per block, and outlier values have to be handled with the kinds of tricks described in the quantization entry. The efficiency claims of this period rest partly on getting that right, and the papers that report them describe the scaling schemes in unusual detail precisely because the details are the contribution.

The general principle is that numerical precision is a resource like memory and bandwidth: you spend it where it affects the result and economize where it does not. The engineering difficulty is that where it affects the result is not knowable in advance and varies between layers and even between parts of a single layer, so the work is measurement-heavy and rarely elegant.

Failure, which is the defining property of large runs

Here is the fact that separates people who have run large training jobs from people who have read about them: at scale, hardware fails constantly, and the run must be engineered to survive it. The BLOOM report (Scao et al.) documented a multi-week training run interrupted dozens of times by hardware faults, with recovery handled by automatic checkpointing and restart. That is not an anomaly; it is the normal operating condition of a cluster with thousands of accelerators, where the mean time between failures across the whole machine is measured in hours.

Three engineering consequences follow, and they explain a great deal about how these systems are built. First, checkpoints must be frequent and cheap, which itself is a distributed problem — you are writing terabytes from thousands of devices simultaneously, and the storage system has to keep up without stalling the compute. Second, the training loop must be deterministic enough to resume correctly, which means the ordering of operations and the seeding of randomness have to be reproducible across a different set of workers than the ones that were running when the failure occurred. Third, elasticity is a live design question: whether a run can continue with fewer devices while a node is repaired is the difference between losing hours and losing a week.

There is a quieter failure mode that matters just as much, which is silent data corruption. A device that returns wrong results without erroring will poison a training run in a way that shows up weeks later as an inexplicable loss anomaly. Detection requires redundancy or periodic validation, and the fact that serious operations invest in it is a good indicator of how much of this work is ordinary reliability engineering rather than anything specific to language models.

Data pipelines, which starve more runs than hardware does

The parallelism machinery gets the attention because it is where the difficulty is visible. In practice, a surprising share of lost throughput comes from the input pipeline: reading tokenized training data fast enough to feed thousands of accelerators without stalling.

The arithmetic is unforgiving. A run at full utilization consumes tokens at a rate determined by the compute budget, and the storage system has to deliver them with enough shuffle randomness that successive batches are not correlated, without the preprocessing becoming the bottleneck. Streaming dataset formats that read sequentially from object storage while shuffling locally exist precisely because the naive approach — random access over a network filesystem — collapses under this load.

Tokenization is usually done offline, for the same reason: converting text to token ids is cheap but not free, and doing it once against the whole corpus removes a per-epoch cost. The tradeoff is that changing the tokenizer invalidates the entire preprocessed corpus, which is a decision that has to be made early and lives for the duration of the project.

The general lesson is that a training cluster is a system with several potential bottlenecks, and whichever one is worst determines throughput. Optimizing the kernel when the data loader is the constraint produces no improvement and a great deal of confusion.

How to tell whether a run is healthy

Because a training run is long, expensive and hard to restart, the discipline that matters most is detecting problems early rather than diagnosing them late.

The primary metric is hardware utilization — the fraction of theoretical peak floating-point throughput actually achieved, usually reported as model FLOPs utilization. A well-tuned run sits in a recognizable band, and a run that drops out of it has a communication, memory or data problem that is worth investigating before more money is spent. Watching that number continuously is the equivalent of watching error rates in a service.

The second practice is instrumenting the quantities that predict divergence: gradient norm, activation magnitudes, attention entropy, and the distribution of routing decisions in sparse layers. Loss spikes are a symptom; by the time one is visible, the cause happened hours earlier, and the only way to find it is to have been recording what changed.

The third is the one that saves the most money. Before committing to a large run, do small ones. Train a series of models at small scale with the same data and recipe, verify that the loss curve behaves as the scaling relationship predicts, and only then scale up. This is the standard practice in serious organizations and it is the reason a failed run is rare relative to a failed project: the failure shows up at the cheap scale first, where it can be fixed.

What this layer decides for everyone else

The reason to care about any of this, if you are not training a frontier model, is that the same techniques determine the cost of everything downstream.

Effective training throughput — the fraction of peak hardware performance actually achieved — is the number that converts a compute budget into a model. A run at 30 percent efficiency costs three times what it should, and most of the difference is in how well the parallelism strategy matches the cluster topology, how much recomputation was traded for memory, and how much time was lost to failures and restarts. Two organizations with the same accelerators and the same budget can produce different results purely on this axis, and the difference compounds across a training campaign.

The second consequence is that infrastructure knowledge is transferable and scarce. The set of people who can debug a hang in a pipeline-parallel training run is small, the debugging tools are immature, and the failure modes are nonlocal in a way that makes them hard to learn from documentation. That is why the same handful of frameworks dominate: PyTorch with FSDP, DeepSpeed, Megatron-LM and their derivatives encode years of accumulated fixes for problems nobody would think to anticipate.

The honest summary is that training infrastructure is the least discussed and most decisive part of the stack. The model architecture gets the paper, the data gets the argument, and the plumbing gets nothing — while quietly determining whether the run finishes, how much it costs, and whether the organization that built it can do it again next quarter.

Works Cited

Chen, Tianqi, et al. "Training Deep Nets with Sublinear Memory Cost." arXiv, 2016, arxiv.org/abs/1604.06174. Accessed 4 Apr. 2024.

Dao, Tri. "FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning." arXiv, 2023, arxiv.org/abs/2307.08691. Accessed 4 Apr. 2024.

Dao, Tri, et al. "FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness." arXiv, 2022, arxiv.org/abs/2205.14135. Accessed 4 Apr. 2024.

Gale, Trevor, et al. "MegaBlocks: Efficient Sparse Training with Mixture-of-Experts." arXiv, 2022, arxiv.org/abs/2211.15841. Accessed 4 Apr. 2024.

Huang, Yanping, et al. "GPipe: Efficient Training of Giant Neural Networks Using Pipeline Parallelism." arXiv, 2018, arxiv.org/abs/1811.06965. Accessed 4 Apr. 2024.

Korthikanti, Vijay Anand, et al. "Reducing Activation Recomputation in Large Transformer Models." arXiv, 2022, arxiv.org/abs/2205.05198. Accessed 4 Apr. 2024.

Lepikhin, Dmitry, et al. "GShard: Scaling Giant Models with Conditional Computation and Automatic Sharding." arXiv, 2020, arxiv.org/abs/2006.16668. Accessed 4 Apr. 2024.

Micikevicius, Paulius, et al. "Mixed Precision Training." arXiv, 2017, arxiv.org/abs/1710.03740. Accessed 4 Apr. 2024.

Narayanan, Deepak, et al. "Efficient Large-Scale Language Model Training on GPU Clusters Using Megatron-LM." arXiv, 2021, arxiv.org/abs/2104.04473. Accessed 4 Apr. 2024.

---. "PipeDream: Generalized Pipeline Parallelism for DNN Training." arXiv, 2018, arxiv.org/abs/1806.03377. Accessed 4 Apr. 2024.

Rajbhandari, Samyam, et al. "DeepSpeed-MoE: Advancing Mixture-of-Experts Inference and Training to Power Next-Generation AI Scale." arXiv, 2022, arxiv.org/abs/2201.05596. Accessed 4 Apr. 2024.

---. "ZeRO: Memory Optimizations Toward Training Trillion Parameter Models." arXiv, 2019, arxiv.org/abs/1910.02054. Accessed 4 Apr. 2024.

Scao, Teven Le, et al. "BLOOM: A 176B-Parameter Open-Access Multilingual Language Model." arXiv, 2022, arxiv.org/abs/2211.05100. Accessed 4 Apr. 2024.

Shoeybi, Mohammad, et al. "Megatron-LM: Training Multi-Billion Parameter Language Models Using Model Parallelism." arXiv, 2019, arxiv.org/abs/1909.08053. Accessed 4 Apr. 2024.

Zhao, Yanli, et al. "PyTorch FSDP: Experiences on Scaling Fully Sharded Data Parallel." arXiv, 2023, arxiv.org/abs/2304.11277. Accessed 4 Apr. 2024.