Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions .github/workflows/prek.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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 \
Expand Down Expand Up @@ -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
1 change: 1 addition & 0 deletions src/art/distributed/host_admission.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
59 changes: 58 additions & 1 deletion src/art/megatron/lora.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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)
Expand Down
181 changes: 181 additions & 0 deletions tests/unit/test_megatron_compile_garbage.py
Original file line number Diff line number Diff line change
@@ -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
Loading