From 5a4650a7351425344e1f8963bb2bd4a92fe1a7ec Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 02:26:24 +0000 Subject: [PATCH 1/5] Collect dynamo's cyclic garbage after compiling a checkpointed layer 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) --- src/art/megatron/lora.py | 31 ++++++++++ tests/unit/test_megatron_compile_garbage.py | 64 +++++++++++++++++++++ 2 files changed, 95 insertions(+) create mode 100644 tests/unit/test_megatron_compile_garbage.py diff --git a/src/art/megatron/lora.py b/src/art/megatron/lora.py index a7fdfb3c0..c3cb76bce 100644 --- a/src/art/megatron/lora.py +++ b/src/art/megatron/lora.py @@ -3,6 +3,7 @@ import contextvars from dataclasses import dataclass, replace import functools +import gc import importlib import json import math @@ -77,11 +78,40 @@ def use_lora_slot(ref: LoRASlotRef | None) -> Iterator[None]: _CURRENT_LORA_SLOT.reset(token) +# Dynamo compilation can leave reference cycles whose frames still hold the +# traced call's real activations. A grad-mode recompile during backward's first +# recompute would otherwise keep one layer's MoE tensors alive through the rest +# of backward, until the cyclic collector happens to run. +_COMPILE_GARBAGE = False + + +def _mark_compile_garbage(_args: Any) -> None: + global _COMPILE_GARBAGE + _COMPILE_GARBAGE = True + + +def _collect_compile_garbage() -> None: + global _COMPILE_GARBAGE + if _COMPILE_GARBAGE and not torch.compiler.is_compiling(): + _COMPILE_GARBAGE = False + gc.collect() + + +def install_compile_garbage_collection() -> None: + """Collect dynamo's cyclic garbage at the next checkpointed call after a compile.""" + if os.environ.get("ART_COLLECT_COMPILE_GARBAGE", "1") in {"0", "false", "False"}: + return + handler = torch._dynamo.callback_handler + if _mark_compile_garbage not in handler.end_callbacks: + handler.register_end_callback(_mark_compile_garbage) + + def _with_captured_lora_slot(function: _F) -> _F: context = _CURRENT_LORA_SLOT.get() @functools.wraps(function) def wrapped(*args: Any, **kwargs: Any) -> Any: + _collect_compile_garbage() token = _CURRENT_LORA_SLOT.set(context) try: return function(*args, **kwargs) @@ -160,6 +190,7 @@ def patch(target: str, name: str, function_index: int) -> None: install_lora_checkpoint_context_hooks() +install_compile_garbage_collection() @dataclass(frozen=True) diff --git a/tests/unit/test_megatron_compile_garbage.py b/tests/unit/test_megatron_compile_garbage.py new file mode 100644 index 000000000..a0c8538a9 --- /dev/null +++ b/tests/unit/test_megatron_compile_garbage.py @@ -0,0 +1,64 @@ +import gc +import weakref + +import pytest +import torch + +lora = pytest.importorskip("art.megatron.lora") + + +class _Cycle: + pass + + +def _garbage_holding_tensor() -> weakref.ref: + # A tensor reachable only through a reference cycle, as dynamo's traced + # frames hold a recompiled call's activations. + first, second = _Cycle(), _Cycle() + first.other, second.other = second, first + first.tensor = torch.empty(1024) + return weakref.ref(first) + + +@pytest.fixture +def no_automatic_gc(monkeypatch): + monkeypatch.setattr(lora, "_COMPILE_GARBAGE", False) + gc.collect() + gc.disable() + try: + yield + finally: + gc.enable() + + +def test_next_checkpointed_call_collects_compile_garbage(no_automatic_gc): + calls = [] + wrapped = lora._with_captured_lora_slot(lambda: calls.append(1)) + held = _garbage_holding_tensor() + wrapped() + assert held() is not None # nothing compiled, nothing collected + lora._mark_compile_garbage(None) + wrapped() + assert held() is None and calls == [1, 1] + assert lora._COMPILE_GARBAGE is False + # Only one collection per compile: later calls do not collect again. + held = _garbage_holding_tensor() + wrapped() + assert held() is not None + + +def test_collection_is_registered_once_and_can_be_disabled(monkeypatch): + handler = torch._dynamo.callback_handler + callbacks = list(handler.end_callbacks) + monkeypatch.setattr(handler, "end_callbacks", []) + try: + monkeypatch.setenv("ART_COLLECT_COMPILE_GARBAGE", "0") + lora.install_compile_garbage_collection() + assert handler.end_callbacks == [] + monkeypatch.delenv("ART_COLLECT_COMPILE_GARBAGE") + lora.install_compile_garbage_collection() + lora.install_compile_garbage_collection() + assert handler.end_callbacks == [lora._mark_compile_garbage] + finally: + monkeypatch.setattr(handler, "end_callbacks", callbacks) + assert lora._mark_compile_garbage in callbacks From 767042a3d93135fa41426186573f64b4789f7b32 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 02:49:01 +0000 Subject: [PATCH 2/5] Collect compile garbage before backward and survive dynamo resets 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) --- src/art/distributed/host_admission.py | 1 + src/art/megatron/lora.py | 38 ++++++++--- tests/unit/test_megatron_compile_garbage.py | 76 +++++++++++++++++++-- 3 files changed, 100 insertions(+), 15 deletions(-) diff --git a/src/art/distributed/host_admission.py b/src/art/distributed/host_admission.py index d99a59b35..7e01c4bab 100644 --- a/src/art/distributed/host_admission.py +++ b/src/art/distributed/host_admission.py @@ -38,6 +38,7 @@ "triton", ) RUNTIME_ENVIRONMENT_KEYS = { + "ART_COLLECT_COMPILE_GARBAGE", "ART_DISABLE_MEGATRON_COMPILE", "ART_MEGATRON_ALLOW_UNVALIDATED_ARCH", "ART_MEGATRON_ENABLE_MOE_ROUTING_REPLAY", diff --git a/src/art/megatron/lora.py b/src/art/megatron/lora.py index c3cb76bce..5530f2d3e 100644 --- a/src/art/megatron/lora.py +++ b/src/art/megatron/lora.py @@ -6,9 +6,11 @@ import gc import importlib import json +import logging import math import os import re +import time from typing import Any, Callable, Literal, NamedTuple, TypeVar, cast from megatron.bridge.models.gpt_provider import GPTModelProvider @@ -78,6 +80,8 @@ def use_lora_slot(ref: LoRASlotRef | None) -> Iterator[None]: _CURRENT_LORA_SLOT.reset(token) +_logger = logging.getLogger(__name__) + # Dynamo compilation can leave reference cycles whose frames still hold the # traced call's real activations. A grad-mode recompile during backward's first # recompute would otherwise keep one layer's MoE tensors alive through the rest @@ -90,20 +94,34 @@ def _mark_compile_garbage(_args: Any) -> None: _COMPILE_GARBAGE = True -def _collect_compile_garbage() -> None: - global _COMPILE_GARBAGE - if _COMPILE_GARBAGE and not torch.compiler.is_compiling(): - _COMPILE_GARBAGE = False - gc.collect() +def install_compile_garbage_collection() -> bool: + """Mark compiles for collection at checkpointed calls; False when disabled. - -def install_compile_garbage_collection() -> None: - """Collect dynamo's cyclic garbage at the next checkpointed call after a compile.""" + Idempotent, and re-registers after ``torch._dynamo.reset()`` clears the + callback handler. + """ if os.environ.get("ART_COLLECT_COMPILE_GARBAGE", "1") in {"0", "false", "False"}: - return + return False handler = torch._dynamo.callback_handler if _mark_compile_garbage not in handler.end_callbacks: handler.register_end_callback(_mark_compile_garbage) + return True + + +def _collect_compile_garbage() -> None: + global _COMPILE_GARBAGE + # Check tracing first: a traced read of the global would guard on it. + if torch.compiler.is_compiling() or not install_compile_garbage_collection(): + return + if _COMPILE_GARBAGE: + _COMPILE_GARBAGE = False + started = time.perf_counter() + collected = gc.collect() + _logger.debug( + "Collected %d objects after a dynamo compile in %.3fs", + collected, + time.perf_counter() - started, + ) def _with_captured_lora_slot(function: _F) -> _F: @@ -117,6 +135,8 @@ def wrapped(*args: Any, **kwargs: Any) -> Any: return function(*args, **kwargs) finally: _CURRENT_LORA_SLOT.reset(token) + # A compile inside this call has unwound; collect before its backward. + _collect_compile_garbage() return cast(_F, wrapped) diff --git a/tests/unit/test_megatron_compile_garbage.py b/tests/unit/test_megatron_compile_garbage.py index a0c8538a9..208fbfa8b 100644 --- a/tests/unit/test_megatron_compile_garbage.py +++ b/tests/unit/test_megatron_compile_garbage.py @@ -3,6 +3,7 @@ import pytest import torch +import torch.utils.checkpoint lora = pytest.importorskip("art.megatron.lora") @@ -22,6 +23,7 @@ def _garbage_holding_tensor() -> weakref.ref: @pytest.fixture def no_automatic_gc(monkeypatch): + monkeypatch.delenv("ART_COLLECT_COMPILE_GARBAGE", raising=False) monkeypatch.setattr(lora, "_COMPILE_GARBAGE", False) gc.collect() gc.disable() @@ -47,18 +49,80 @@ def test_next_checkpointed_call_collects_compile_garbage(no_automatic_gc): assert held() is not None -def test_collection_is_registered_once_and_can_be_disabled(monkeypatch): +def test_compile_inside_a_call_is_collected_before_it_returns(no_automatic_gc): + # The compiling call's own garbage is freed before its backward runs. + held = [] + + def compiling_call(): + held.append(_garbage_holding_tensor()) + lora._mark_compile_garbage(None) + + lora._with_captured_lora_slot(compiling_call)() + assert held[0]() is None and lora._COMPILE_GARBAGE is False + + +def test_tracing_keeps_the_mark_for_a_later_call(no_automatic_gc, monkeypatch): + wrapped = lora._with_captured_lora_slot(lambda: None) + held = _garbage_holding_tensor() + lora._mark_compile_garbage(None) + with monkeypatch.context() as tracing: + tracing.setattr(torch.compiler, "is_compiling", lambda: True) + wrapped() + assert held() is not None and lora._COMPILE_GARBAGE is True + lora._with_captured_lora_slot(lambda: None)() + assert held() is None + + +def test_registration_survives_dynamo_reset_and_can_be_disabled(monkeypatch): handler = torch._dynamo.callback_handler callbacks = list(handler.end_callbacks) - monkeypatch.setattr(handler, "end_callbacks", []) + monkeypatch.setattr(handler, "end_callbacks", []) # as torch._dynamo.reset() try: monkeypatch.setenv("ART_COLLECT_COMPILE_GARBAGE", "0") - lora.install_compile_garbage_collection() + assert lora.install_compile_garbage_collection() is False + lora._with_captured_lora_slot(lambda: None)() assert handler.end_callbacks == [] monkeypatch.delenv("ART_COLLECT_COMPILE_GARBAGE") - lora.install_compile_garbage_collection() - lora.install_compile_garbage_collection() + lora._with_captured_lora_slot(lambda: None)() + lora._with_captured_lora_slot(lambda: None)() assert handler.end_callbacks == [lora._mark_compile_garbage] finally: monkeypatch.setattr(handler, "end_callbacks", callbacks) - assert lora._mark_compile_garbage in callbacks + + +def test_real_compile_marks_garbage(no_automatic_gc): + torch._dynamo.reset() + lora.install_compile_garbage_collection() + + def double(x): + return x * 2 + 1 + + torch.compile(double, backend="eager")(torch.ones(3)) + assert lora._COMPILE_GARBAGE is True + + +@pytest.mark.parametrize("use_reentrant", [False, True]) +def test_checkpoint_backward_with_compile_keeps_gradients( + no_automatic_gc, use_reentrant +): + torch._dynamo.reset() + lora.install_compile_garbage_collection() + torch.manual_seed(0) + weight = torch.randn(8, 8) + + def layer(x, w): + return torch.tanh(x @ w).square() + + compiled = torch.compile(layer, backend="eager") + x = torch.randn(4, 8, requires_grad=True) + w = weight.clone().requires_grad_() + reference_x = x.detach().clone().requires_grad_() + reference_w = weight.clone().requires_grad_() + layer(reference_x, reference_w).sum().backward() + # The patched checkpoint wraps the function, so recompute runs the + # collection hook, and compiles (forward and recompute) set the mark. + out = torch.utils.checkpoint.checkpoint(compiled, x, w, use_reentrant=use_reentrant) + out.sum().backward() + torch.testing.assert_close(x.grad, reference_x.grad) + torch.testing.assert_close(w.grad, reference_w.grad) + assert lora._COMPILE_GARBAGE is False From 55e72d3fa3004e6de899ba8d142415a699472744 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 03:09:18 +0000 Subject: [PATCH 3/5] Keep a failed call's error and defer collection during graph capture 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) --- src/art/megatron/lora.py | 30 +++++++++++------- tests/unit/test_megatron_compile_garbage.py | 35 +++++++++++++++++++-- 2 files changed, 50 insertions(+), 15 deletions(-) diff --git a/src/art/megatron/lora.py b/src/art/megatron/lora.py index 5530f2d3e..f3baead14 100644 --- a/src/art/megatron/lora.py +++ b/src/art/megatron/lora.py @@ -113,15 +113,19 @@ def _collect_compile_garbage() -> None: # Check tracing first: a traced read of the global would guard on it. if torch.compiler.is_compiling() or not install_compile_garbage_collection(): return - if _COMPILE_GARBAGE: - _COMPILE_GARBAGE = False - started = time.perf_counter() - collected = gc.collect() - _logger.debug( - "Collected %d objects after a dynamo compile in %.3fs", - collected, - time.perf_counter() - started, - ) + # Finalizers must not run inside a CUDA graph capture; keep the mark. + if not _COMPILE_GARBAGE or ( + torch.cuda.is_initialized() and torch.cuda.is_current_stream_capturing() + ): + return + _COMPILE_GARBAGE = False + started = time.perf_counter() + collected = gc.collect() + _logger.debug( + "Collected %d objects after a dynamo compile in %.3fs", + collected, + time.perf_counter() - started, + ) def _with_captured_lora_slot(function: _F) -> _F: @@ -132,11 +136,13 @@ def wrapped(*args: Any, **kwargs: Any) -> Any: _collect_compile_garbage() token = _CURRENT_LORA_SLOT.set(context) try: - return function(*args, **kwargs) + result = function(*args, **kwargs) finally: _CURRENT_LORA_SLOT.reset(token) - # A compile inside this call has unwound; collect before its backward. - _collect_compile_garbage() + # A compile inside this call has unwound; collect before its backward. + # Failed calls leave the mark for the next call, keeping their error. + _collect_compile_garbage() + return result return cast(_F, wrapped) diff --git a/tests/unit/test_megatron_compile_garbage.py b/tests/unit/test_megatron_compile_garbage.py index 208fbfa8b..868d29957 100644 --- a/tests/unit/test_megatron_compile_garbage.py +++ b/tests/unit/test_megatron_compile_garbage.py @@ -61,12 +61,41 @@ def compiling_call(): assert held[0]() is None and lora._COMPILE_GARBAGE is False -def test_tracing_keeps_the_mark_for_a_later_call(no_automatic_gc, monkeypatch): +def test_failed_call_raises_its_own_error_and_keeps_the_mark(no_automatic_gc): + held = [] + + def failing_call(): + held.append(_garbage_holding_tensor()) + lora._mark_compile_garbage(None) + raise ValueError("model error") + + with pytest.raises(ValueError, match="model error"): + lora._with_captured_lora_slot(failing_call)() + assert held[0]() is not None and lora._COMPILE_GARBAGE is True + lora._with_captured_lora_slot(lambda: None)() + assert held[0]() is None + + +@pytest.mark.parametrize( + "deferring", + [ + {(torch.compiler, "is_compiling"): True}, + { + (torch.cuda, "is_initialized"): True, + (torch.cuda, "is_current_stream_capturing"): True, + }, + ], + ids=["tracing", "cuda-graph-capture"], +) +def test_deferred_collection_keeps_the_mark_for_a_later_call( + no_automatic_gc, monkeypatch, deferring +): wrapped = lora._with_captured_lora_slot(lambda: None) held = _garbage_holding_tensor() lora._mark_compile_garbage(None) - with monkeypatch.context() as tracing: - tracing.setattr(torch.compiler, "is_compiling", lambda: True) + with monkeypatch.context() as patch: + for (module, name), value in deferring.items(): + patch.setattr(module, name, lambda value=value: value) wrapped() assert held() is not None and lora._COMPILE_GARBAGE is True lora._with_captured_lora_slot(lambda: None)() From 10e32d791a1c8f7db0e51a5081a8bacc659c6e50 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 03:45:39 +0000 Subject: [PATCH 4/5] Declare the test cycle's attributes and cover early-stopped recompute 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) --- tests/unit/test_megatron_compile_garbage.py | 25 ++++++++++++++++++++- 1 file changed, 24 insertions(+), 1 deletion(-) diff --git a/tests/unit/test_megatron_compile_garbage.py b/tests/unit/test_megatron_compile_garbage.py index 868d29957..25421635c 100644 --- a/tests/unit/test_megatron_compile_garbage.py +++ b/tests/unit/test_megatron_compile_garbage.py @@ -9,7 +9,8 @@ class _Cycle: - pass + other: "_Cycle" + tensor: torch.Tensor def _garbage_holding_tensor() -> weakref.ref: @@ -155,3 +156,25 @@ def layer(x, w): torch.testing.assert_close(x.grad, reference_x.grad) torch.testing.assert_close(w.grad, reference_w.grad) assert lora._COMPILE_GARBAGE is False + + +def test_early_stopped_recompute_leaves_the_mark_for_the_next_call(no_automatic_gc): + # Non-reentrant recompute stops by raising out of the wrapped call, so a + # compile during it is collected at the next checkpointed call instead. + torch._dynamo.reset() + lora.install_compile_garbage_collection() + + def layer(x, w): + return torch.tanh(x @ w).square() + + compiled = torch.compile(layer, backend="eager") + x = torch.randn(4, 8, requires_grad=True) + w = torch.randn(8, 8, requires_grad=True) + out = torch.utils.checkpoint.checkpoint(compiled, x, w, use_reentrant=False) + assert lora._COMPILE_GARBAGE is False + torch._dynamo.reset() # the recompute compiles again + with torch.utils.checkpoint.set_checkpoint_early_stop(True): + out.sum().backward() + assert x.grad is not None and lora._COMPILE_GARBAGE is True + lora._with_captured_lora_slot(lambda: None)() + assert lora._COMPILE_GARBAGE is False From 336f940551419b7aed02271ed2033e657d86f966 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 03:55:33 +0000 Subject: [PATCH 5/5] Run the compile-garbage tests in the Megatron CI lane 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) --- .github/workflows/prek.yml | 2 ++ tests/unit/test_megatron_compile_garbage.py | 7 ++++--- 2 files changed, 6 insertions(+), 3 deletions(-) diff --git a/.github/workflows/prek.yml b/.github/workflows/prek.yml index 17e6913da..090a61919 100644 --- a/.github/workflows/prek.yml +++ b/.github/workflows/prek.yml @@ -238,6 +238,7 @@ jobs: tests/unit/test_trainer_rank_shared_memory.py \ tests/unit/test_trainer_rank_converted_memory.py \ tests/unit/test_trainer_rank_split.py \ + tests/unit/test_megatron_compile_garbage.py \ tests/unit/test_trainer_rank_cache_recovery.py::test_dense_cp_exact_demand_fits_after_recovery \ tests/unit/test_trainer_rank_cache_recovery.py::test_dense_cp_exact_demand_refuses_after_recovery \ tests/acceptance/trainer_rank_planner \ @@ -277,4 +278,5 @@ jobs: --ignore=tests/unit/test_trainer_rank_ignored_mixed_head.py \ --ignore=tests/unit/test_trainer_rank_pending_memory.py \ --ignore=tests/unit/test_trainer_rank_shared_memory.py \ + --ignore=tests/unit/test_megatron_compile_garbage.py \ --ignore=tests/unit/test_trainer_rank_converted_memory.py diff --git a/tests/unit/test_megatron_compile_garbage.py b/tests/unit/test_megatron_compile_garbage.py index 25421635c..375981dba 100644 --- a/tests/unit/test_megatron_compile_garbage.py +++ b/tests/unit/test_megatron_compile_garbage.py @@ -170,11 +170,12 @@ def layer(x, w): compiled = torch.compile(layer, backend="eager") x = torch.randn(4, 8, requires_grad=True) w = torch.randn(8, 8, requires_grad=True) - out = torch.utils.checkpoint.checkpoint(compiled, x, w, use_reentrant=False) + # Checkpoint reads the early-stop setting when it runs forward. + with torch.utils.checkpoint.set_checkpoint_early_stop(True): + out = torch.utils.checkpoint.checkpoint(compiled, x, w, use_reentrant=False) assert lora._COMPILE_GARBAGE is False torch._dynamo.reset() # the recompute compiles again - with torch.utils.checkpoint.set_checkpoint_early_stop(True): - out.sum().backward() + out.sum().backward() assert x.grad is not None and lora._COMPILE_GARBAGE is True lora._with_captured_lora_slot(lambda: None)() assert lora._COMPILE_GARBAGE is False