[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
Conversation
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
Author
|
@misko Hi Misko, anything you would like me to change here? :) |
This branch has not been deployed
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
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_checkconfig runs 15 tasks per step.DDPMTLoss._reductioncounts the valid samples withloss[mult_mask].numel()(modules/loss.py:129) and_ddp_meanall-reduces that Python int (:87), whichdistutils.all_reducecopies to the device and reads back with.item();per_structureaddstorch.nonzero(...).numel()(:125), abincountover a masked index (:117) and two asserts on device tensors;forwardthen branches ontorch.all(loss.isfinite())(:163). One function up,get_output_maskbuilds a mask on the host and copies it to the device once per task and per dataset (units/mlip_unit/mlip_unit.py:207), andcompute_lossusesbatch.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 onmainat473513e15, so none of it has been fixed upstream._reductioncomputes the same integer asmult_mask.sum()times the loss's elements per mask entry, and_ddp_meanall-reduces that 0-dim tensor and clamps it withtorch.clamp(min=1)instead of takingmax()of a Python int, so the all-reduce is enqueued and nothing waits on its result.num_samplesis a tensor now rather than an int; inside the tree it only ever reaches_ddp_mean, but an out-of-tree subclass ofDDPMTLossthat overrode a reduction or_ddp_meanand used it as a number would see the difference.per_structurekeeps its sizes static.repeat_interleavegets itsoutput_size, ascatter_add_over the structure index replaces thebincountover a masked index,torch.wherereplaces the index assignment that filled in structures with no free atoms, and the count of contributing structures is(per_struct_loss != 0).sum()instead ofnonzero().numel(). The three asserts go with them: each one read the device, and each holds by construction, sincestruct_idxcovers0..natoms.numel()-1exactly andfree_natomscounts set mask entries per structure, so after thewhereit lies between 1 andnatoms.nan_to_numis 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.get_dataset_masksbuilds one boolean mask over the systems of the batch per dataset in it, andget_output_masktakes them as an optional argument, building them itself when it is not given, so its other caller atmlip_unit.py:325is unaffected. The atom-levelrepeat_interleavegets itsoutput_size, andcompute_losstakes the atom count frombatch.pos.shape[0]instead ofbatch.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.473513e15nvidia-smiThe 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
DDPMTLossforward per reduction undertorch.cuda.set_sync_debug_mode("warn"), in a single process on one GPU with synthetic tensors shaped like this config's tasks.mean, energy with a per-system masksum, structure levelper_structure, forces with a per-atom masktorch 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 == 1distutils.all_reducehands 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_pathrecorder 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 inindex_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
meanreduction reads 1.5249106884002686 unmodified against 1.524910569190979 here, one ulp and 7.8e-8 relative, whilesumandper_structurecome out bit-identical. The table below is what bounds that end to end. Each cell is the worst over steps 0 to 4 ofmax |a - b|over the tensor divided bymax |a|, and the tolerance is the one that applies at the step where that worst value falls.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-baseandperf/loss-path-no-host-syncs-verifyof my fork.ruff checkandruff format --checkpass on both changed files with ruff 0.5.1, the version.pre-commit-config.yamlpins.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
PYTHONPATHand each run logs which tree it imported.Model and job. UMA-S style eSCN-MD MoLE backbone from
backbone/K4L2.yamlwith theperf_checkoverrides: 4 blocks, lmax = mmax = 2, 64 experts, about 290 M parameters, fp32 with TF32 off, trained from random init.MLP_EFS_Headper dataset (omol, oc20, omat, odac, omc),regress_stress: Trueanddirect_forces: False, so forces and stress are autograd of the energy and every step runs a second-order backward. 15 tasks fromtasks/oc20_omol_conserving_all.yaml, AdamW at lr 0.1, EMA 0.999, DDP over 8 ranks. Batches come fromMaxAtomDistributedBatchSamplerwithmax_atoms: 350. The data is the synthetic 5-dataset aselmdb corpus the repo'sperf_checkfixtures build, 1.1 MB, copied to node-local disk before each run.Measurement. Against the unmodified
473513e15, 250 steps, seed 42, 8 ranks withjob.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 istime.perf_counter()around each optimizer step from the repo's ownBenchmarkTrainCallback, 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:
max_epochs=nullbecausetraining_inner.yamlshipsmax_epochs: 1and an epoch of this corpus is 6 steps per rank at 8 ranks;++job.seedand++job.run_dirbecause the yaml'sjob:block does not declare them. Nothing is set in the environment beyond node-localTMPDIR, caches and data, HF offline, W&B disabled,CUDA_VISIBLE_DEVICES=0-7andPYTHONPATH=<tree>/src.The correctness runs use the same command with
max_steps=12and+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 thatruff checkpasses 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=1and=16,NCCL_CUMEM_ENABLE=0,TORCH_NCCL_AVOID_RECORD_STREAMS=1, andrunner.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_checkcorpus, 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, sincegp_utils.initialized()is false in this workload;world_size == 1outside the single-process sync count above; other trainers, althoughDDPMTLossis 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.