Skip to content
1 change: 1 addition & 0 deletions qa/L0_pytorch_unittest/test.sh
Original file line number Diff line number Diff line change
Expand Up @@ -54,6 +54,7 @@ python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_multi_tensor.xml
python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_fusible_ops.xml $TE_PATH/tests/pytorch/test_fusible_ops.py || test_fail "test_fusible_ops.py"
python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_selective_activation_checkpoint.xml $TE_PATH/tests/pytorch/layernorm_mlp/test_selective_activation_checkpoint.py || test_fail "test_selective_activation_checkpoint.py"
python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_distributed_weight.xml $TE_PATH/tests/pytorch/test_distributed_weight.py || test_fail "test_distributed_weight.py"
python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_module_distributed_weight_saved_tensor_hooks.xml $TE_PATH/tests/pytorch/test_module_distributed_weight_saved_tensor_hooks.py || test_fail "test_module_distributed_weight_saved_tensor_hooks.py"
python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_backward_override.xml $TE_PATH/tests/pytorch/test_backward_override.py || test_fail "test_backward_override.py"
python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_permutation.xml $TE_PATH/tests/pytorch/test_permutation.py || test_fail "test_permutation.py"
python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_cross_entropy.xml $TE_PATH/tests/pytorch/test_cross_entropy.py || test_fail "test_cross_entropy.py"
Expand Down
237 changes: 237 additions & 0 deletions tests/pytorch/test_module_distributed_weight_saved_tensor_hooks.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,237 @@
# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
#
# See LICENSE for license information.

"""DistributedWeight dispatch in ``Linear`` / ``GroupedLinear`` under saved-tensor hooks.

Both modules pass the DistributedWeight parameter itself through ``save_for_backward`` and
used to gate their backward on ``is_distributed_weight(saved_weight)``. That only holds while
no ``torch.autograd.graph.saved_tensors_hooks`` are installed: with hooks active, autograd
unpacks a saved leaf as a *fresh plain tensor* (whatever the unpack hook returns), dropping the
Python subclass and its ``is_distributed_weight`` marker. Backward then took the plain-parameter
branch, whose weakrefs point at the transient all-gathered weights, and failed with
"weight was removed while fuse_wgrad_accumulation=True" (or silently used the wrong weights).

Megatron-LM's fine-grained activation offloading installs exactly such hooks around expert
layers, which is how this surfaced. The fix keeps the DistributedWeight objects on the autograd
context and prefers them in backward. These tests pin that: with identity hooks installed, the
distributed path must still be taken and produce the same results as without hooks.

TE ships no DistributedWeight implementer (they live in the caller, e.g. Megatron-LM's GTP), so
an in-repo fake is used. It applies distinct scales in the forward / backward materialize hooks,
so taking the wrong branch fails a specific assertion rather than drifting numerically.
"""

import pytest
import torch
from torch.autograd.graph import saved_tensors_hooks

import transformer_engine.pytorch as te
from transformer_engine.pytorch.module.grouped_linear import (
is_module_grouped_tensor_path_supported,
)

FWD_SCALE = 2.0
BWD_SCALE = 4.0
IN_F, OUT_F, TOKENS = 256, 512, 128
DTYPE, DEVICE = torch.bfloat16, "cuda"

FUSE_WGRAD = pytest.mark.parametrize("fuse_wgrad", [False, True], ids=["unfused", "fused-wgrad"])


class _FakeDistWeight(torch.nn.Parameter):
"""Fake distributed weight; the group leader holds every member in ``_group``."""

is_distributed_weight = True

def materialize_group_for_forward(self):
self.calls["fwd"] += 1
return [(w.detach() * FWD_SCALE).requires_grad_(True) for w in self._group]

def materialize_group_for_backward(self, **kwargs):
self.calls["bwd"] += 1
return [w.detach() * BWD_SCALE for w in self._group]

def finalize_group_grads(self, wgrads, **kwargs):
self.calls["finalize"] += 1
wl = list(wgrads) if isinstance(wgrads, (list, tuple)) else [wgrads]
if not self.fuse_wgrad_accumulation:
return [g.clone() for g in wl]
for w, g in zip(self._group, wl):
w.main_grad.add_(g.to(w.main_grad.dtype)) # stands in for the reduce-scatter
w.grad_added_to_main_grad = True
return [torch.zeros_like(g) for g in wl] # real grad is in main_grad

def grad_buffer(self):
return self.wgrad_scratch


def _identity_hooks():
"""Hooks that change nothing; their mere presence makes autograd re-wrap saved leaves."""
return saved_tensors_hooks(lambda t: t, lambda t: t)


def _install_fakes(module, weight_names, fuse_wgrad):
"""Replace the module's weights with fakes sharing one group; return the leader."""
fakes = []
for name in weight_names:
fake = _FakeDistWeight(getattr(module, name).data)
fake.calls = {"fwd": 0, "bwd": 0, "finalize": 0}
fake.fuse_wgrad_accumulation = fuse_wgrad
fake.main_grad = torch.zeros((OUT_F, IN_F), dtype=torch.float32, device=DEVICE)
fake.wgrad_scratch = torch.zeros_like(fake.main_grad)
fake.grad_added_to_main_grad = False
setattr(module, name, fake)
fakes.append(fake)
for fake in fakes:
fake._group = fakes
return fakes[0]


def _check_dist_path_kept(
module, weight_names, leader, reference, ref_leader, out, ref_out, x, ref_x, fuse_wgrad
):
"""Hooked run must have used every DistributedWeight hook and match the unhooked run."""
assert leader.calls["fwd"] > 0, "materialize_group_for_forward never called"
assert leader.calls["bwd"] > 0, (
"materialize_group_for_backward never called: backward lost the DistributedWeight "
"after saved_tensors_hooks unpacked the saved weight as a plain tensor"
)
assert leader.calls["finalize"] > 0, "finalize_group_grads never called"
assert leader.calls == ref_leader.calls
# Same kernels on the same data: results are bitwise equal to the unhooked run.
torch.testing.assert_close(out, ref_out, rtol=0, atol=0)
torch.testing.assert_close(x.grad, ref_x.grad, rtol=0, atol=0)
for name in weight_names:
w, ref_w = getattr(module, name), getattr(reference, name)
if fuse_wgrad:
assert w.grad_added_to_main_grad is True
torch.testing.assert_close(w.main_grad, ref_w.main_grad, rtol=0, atol=0)
assert torch.count_nonzero(w.main_grad) > 0
else:
assert w.grad is not None and ref_w.grad is not None
torch.testing.assert_close(w.grad, ref_w.grad, rtol=0, atol=0)


def _skip_without_cuda():
if not torch.cuda.is_available():
pytest.skip("requires CUDA")


def test_saved_tensor_hooks_unpack_subclass_as_plain_tensor():
"""Premise: with hooks installed, a saved DistributedWeight comes back as a plain tensor."""
_skip_without_cuda()
seen = {}

class _Probe(torch.autograd.Function):
@staticmethod
def forward(ctx, x, w):
ctx.save_for_backward(w)
return x * 2

@staticmethod
def backward(ctx, g):
(w,) = ctx.saved_tensors
seen["type"] = type(w)
seen["marked"] = bool(getattr(w, "is_distributed_weight", False))
return g * 2, None

w = _FakeDistWeight(torch.ones(4, device=DEVICE))
x = torch.ones(4, device=DEVICE, requires_grad=True)

_Probe.apply(x, w).sum().backward()
assert seen["type"] is _FakeDistWeight and seen["marked"]

with _identity_hooks():
_Probe.apply(x, w).sum().backward()
assert seen["type"] is torch.Tensor and not seen["marked"]


@FUSE_WGRAD
def test_linear_distributed_weight_under_saved_tensor_hooks(fuse_wgrad):
"""``Linear`` must keep the distributed path in backward when saved-tensor hooks are on."""
_skip_without_cuda()
torch.manual_seed(0)
module, reference = (
te.Linear(
IN_F,
OUT_F,
bias=False,
device=DEVICE,
params_dtype=DTYPE,
fuse_wgrad_accumulation=fuse_wgrad,
)
for _ in range(2)
)
reference.load_state_dict(module.state_dict())
leader = _install_fakes(module, ["weight"], fuse_wgrad)
ref_leader = _install_fakes(reference, ["weight"], fuse_wgrad)

x = torch.randn(TOKENS, IN_F, dtype=DTYPE, device=DEVICE, requires_grad=True)
ref_x = x.detach().clone().requires_grad_(True)

ref_out = reference(ref_x)
ref_out.sum().backward()
with _identity_hooks():
out = module(x)
out.sum().backward()

_check_dist_path_kept(
module, ["weight"], leader, reference, ref_leader, out, ref_out, x, ref_x, fuse_wgrad
)
# The scales prove which weights were used: FWD_SCALE in forward, BWD_SCALE in dgrad.
plain = te.Linear(IN_F, OUT_F, bias=False, device=DEVICE, params_dtype=DTYPE)
plain.load_state_dict({"weight": leader.data}, strict=False)
px = x.detach().clone().requires_grad_(True)
pout = plain(px)
pout.sum().backward()
torch.testing.assert_close(out.float(), FWD_SCALE * pout.float(), rtol=1e-5, atol=1e-5)
torch.testing.assert_close(x.grad.float(), BWD_SCALE * px.grad.float(), rtol=1e-5, atol=1e-5)


@FUSE_WGRAD
@pytest.mark.parametrize("num_gemms", [2, 4])
@pytest.mark.parametrize(
"use_grouped_tensor", [False, True], ids=["split-quantize", "grouped-tensor"]
)
Comment thread
greptile-apps[bot] marked this conversation as resolved.
def test_grouped_linear_distributed_weight_under_saved_tensor_hooks(
num_gemms, fuse_wgrad, use_grouped_tensor
):
"""``GroupedLinear`` (both GEMM paths) must keep the distributed path under hooks."""
_skip_without_cuda()
if use_grouped_tensor and not is_module_grouped_tensor_path_supported(None, DTYPE):
# GroupedLinear would silently fall back to split-quantize, which the other case
# already covers; skip rather than report coverage of the native backward.
pytest.skip("native grouped-tensor path unsupported on this device / cuBLASLt")
torch.manual_seed(0)
names = [f"weight{i}" for i in range(num_gemms)]
module, reference = (
te.GroupedLinear(
num_gemms,
IN_F,
OUT_F,
bias=False,
device=DEVICE,
params_dtype=DTYPE,
fuse_wgrad_accumulation=fuse_wgrad,
use_grouped_tensor=use_grouped_tensor,
)
for _ in range(2)
)
reference.load_state_dict(module.state_dict())
leader = _install_fakes(module, names, fuse_wgrad)
ref_leader = _install_fakes(reference, names, fuse_wgrad)

m_splits = torch.full((num_gemms,), TOKENS, dtype=torch.int64, device=DEVICE)
x = torch.randn(num_gemms * TOKENS, IN_F, dtype=DTYPE, device=DEVICE, requires_grad=True)
ref_x = x.detach().clone().requires_grad_(True)

ref_out = reference(ref_x, m_splits)
ref_out.sum().backward()
with _identity_hooks():
out = module(x, m_splits)
out.sum().backward()

_check_dist_path_kept(
module, names, leader, reference, ref_leader, out, ref_out, x, ref_x, fuse_wgrad
)
22 changes: 21 additions & 1 deletion transformer_engine/pytorch/module/grouped_linear.py
Original file line number Diff line number Diff line change
Expand Up @@ -594,6 +594,9 @@ def _forward_grouped_tensor(
ctx.single_grouped_weight = single_grouped_weight
ctx.single_grouped_bias = single_grouped_bias
ctx.is_dist_weight = is_dist_weight
# Same reason as in ``forward``: saved-tensor hooks unpack the saved
# DistributedWeight shards as plain tensors, so keep the originals on ctx.
ctx.dist_weights = list(origin_weights) if is_dist_weight else None
ctx.fp8_for_weight_prep = fp8
if fuse_wgrad_accumulation and ctx.weights_requires_grad:
# Weakref the parameters, not the gathered copies: those can be dead by backward.
Expand Down Expand Up @@ -974,6 +977,12 @@ def forward(
# weight_requires_grad was read from the parameters, before materialization
# replaced ``weights`` with gathered copies that need not carry requires_grad.
ctx.weights_requires_grad = weight_requires_grad
# DistributedWeight objects must survive to backward as Python objects. They are
# also passed through save_for_backward (as saved_weights), but when saved-tensor
# hooks are active (e.g. activation offloading) autograd unpacks saved leaves as
# fresh plain tensors, dropping the subclass and its ``is_distributed_weight``
# marker. Keep the originals on ctx and prefer them in backward.
ctx.dist_weights = list(origin_weights) if is_dist_weight else None
Comment thread
xrennvidia marked this conversation as resolved.
if fuse_wgrad_accumulation and ctx.weights_requires_grad:
# Keep weakrefs to weights to preserve attributes like main_grad
# when we need to modify the weight python objects. Target the parameters:
Expand Down Expand Up @@ -1089,7 +1098,12 @@ def _backward_grouped_tensor(
main_grads = [None] * num_weight_args
is_dist_weight = getattr(ctx, "is_dist_weight", False)
if is_dist_weight:
# Forward saved the shards, not the gathered copies.
# Forward saved the shards, not the gathered copies. Prefer the originals kept
# on ctx: saved-tensor hooks may have unpacked the shards as plain tensors.
dist_weights = getattr(ctx, "dist_weights", None)
if dist_weights is not None:
weight_tensors = dist_weights
ctx.dist_weights = None

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Second backward loses distributed weights

When GroupedLinear runs under saved-tensor hooks and its output is backpropagated twice with retain_graph=True, the first backward clears ctx.dist_weights here; the split backward path clears it too. On the second backward, the saved weights unpack as plain tensors, so the distributed-weight path is lost again. Backward can then fail when it tries to use the unsaved gathered weights, or skip distributed gradient finalization. Keep the original weights available for each backward while the graph is retained.

origin_weights = list(weight_tensors)
if ctx.fuse_wgrad_accumulation and ctx.weights_requires_grad:
main_grads = [main_grad_func() for main_grad_func in ctx.main_grad_funcs]
Expand Down Expand Up @@ -1351,6 +1365,12 @@ def backward(
weights = saved_tensors[N : 2 * N]
saved_weights = saved_tensors[2 * N : 3 * N]
biases = saved_tensors[3 * N : 4 * N]
dist_weights = getattr(ctx, "dist_weights", None)
if dist_weights is not None:
# See forward: saved-tensor hooks may have replaced the DistributedWeight
# objects with plain aliases; the originals were kept on ctx.
saved_weights = dist_weights
ctx.dist_weights = None

# Restore from weakrefs to get original weight python objects
# (preserves attributes like main_grad, grad_added_to_main_grad, etc.)
Expand Down
9 changes: 9 additions & 0 deletions transformer_engine/pytorch/module/linear.py
Original file line number Diff line number Diff line change
Expand Up @@ -291,6 +291,10 @@ class LinearBwdArgs:
fuse_wgrad_accumulation: bool = False
wgrad_store: Optional[Any] = None
origin_weight_ref: Optional[Any] = None
# DistributedWeight object kept as a Python reference: saved-tensor hooks (e.g. activation
# offloading) make autograd unpack ``saved_weight`` as a fresh plain tensor, dropping the
# subclass and its ``is_distributed_weight`` marker.
dist_weight: Optional[Any] = None
origin_weight_overwrites_main_grad: bool = False
main_grad_func: Optional[Callable[[], torch.Tensor]] = None

Expand Down Expand Up @@ -1060,6 +1064,7 @@ def _linear_setup_ctx(
bwd_args.is_first_microbatch = fwd_args.is_first_microbatch
bwd_args.fuse_wgrad_accumulation = fuse_wgrad_accumulation
bwd_args.wgrad_store = fwd_args.wgrad_store
bwd_args.dist_weight = weight if is_distributed_weight(weight) else None
if fuse_wgrad_accumulation and fwd_args.weight_requires_grad:
bwd_args.origin_weight_ref = weakref.ref(weight)
bwd_args.origin_weight_overwrites_main_grad = getattr(weight, "overwrite_main_grad", False)
Expand Down Expand Up @@ -1116,6 +1121,10 @@ def _linear_backward_impl(args: LinearBwdArgs) -> Tuple[Union[torch.Tensor, None
inputmat = args.inputmat
weight_fp8 = args.weight_fp8
saved_weight = args.saved_weight
if args.dist_weight is not None:
# Prefer the object kept in forward; see LinearBwdArgs.dist_weight.
saved_weight = args.dist_weight
args.dist_weight = None
is_dist_weight = is_distributed_weight(saved_weight)
bias = args.bias
input_quantizer = args.input_quantizer
Expand Down
Loading