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/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 a7fdfb3c0..f3baead14 100644 --- a/src/art/megatron/lora.py +++ b/src/art/megatron/lora.py @@ -3,11 +3,14 @@ import contextvars from dataclasses import dataclass, replace import functools +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 @@ -77,16 +80,69 @@ 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 +# 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 install_compile_garbage_collection() -> bool: + """Mark compiles for collection at checkpointed calls; False when disabled. + + 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 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 + # 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: 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) + result = function(*args, **kwargs) finally: _CURRENT_LORA_SLOT.reset(token) + # 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) @@ -160,6 +216,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..375981dba --- /dev/null +++ b/tests/unit/test_megatron_compile_garbage.py @@ -0,0 +1,181 @@ +import gc +import weakref + +import pytest +import torch +import torch.utils.checkpoint + +lora = pytest.importorskip("art.megatron.lora") + + +class _Cycle: + other: "_Cycle" + tensor: torch.Tensor + + +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.delenv("ART_COLLECT_COMPILE_GARBAGE", raising=False) + 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_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_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 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)() + 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", []) # as torch._dynamo.reset() + try: + monkeypatch.setenv("ART_COLLECT_COMPILE_GARBAGE", "0") + 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._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) + + +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 + + +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) + # 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 + 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