Skip to content

[perf] UMA-S conserving training on 8x B200: 152.2 ms to 144.6 ms per step, from the host synchronizations in the loss - #2203

Open
TarzanZhao wants to merge 1 commit into
facebookresearch:mainfrom
TarzanZhao:perf/loss-path-no-host-syncs
Open

TarzanZhao wants to merge 1 commit into
facebookresearch:mainfrom
TarzanZhao:perf/loss-path-no-host-syncs

Conversation

@TarzanZhao

@TarzanZhao TarzanZhao commented Sep 18, 2026 •

Copy link
Copy Markdown

Summary

In UMA-S conserving training every task's loss reads a count back from the GPU before it can be computed, and the perf_check config runs 15 tasks per step. DDPMTLoss._reduction counts the valid samples with loss[mult_mask].numel() (modules/loss.py:129) and _ddp_mean all-reduces that Python int (:87), which distutils.all_reduce copies to the device and reads back with .item(); per_structure adds torch.nonzero(...).numel() (:125), a bincount over a masked index (:117) and two asserts on device tensors; forward then branches on torch.all(loss.isfinite()) (:163). One function up, get_output_mask builds a mask on the host and copies it to the device once per task and per dataset (units/mlip_unit/mlip_unit.py:207), and compute_loss uses batch.natoms.sum(), a device scalar, as a view size (:253). Each one stops the CPU until the GPU has drained the queue the CPU is filling. This PR computes the same quantities on the device. The steady step goes from 152.18 ms to 144.57 ms (1.053x) on the 8 B200 GPUs I measured; I have not measured other hardware. Nothing is behind an option, and no config file or dependency changes: two files, +61 / -28 lines. Every line quoted above reads the same on main at 473513e15, so none of it has been fixed upstream.

  1. The valid-sample count stays on the device. _reduction computes the same integer as mult_mask.sum() times the loss's elements per mask entry, and _ddp_mean all-reduces that 0-dim tensor and clamps it with torch.clamp(min=1) instead of taking max() of a Python int, so the all-reduce is enqueued and nothing waits on its result. num_samples is a tensor now rather than an int; inside the tree it only ever reaches _ddp_mean, but an out-of-tree subclass of DDPMTLoss that overrode a reduction or _ddp_mean and used it as a number would see the difference.
  2. per_structure keeps its sizes static. repeat_interleave gets its output_size, a scatter_add_ over the structure index replaces the bincount over a masked index, torch.where replaces the index assignment that filled in structures with no free atoms, and the count of contributing structures is (per_struct_loss != 0).sum() instead of nonzero().numel(). The three asserts go with them: each one read the device, and each holds by construction, since struct_idx covers 0..natoms.numel()-1 exactly and free_natoms counts set mask entries per structure, so after the where it lies between 1 and natoms.
  3. The NaN guard runs unconditionally. nan_to_num is the identity on finite values, so applying it always costs one elementwise kernel and removes a third read of the loss. A NaN loss is still zeroed exactly as before; what goes is the "Found nans while computing loss" warning, because deciding whether to log it is what read the device.
  4. The per-dataset masks are built once per batch. get_dataset_masks builds one boolean mask over the systems of the batch per dataset in it, and get_output_mask takes them as an optional argument, building them itself when it is not given, so its other caller at mlip_unit.py:325 is unaffected. The atom-level repeat_interleave gets its output_size, and compute_loss takes the atom count from batch.pos.shape[0] instead of batch.natoms.sum().

Speedup Result

Measured on 8 NVIDIA B200 GPUs in one node, four runs per version, interleaved as base, PR, base, PR and so on, each run alone on the node. The step time is the median of steps 51 to 250 from the repo's own BenchmarkTrainCallback; the figures below average that median over the four runs.

unmodified 473513e15 this PR change
step time, median of steps 51 to 250 152.18 ms 144.57 ms 1.053x, -7.61 ms
run-to-run spread of that median, 4 runs 2.24 ms (1.5 %) 2.93 ms (2.0 %)
slowest and fastest of the four runs 150.87 to 153.11 ms 143.64 to 146.57 ms the two do not overlap
step time, mean of steps 51 to 250 151.80 ms 144.36 ms 1.052x
steps 51 to 250, summed 30.36 s 28.87 s 1.052x
whole run, launch to exit 137.4 s 134.6 s
peak allocated memory, writing rank 12.87 GB 12.87 GB unchanged
peak GPU memory per GPU, nvidia-smi 15.7 to 16.0 GB 15.7 to 16.0 GB unchanged
host RSS, summed over 8 ranks 25.1 to 25.3 GB 24.9 to 25.0 GB unchanged

The gain is 2.6 times the larger of the two spreads, and the arms do not overlap over the eight interleaved runs: the slowest run of this branch, 146.57 ms, is 4.3 ms below the fastest unmodified run, 150.87 ms. Inside a single run the per-step spread is 5 to 6 % on both versions, which is this workload's variable-size batches rather than the measurement.

The synchronizations themselves, counted at the call site: one DDPMTLoss forward per reduction under torch.cuda.set_sync_debug_mode("warn"), in a single process on one GPU with synthetic tensors shaped like this config's tasks.

reduction, as this config's tasks use it unmodified this PR
mean, energy with a per-system mask 2 0
sum, structure level 2 0
per_structure, forces with a per-atom mask 12 0

torch prints that sync debug mode "does not yet detect all synchronizing operations", so those are lower bounds, and one process does not exercise the all-reduce leg at all: at world_size == 1 distutils.all_reduce hands its argument straight back. Under 8-rank DDP each of the 15 tasks also pays the blocking round trip of all-reducing a Python int, which this PR removes and which I measured only through the step time.

Whole-run wall time is 137.4 s to 134.6 s, and the training part of it 117.3 s to 114.0 s. About 71 s of each run goes on elastic launch and loading five datasets before step 0, which this PR does not touch and which dominates the wall figure, so the claim here is per-step time.

Correctness Verification

I trained the unmodified code and this branch on the same 12 steps, seed 42, 8 ranks, and compared what the repo's own debug_checksums_save_path recorder writes, extended to also save the per-task losses, the per-head energy, forces, stress and node embeddings, and the shapes and dtypes of the batch. The comparison covers steps 0 to 4 on all 8 ranks, 40 rank-and-step records per pair. I recorded three runs of the unmodified code and two of this branch, and compared all six cross-version pairs, plus the three unmodified-against-unmodified pairs and the one branch-against-branch pair as controls. The tolerances were set before the runs: per quantity and step, three times the spread the unmodified code showed over three earlier repeats of itself, with a floor of 1e-5. Only steps 0 to 4 carry a signal, because at lr 0.1 from random init, with atomics in index_add_, the unmodified code differs from itself past step 4.

One of the six cross-version pairs failed those tolerances. The gradient-magnitude check at step 2 recorded 7.55e-4 against a budget of 6.26e-4, 1.21x over. It is one number rather than eight: DDP has averaged the gradients before the recorder sees them, so the same value is recorded on all 8 ranks and it marks 8 of that pair's 40 rank-and-step records. The other five cross-version pairs are within tolerance on all 40, as are the three unmodified-against-unmodified pairs and the branch-against-branch pair.

That budget is tighter than what the unmodified code does against itself at that step today. Derived the same way from the three unmodified runs of this round, the run-to-run spread of the gradient magnitudes at step 2 is 6.21e-4, so the budget of 6.26e-4 that the check used is 1.01x that spread rather than the 3x it was meant to be; the same recipe on today's unmodified runs gives 1.86e-3. Re-running all six cross-version pairs against that freshly derived file, which comes from the unmodified runs only, puts 6 of 6 within tolerance. Both budgets are above so that you can judge it yourself: the declared check failed, and the file it passes against was derived after the runs rather than before them.

One value change is real, and measured rather than assumed. Dividing a float32 tensor by a 0-dim int64 tensor is not bit-identical to dividing it by a Python int, so each task's loss can move in its last bit: on the synthetic case above, the mean reduction reads 1.5249106884002686 unmodified against 1.524910569190979 here, one ulp and 7.8e-8 relative, while sum and per_structure come out bit-identical. The table below is what bounds that end to end. Each cell is the worst over steps 0 to 4 of max |a - b| over the tensor divided by max |a|, and the tolerance is the one that applies at the step where that worst value falls.

recorded unmodified vs this PR, 6 pairs unmodified vs itself, 3 pairs tolerance
loss 6.0e-6 to 2.0e-5 1.0e-5 to 2.0e-5 4.3e-5 at step 4
per-task losses, 15 tasks 5.6e-5 to 1.6e-4 1.1e-4 to 2.2e-4 2.3e-4 at step 4
parameters, mean of |w| per tensor 1.9e-5 to 4.2e-5 2.7e-5 to 6.3e-5 1.3e-4 at step 2
gradients, mean of |g| per tensor 7.7e-3 to 5.2e-2 at step 4, and the one exceedance above, 7.55e-4 against 6.26e-4 at step 2 on one pair 2.3e-2 to 4.7e-2 9.7e-2 at step 4, 6.3e-4 at step 2
energy per system, all 5 heads 1.3e-4 to 4.8e-4 1.4e-4 to 4.6e-4 6.8e-4 at step 0
forces per atom, all 5 heads 5.4e-3 to 1.3e-2 6.7e-3 to 1.2e-2 1.9e-2 at step 4
stress per system, all 5 heads 2.8e-3 to 7.6e-3 5.6e-3 to 1.1e-2 2.4e-2 at step 4
node embeddings 5.7e-4 to 2.0e-3 8.0e-4 to 2.1e-3 2.6e-3 at step 4
batch shapes and dtypes, atoms per system, dataset per system, parameter and gradient names identical in every pair identical in every pair exact

Outside the cell named above, every cross-version figure is within its tolerance and is the same size as what the unmodified code shows against itself.

The recording code is not in this PR: one commit, applied identically to both sides, on the branches perf/loss-path-no-host-syncs-verify-base and perf/loss-path-no-host-syncs-verify of my fork.

ruff check and ruff format --check pass on both changed files with ruff 0.5.1, the version .pre-commit-config.yaml pins.

Details: hardware, model, full command, scope of the measurement

Hardware and environment. One node, 8 x NVIDIA B200 (sm_100, 183 GB each), 2 x Xeon Platinum 8562Y+ (128 threads), driver 580.126.20. Python 3.12.14, torch 2.13.0+cu130 (cuDNN 9.2, NCCL 2.29.7), e3nn 0.6.0, torchtnt 0.2.4, ase 3.29.0, hydra-core 1.3.6, numpy 2.4.6, pymatgen 2026.5.4, pandas 3.0.5, pyarrow 25.0.1. I installed nothing for either version; the tree under test goes first on PYTHONPATH and each run logs which tree it imported.

Model and job. UMA-S style eSCN-MD MoLE backbone from backbone/K4L2.yaml with the perf_check overrides: 4 blocks, lmax = mmax = 2, 64 experts, about 290 M parameters, fp32 with TF32 off, trained from random init. MLP_EFS_Head per dataset (omol, oc20, omat, odac, omc), regress_stress: True and direct_forces: False, so forces and stress are autograd of the energy and every step runs a second-order backward. 15 tasks from tasks/oc20_omol_conserving_all.yaml, AdamW at lr 0.1, EMA 0.999, DDP over 8 ranks. Batches come from MaxAtomDistributedBatchSampler with max_atoms: 350. The data is the synthetic 5-dataset aselmdb corpus the repo's perf_check fixtures build, 1.1 MB, copied to node-local disk before each run.

Measurement. Against the unmodified 473513e15, 250 steps, seed 42, 8 ranks with job.scheduler.mode=LOCAL, four runs per version interleaved. Each run starts with all 8 GPUs at 0 MiB, no leftover trainer process and a clean tree; the script refuses to start otherwise, and all eight runs exited 0. Both versions share one Hydra override that is not part of this PR, ++optimizer.fused=true, which swaps AdamW's foreach path for its fused kernels. Step time is time.perf_counter() around each optimizer step from the repo's own BenchmarkTrainCallback, read from <run_dir>/benchmark_results.pkl; steps 1 to 50 are warm-up. Every rank writes that path and the last one to finish wins, which is also why the peak allocated memory row varies by which rank wrote the file.

The measured job, both versions, run from each tree:

fairchem -c configs/uma/benchmark/perf_check/training_inner.yaml \
  datasets.data_root_dir=<data_root> job.device_type=CUDA bf16=False \
  max_steps=250 max_epochs=null ++job.seed=42 ++job.run_dir=<run_dir> \
  job.scheduler.mode=LOCAL ++job.scheduler.ranks_per_node=8 \
  runner.callbacks.0.benchmark_results_path=<run_dir>/benchmark_results.pkl \
  ++optimizer.fused=true

max_epochs=null because training_inner.yaml ships max_epochs: 1 and an epoch of this corpus is 6 steps per rank at 8 ranks; ++job.seed and ++job.run_dir because the yaml's job: block does not declare them. Nothing is set in the environment beyond node-local TMPDIR, caches and data, HF offline, W&B disabled, CUDA_VISIBLE_DEVICES=0-7 and PYTHONPATH=<tree>/src.

The correctness runs use the same command with max_steps=12 and +runner.train_eval_unit.debug_checksums_save_path=<dir>, on the two verify branches named above, which carry the recorder commit on top of the unmodified base and on top of this change. The branch that was recorded and timed differs from the head of this PR by one local variable rename, made so that ruff check passes on the loop that now unpacks the dataset masks.

Settings tried and rejected on the unmodified code, all inside the run-to-run spread: OMP_NUM_THREADS=1 and =16, NCCL_CUMEM_ENABLE=0, TORCH_NCCL_AVOID_RECORD_STREAMS=1, and runner.train_eval_unit.print_every=1000. Not tried, because they change the numerics or the job: runner.train_eval_unit.tf32=true, bf16=True, ema_decay, clip_grad_norm, datasets.max_atoms, job.deterministic.

Scope of the measurement, and what is not covered. One node, 8 ranks, fp32 with TF32 off, the synthetic perf_check corpus, 250 steps. What the change is worth depends on the host as much as on the GPU, since what it removes is CPU stall. Not exercised: graph parallel, since gp_utils.initialized() is false in this workload; world_size == 1 outside the single-process sync count above; other trainers, although DDPMTLoss is shared and every trainer that uses it gets the device-side form; steps past 4 in the numerical comparison; other world sizes, longer runs and real datasets.

Every task's loss reads a count back to the host before it can be
computed. DDPMTLoss counts the valid samples with loss[mult_mask].numel(),
all-reduces that Python int -- which sends it to the device and fetches
it back -- and per_structure adds nonzero().numel(), a bincount over a
masked index and two asserts on device tensors. With fifteen tasks that
is fifteen round trips per step, each one draining the queue the GPU is
working from.

The count is a device scalar now: mult_mask.sum() times the loss's
elements per mask entry, all-reduced as a tensor and clamped instead of
max()'d, so nothing waits on the host. per_structure keeps its sizes
static (output_size on repeat_interleave, scatter_add_ instead of
bincount over a masked index, torch.where instead of index assignment)
and drops three asserts that hold by construction. nan_to_num is applied
unconditionally: it is the identity on finite values, and the isfinite()
test it replaces read the loss back (the warning log goes with it).

On the mask side, the per-dataset system masks are built on the device
once per batch instead of once per task and dataset, repeat_interleave
gets an output_size, and compute_loss takes the atom count from
batch.pos.shape[0] rather than natoms.sum(), a device scalar used as a
view size.

The same integer counts, the same divisions, the same fp32 ops. One
caveat measured rather than assumed: dividing a float32 tensor by a
0-dim int64 tensor is not bit-identical to dividing it by a Python int,
so each task's loss can move in its last bit (1 ulp, 8e-8 relative, on a
synthetic case). Over five steps of the perf_check training that stays
well inside the run-to-run spread of the unmodified code.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01ATw7g6cxq9N2LtnMratsru
@TarzanZhao

Copy link
Copy Markdown
Author

@misko Hi Misko, anything you would like me to change here? :)

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant