diff --git a/src/art/megatron/lora.py b/src/art/megatron/lora.py index f3baead14..6d0e32644 100644 --- a/src/art/megatron/lora.py +++ b/src/art/megatron/lora.py @@ -1694,6 +1694,7 @@ def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor | None]: ) alpha = beta.clone() adapter_output = torch.cat([qkv, z, beta, alpha], dim=-1) + del qkv, z, beta, alpha return linear_output + adapter_output, bias diff --git a/tests/integration/megatron/gdn_shared_prefix/test_real_gdn_tp_lora.py b/tests/integration/megatron/gdn_shared_prefix/test_real_gdn_tp_lora.py index 0314fd681..191ff5056 100644 --- a/tests/integration/megatron/gdn_shared_prefix/test_real_gdn_tp_lora.py +++ b/tests/integration/megatron/gdn_shared_prefix/test_real_gdn_tp_lora.py @@ -19,6 +19,8 @@ ) from torch.distributed import destroy_process_group, init_process_group # noqa: E402 import torch.multiprocessing as mp # noqa: E402 +from torch.multiprocessing.reductions import StorageWeakRef # noqa: E402 +from torch.utils._python_dispatch import TorchDispatchMode # noqa: E402 from art.megatron.lora import apply_lora_adapters # noqa: E402 from art.megatron.model_support import QWEN3_5_MOE_SPEC # noqa: E402 @@ -63,6 +65,70 @@ def test_real_qwen35_gdn_lora_gradients_match_flattened() -> None: assert _gdn_lora_grad_names(packed_gdn) +class _AdapterOutputLifetime(TorchDispatchMode): + """Observe storage ownership without retaining projection tensors.""" + + def __init__(self): + super().__init__() + self.parts = () + self.adapter = None + self.at_add = [] + + def __torch_dispatch__(self, func, types, args: tuple = (), kwargs=None): + if ( + func is torch.ops.aten.add.Tensor + and self.adapter is not None + and not self.adapter.expired() + and args[1].untyped_storage()._cdata == self.adapter.cdata + ): + self.at_add.append(tuple(part.expired() for part in self.parts)) + output = func(*args, **(kwargs or {})) + if func is torch.ops.aten.cat.default and len(args[0]) == 4: + self.parts = tuple(StorageWeakRef(t.untyped_storage()) for t in args[0]) + self.adapter = StorageWeakRef(output.untyped_storage()) + return output + + +@pytest.mark.skipif( + not torch.cuda.is_available(), + reason="CUDA is required for real Megatron/FLA GDN LoRA coverage.", +) +@pytest.mark.parametrize("grad_enabled", (False, True), ids=("no_grad", "grad")) +def test_real_qwen35_gdn_lora_releases_components_before_add( + grad_enabled: bool, +) -> None: + with _single_rank_model_parallel(), torch.set_grad_enabled(grad_enabled): + gdn, _ = _make_matching_gdn_pair(tp_size=1, lora=True) + projection = gdn.in_proj + x = torch.randn( + 16, + 1, + 64, + device="cuda", + dtype=GDN_CORRECTNESS_DTYPE, + requires_grad=grad_enabled, + ) + with _AdapterOutputLifetime() as lifetime: + output, bias = projection(x) + # One observed add is required; an empty observation must fail. + assert lifetime.at_add == [(True, True, True, True)] + assert bias is None + assert torch.isfinite(output).all() + assert output.requires_grad == grad_enabled + if grad_enabled: + output.square().mean().backward() + gradients = [ + parameter.grad + for lora in (projection.qkv_lora, projection.z_lora) + for parameter in (lora.A_T, lora.B_T) + ] + assert len(gradients) == 4 + for gradient in [x.grad, *gradients]: + assert gradient is not None + assert torch.isfinite(gradient).all() + assert torch.count_nonzero(gradient) > 0 + + @pytest.mark.skipif( not torch.cuda.is_available() or torch.cuda.device_count() < 2, reason="At least two CUDA devices are required for TP2 GDN coverage.",