NotesDraft/Technical Notes
The FSDP throughput cliff nobody warns you about
Throughput dropped 40% going from 4 GPUs to 8. The network was fine. The problem was where the shards landed.
Symptom
Scaling was clean up to four GPUs, then fell off a cliff at eight and partially recovered at sixteen. A dip that recovers is not a bandwidth problem. Bandwidth problems get monotonically worse.
What it was not
- Interconnect saturation. NCCL bandwidth tests came back at expected numbers.
- Data loading. The pipeline was prefetching well ahead of the model.
- Stragglers. All eight devices were identical and their step times matched.
What it was
The auto wrap policy was sharding on parameter count, not on module boundaries. At eight GPUs the split landed inside a transformer block, so a single forward pass needed two all-gathers where the four GPU configuration needed one.
# Before: shards wherever the parameter count says to
auto_wrap_policy = size_based_auto_wrap_policy
# After: shards on block boundaries, always
auto_wrap_policy = functools.partial(
transformer_auto_wrap_policy,
transformer_layer_cls={TransformerBlock},
)
Result
| GPUs | Before | After | Change |
|---|---|---|---|
| 4 | 100% | 100% | none |
| 8 | 61% | 94% | +33 pts |
| 16 | 78% | 91% | +13 pts |
Scaling efficiency against the 4 GPU baseline. Placeholder numbers.
Takeaway
If your scaling curve dips and then recovers, stop looking at the network and go look at where your sharding boundary falls relative to your model architecture. A size based policy knows nothing about your blocks and will happily cut one in half.