Collect dynamo's cyclic garbage after compiling a checkpointed layer - #955
Merged
Merged
Conversation
A grad-mode recompile during backward's first recompute leaves reference cycles whose Python frames still hold that call's real activations. On an 8-layer CP1 Qwen3.6 fixture, four 6.38 GB MoE permute/sort/FC2 tensors stayed alive through the rest of backward, raising the cold peak from 69.1 GB (eager) to 95.3 GB until the cyclic collector ran. Mark garbage from a dynamo end-of-compile callback and run one gc.collect() at the next checkpointed call (skipped while tracing). ART_COLLECT_COMPILE_GARBAGE=0 disables it. Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
bradhilton
deployed
to
trainer-rank-gpu-validation
September 25, 2026 02:27 — with
GitHub Actions
Active
Collect at entry and after each checkpointed call, so a compile inside a call is freed before its backward. Re-register the dynamo end callback on use, check tracing before reading the mark, forward the toggle to remote hosts, and log collection time at debug level. Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
bradhilton
had a problem deploying
to
trainer-rank-gpu-validation
September 25, 2026 03:06 — with
GitHub Actions
Error
Collect after a checkpointed call only when it returns, so a failure keeps its own exception and leaves the mark for the next call. Skip collection while the current CUDA stream is capturing a graph. Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
bradhilton
had a problem deploying
to
trainer-rank-gpu-validation
September 25, 2026 03:27 — with
GitHub Actions
Error
Non-reentrant recompute ends by raising out of the wrapped call, so a compile during it leaves the mark for the next checkpointed call. Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
bradhilton
had a problem deploying
to
trainer-rank-gpu-validation
September 25, 2026 03:46 — with
GitHub Actions
Error
The root unit-test environment lacks Megatron, so importorskip skipped the whole module there. Also enable early stop around the checkpointed forward, where PyTorch reads the setting. Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
bradhilton
deployed
to
trainer-rank-gpu-validation
September 25, 2026 03:56 — with
GitHub Actions
Active
This was referenced Sep 25, 2026
This was referenced Sep 25, 2026
This branch was successfully 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.
Part of #949. With ART's per-layer
torch.compile, the first recompute in backward triggers a dynamo recompile for grad mode. That compilation leaves reference cycles whose Python frames still hold the traced call's real activations. Those tensors stay allocated through the rest of backward, until Python's cyclic collector happens to run. In our traces that was when the next call started compiling.Evidence
One H200, Qwen3.6-35B-A3B with random weights, 8 layers, CP1, ten unique-token requests (194,753 tokens), forward plus backward.
Memory numbers are incremental unless marked absolute:
max_memory_allocatedduring the call minusmemory_allocatedjust before it.predicted_peak_bytes: the bytes it requires beyond what is already allocated, admitted against free device memory minus a reserve. That is the same incremental basis.aab5d5416aab5d5416ART_DISABLE_MEGATRON_COMPILE=1)ART_COLLECT_COMPILE_GARBAGE=0The base rows ran on #954's head
587ae7e6b, whose tree is identical toaab5d5416. The compile-off run recorded its cold peak, then exited 1 in a later diagnostic probe step. The first revision of this PR collected only at the next checkpointed call's entry. It measured 69.83 GB incremental. Collecting when the compiling call returns frees that call's garbage before its own backward, which removes another 5.7 GB.CP2/EP1: the same 8-layer shape on two H200s, seeded, on this PR's head. The planner predicts 29.41 GB incremental per rank.
ART_COLLECT_COMPILE_GARBAGE=0Base match. The toggle-off arm matches two earlier base-tree CP2 runs byte for byte.
Numerics. Both arms ran one subforward per call. The summed loss is bitwise identical between arms in all three calls.
Remaining gap. The remaining excess over the planner, 2.2 GB and 6.8 GB, is the CP planner gap still tracked in TrainerRank CP2/EP1 normal admission reaches CUDA OOM in compiled MoE LoRA forward #949.
What the extra memory is. Allocator traces show the compiled run's extra memory is four
[1558024, 2048]BF16 tensors (6.38 GB each): Transformer Engine permute/sort-chunks outputs and the expert FC2 output. They are allocated in backward's first recompute, when dynamo logs agrad_modeguard failure and recompiles, and are never freed in that call. In the next call all four are freed insidetorch._dynamoOutputGraph.__init__(guard creation →torch.utils._traceback.extract).Confirmed as cyclic garbage. After the cold call, a
gc.collect()undergc.DEBUG_SAVEALLmoved the unreachable objects intogc.garbagefor inspection. Clearing that list and collecting again freed 32.02 GB of CUDA memory. The unreachable objects included 357frameobjects plus dynamo tracing state (Instruction,FrameSummary,StackSummary,SpeculationEntry,TensorVariable), whose frame locals held these tensors and hidden states. The warm call left none.Change
A dynamo end-of-compile callback marks compile garbage as pending. ART's checkpoint wrapper runs at every checkpointed layer call, including recompute. It now runs one
gc.collect()when a mark is pending:Behavior:
torch._dynamo.reset()does not silently disable it.ART_COLLECT_COMPILE_GARBAGE=0disables it. The key is inRUNTIME_ENVIRONMENT_KEYS, so SSH-launched ranks receive it and host admission reports it. Admission does not require it to match across hosts.art.megatron.lora.Validation
Actual peak: the table above. The cold call includes the compiles. The comparison is limited to this one shape. The seeded runs are two fix-on runs and one fix-off run.
Numerics. These come from seeded runs (
ART_MEGATRON_RANDOM_STATE=0). A SHA-256 fingerprint of every parameter after the checkpoint loads, and of each call's input tokens, matched across all three arms. The cold-call summed loss is bitwise identical in all three (2499687.75). Gradients of all 132 trained adapter tensors were compared elementwise; the other 132 LoRA tensors are unused and zero in every run.DEBUG_SAVEALLprobe. By call 2 that memory had been freed, but PyTorch's allocator kept it cached, and the planner does not count cached memory as available. For calls 1–2, the off arm's usable limit (free device memory minus the reserve, incremental) was 25.0 GB, against 58.1 GB in the fix-on arm; both arms had 128.5 GB on the cold call. The off arm therefore split each call into 10 subforwards instead of 2. Their difference reflects that different split, and is not evidence either way.Cost (same fixture, one GPU, seeded runs):
GC time counts every collection, automatic and explicit, timed through
gc.callbacks. The off arm's calls 1–2 ran 10 subforwards instead of 2 (see Numerics). That is why its warm call is slower, and why its call 1 is not like-for-like either.Collections happen only around compiles and recompiles: 7 in the cold call and 8 in the second call, whose new shapes trigger dynamic recompiles. Each takes 0.47–0.96 s, dominated by heap size rather than garbage. None happened in the warm call. The counts were the same in all four fix-on runs.
Unit tests (CPU,
tests/unit/test_megatron_compile_garbage.py, now run in CI's Megatron lightweight lane, since the root unit-test environment lacks Megatron and would skip them):torch.compilesets the mark;tests/integration/megatron/model_support/test_compile_flags.py, which callstorch._dynamo.reset(), passes when run before these tests.Correction
Earlier versions of this description cited summed-loss and gradient-norm digests as a correctness check. Those runs did not seed the random base weights, so each process drew different weights. Elementwise, the fix-on and fix-off gradients from those unseeded runs have cosine 0.098, even though the cold-call digests agree within 1%. The digests were not evidence of equivalence and are withdrawn. The numerics above come from seeded runs whose weight and input fingerprints match.
Limits
gc.collect()andgc.freeze()after its first job when layers are compiled, so later collections there scan fewer objects. A larger Python heap, more compiled regions or recompile churn would scale it, and there is no throttle. The measured warm call made no explicit collections. Every checkpointed call still reads the toggle and checks the callback registration at entry and return; that overhead was not timed separately.recompute_granularityunset orselective, no checkpointed layer call exists to collect at, and behavior is unchanged from today.torch._dynamo.callback_handler, the same API ART's compile telemetry already uses. The underlying retention (guard tracebacks capturing live frames) is dynamo behavior. This is a workaround.🤖 Generated with Claude Code