From b5b9cc2afc3f49cfa4f165cdfe620bcc99144d18 Mon Sep 17 00:00:00 2001 From: 0z5a Date: Mon, 21 Sep 2026 04:50:56 +0000 Subject: [PATCH 1/4] Accumulate delayed wgrads into shared (tied) parameters backward_dw() assigned the popped delayed wgrad to Parameter.grad. When two modules share one Parameter, each pops its own contribution and the later one silently replaced the earlier, so the shared gradient came out too small -- on the issue's own 16x16 / batch=2 BF16 repro it was 0.6015x the regular path, with 256/256 elements differing. The bias takes the same pop path and hits the same bug, with an extra twist: the second module's bgrad is numerically ~0, so it wiped out a contribution the first module had already written correctly (1.0984 delivered vs 33.0859 expected). Initialise on the first contribution of the step and add on every later one, for both the weight and the bias. dtype/shape/buffer-alias contracts are unchanged, and the fused wgrad-accumulation branch is untouched. tests/pytorch/test_tied_delayed_wgrad.py covers the matrix: tied and non-tied, both backward_dw() orders, multiple microbatches, two optimizer steps with set_to_none True/False, the bias case, and parity with a plain nn.Linear oracle. It uses only the public te.Linear API so it cannot be satisfied by a shared implementation detail of the fix. Without the fix it fails 5 and passes 3; with it, 8 pass. Signed-off-by: 0z5a --- tests/pytorch/test_tied_delayed_wgrad.py | 177 ++++++++++++++++++++++ transformer_engine/pytorch/module/base.py | 19 ++- 2 files changed, 194 insertions(+), 2 deletions(-) create mode 100644 tests/pytorch/test_tied_delayed_wgrad.py diff --git a/tests/pytorch/test_tied_delayed_wgrad.py b/tests/pytorch/test_tied_delayed_wgrad.py new file mode 100644 index 00000000000..0b9ad722ee7 --- /dev/null +++ b/tests/pytorch/test_tied_delayed_wgrad.py @@ -0,0 +1,177 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. +"""Tied parameters must accumulate delayed wgrads, not overwrite them (#3437). + +``backward_dw()`` used to *assign* the popped delayed wgrad to +``Parameter.grad``. When two modules share one Parameter (tied weights), each +pops its own contribution and the later one silently replaced the earlier, so +the shared gradient came out too small. The bias takes the same path and hit the +same bug, with the extra twist that the second module's bgrad is often ~0 and so +wiped out a correct contribution. + +These tests use only the public ``te.Linear`` API, so they cannot be satisfied by +a shared implementation detail of the fix. +""" + +from __future__ import annotations + +import pytest +import torch +from torch import nn + +import transformer_engine.pytorch as te + +SIZE = 16 +BATCH = 2 +DTYPE = torch.bfloat16 +DEVICE = "cuda" + +# The delayed path sums contributions in arrival order while the regular path +# lets autograd accumulate, so with several microbatches the two orders can +# differ by a BF16 ULP. That is not something the fix promises to remove, so the +# tolerance is declared here and applies to the multi-microbatch case only; +# every single-microbatch comparison below stays bit-exact. +MULTI_MICROBATCH_RTOL = 1e-2 +MULTI_MICROBATCH_ATOL = 1e-2 + + +def _tie(model, bias): + """Re-tie after load_state_dict, which rebinds parameters.""" + model[1].weight = model[0].weight + if bias: + model[1].bias = model[0].bias + return model + + +def _build(delay, tied, bias=False): + model = nn.Sequential( + te.Linear( + SIZE, + SIZE, + bias=bias, + params_dtype=DTYPE, + device=DEVICE, + delay_wgrad_compute=delay, + fuse_wgrad_accumulation=False, + ), + te.Linear( + SIZE, + SIZE, + bias=bias, + params_dtype=DTYPE, + device=DEVICE, + delay_wgrad_compute=delay, + fuse_wgrad_accumulation=False, + ), + ) + return _tie(model, bias) if tied else model + + +def _run(model, x, microbatches=1, delay=False, order=(0, 1)): + """Run the microbatches, then drain the delayed wgrad queues.""" + for xb in torch.chunk(x, microbatches, dim=0): + model(xb).float().sum().backward() + if delay: + for _ in range(microbatches): + for i in order: + model[i].backward_dw() + + +def _x(batch=BATCH): + return torch.randn(batch, SIZE, device=DEVICE, dtype=DTYPE, requires_grad=True) + + +@pytest.fixture(autouse=True) +def _seed(): + torch.manual_seed(9) + + +def test_tied_delayed_wgrad_accumulates(): + """The issue's own case: a tied delayed grad must match the regular path.""" + delayed, regular = _build(True, True), _build(False, True) + regular.load_state_dict(delayed.state_dict()) + _tie(regular, False) + xd = _x() + _run(delayed, xd, delay=True) + _run(regular, xd.detach().clone().requires_grad_(True)) + torch.testing.assert_close(delayed[0].weight.grad, regular[0].weight.grad, rtol=0, atol=0) + + +def test_non_tied_control_still_passes(): + """Independent parameters keep working.""" + delayed, regular = _build(True, False), _build(False, False) + regular.load_state_dict(delayed.state_dict()) + xd = _x() + _run(delayed, xd, delay=True) + _run(regular, xd.detach().clone().requires_grad_(True)) + for i in (0, 1): + torch.testing.assert_close(delayed[i].weight.grad, regular[i].weight.grad, rtol=0, atol=0) + + +def test_both_backward_dw_orders_match(): + """The shared gradient must not depend on which module pops first.""" + a, b = _build(True, True), _build(True, True) + b.load_state_dict(a.state_dict()) + _tie(b, False) + xa = _x() + _run(a, xa, delay=True, order=(0, 1)) + _run(b, xa.detach().clone().requires_grad_(True), delay=True, order=(1, 0)) + torch.testing.assert_close(a[0].weight.grad, b[0].weight.grad, rtol=0, atol=0) + + +def test_multiple_microbatches_accumulate(): + delayed, regular = _build(True, True), _build(False, True) + regular.load_state_dict(delayed.state_dict()) + _tie(regular, False) + x = torch.randn(4, SIZE, device=DEVICE, dtype=DTYPE) + _run(delayed, x.clone().requires_grad_(True), microbatches=2, delay=True) + _run(regular, x.clone().requires_grad_(True), microbatches=2) + torch.testing.assert_close( + delayed[0].weight.grad, + regular[0].weight.grad, + rtol=MULTI_MICROBATCH_RTOL, + atol=MULTI_MICROBATCH_ATOL, + ) + + +@pytest.mark.parametrize("set_to_none", [True, False]) +def test_two_optimizer_steps_do_not_pollute_each_other(set_to_none): + """After a reset, the second step's grad must not inherit the first's.""" + model = _build(True, True) + params = list({id(p): p for p in model.parameters()}.values()) + opt = torch.optim.SGD(params, lr=0.0) + + _run(model, torch.randn(BATCH, SIZE, device=DEVICE, dtype=DTYPE), delay=True) + first = model[0].weight.grad.detach().clone() + opt.zero_grad(set_to_none=set_to_none) + + _run(model, torch.randn(BATCH, SIZE, device=DEVICE, dtype=DTYPE), delay=True) + second = model[0].weight.grad.detach().clone() + + assert not torch.equal(first, second) + assert second.abs().sum() > 0 + + +def test_tied_bias_delayed_wgrad_accumulates(): + """The bias shares the pop path, so it must accumulate too.""" + delayed, regular = _build(True, True, bias=True), _build(False, True, bias=True) + regular.load_state_dict(delayed.state_dict()) + _tie(regular, True) + xd = _x() + _run(delayed, xd, delay=True) + _run(regular, xd.detach().clone().requires_grad_(True)) + torch.testing.assert_close(delayed[0].bias.grad, regular[0].bias.grad, rtol=0, atol=0) + + +def test_matches_pure_pytorch_reference(): + """Independent oracle: plain nn.Linear over the same weight and input.""" + model = _build(True, True) + x = _x() + _run(model, x, delay=True) + + lin = nn.Linear(SIZE, SIZE, bias=False).to(device=DEVICE, dtype=DTYPE) + with torch.no_grad(): + lin.weight.copy_(model[0].weight) + lin(lin(x.detach().clone().requires_grad_(True))).float().sum().backward() + torch.testing.assert_close(model[0].weight.grad, lin.weight.grad, rtol=0, atol=0) diff --git a/transformer_engine/pytorch/module/base.py b/transformer_engine/pytorch/module/base.py index a3131f7436d..94118761149 100644 --- a/transformer_engine/pytorch/module/base.py +++ b/transformer_engine/pytorch/module/base.py @@ -1960,11 +1960,26 @@ def backward_dw(self): (wgrad, bgrad), _ = self.wgrad_store.pop() if not self.fuse_wgrad_accumulation: weight_tensor = noop_cat(self._get_weight_tensors()) - weight_tensor.grad = wgrad.to(weight_tensor.dtype) + wgrad = wgrad.to(weight_tensor.dtype) + # Multiple modules can share one Parameter (tied weights), and each + # of them pops its own delayed wgrad. Assigning would let the last + # module overwrite the contributions of the earlier ones, so + # accumulate instead: initialise on the first contribution of the + # step and add on every later one. + if weight_tensor.grad is None: + weight_tensor.grad = wgrad + else: + weight_tensor.grad.add_(wgrad) if self.use_bias and bgrad is not None and bgrad.numel() != 0: bias_tensor = noop_cat([getattr(self, name) for name in self.bias_names]) + bgrad = bgrad.to(bias_tensor.dtype) + # Same reasoning as for the weight above. A tied bias is shared, + # and a module whose delayed bgrad is (numerically) zero must not + # overwrite the contribution another module already wrote. if bias_tensor.grad is None: - bias_tensor.grad = bgrad.to(bias_tensor.dtype) + bias_tensor.grad = bgrad + else: + bias_tensor.grad.add_(bgrad) del wgrad del bgrad self._trigger_wgrad_accumulation_and_reduce_hooks() From dd805bd01de96df4709682831c022eadc01c44ad Mon Sep 17 00:00:00 2001 From: 0z5a <192209249+0z5a@users.noreply.github.com> Date: Mon, 21 Sep 2026 13:23:00 +0000 Subject: [PATCH 2/4] Clarify the accumulate path for zero_grad(set_to_none=False) Review question: is the `grad is None` first-contribution test safe when the optimizer was built with `zero_grad(set_to_none=False)`? It is: that setting leaves a zeroed tensor where `set_to_none=True` leaves `None`, so the first contribution takes the add path and lands on the same value, and nothing in the branch depends on `grad` being `None`. Say so where the branch is, and pin it in the test. `test_two_optimizer_steps_do_not_pollute_each_other` already ran both settings, but `second != first` cannot show pollution - a first-step value carried into the second still differs from the first. It now compares the second step's gradient against a fresh model on the same weights and input, bit-exactly. Verified with the new test over an out-of-tree copy of an installed TransformerEngine: 5 failed, 3 passed against the unfixed `backward_dw` (the same five this PR describes), and 8 passed with this PR's `backward_dw` in place. Signed-off-by: 0z5a <192209249+0z5a@users.noreply.github.com> Signed-off-by: 0z5a --- tests/pytorch/test_tied_delayed_wgrad.py | 21 ++++++++++++++++++--- transformer_engine/pytorch/module/base.py | 6 ++++++ 2 files changed, 24 insertions(+), 3 deletions(-) diff --git a/tests/pytorch/test_tied_delayed_wgrad.py b/tests/pytorch/test_tied_delayed_wgrad.py index 0b9ad722ee7..3b289cdc216 100644 --- a/tests/pytorch/test_tied_delayed_wgrad.py +++ b/tests/pytorch/test_tied_delayed_wgrad.py @@ -137,18 +137,33 @@ def test_multiple_microbatches_accumulate(): @pytest.mark.parametrize("set_to_none", [True, False]) def test_two_optimizer_steps_do_not_pollute_each_other(set_to_none): - """After a reset, the second step's grad must not inherit the first's.""" + """After a reset, the second step's grad must be the second step's alone. + + ``zero_grad(set_to_none=False)`` leaves a zeroed tensor where ``set_to_none=True`` leaves + ``None``, so the first contribution of the step takes the add path instead of the assign + path. Both settings have to deliver the same value, and ``second != first`` on its own + would not show pollution: a first-step value carried into the second still differs from + the first. + """ model = _build(True, True) params = list({id(p): p for p in model.parameters()}.values()) opt = torch.optim.SGD(params, lr=0.0) - _run(model, torch.randn(BATCH, SIZE, device=DEVICE, dtype=DTYPE), delay=True) + _run(model, _x(), delay=True) first = model[0].weight.grad.detach().clone() opt.zero_grad(set_to_none=set_to_none) - _run(model, torch.randn(BATCH, SIZE, device=DEVICE, dtype=DTYPE), delay=True) + second_x = _x() + _run(model, second_x, delay=True) second = model[0].weight.grad.detach().clone() + # What the second step owes on its own: same weights, same input, first backward only. + fresh = _build(True, True) + fresh.load_state_dict(model.state_dict()) + _tie(fresh, True) + _run(fresh, second_x.detach().clone().requires_grad_(True), delay=True) + + torch.testing.assert_close(second, fresh[0].weight.grad, rtol=0, atol=0) assert not torch.equal(first, second) assert second.abs().sum() > 0 diff --git a/transformer_engine/pytorch/module/base.py b/transformer_engine/pytorch/module/base.py index 94118761149..016d52d1591 100644 --- a/transformer_engine/pytorch/module/base.py +++ b/transformer_engine/pytorch/module/base.py @@ -1966,6 +1966,12 @@ def backward_dw(self): # module overwrite the contributions of the earlier ones, so # accumulate instead: initialise on the first contribution of the # step and add on every later one. + # + # The first contribution is assigned only when zero_grad left grad as + # None, which is what set_to_none=True does. set_to_none=False leaves a + # zeroed tensor instead, so that same contribution takes the add path and + # lands on the same value; nothing here needs grad to be None to start a + # step. if weight_tensor.grad is None: weight_tensor.grad = wgrad else: From 557a35bb80245e09363a9449489dee596f5b3f2e Mon Sep 17 00:00:00 2001 From: 0z5a <192209249+0z5a@users.noreply.github.com> Date: Mon, 21 Sep 2026 14:36:10 +0000 Subject: [PATCH 3/4] Address review feedback on the tied-wgrad test file - `_tie(model, *, bias)` and `_build(delay, *, bias=False, device="cuda")`: the boolean arguments are keyword-only, so a call site says what it means. - `_build` only builds. Tying moved out to the call sites, which is where a test decides whether it wants a tied pair or two independent ones. - `order` is now `backward_dw_order`, and the docstring says what it orders. - The module-level `DEVICE` and tolerance globals are gone: the device is a default on the two helpers that need it, and the multi-microbatch tolerance is local to the test that needs it. Signed-off-by: 0z5a <192209249+0z5a@users.noreply.github.com> Signed-off-by: 0z5a --- tests/pytorch/test_tied_delayed_wgrad.py | 88 ++++++++++++------------ 1 file changed, 45 insertions(+), 43 deletions(-) diff --git a/tests/pytorch/test_tied_delayed_wgrad.py b/tests/pytorch/test_tied_delayed_wgrad.py index 3b289cdc216..7ad756f90bb 100644 --- a/tests/pytorch/test_tied_delayed_wgrad.py +++ b/tests/pytorch/test_tied_delayed_wgrad.py @@ -25,18 +25,9 @@ SIZE = 16 BATCH = 2 DTYPE = torch.bfloat16 -DEVICE = "cuda" -# The delayed path sums contributions in arrival order while the regular path -# lets autograd accumulate, so with several microbatches the two orders can -# differ by a BF16 ULP. That is not something the fix promises to remove, so the -# tolerance is declared here and applies to the multi-microbatch case only; -# every single-microbatch comparison below stays bit-exact. -MULTI_MICROBATCH_RTOL = 1e-2 -MULTI_MICROBATCH_ATOL = 1e-2 - -def _tie(model, bias): +def _tie(model, *, bias): """Re-tie after load_state_dict, which rebinds parameters.""" model[1].weight = model[0].weight if bias: @@ -44,14 +35,15 @@ def _tie(model, bias): return model -def _build(delay, tied, bias=False): +def _build(delay, *, bias=False, device="cuda"): + """Two Linear layers over the same shapes. Tying is the caller's business.""" model = nn.Sequential( te.Linear( SIZE, SIZE, bias=bias, params_dtype=DTYPE, - device=DEVICE, + device=device, delay_wgrad_compute=delay, fuse_wgrad_accumulation=False, ), @@ -60,26 +52,30 @@ def _build(delay, tied, bias=False): SIZE, bias=bias, params_dtype=DTYPE, - device=DEVICE, + device=device, delay_wgrad_compute=delay, fuse_wgrad_accumulation=False, ), ) - return _tie(model, bias) if tied else model + return model + +def _run(model, x, microbatches=1, delay=False, backward_dw_order=(0, 1)): + """Run the microbatches, then drain the delayed wgrad queues. -def _run(model, x, microbatches=1, delay=False, order=(0, 1)): - """Run the microbatches, then drain the delayed wgrad queues.""" + ``backward_dw_order`` is the module order used for draining, which decides which tied + module pops its delayed wgrad first. + """ for xb in torch.chunk(x, microbatches, dim=0): model(xb).float().sum().backward() if delay: for _ in range(microbatches): - for i in order: + for i in backward_dw_order: model[i].backward_dw() -def _x(batch=BATCH): - return torch.randn(batch, SIZE, device=DEVICE, dtype=DTYPE, requires_grad=True) +def _x(batch=BATCH, device="cuda"): + return torch.randn(batch, SIZE, device=device, dtype=DTYPE, requires_grad=True) @pytest.fixture(autouse=True) @@ -89,9 +85,10 @@ def _seed(): def test_tied_delayed_wgrad_accumulates(): """The issue's own case: a tied delayed grad must match the regular path.""" - delayed, regular = _build(True, True), _build(False, True) + delayed = _tie(_build(True), bias=False) + regular = _build(False) regular.load_state_dict(delayed.state_dict()) - _tie(regular, False) + _tie(regular, bias=False) xd = _x() _run(delayed, xd, delay=True) _run(regular, xd.detach().clone().requires_grad_(True)) @@ -100,7 +97,7 @@ def test_tied_delayed_wgrad_accumulates(): def test_non_tied_control_still_passes(): """Independent parameters keep working.""" - delayed, regular = _build(True, False), _build(False, False) + delayed, regular = _build(True), _build(False) regular.load_state_dict(delayed.state_dict()) xd = _x() _run(delayed, xd, delay=True) @@ -111,28 +108,31 @@ def test_non_tied_control_still_passes(): def test_both_backward_dw_orders_match(): """The shared gradient must not depend on which module pops first.""" - a, b = _build(True, True), _build(True, True) + a = _tie(_build(True), bias=False) + b = _build(True) b.load_state_dict(a.state_dict()) - _tie(b, False) + _tie(b, bias=False) xa = _x() - _run(a, xa, delay=True, order=(0, 1)) - _run(b, xa.detach().clone().requires_grad_(True), delay=True, order=(1, 0)) + _run(a, xa, delay=True, backward_dw_order=(0, 1)) + _run(b, xa.detach().clone().requires_grad_(True), delay=True, backward_dw_order=(1, 0)) torch.testing.assert_close(a[0].weight.grad, b[0].weight.grad, rtol=0, atol=0) def test_multiple_microbatches_accumulate(): - delayed, regular = _build(True, True), _build(False, True) + # The delayed path sums contributions in arrival order while the regular path lets + # autograd accumulate, so with several microbatches the two orders can differ by a BF16 + # ULP. That is not something the fix promises to remove, so the tolerance belongs to this + # test only; every single-microbatch comparison below stays bit-exact. + rtol = atol = 1e-2 + + delayed = _tie(_build(True), bias=False) + regular = _build(False) regular.load_state_dict(delayed.state_dict()) - _tie(regular, False) - x = torch.randn(4, SIZE, device=DEVICE, dtype=DTYPE) + _tie(regular, bias=False) + x = _x(4) _run(delayed, x.clone().requires_grad_(True), microbatches=2, delay=True) _run(regular, x.clone().requires_grad_(True), microbatches=2) - torch.testing.assert_close( - delayed[0].weight.grad, - regular[0].weight.grad, - rtol=MULTI_MICROBATCH_RTOL, - atol=MULTI_MICROBATCH_ATOL, - ) + torch.testing.assert_close(delayed[0].weight.grad, regular[0].weight.grad, rtol=rtol, atol=atol) @pytest.mark.parametrize("set_to_none", [True, False]) @@ -145,7 +145,7 @@ def test_two_optimizer_steps_do_not_pollute_each_other(set_to_none): would not show pollution: a first-step value carried into the second still differs from the first. """ - model = _build(True, True) + model = _tie(_build(True), bias=False) params = list({id(p): p for p in model.parameters()}.values()) opt = torch.optim.SGD(params, lr=0.0) @@ -158,9 +158,9 @@ def test_two_optimizer_steps_do_not_pollute_each_other(set_to_none): second = model[0].weight.grad.detach().clone() # What the second step owes on its own: same weights, same input, first backward only. - fresh = _build(True, True) + fresh = _build(True) fresh.load_state_dict(model.state_dict()) - _tie(fresh, True) + _tie(fresh, bias=False) _run(fresh, second_x.detach().clone().requires_grad_(True), delay=True) torch.testing.assert_close(second, fresh[0].weight.grad, rtol=0, atol=0) @@ -170,9 +170,10 @@ def test_two_optimizer_steps_do_not_pollute_each_other(set_to_none): def test_tied_bias_delayed_wgrad_accumulates(): """The bias shares the pop path, so it must accumulate too.""" - delayed, regular = _build(True, True, bias=True), _build(False, True, bias=True) + delayed = _tie(_build(True, bias=True), bias=True) + regular = _build(False, bias=True) regular.load_state_dict(delayed.state_dict()) - _tie(regular, True) + _tie(regular, bias=True) xd = _x() _run(delayed, xd, delay=True) _run(regular, xd.detach().clone().requires_grad_(True)) @@ -181,12 +182,13 @@ def test_tied_bias_delayed_wgrad_accumulates(): def test_matches_pure_pytorch_reference(): """Independent oracle: plain nn.Linear over the same weight and input.""" - model = _build(True, True) + model = _tie(_build(True), bias=False) x = _x() _run(model, x, delay=True) - lin = nn.Linear(SIZE, SIZE, bias=False).to(device=DEVICE, dtype=DTYPE) + weight = model[0].weight + lin = nn.Linear(SIZE, SIZE, bias=False).to(device=weight.device, dtype=weight.dtype) with torch.no_grad(): - lin.weight.copy_(model[0].weight) + lin.weight.copy_(weight) lin(lin(x.detach().clone().requires_grad_(True))).float().sum().backward() torch.testing.assert_close(model[0].weight.grad, lin.weight.grad, rtol=0, atol=0) From c5cf151827b0506a0d5b24f9608ff9afcb1915f2 Mon Sep 17 00:00:00 2001 From: 0z5a Date: Wed, 7 Oct 2026 22:15:49 +0200 Subject: [PATCH 4/4] fix(pytorch): accumulate all delayed tied-parameter gradients Signed-off-by: 0z5a --- qa/L0_pytorch_unittest/test.sh | 1 + tests/pytorch/test_tied_delayed_wgrad.py | 468 ++++++++++++------ transformer_engine/pytorch/module/base.py | 15 +- .../pytorch/module/grouped_linear.py | 20 +- .../pytorch/module/layernorm_mlp.py | 45 +- .../pytorch/ops/basic/grouped_linear.py | 14 +- 6 files changed, 353 insertions(+), 210 deletions(-) diff --git a/qa/L0_pytorch_unittest/test.sh b/qa/L0_pytorch_unittest/test.sh index ea2d58825f0..60d3d4d594e 100644 --- a/qa/L0_pytorch_unittest/test.sh +++ b/qa/L0_pytorch_unittest/test.sh @@ -34,6 +34,7 @@ python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_recipe.xml $TE_P python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_custom_recipe.xml $TE_PATH/tests/pytorch/test_custom_recipe.py || test_fail "test_custom_recipe.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_deferred_init.xml $TE_PATH/tests/pytorch/test_deferred_init.py || test_fail "test_deferred_init.py" PYTORCH_JIT=0 NVTE_TORCH_COMPILE=0 NVTE_ALLOW_NONDETERMINISTIC_ALGO=0 NVTE_FUSED_ATTN=0 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_numerics.xml $TE_PATH/tests/pytorch/test_numerics.py || test_fail "test_numerics.py" +PYTORCH_JIT=0 NVTE_TORCH_COMPILE=0 NVTE_ALLOW_NONDETERMINISTIC_ALGO=0 NVTE_FUSED_ATTN=0 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_tied_delayed_wgrad.xml $TE_PATH/tests/pytorch/test_tied_delayed_wgrad.py || test_fail "test_tied_delayed_wgrad.py" PYTORCH_JIT=0 NVTE_TORCH_COMPILE=0 NVTE_ALLOW_NONDETERMINISTIC_ALGO=0 NVTE_FUSED_ATTN=0 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_cuda_graphs.xml $TE_PATH/tests/pytorch/test_cuda_graphs.py || test_fail "test_cuda_graphs.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_jit.xml $TE_PATH/tests/pytorch/test_jit.py || test_fail "test_jit.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_fused_rope.xml $TE_PATH/tests/pytorch/test_fused_rope.py || test_fail "test_fused_rope.py" diff --git a/tests/pytorch/test_tied_delayed_wgrad.py b/tests/pytorch/test_tied_delayed_wgrad.py index 7ad756f90bb..75a53a9b695 100644 --- a/tests/pytorch/test_tied_delayed_wgrad.py +++ b/tests/pytorch/test_tied_delayed_wgrad.py @@ -1,194 +1,334 @@ # Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. # # See LICENSE for license information. -"""Tied parameters must accumulate delayed wgrads, not overwrite them (#3437). - -``backward_dw()`` used to *assign* the popped delayed wgrad to -``Parameter.grad``. When two modules share one Parameter (tied weights), each -pops its own contribution and the later one silently replaced the earlier, so -the shared gradient came out too small. The bias takes the same path and hit the -same bug, with the extra twist that the second module's bgrad is often ~0 and so -wiped out a correct contribution. - -These tests use only the public ``te.Linear`` API, so they cannot be satisfied by -a shared implementation detail of the fix. -""" - -from __future__ import annotations +"""Delayed weight gradients accumulate for tied parameters and microbatches.""" import pytest import torch from torch import nn import transformer_engine.pytorch as te - -SIZE = 16 -BATCH = 2 -DTYPE = torch.bfloat16 - - -def _tie(model, *, bias): - """Re-tie after load_state_dict, which rebinds parameters.""" - model[1].weight = model[0].weight - if bias: - model[1].bias = model[0].bias - return model - - -def _build(delay, *, bias=False, device="cuda"): - """Two Linear layers over the same shapes. Tying is the caller's business.""" - model = nn.Sequential( +from transformer_engine.common.recipe import Float8BlockScaling +from transformer_engine.pytorch.module.grouped_linear import is_module_grouped_tensor_path_supported +from transformer_engine.pytorch.ops.basic.grouped_linear import ( + is_op_fuser_grouped_tensor_path_supported, +) +from transformer_engine.pytorch.quantization import is_fp8_block_scaling_available + + +def _build_two_linears(delay, *, bias=False, device="cuda"): + """Build two BF16 Linear layers with independent parameters.""" + hidden_size = 16 + return nn.Sequential( te.Linear( - SIZE, - SIZE, + hidden_size, + hidden_size, bias=bias, - params_dtype=DTYPE, + params_dtype=torch.bfloat16, device=device, delay_wgrad_compute=delay, fuse_wgrad_accumulation=False, ), te.Linear( - SIZE, - SIZE, + hidden_size, + hidden_size, bias=bias, - params_dtype=DTYPE, + params_dtype=torch.bfloat16, device=device, delay_wgrad_compute=delay, fuse_wgrad_accumulation=False, ), ) - return model -def _run(model, x, microbatches=1, delay=False, backward_dw_order=(0, 1)): - """Run the microbatches, then drain the delayed wgrad queues. - - ``backward_dw_order`` is the module order used for draining, which decides which tied - module pops its delayed wgrad first. - """ - for xb in torch.chunk(x, microbatches, dim=0): - model(xb).float().sum().backward() - if delay: +@pytest.mark.parametrize("bias", [False, True]) +@pytest.mark.parametrize("set_to_none", [False, True]) +@pytest.mark.parametrize("microbatches", [1, 2]) +@pytest.mark.parametrize("backward_dw_order", [(0, 1), (1, 0)]) +def test_tied_linear_gradients(bias, set_to_none, microbatches, backward_dw_order): + """Tied weight and bias gradients match regular backward after either reset.""" + torch.manual_seed(9) + delayed = _build_two_linears(True, bias=bias) + regular = _build_two_linears(False, bias=bias) + regular.load_state_dict(delayed.state_dict()) + for model in (delayed, regular): + model[1].weight = model[0].weight + if bias: + model[1].bias = model[0].bias + hidden_size, batch_size = 16, 2 + optimizers = [torch.optim.SGD(model.parameters(), lr=0.0) for model in (delayed, regular)] + for _ in range(2): + for optimizer in optimizers: + optimizer.zero_grad(set_to_none=set_to_none) + x = torch.randn(batch_size * microbatches, hidden_size, device="cuda", dtype=torch.bfloat16) + for model in (delayed, regular): + for microbatch in x.chunk(microbatches): + model(microbatch.detach().clone().requires_grad_(True)).float().sum().backward() for _ in range(microbatches): - for i in backward_dw_order: - model[i].backward_dw() - - -def _x(batch=BATCH, device="cuda"): - return torch.randn(batch, SIZE, device=device, dtype=DTYPE, requires_grad=True) - - -@pytest.fixture(autouse=True) -def _seed(): + for index in backward_dw_order: + delayed[index].backward_dw() + # Microbatch contributions can accumulate in a different BF16 rounding order. + tolerance = 1e-2 if microbatches > 1 else 0.0 + torch.testing.assert_close( + delayed[0].weight.grad, regular[0].weight.grad, rtol=tolerance, atol=tolerance + ) + assert torch.count_nonzero(delayed[0].weight.grad) > 0 + if bias: + torch.testing.assert_close( + delayed[0].bias.grad, regular[0].bias.grad, rtol=tolerance, atol=tolerance + ) + for optimizer in optimizers: + optimizer.step() + + +def test_non_tied_linear_gradients(): + """Independent delayed parameters match regular backward.""" torch.manual_seed(9) - - -def test_tied_delayed_wgrad_accumulates(): - """The issue's own case: a tied delayed grad must match the regular path.""" - delayed = _tie(_build(True), bias=False) - regular = _build(False) + delayed, regular = _build_two_linears(True), _build_two_linears(False) regular.load_state_dict(delayed.state_dict()) - _tie(regular, bias=False) - xd = _x() - _run(delayed, xd, delay=True) - _run(regular, xd.detach().clone().requires_grad_(True)) - torch.testing.assert_close(delayed[0].weight.grad, regular[0].weight.grad, rtol=0, atol=0) - - -def test_non_tied_control_still_passes(): - """Independent parameters keep working.""" - delayed, regular = _build(True), _build(False) + x = torch.randn(2, 16, device="cuda", dtype=torch.bfloat16) + for model in (delayed, regular): + model(x.detach().clone().requires_grad_(True)).float().sum().backward() + for index in (0, 1): + delayed[index].backward_dw() + torch.testing.assert_close( + delayed[index].weight.grad, regular[index].weight.grad, rtol=0, atol=0 + ) + + +def test_tied_linear_matches_pytorch(): + """A plain PyTorch Linear provides an independent tied-weight reference.""" + torch.manual_seed(9) + model = _build_two_linears(True) + model[1].weight = model[0].weight + reference = nn.Linear(16, 16, bias=False, device="cuda", dtype=torch.bfloat16) + with torch.no_grad(): + reference.weight.copy_(model[0].weight) + x = torch.randn(2, 16, device="cuda", dtype=torch.bfloat16) + model(x.detach().clone().requires_grad_(True)).float().sum().backward() + for layer in model: + layer.backward_dw() + reference_x = x.detach().clone().requires_grad_(True) + reference(reference(reference_x)).float().sum().backward() + torch.testing.assert_close(model[0].weight.grad, reference.weight.grad, rtol=0, atol=0) + + +@pytest.mark.parametrize("set_to_none", [False, True]) +@pytest.mark.parametrize("microbatches", [1, 2]) +@pytest.mark.parametrize("backward_dw_order", [(0, 1), (1, 0)]) +@pytest.mark.parametrize("fp8_block_scaling", [False, True]) +def test_tied_layernorm_mlp_gradients( + set_to_none, microbatches, backward_dw_order, fp8_block_scaling +): + """Both FCs accumulate delayed gradients without duplicating unfused bias grads.""" + if fp8_block_scaling: + supported, reason = is_fp8_block_scaling_available(return_reason=True) + if not supported: + pytest.skip(reason) + torch.manual_seed(9) + hidden_size, ffn_hidden_size, batch_size = 128, 256, 128 + delayed = nn.Sequential( + te.LayerNormMLP( + hidden_size, + ffn_hidden_size, + params_dtype=torch.bfloat16, + delay_wgrad_compute=True, + fuse_wgrad_accumulation=False, + ), + te.LayerNormMLP( + hidden_size, + ffn_hidden_size, + params_dtype=torch.bfloat16, + delay_wgrad_compute=True, + fuse_wgrad_accumulation=False, + ), + ) + regular = nn.Sequential( + te.LayerNormMLP( + hidden_size, + ffn_hidden_size, + params_dtype=torch.bfloat16, + delay_wgrad_compute=False, + fuse_wgrad_accumulation=False, + ), + te.LayerNormMLP( + hidden_size, + ffn_hidden_size, + params_dtype=torch.bfloat16, + delay_wgrad_compute=False, + fuse_wgrad_accumulation=False, + ), + ) regular.load_state_dict(delayed.state_dict()) - xd = _x() - _run(delayed, xd, delay=True) - _run(regular, xd.detach().clone().requires_grad_(True)) - for i in (0, 1): - torch.testing.assert_close(delayed[i].weight.grad, regular[i].weight.grad, rtol=0, atol=0) - - -def test_both_backward_dw_orders_match(): - """The shared gradient must not depend on which module pops first.""" - a = _tie(_build(True), bias=False) - b = _build(True) - b.load_state_dict(a.state_dict()) - _tie(b, bias=False) - xa = _x() - _run(a, xa, delay=True, backward_dw_order=(0, 1)) - _run(b, xa.detach().clone().requires_grad_(True), delay=True, backward_dw_order=(1, 0)) - torch.testing.assert_close(a[0].weight.grad, b[0].weight.grad, rtol=0, atol=0) - - -def test_multiple_microbatches_accumulate(): - # The delayed path sums contributions in arrival order while the regular path lets - # autograd accumulate, so with several microbatches the two orders can differ by a BF16 - # ULP. That is not something the fix promises to remove, so the tolerance belongs to this - # test only; every single-microbatch comparison below stays bit-exact. - rtol = atol = 1e-2 - - delayed = _tie(_build(True), bias=False) - regular = _build(False) + names = ("fc1_weight", "fc2_weight", "fc1_bias", "fc2_bias") + for model in (delayed, regular): + for name in names: + setattr(model[1], name, getattr(model[0], name)) + optimizers = [torch.optim.SGD(model.parameters(), lr=0.0) for model in (delayed, regular)] + for _ in range(2): + for optimizer in optimizers: + optimizer.zero_grad(set_to_none=set_to_none) + x = torch.randn(batch_size * microbatches, hidden_size, device="cuda", dtype=torch.bfloat16) + for model in (delayed, regular): + for microbatch in x.chunk(microbatches): + with te.autocast(enabled=fp8_block_scaling, recipe=Float8BlockScaling()): + output = model(microbatch.detach().clone().requires_grad_(True)) + output.float().sum().backward() + for _ in range(microbatches): + for index in backward_dw_order: + delayed[index].backward_dw() + for name in names: + torch.testing.assert_close( + getattr(delayed[0], name).grad, + getattr(regular[0], name).grad, + rtol=1e-2, + atol=1e-2, + ) + assert torch.count_nonzero(delayed[0].fc1_weight.grad) > 0 + assert torch.count_nonzero(delayed[0].fc2_weight.grad) > 0 + for optimizer in optimizers: + optimizer.step() + + +@pytest.mark.parametrize("set_to_none", [False, True]) +@pytest.mark.parametrize("backward_dw_order", [(0, 1), (1, 0)]) +@pytest.mark.parametrize("packed", [False, True]) +def test_tied_grouped_linear_gradients(set_to_none, backward_dw_order, packed, monkeypatch): + """Module grouped weights and deferred discrete biases accumulate across two microbatches.""" + dtype = torch.bfloat16 + if packed and not is_module_grouped_tensor_path_supported(None, dtype): + pytest.skip("packed grouped weights require a supported GPU and cuBLASLt") + monkeypatch.setenv("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "1" if packed else "0") + torch.manual_seed(9) + num_groups, hidden_size, tokens_per_group, microbatches = 2, 16, 16, 2 + delayed = nn.ModuleList( + [ + te.GroupedLinear( + num_groups, + hidden_size, + hidden_size, + bias=not packed, + params_dtype=dtype, + fuse_wgrad_accumulation=False, + use_grouped_tensor=packed, + delay_wgrad_compute=True, + single_grouped_weight=packed, + ) + for _ in range(2) + ] + ) + regular = nn.ModuleList( + [ + te.GroupedLinear( + num_groups, + hidden_size, + hidden_size, + bias=not packed, + params_dtype=dtype, + fuse_wgrad_accumulation=False, + use_grouped_tensor=packed, + delay_wgrad_compute=False, + single_grouped_weight=packed, + ) + for _ in range(2) + ] + ) regular.load_state_dict(delayed.state_dict()) - _tie(regular, bias=False) - x = _x(4) - _run(delayed, x.clone().requires_grad_(True), microbatches=2, delay=True) - _run(regular, x.clone().requires_grad_(True), microbatches=2) - torch.testing.assert_close(delayed[0].weight.grad, regular[0].weight.grad, rtol=rtol, atol=atol) - - -@pytest.mark.parametrize("set_to_none", [True, False]) -def test_two_optimizer_steps_do_not_pollute_each_other(set_to_none): - """After a reset, the second step's grad must be the second step's alone. - - ``zero_grad(set_to_none=False)`` leaves a zeroed tensor where ``set_to_none=True`` leaves - ``None``, so the first contribution of the step takes the add path instead of the assign - path. Both settings have to deliver the same value, and ``second != first`` on its own - would not show pollution: a first-step value carried into the second still differs from - the first. - """ - model = _tie(_build(True), bias=False) - params = list({id(p): p for p in model.parameters()}.values()) - opt = torch.optim.SGD(params, lr=0.0) - - _run(model, _x(), delay=True) - first = model[0].weight.grad.detach().clone() - opt.zero_grad(set_to_none=set_to_none) - - second_x = _x() - _run(model, second_x, delay=True) - second = model[0].weight.grad.detach().clone() - - # What the second step owes on its own: same weights, same input, first backward only. - fresh = _build(True) - fresh.load_state_dict(model.state_dict()) - _tie(fresh, bias=False) - _run(fresh, second_x.detach().clone().requires_grad_(True), delay=True) - - torch.testing.assert_close(second, fresh[0].weight.grad, rtol=0, atol=0) - assert not torch.equal(first, second) - assert second.abs().sum() > 0 - - -def test_tied_bias_delayed_wgrad_accumulates(): - """The bias shares the pop path, so it must accumulate too.""" - delayed = _tie(_build(True, bias=True), bias=True) - regular = _build(False, bias=True) + names = ("weight",) if packed else ("weight0", "weight1", "bias0", "bias1") + for model in (delayed, regular): + for name in names: + setattr(model[1], name, getattr(model[0], name)) + splits = [tokens_per_group] * num_groups + if packed: + splits = torch.tensor(splits, device="cuda", dtype=torch.int64) + for _ in range(2): + for model in (delayed, regular): + model.zero_grad(set_to_none=set_to_none) + x = torch.randn( + microbatches * tokens_per_group * num_groups, hidden_size, device="cuda", dtype=dtype + ) + for model in (delayed, regular): + for microbatch in x.chunk(microbatches): + output = model[0](microbatch.detach().clone().requires_grad_(True), splits) + model[1](output, splits).float().sum().backward() + for _ in range(microbatches): + for index in backward_dw_order: + delayed[index].backward_dw() + for name in names: + torch.testing.assert_close( + getattr(delayed[0], name).grad, + getattr(regular[0], name).grad, + rtol=1e-2, + atol=1e-2, + ) + assert torch.count_nonzero(getattr(delayed[0], names[0]).grad) > 0 + + +@pytest.mark.parametrize("set_to_none", [False, True]) +@pytest.mark.parametrize("backward_dw_order", [(0, 1), (1, 0)]) +@pytest.mark.parametrize("packed", [False, True]) +def test_tied_grouped_linear_ops_gradients(set_to_none, backward_dw_order, packed, monkeypatch): + """Fusible grouped operations accumulate delayed packed and discrete weight grads.""" + dtype = torch.bfloat16 + if packed and not is_op_fuser_grouped_tensor_path_supported(None, dtype): + pytest.skip("packed grouped weights require a supported GPU and cuBLASLt") + monkeypatch.setenv("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "1" if packed else "0") + torch.manual_seed(9) + num_groups, hidden_size, tokens_per_group, microbatches = 2, 16, 16, 2 + delayed = nn.ModuleList( + [ + te.ops.GroupedLinear( + num_groups, + hidden_size, + hidden_size, + bias=False, + dtype=dtype, + accumulate_into_main_grad=False, + delay_wgrad_compute=True, + single_grouped_weight=packed, + ) + for _ in range(2) + ] + ) + regular = nn.ModuleList( + [ + te.ops.GroupedLinear( + num_groups, + hidden_size, + hidden_size, + bias=False, + dtype=dtype, + accumulate_into_main_grad=False, + delay_wgrad_compute=False, + single_grouped_weight=packed, + ) + for _ in range(2) + ] + ) regular.load_state_dict(delayed.state_dict()) - _tie(regular, bias=True) - xd = _x() - _run(delayed, xd, delay=True) - _run(regular, xd.detach().clone().requires_grad_(True)) - torch.testing.assert_close(delayed[0].bias.grad, regular[0].bias.grad, rtol=0, atol=0) - - -def test_matches_pure_pytorch_reference(): - """Independent oracle: plain nn.Linear over the same weight and input.""" - model = _tie(_build(True), bias=False) - x = _x() - _run(model, x, delay=True) - - weight = model[0].weight - lin = nn.Linear(SIZE, SIZE, bias=False).to(device=weight.device, dtype=weight.dtype) - with torch.no_grad(): - lin.weight.copy_(weight) - lin(lin(x.detach().clone().requires_grad_(True))).float().sum().backward() - torch.testing.assert_close(model[0].weight.grad, lin.weight.grad, rtol=0, atol=0) + names = ("weight",) if packed else ("weight0", "weight1") + for model in (delayed, regular): + for name in names: + setattr(model[1], name, getattr(model[0], name)) + splits = torch.tensor([tokens_per_group] * num_groups, device="cuda", dtype=torch.int64) + for _ in range(2): + for model in (delayed, regular): + model.zero_grad(set_to_none=set_to_none) + x = torch.randn( + microbatches * tokens_per_group * num_groups, hidden_size, device="cuda", dtype=dtype + ) + for model in (delayed, regular): + for microbatch in x.chunk(microbatches): + output = model[0](microbatch.detach().clone().requires_grad_(True), splits) + model[1](output, splits).float().sum().backward() + for _ in range(microbatches): + for index in backward_dw_order: + delayed[index].backward_dw() + for name in names: + torch.testing.assert_close( + getattr(delayed[0], name).grad, + getattr(regular[0], name).grad, + rtol=1e-2, + atol=1e-2, + ) + assert torch.count_nonzero(getattr(delayed[0], names[0]).grad) > 0 diff --git a/transformer_engine/pytorch/module/base.py b/transformer_engine/pytorch/module/base.py index 016d52d1591..a330eac4339 100644 --- a/transformer_engine/pytorch/module/base.py +++ b/transformer_engine/pytorch/module/base.py @@ -1961,17 +1961,7 @@ def backward_dw(self): if not self.fuse_wgrad_accumulation: weight_tensor = noop_cat(self._get_weight_tensors()) wgrad = wgrad.to(weight_tensor.dtype) - # Multiple modules can share one Parameter (tied weights), and each - # of them pops its own delayed wgrad. Assigning would let the last - # module overwrite the contributions of the earlier ones, so - # accumulate instead: initialise on the first contribution of the - # step and add on every later one. - # - # The first contribution is assigned only when zero_grad left grad as - # None, which is what set_to_none=True does. set_to_none=False leaves a - # zeroed tensor instead, so that same contribution takes the add path and - # lands on the same value; nothing here needs grad to be None to start a - # step. + # Tied parameters and microbatches contribute to the same gradient. if weight_tensor.grad is None: weight_tensor.grad = wgrad else: @@ -1979,9 +1969,6 @@ def backward_dw(self): if self.use_bias and bgrad is not None and bgrad.numel() != 0: bias_tensor = noop_cat([getattr(self, name) for name in self.bias_names]) bgrad = bgrad.to(bias_tensor.dtype) - # Same reasoning as for the weight above. A tied bias is shared, - # and a module whose delayed bgrad is (numerically) zero must not - # overwrite the contribution another module already wrote. if bias_tensor.grad is None: bias_tensor.grad = bgrad else: diff --git a/transformer_engine/pytorch/module/grouped_linear.py b/transformer_engine/pytorch/module/grouped_linear.py index 8b078944f17..e305c9664ba 100644 --- a/transformer_engine/pytorch/module/grouped_linear.py +++ b/transformer_engine/pytorch/module/grouped_linear.py @@ -2342,12 +2342,20 @@ def backward_dw(self): weight_params = self._get_weight_tensors() if not self.fuse_wgrad_accumulation: if self.single_grouped_weight: - weight_params[0].grad = wgrad_output.rowwise_data.view( + wgrad = wgrad_output.rowwise_data.view( self.num_gemms, self.out_features, self.in_features ).to(weight_params[0].dtype) + if weight_params[0].grad is None: + weight_params[0].grad = wgrad + else: + weight_params[0].grad.add_(wgrad) else: for i in range(self.num_gemms): - weight_params[i].grad = wgrad_output[i].to(weight_params[i].dtype) + wgrad = wgrad_output[i].to(weight_params[i].dtype) + if weight_params[i].grad is None: + weight_params[i].grad = wgrad + else: + weight_params[i].grad.add_(wgrad) has_grad_biases = [ grad_bias is not None and grad_bias.numel() != 0 for grad_bias in grad_biases_ ] @@ -2362,8 +2370,12 @@ def backward_dw(self): ) bias_params = [getattr(self, f"bias{i}") for i in range(self.num_gemms)] for i in range(self.num_gemms): - if has_grad_biases[i] and bias_params[i].grad is None: - bias_params[i].grad = grad_biases_[i].to(bias_params[i].dtype) + if has_grad_biases[i]: + bgrad = grad_biases_[i].to(bias_params[i].dtype) + if bias_params[i].grad is None: + bias_params[i].grad = bgrad + else: + bias_params[i].grad.add_(bgrad) del grad_biases_ del wgrad_output del tensor_list diff --git a/transformer_engine/pytorch/module/layernorm_mlp.py b/transformer_engine/pytorch/module/layernorm_mlp.py index 8d2dc192159..1287e3d23fd 100644 --- a/transformer_engine/pytorch/module/layernorm_mlp.py +++ b/transformer_engine/pytorch/module/layernorm_mlp.py @@ -3148,36 +3148,31 @@ def _get_weight_quantizers(self) -> List[Quantizer]: return [fc1_weight_quantizer, fc2_weight_quantizer] def backward_dw(self): - """ - Execute the delayed weight gradient computation. - This method is called after the main backward pass to compute weight gradients. - """ + """Execute delayed weight-gradient GEMMs and accumulate their contributions.""" if not self.need_backward_dw(): return with get_nvtx_range_context("_LayerNormMLP_wgrad"): - (fc2_wgrad, fc2_bias_grad_, *_), tensor_list_fc2 = self.wgrad_store.pop() - if self.use_bias and self.fc1_bias.grad is None: - (fc1_wgrad, fc1_bias_grad, *_), _ = self.wgrad_store.pop() - else: - (fc1_wgrad, *_), _ = self.wgrad_store.pop() - fc1_bias_grad = None + (fc2_wgrad, fc2_bias_grad_, *_), _ = self.wgrad_store.pop() + (fc1_wgrad, fc1_bias_grad, *_), _ = self.wgrad_store.pop() if self.use_bias: - if self.fc2_bias.grad is None: - if ( - self.fp8 - and FP8GlobalStateManager.get_fp8_recipe().float8_block_scaling() - and self.apply_bias - and not self.gemm_bias_unfused_add - ): - act_out = tensor_list_fc2[0] - # BGRAD not fused with GEMM for float8 blockwise gemm. - fc2_bias_grad_ = act_out.view(-1, act_out.shape[-1]).sum(dim=0) - self.fc2_bias.grad = fc2_bias_grad_.to(self.fc2_bias.dtype) - if self.fc1_bias.grad is None: - self.fc1_bias.grad = fc1_bias_grad.to(self.fc1_bias.dtype) + # Unfused bias gradients are already returned by the main backward. + for bias, bgrad in ( + (self.fc2_bias, fc2_bias_grad_), + (self.fc1_bias, fc1_bias_grad), + ): + if bgrad is not None and bgrad.numel() != 0: + bgrad = bgrad.to(bias.dtype) + if bias.grad is None: + bias.grad = bgrad + else: + bias.grad.add_(bgrad) if not self.fuse_wgrad_accumulation: - self.fc2_weight.grad = fc2_wgrad.to(self.fc2_weight.dtype) - self.fc1_weight.grad = fc1_wgrad.to(self.fc1_weight.dtype) + for weight, wgrad in ((self.fc2_weight, fc2_wgrad), (self.fc1_weight, fc1_wgrad)): + wgrad = wgrad.to(weight.dtype) + if weight.grad is None: + weight.grad = wgrad + else: + weight.grad.add_(wgrad) del fc2_bias_grad_ del fc2_wgrad del fc1_wgrad diff --git a/transformer_engine/pytorch/ops/basic/grouped_linear.py b/transformer_engine/pytorch/ops/basic/grouped_linear.py index 95516500450..4bc1cd58ad9 100644 --- a/transformer_engine/pytorch/ops/basic/grouped_linear.py +++ b/transformer_engine/pytorch/ops/basic/grouped_linear.py @@ -363,17 +363,25 @@ def backward_dw(self) -> None: return if self.single_grouped_weight: if isinstance(grad_weights, list): - self.weight.grad = torch.stack(grad_weights, dim=0).to(self.weight.dtype) + wgrad = torch.stack(grad_weights, dim=0).to(self.weight.dtype) else: - self.weight.grad = grad_weights.rowwise_data.view( + wgrad = grad_weights.rowwise_data.view( self.num_groups, self.out_features, self.in_features, ).to(self.weight.dtype) + if self.weight.grad is None: + self.weight.grad = wgrad + else: + self.weight.grad.add_(wgrad) else: for group_idx in range(self.num_groups): w = getattr(self, f"weight{group_idx}") - w.grad = grad_weights[group_idx].to(w.dtype) + wgrad = grad_weights[group_idx].to(w.dtype) + if w.grad is None: + w.grad = wgrad + else: + w.grad.add_(wgrad) self._trigger_wgrad_accumulation_and_reduce_hooks() def _get_discrete_bias_tensors(self, dtype: torch.dtype) -> list[torch.Tensor]: