Skip to content

Collect dynamo's cyclic garbage after compiling a checkpointed layer - #955

Merged
bradhilton merged 5 commits into
mainfrom
dalinar/compile-gc
Sep 25, 2026
Merged

bradhilton merged 5 commits into
mainfrom
dalinar/compile-gc

Conversation

@bradhilton

@bradhilton bradhilton commented Sep 25, 2026 •

Copy link
Copy Markdown
Collaborator

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:

  • Incremental peak. max_memory_allocated during the call minus memory_allocated just before it.
  • Planner prediction. Its 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.
  • Absolute peak. Before the cold call, 15.80 GB is allocated (weights, adapters, optimizer state) in every run below. The absolute peak adds that.
run layer compile cold peak, incremental cold peak, absolute planner prediction, incremental
base aab5d5416 on (default) 95.34 GB (byte-identical in 2 runs) 111.13 GB 69.5 GB
base aab5d5416 off (ART_DISABLE_MEGATRON_COMPILE=1) 69.13 GB 84.92 GB 69.5 GB
this PR on 64.14 GB (byte-identical in 4 runs) 79.94 GB 69.5 GB
this PR, ART_COLLECT_COMPILE_GARBAGE=0 on 95.34 GB (byte-identical in 3 runs: 2 on this head, 1 on the first revision) 111.13 GB 69.5 GB

The base rows ran on #954's head 587ae7e6b, whose tree is identical to aab5d5416. 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.

rank this PR, cold peak incremental (absolute) ART_COLLECT_COMPILE_GARBAGE=0 planner prediction
0 31.59 GB (47.38 GB) 45.98 GB (61.77 GB) 29.41 GB
1 36.17 GB (51.97 GB) 51.85 GB (67.65 GB) 29.41 GB
  • Base 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 a grad_mode guard failure and recompiles, and are never freed in that call. In the next call all four are freed inside torch._dynamo OutputGraph.__init__ (guard creation → torch.utils._traceback.extract).

  • Confirmed as cyclic garbage. After the cold call, a gc.collect() under gc.DEBUG_SAVEALL moved the unreachable objects into gc.garbage for inspection. Clearing that list and collecting again freed 32.02 GB of CUDA memory. The unreachable objects included 357 frame objects 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:

  • At entry. This covers compiles that happened outside a checkpointed call.
  • When the call returns. The compile frames have unwound by then, so a compile inside the call is freed before that layer's backward.
    • A call that raises keeps its own exception. Its mark waits for the next call.

Behavior:

  • Skips. Nothing is collected while dynamo is tracing (checked first, so a traced read cannot guard on the mark) or while the current CUDA stream is capturing a graph. The mark is kept for later.
  • Frequency. Each mark is collected once.
  • Re-registration. Every checkpointed call checks the callback's registration, so a torch._dynamo.reset() does not silently disable it.
  • Toggle. ART_COLLECT_COMPILE_GARBAGE=0 disables it. The key is in RUNTIME_ENVIRONMENT_KEYS, so SSH-launched ranks receive it and host admission reports it. Admission does not require it to match across hosts.
  • Logging. Each collection's object count and duration are logged at debug level on 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.

    comparison cold call: rel. L2 / cosine calls 1–2: rel. L2
    fix on vs fix on (noise floor) 1.305e-2 / 0.999915 1.336e-2, 1.329e-2
    fix on vs fix off (each on run) 1.311e-2, 1.309e-2 / 0.999914 5.43e-2–5.44e-2
    • Noise floor. Backward is not deterministic here, so the floor is non-zero.
    • Cold call. Both arms run the same plan, and fix on vs fix off is indistinguishable from the floor.
    • Calls 1–2. These are not a matched comparison. Entering call 1, the off arm still had the cold call's garbage allocated: 50.86 GB absolute before the call, versus 18.84 GB absolute with the fix. The 32.02 GB difference matches the DEBUG_SAVEALL probe. 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):

    call fix on (2 runs) fix off
    0, cold (compiles) 28.3 s / 26.5 s, of which GC 5.5 s / 4.3 s 23.8 s, GC 1.3 s
    1, new shapes (recompiles) 16.7 s / 16.1 s, GC 6.2 s / 5.5 s 12.4 s, GC 0.7 s
    2, warm 5.3 s / 5.3 s, GC 0.01 s 5.9 s, GC 0.03 s

    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):

    • the next call after a mark collects a tensor held only by a cycle, once per mark;
    • a compile inside a call is collected before the call returns;
    • a failing call raises its own error and keeps the mark;
    • tracing and CUDA graph capture defer collection and keep the mark;
    • registration survives a cleared callback handler, and the toggle disables it;
    • a real torch.compile sets the mark;
    • checkpointed compiled backward, reentrant and non-reentrant, matches uncheckpointed gradients;
    • a compile during early-stopped non-reentrant recompute leaves the mark for the next call.

    tests/integration/megatron/model_support/test_compile_flags.py, which calls torch._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

  • Coverage. One fixture shape (8 layers, CP1, random weights), PyTorch-allocated peaks. Calls 1–2 are confounded by the off arm's different plan. This is not end-to-end training qualification.
  • Other topologies. CP2/EP1 was re-run with the fix (above). EP2 snapshots show the same one-layer retention pattern, but were not re-run with the fix.
  • Collection cost. The cost is 3–5.5 s of extra collection time on each of this fixture's first two calls: 15 full collections of 0.47–0.96 s each. The harness does not freeze the heap. ART's Megatron executor runs gc.collect() and gc.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.
  • Cross-thread marks. The mark is process-global, so another thread's checkpointed call can consume it first. Prompt collection is only guaranteed on the compiling thread. Otherwise the garbage waits for the next mark or the automatic collector, which is today's behavior.
  • Needs full-layer recompute. Collection runs in ART's checkpoint wrapper. With recompute_granularity unset or selective, no checkpointed layer call exists to collect at, and behavior is unchanged from today.
  • Failed calls. A call that fails after a compile, such as an OOM in the recompiling layer, leaves its garbage allocated until the next checkpointed call. A retry's admission may therefore see less free memory than it should. Collecting in the planner's OOM path would close this; that is left as a follow-up.
  • Early-stopped recompute. Non-reentrant checkpointing ends recompute by raising out of the wrapped call. A compile during that recompute is therefore collected at the next checkpointed call rather than before the layer's backward (tested). ART's layer recompute (Megatron and Transformer Engine checkpointing) is reentrant, and that is where the measured grad-mode recompile happens. ART's non-reentrant use is the output-head chunks.
  • What it cannot free. Garbage from a compile that no checkpointed call follows, or compiler state that is still reachable, is not freed by this change.
  • Dynamo API. The callback uses 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

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
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
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
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
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>

This branch was successfully deployed

1 active deployment
trainer-rank-gpu-validation — 336f9405 Deployed Sep 25, 2026 by bradhilton via Run on 2x H200 #667
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant