Jordan Sassoon
Scaling Pre-Training in Practice: A Hierarchical Approach
Taking a 30B-A3B MoE model from 16 to 512 B200 GPUs while achieving near-linear scaling and 35.3% MFU.
In a GPU-hungry world, one must make use of their compute as best they can. It is tempting to speed up model training by simply scaling up and “buying more GPUs”, crunching more data in less time. However, naively tossing more GPUs in the mix is not a solution: training efficiency quickly degrades if the software is not carefully tuned. And frontier labs care deeply about training efficiency. It’s no secret they invest substantial human- and agent-hours into optimising their training stack, fiending for a couple of percentage-point improvements in their efficiency metrics. There are many efficiency-related hyperparameters to tune, and a naive grid search at scale is prohibitively expensive, so how do you narrow down the search space? While the final training configurations are (at times) disclosed, the choices that lead to them rarely are. So we’re pulling back the curtain.
In this blog post, we will break down how we, in practice, efficiently scale up and pre-train a 30B-A3B1 MoE2 model on 512 NVIDIA B200 GPUs. We use a hierarchical approach to find our training setup, which entails shrinking the search space before increasing the GPU count. With this, we achieve 35.3% MFU3 on 512 GPUs, only 6% below perfect linear scaling – efficient pre-training at scale!
And why should you care about efficiency if you don’t have hundreds of SOTA GPUs? As we will see, many of the tools and ideas we use are scale-agnostic. If you don’t care about saving tonnes of cash4, increasing training efficiency cuts down time, and we only ever have so much of that.
We train with our own optimised fork of torchtitan. The target compute setup consists of 64 nodes connected via InfiniBand, and 8 GPUs per node connected via NVLink. We want to make all these GPUs hum together in perfect synchrony, orchestrating a perfect pre-training symphony.
Note that we focus here on the hierarchical scaling methodology itself. The system- and kernel-level optimisations that make 35.3% MFU possible in the first place are (mostly) out of scope for this post.
Before diving into experiments, we must first understand what training efficiency is and how to measure it – section “MFU is good; TPS is better”. Then, we will look at how to tweak our training setup to achieve maximal efficiency in “Three degrees of freedom to make us compute-bound”. At this point, we can start getting hierarchical. In section “16 GPUs are enough to decimate our search space”, we learn to use small-scale experiments to reduce the hyperparameter space before scaling up. We will use these first experiments to see exactly which operations our model is performing, and if they are efficient, in section “A profile is worth a thousand experiments”. Finally, we’ll take all the learnings and scale up in “From 16 to 512 GPUs, scaling cost is (nearly) free”. But first, metrics!
MFU is good; TPS is better
We use three metrics in our experiments: Tokens Per Second per GPU (TPS), Model FLOPS Utilisation (MFU), and GPU memory usage (mem). TPS is the number of tokens processed by a single GPU divided by wall time:
It is a crude but honest measure of training speed, and it is the one we optimise against. MFU is instead an efficiency metric. It expresses the same measurement as TPS but as a fraction of the hardware’s ceiling – the theoretical limit of how many floating-point operations per second (FLOPS) our GPUs can do. Torchtitan computes it as:
Where is the number of floating-point operations done per token, and is the hardware’s peak FLOPS. There are various caveats with MFU which make training efficiency comparisons between models quite murky. For example, mixed-precision training changes significantly what the peak FLOPS are. We train in BF16 (apart from FP32 gradients, though the vast majority of compute is still BF16), so all MFU figures in this post use BF16 peak FLOPS as the denominator. If you are interested in the exact MFU calculations, see the torchtitan code. The variable is a nasty one to compute, so our north star objective function is TPS. Nevertheless, we always show MFU as an indicator for efficiency, also because it is the industry standard. Finally, memory usage is the ratio of peak reserved memory during training to total device memory:
We use PyTorch’s reserved memory (not allocated memory) since it dictates hard failures – out-of-memory (OOM) errors. Now that our metrics of speed and efficiency have been explained, let’s look at how we can modify our training configuration to achieve blazingly fast pre-training.
Three degrees of freedom to make us compute-bound
We know how to measure efficiency, but what does efficient GPU use look like in practice? A recurring theme in efficiency engineering is the communication-to-computation ratio: how much time the GPU spends performing the mathematical operations needed to run the model (compute kernels), such as calculating attention, versus moving data between GPUs (communication kernels), such as gathering model weights before a forward pass. Ideally, we are always “compute-bound”, meaning that we are limited by how quickly the GPUs can perform these mathematical operations rather than by how quickly data can be moved around. This matters because communication does not directly advance the model’s computation: while the GPU is waiting for data to arrive, its computational resources may sit idle. We want to avoid stalling our compute kernels and throttling our GPUs – “I paid for the whole GPU, I am going to use the whole GPU”. There are many tools we can leverage to tilt the communication-to-computation ratio in our favour. In this blog post, we will take a look at:
- the parallelism scheme
- the activation checkpointing technique
- the local batch size
These are the three degrees of freedom we will explore to optimise the pre-training of our 30B-A3B model. Let’s dive deeper and understand how they change what is computed, what is communicated, how memory usage is affected, and how they interact. We will borrow torchtitan naming throughout.
1 – Parallelism scheme
We base our initial understanding of parallelism schemes on Hugging Face’s Ultra-Scale Playbook, which we highly recommend reading if you haven’t already done so. Roughly speaking, pre-training involves four kinds of model-related data: weights, gradients, optimiser states, and activations. Weights are the model’s trainable parameters. Gradients are partial derivatives of the loss w.r.t. the parameters. They tell us how changing each parameter would affect the loss, and therefore how the weights should be adjusted to reduce it. Optimiser states store information from previous gradients to help determine the size and direction of weight updates. Activations are intermediate results produced during the forward pass that are needed to compute gradients during the backward pass. Parallelism schemes determine how weights, gradients, optimiser states, and activations are distributed across GPUs. Nowadays, multi-trillion parameter models are being trained on tens of thousands of GPUs concurrently. We, however, have a 30B model on 512 GPUs – we can be compute-bound with just a few tricks. It will suffice to look into these two parallelism dimensions:
-
Fully Sharded Data Parallelism (FSDP or ZeRO-3): when a single GPU cannot hold
the full model in memory, the weights, gradients, and optimiser states are split (or “sharded”)
evenly across multiple GPUs (i.e., the FSDP group). Together, the GPUs in an FSDP group make
up one full model replica. For example, if there are 64 GPUs in one group, each GPU holds 1/64th
of the replica. During the forward pass, each GPU collects the model weights with an
all_gatheroperation. During the backward pass, each GPU computes its local gradients. Theall_gatheris once more present to compute gradients of the model weights. In addition, thereduce_scatteroperation reduces (averages) and shards back the gradients between all GPUs in the FSDP group. -
Data Parallelism (DP): this refers to the number of model replicas our GPUs
hold, which is equivalent to the number of independent FSDP groups. DP introduces the
all_reduceoperation, which reduces gradients between FSDP groups.
Each GPU processes a different local batch of sequences, thus creating different local gradients. It’s important to reduce gradients across all GPUs before doing an optimiser step so that each GPU has global information of how its model parameters should change.
Mixture of Experts (MoE) models present specific challenges and opportunities for parallelism. Specifically, MoE models are sparse: only a fraction of parameters (experts) are activated per token, meaning the compute needed for a forward pass is small compared to the total model size. Expert parallelism (EP) is another technique commonly used during MoE pre-training, though it is out of scope for this blog post. If our model were larger or sparser, we would likely need EP, as the communication-to-computation ratio would not tilt in our favour.
2 – Activation checkpointing
During the backward pass, each GPU needs to compute the gradient of the loss w.r.t. its trainable parameters. To do this, the input (activation) of each operation that involves trainable parameters has to be available in memory. We already computed activations during the forward pass, so storing them is a natural first instinct. However, this is very memory expensive, especially for large models – after all, large models have large intermediate computations! Instead, one can introduce “checkpoints”, selected activations from which we can replay the forward pass and recompute the missing activations. Activation Checkpointing (AC) is a great tool to regulate memory usage.
There are various strategies for selecting these checkpoints; it’s a careful trade-off between spending time in recomputation versus increasing memory consumption. We consider three strategies:
- Full checkpointing (full-AC): During the forward pass, only the inputs to checkpointed “regions” are stored, while all intermediate activations within those regions are discarded. In torchtitan terms, each transformer layer is the checkpointing region; therefore, we store only the input to each layer. During the backward pass, the whole transformer layer is recomputed. This provides the largest reduction in activation memory, but incurs the highest recomputation overhead.
- Selective checkpointing, per operation (selective-AC): During the forward pass, activations for selected operations are stored, typically prioritising operations whose recomputation is expensive (e.g. attention or linear layers, which require significant forward compute time). Activations for cheaper operations are discarded and recomputed during the backward pass. This provides a middle ground between memory consumption and recomputation overhead.
- No checkpointing (no-AC): During the forward pass, all activations required for the backward pass are stored. This results in the highest memory consumption, but minimises backward pass recomputation.
3 – Local batch size
Parallelism schemes and activation checkpointing are great tools to tweak memory usage: we can choose how much memory to assign to activations storage or model parameters, etc. This is useful since local batch size, our third degree of freedom, determines the most memory-hungry component: the activations themselves.
Local batch size refers to the number of sequences of tokens each GPU processes simultaneously during one forward-backward pass. It has a significant impact on training efficiency, since usually processing more tokens at once means we will consume our token budget more quickly. We aim to fill up our memory and use the whole GPU (since we paid for it), and the easiest way to achieve this is increasing the local batch size: the size of intermediate computations grows with local batch size. We aim for maximal token throughput, so if we can fit a higher local batch size, we will do so.
The global batch size is, in torchtitan terms, the total number of sequences processed by all GPUs before the optimiser step. Contrary to the local batch size, the global batch size dictates the training dynamics of the model, and thus represents an external requirement that our efficiency engineers do not choose. However, we have an important constraint to follow, namely:
which states that the global batch size () is a multiple of the local batch size () and the number of GPUs we are training on (). A key element of achieving this global batch size is gradient accumulation (), which is the number of forward-backward passes between optimiser steps. We must respect this constraint as we perform local batch size sweeps at different GPU counts. We set the target global batch size for that GPU count and select the closest global batch size that satisfies the constraint for each experiment in the sweep.
How it all fits together
As stated earlier, we aim to make full use of our GPU memory. All three dimensions we defined affect the memory usage of model training. Therefore, the key question we need to answer is: “How should we allocate our memory?” Given that we are in a constrained problem, a change in one of the three dimensions impacts the other two. To illustrate how they are intertwined, let’s take a look at a few examples.
Local batch size interacts with activation checkpointing. For example, in memory terms, selective-AC sits between full-AC and no-AC. It requires a lower local batch size than full-AC, since it keeps more activations in memory, but it avoids recomputing all of a layer’s intermediate activations, which makes it significantly faster. This battle between full-AC and selective-AC will resurface throughout the post.
Local batch size also interacts with FSDP. For example, increasing the number of GPUs in one FSDP group reduces the per-GPU memory footprint of weights, gradients, and optimiser state, freeing up room for a higher local batch size. However, this comes at the cost of communicating with a larger group, and therefore spending more time on communication kernels.
The knobs are not independent: every one of them moves the others. Striking a balance in this push-and-pull game requires careful adjustments, which is why sweeping is so important. One thing stays true: we want to maximise throughput, and staying compute-bound is a great indicator that we are making the right trade-offs.
16 GPUs are enough to decimate our search space
We can now start running our 30B-A3B model and gather the first pre-training efficiency results. We begin to explore the search space at small scale: 8–16 GPUs. The model needs several B200s just to hold weights, optimiser state and gradients, so we treat one node (8×B200) as the basic unit to scale up from.
In the 1–2 node regime, experiments are cheap in cluster resources and time. Therefore, we run the largest sweep at this scale, and transfer the learnings when scaling up. Specifically, we sweep over feasible combinations of:
- FSDP ∈ {8, 16}
- DP ∈ {1, 2}
- Number of GPUs ∈ {8, 16}
- AC ∈ {no-AC, selective-AC, full-AC}
- Local batch size ∈ {2, 6, 10, 14, …}
All experiments use a fixed sequence length of 4096 tokens.
The heatmap shows model efficiency under different parallelism and activation checkpointing setups, across local batch sizes until OOM. We observe memory fragmentation with full-AC: increasing local batch size oscillates memory usage as our tensors are more or less efficiently packed into reserved memory.
Local batch size and memory usage are happily coupled: where one goes, the other goes too.
Training speed, on the other hand, has a more complex relationship with the couple. If the
local batch size is too low, the GPUs are underutilised, and TPS is low. If the local batch
size is too high, training efficiency starts to degrade – the memory is saturating (we see CUDA allocation retry warnings from PyTorch). We observe peak training efficiency at 92% memory usage for selective-AC.
Full-AC, on the other hand, seems to be performant on a wide range of local batch sizes, ranging
from 60 to 94% memory usage. Those are the target memory usages we’ll aim for when scaling up,
ensuring the local batch size is high enough to reach them.
The OOM boundary and training efficiency drastically change depending on the activation checkpointing technique. No-AC is the clear winner for a fixed, small batch size, but its high memory consumption only allows for batch sizes too small to truly compete for best efficiency. Selective-AC introduces recomputation overhead, and therefore requires a higher local batch size to amortise the cost. Full-AC has the highest recomputation overhead and is competitive at large local batch sizes.
Training efficiency on single- and multi-node shows there are no issues when scaling from 8 to 16 GPUs with InfiniBand. The communication medium is fast enough to keep the training compute-bound; otherwise, we would see a significant degradation in TPS on multi-node runs.
By running on just 16 GPUs (1/32 of the target scale), we have shrunk our search space considerably: no-AC runs OOM too early, selective-AC has peak TPS at a local batch size of ~18, and full-AC has peak TPS at a local batch size of ~38. This information is crucial when scaling up, since we will only have to test efficiency around these peak regions instead of sweeping the whole space again.
A profile is worth a thousand experiments
MFU, TPS and memory usage metrics are great indicators to quickly reduce our search space, but to really understand what is happening under the hood, we need to analyse PyTorch profiles. Which kernels are running? When are they running? How do communication and computation kernels overlap? What kind of data and how much is passing through the network? In-depth analysis like this also gives us a good mental model of how our model’s behaviour will change when scaling up.
Let’s take a look at the forward pass of the fastest configuration found so far: fully sharded on 16 GPUs, selective-AC, and local batch size 18.
The trace shows the compute stream at the top, which contains the torch-compiled graph calls
of the transformer layers and all the compute kernels within them. The bottom stream instead
contains the communication kernels. Computation and communication are happening
simultaneously: the weights for the next layer in the forward pass are pre-fetched with the all_gather operation, so that the compute kernels do not have to wait to get started. The trace shows the
compute kernels are running back to back; thus we are bound by the speed at which we can do these
computations – we are compute-bound. Does this hold also for the backward pass?
As detailed above, the backward pass contains reduce_scatter operations in addition to the all_gathers. Both
operations are fully overlapped with compute. Recalling how FSDP works, gradients, as soon
as they are computed for the whole layer, get reduced between GPUs and sharded across the
FSDP group. This is why we see reduce_scatter calls. The all_gather calls are instead needed to gather weights to compute the gradients and for recomputing activations.
All in all, we are compute-bound also in the backward pass. Let’s take a look at a trace that
includes DP:
We have a third communication operation: the all_reduce.
Compared to a reduce_scatter, the all_reduce is used to average gradients across model replicas. This third type of communication scales
with the number of DP replicas, thus creating a push-and-pull dynamic between DP and FSDP. With
a fixed number of GPUs, we can decide how to allocate our FSDP and DP groups. If we increase the
DP degree, FSDP kernel runtime lowers since we have a smaller group, and vice versa. Striking
a balance here is crucial.
A keen reader would point out that we don’t actually need to all_reduce gradients in every backward pass, but only the final one before the optimiser. For the exact
reason, we refer to this section in the Ultra-Scale Playbook. Let’s try the same profile again, but with the all_reduce optimisation:
Trace of our model with DP during the backward pass of a non-final gradient accumulation step. DP all_reduce operations are not needed and therefore removed.
Trace of our model with DP during the backward pass of the final gradient accumulation step, containing the all_reduce operations to ensure all GPUs have the same gradients before the optimiser step.
This is a simple example of how analysing profiles is the cornerstone of efficiency engineering, and showcases the profile → optimise → profile approach we are so fond of. If you really want to get to the bottom of why a model is underperforming, profiling is your best tool, and we cannot recommend it enough. Other system-level optimisations are also enabled at larger scales, but their individual impact is limited at this model size. The most impactful is keeping weights resident in memory between a backward pass and the next forward pass, rather than sharding and re-gathering them, worth roughly a 2–3% throughput speedup at target scale. Even so, the hierarchical scaling approach remains the main driver of the efficiency we report.
From 16 to 512 GPUs, scaling cost is (nearly) free
We have narrowed down our search space significantly with a sweep at the 16 GPU scale. Before running experiments on more GPUs, let’s think about our findings and what we can expect from a scale-up.
Our results show the communication ends earlier than the compute. This means we can increase
the FSDP group size until we are nearly communication-bound. Furthermore, all but the last
forward-backward pass do not contain any all_reduce operations,
meaning DP can eat some of the scale-up cost too, as we would only pay it in the last gradient
accumulation step. The game plan is to increase FSDP as long as we are compute-bound and scale
DP until we get to 512 GPUs.
As a first step, let’s scale to 128 GPUs and profile the runs to see how the communication ratio goes up. Since we don’t know yet if selective-AC or full-AC is the faster choice, we will gather metrics for both. We begin with a run on 128 GPUs with selective-AC, where we can increase local batch size as each GPU holds less of the model in memory.
Forward pass on 128 GPUs, FSDP=128, selective-AC.
Backward pass on 128 GPUs, FSDP=128, selective-AC.
This configuration achieves 28.0k TPS and 37.0% MFU with 93% memory usage. These are great numbers given that we had 28.4k TPS at 1/8th of this scale. It’s a relative drop in per-GPU training efficiency of only 1.4%. We are still compute-bound in both the forward and the backward pass.
With full-AC, the picture looks similar:
Forward pass on 128 GPUs, FSDP=128, full-AC.
Backward pass on 128 GPUs, FSDP=128, full-AC.
Once again compute-bound, though this time we train at 26.3k TPS, 34.7% MFU and 89% memory usage. The communication kernels in the forward passes of both traces indicate we are right at the boundary of being communication-bound. The backward passes, on the other hand, show we are still clearly compute-bound. Increasing the FSDP degree would make us communication-bound during the whole training. We can conclude that selective-AC is the preferred AC technique, and 128 is the largest FSDP degree we should use.
Let’s scale up to the target count of 512 GPUs. We will use the remaining factor of 4 (512/128 = 4) as the DP degree. In accordance with our hierarchical scaling strategy, we keep the rest of the setup the same. The final configuration is thus: selective-AC, local batch size = 22, FSDP = 128, DP = 4. The traces follow:
Forward pass on 512 GPUs, FSDP=128, DP=4, selective-AC.
Backward pass of a non-final gradient accumulation step on 512 GPUs: no all_reduce, still compute-bound.
Backward pass of the final gradient accumulation step on 512 GPUs: the DP all_reduce bleeds past the compute kernels.
We are compute-bound in all but the last backward pass, exactly where we introduced the DP all_reduces. Our final TPS is 26.7k, and MFU is 35.3%. This means we could train this model for 20T
tokens on 512 GPUs within 17 days. The big win comes from scaling: for example, adding two
more DP replicas cuts our training down to just 11 days. Careful balance of hardware
resources and FSDP and DP degrees, along with cheap DP scaling, grants us a near-linear
scaling effect.
Our 30B-A3B model has been successfully optimised to run on 512 GPUs.
Let’s take a look at all the optimums we found so far:
| Number of GPUs | FSDP/DP | TPS | MFU | HBM usage |
|---|---|---|---|---|
| 16 | 16/1 | 28.4k | 37.5% | 92% |
| 128 | 128/1 | 28.0k | 37.0% | 93% |
| 512 | 128/4 | 26.7k | 35.3% | 94% |
We are running at very high MFU and TPS on all scales, with only a 6% drop in per-GPU efficiency from 16 to 512 GPUs. This indicates we have balanced the usage of all hardware components well. We achieved efficient pre-training by hierarchically transferring the findings at lower GPU counts to higher GPU counts. With these efficiency numbers, we would be more than happy to roll out an actual pre-training run.
Limitations
We used a hierarchical approach to scaling up pre-training. Notably, we use this approach because we are in a constrained environment; each GPU hour is extremely valuable. While we might not be running the overall fastest configuration possible, finding such a configuration would require a significant amount of experiments at the target scale, which is not desirable for us. Those GPU hours can be used for functional performance ablations (actually producing a good model).
Furthermore, the experiments presented here are limited to a single architecture and a single sequence length, with FSDP and DP only, and a target scale of 512 GPUs.
Of course, it takes more than “just” a well-designed scale-up to achieve 35.3% MFU. We do system- and kernel-level optimisations in conjunction with the hierarchical scale-up. They are the foundation and the reason that our scale-up works well. Those system- and kernel-level optimisations are, however, out of scope for the current blog post.
We deliberately don’t compare this MFU figure against other published MoE runs. Architecture, precision, and recipe differences make such comparisons complex and prone to mistakes. Instead, judge us on the hierarchical scale-up itself: how little efficiency we lose as GPU count grows.
Conclusion
We achieved 35.3% MFU in the pre-training of a 30B-A3B model on 512 NVIDIA B200 GPUs, successfully parallelising at scale using a hierarchical scaling approach. We showed how we removed most of the search space with just 16 GPUs by finding local batch size bounds, ruling out underperforming configurations, and understanding the behaviour of our model under the hood with PyTorch profiles. By carrying over learnings from small-scale experiments to a 32× larger target scale, we found a highly efficient configuration that is only 6% away from perfect linear scaling.
We hope you found this blog post fun to read! More optimisation posts will follow; if you have any requests, feel free to reach out to us. And if this sort of work interests you, we’d love to talk to you about opportunities to work together!
Acknowledgements
A very special thank you to Steffen Hirschmann, who helped shape and write this post from the ground up. This piece would not have been possible without his invaluable contributions.
I’m also grateful to Samuel Weinbach, Yasser Jadidi, Max Höth, Fabien Benureau, Alessio Serra, Vale Cofer-Shabica, Matteo Silla, Pablo Schumacher, and Helena Treeck for their thoughtful reviews, and to Noé Beckerle Vallejo and Alexander Wortmeier for bringing the piece to life.