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
1 change: 1 addition & 0 deletions src/art/megatron/lora.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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.",
Expand Down
Loading