the model fit. the work piled up.
Our five-Mac training run was spending far longer on each update than its communication timings could explain. The problem was how the backward computation held its intermediate work in memory.
The trainer compiled four accumulation microsteps into one lazy computation. On the worker handling the largest local batch, peak active MLX allocations reached 87.436 GiB. With 96 GiB of physical memory, that left little headroom for streamed datasets, the operating system and other applications. Memory compression became substantial.
In the initial synthetic cluster test, that worker took about 15.2 seconds to compute. Its peers spent roughly 11–12 seconds waiting. Gradient transfer and averaging took about 0.2 seconds after synchronization. Local computation under memory pressure was the dominant bottleneck in that test.
- Parameters
- 175,611,648
- Context length
- 2,048 tokens
- Accumulation
- 4 microsteps
- Global update
- 131,072 token positions
finish one microbatch before the next.
We changed the trainer to compile one microbatch’s loss and gradient computation, then materialize the FP32 accumulated gradients after each microstep. The optimizer still updated once after four microsteps. The model, context, local batch sizes and total token positions per update stayed the same.
That change reduced the overlap of temporary allocations. Peak active memory on the largest-batch worker fell to 26.839 GiB, about 69% lower. All five workers showed lower peaks.
| Worker group | Before | After |
|---|---|---|
| Largest local batch | 87.436 GiB | 26.839 GiB |
| Next largest local batch | 59.532 GiB | 19.472 GiB |
| Each of three smaller workers | 31.119 GiB | 11.596 GiB |
The cluster comparison paired sequential accumulation with a 4 GiB allocation-cache cap. Separate local tests also reached 26.839 GiB with the default cache; a cache cap alone left the unrolled computation’s peak at 87.436 GiB.
The slowest worker’s median update time fell from 15.576 to 4.037 seconds in the short synthetic cluster comparison. Each variant ran three updates, excluding the first from median timing. These tests used synthetic inputs and fresh model initialization, so their timings exclude streaming, validation, sample generation and checkpoint writes.
Active MLX memory is not total system memory. These figures exclude the reusable allocation cache, Python and Arrow dataset memory, and the operating system. Post-update active memory was only about 2.94 GiB per worker, which would have hidden the much larger temporary peaks.
the next bottleneck was workload balance.
After reducing memory pressure, another worker’s compute time became the limiting factor. We tested redistributing the local batches and changing allocation-cache limits while keeping the global batch fixed.
Moving from local batch sizes of 6 / 4 / 2 / 2 / 2 to 7 / 3 / 2 / 2 / 2 shifted work toward the worker that could finish it faster. Increasing the cache caps on the two 96 GiB machines to 16 GiB, while retaining 4 GiB caps on the three 36 GiB machines, improved throughput further.
| Configuration | Local batch sizes | Cache caps · GiB | Seconds/update | Tokens/s |
|---|---|---|---|---|
| Baseline | 6 / 4 / 2 / 2 / 2 | 4 / 4 / 4 / 4 / 4 | 4.047 | 32,391 |
| Baseline repeat | 6 / 4 / 2 / 2 / 2 | 4 / 4 / 4 / 4 / 4 | 4.031 | 32,516 |
| Rebalanced workload | 7 / 3 / 2 / 2 / 2 | 4 / 4 / 4 / 4 / 4 | 3.509 | 37,351 |
| Rebalanced + 8 GiB caches | 7 / 3 / 2 / 2 / 2 | 8 / 8 / 4 / 4 / 4 | 3.353 | 39,092 |
| Rebalanced + 16 GiB caches | 7 / 3 / 2 / 2 / 2 | 16 / 16 / 4 / 4 / 4 | 3.160 | 41,483 |
| 16 GiB repeat | 7 / 3 / 2 / 2 / 2 | 16 / 16 / 4 / 4 / 4 | 3.183 | 41,185 |
| 16 GiB + compiled optimizer | 7 / 3 / 2 / 2 / 2 | 16 / 16 / 4 / 4 / 4 | 3.153 | 41,569 |
| More work on the largest worker | 8 / 2 / 2 / 2 / 2 | 4 / 4 / 4 / 4 / 4 | 3.985 | 32,888 |
| More work + 16 GiB caches | 8 / 2 / 2 / 2 / 2 | 16 / 16 / 4 / 4 / 4 | 3.368 | 38,914 |
Each variant ran eight synthetic updates with two warmup updates excluded. Seconds/update is the median of the six per-update maximum times across all workers. Tokens/s is 131,072 divided by that median. Cache caps refer to reusable allocations, not a total memory limit.
The baseline controls measured 32.4k and 32.5k tokens/s. The selected configuration repeated at 41.5k and 41.2k tokens/s: approximately 27% greater throughput, or 21% shorter updates. Giving the largest worker eight samples instead of seven made the cluster slower than the selected balance.
Compiling the optimizer shortened its own phase, but overall throughput was too close to the noncompiled repeats to establish a meaningful gain. We retained the existing optimizer execution path.
check the improvement with real data.
We then ran isolated resume pilots from copies of saved checkpoints. The first used sequential accumulation; the later pilot also used the selected workload balance and cache caps. The original checkpoints were preserved.
| Pilot | Timed updates | Seconds/update | Tokens/s |
|---|---|---|---|
| Sequential accumulation | 3 | 4.042 | 32,424 |
| Balanced workloads + caches | 6 | 3.185 | 41,158 |
The later interval measured about 27% more throughput than the earlier one, consistent with the controlled workload tests. The pilots resumed different checkpoints, so this real-data comparison is supporting evidence rather than a same-checkpoint A/B experiment.
The first resumed updates took about 104 and 113 seconds, including cold stream preparation and compilation; the later start also included a transition save. We excluded those updates from the warm interval rates above. Validation, checkpoints and later dataset stalls can still reduce the average speed of a longer run.
The balanced pilot’s peak active allocations were approximately 30.67 / 15.54 / 11.60 / 11.60 / 11.60 GiB. Its largest-batch worker now handled seven samples rather than six, so this peak should not be confused with the 26.839 GiB accumulation comparison.
faster execution still needs verification.
A numerical fixture compared the old and new accumulation paths, checking loss and every FP32 gradient with absolute tolerance 1e-5 and relative tolerance 1e-4. Distributed checks covered token-weighted averaging with unequal local batches, matching model and optimizer replicas, and checkpoint round trips.
Both real-data pilots retained the validation identity and passed checkpoint manifest checks. All five model and optimizer replicas agreed at the pilot checkpoint saves. The transitions preserved the architecture, tokenizer, data recipe, global batch and token-based learning-rate schedule.
These checks establish numerical agreement within the tested tolerances and consistency across workers. They do not assert bitwise identity between execution variants or improved model quality. The synthetic losses were diagnostic values, not evaluations of a trained model.
measure the work, then change its execution.
The first useful measurement was peak active memory during the backward pass. The next was each worker’s compute time. Together, they showed where the cluster was losing time and which changes were worth testing.
Sequential accumulation addressed the temporary memory pressure. Workload redistribution and cache tuning then reduced the wait for the slowest worker. The model stayed at 175.6 million parameters throughout these comparisons.
The results apply to this model, hardware and execution path. The tests were bounded diagnostics and short pilots; they do not establish a universal speedup for MLX training or a sustained throughput guarantee.
inspect the measurements.
The accompanying data contains all nine workload variants, the memory comparison and the two streamed-data pilot intervals. It identifies the fixed training dimensions and the scope of each measurement.