Repository navigation
Fix flops profiler FLOPS and samples/s under gradient accumulation - #8763
Draft
vineethsaivs wants to merge 1 commit into
Draft
vineethsaivs wants to merge 1 commit into
vineethsaivs wants to merge 1 commit into
Conversation
print_model_profile divides one forward's flops by the fwd/bwd/step global timers, which hold every micro-batch since the last optimizer step, so FLOPS and samples/s were low by the number of micro-batches. Average the timers over the recorded forwards, and count samples per iteration as micro-batch times data-parallel size, not world size. Signed-off-by: Vineeth Sai <vineethsai4444@gmail.com>
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.
With gradient accumulation, the flops profiler summary reports latencies too high and FLOPS and samples/second too low, by the number of accumulated micro-batches. This is what #2565 and #4385 ran into.
print_model_profiledivides one forward's flops by the fwd/bwd/step global timers, which hold every micro-batch since the last optimizer step. samples/second also usedmicro_batch * world_size, counting model-parallel ranks as extra samples.Fix: average the timers over the recorded forwards, and count samples as
train_batch_size() // gradient_accumulation_steps(). Step latency is now amortized per micro-batch, so iter = fwd + bwd + step still holds. GAS 1 without model parallelism prints the same as before.Test: with fixed timer records (GAS 4, mp 2) upstream prints 400 ms fwd and 12.5 samples/s; this prints 100 ms and 25. On a real CPU engine at GAS 4, samples/s goes from 33 to 133 and fwd latency matches one forward. Replaying #2565's timings gives 13.5 samples/s against 13.23 in Megatron's own log (upstream prints 4.5). Profiler tests: 19 passed, 1 skipped.