Repository navigation
Accumulate delayed wgrads into shared (tied) parameters #3552
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Open
0z5a
wants to merge
5
commits into
NVIDIA:main
Choose a base branch
from
0z5a:tied-delayed-wgrad-accumulate
base: main
Could not load branches
Branch not found: {{ refName }}
Loading
Could not load tags
Nothing to show
Loading
Are you sure you want to change the base?
Some commits from the old base branch may be removed from the timeline,
and old review comments may become outdated.
+392
−34
Open
Changes from all commits
Commits
Show all changes
5 commits
Select commit
Hold shift + click to select a range
b5b9cc2
Accumulate delayed wgrads into shared (tied) parameters
0z5a dd805bd
Clarify the accumulate path for zero_grad(set_to_none=False)
0z5a 557a35b
Address review feedback on the tied-wgrad test file
0z5a c5cf151
fix(pytorch): accumulate all delayed tied-parameter gradients
0z5a 67a310d
Merge main and preserve delayed tied-gradient accumulation
0z5a File filter
Filter by extension
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
There are no files selected for viewing
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,334 @@ | ||
| # Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. | ||
| # | ||
| # See LICENSE for license information. | ||
| """Delayed weight gradients accumulate for tied parameters and microbatches.""" | ||
|
|
||
| import pytest | ||
| import torch | ||
| from torch import nn | ||
|
|
||
| import transformer_engine.pytorch as te | ||
| 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( | ||
| hidden_size, | ||
| hidden_size, | ||
| bias=bias, | ||
| params_dtype=torch.bfloat16, | ||
| device=device, | ||
| delay_wgrad_compute=delay, | ||
| fuse_wgrad_accumulation=False, | ||
| ), | ||
| te.Linear( | ||
| hidden_size, | ||
| hidden_size, | ||
| bias=bias, | ||
| params_dtype=torch.bfloat16, | ||
| device=device, | ||
| delay_wgrad_compute=delay, | ||
| fuse_wgrad_accumulation=False, | ||
| ), | ||
| ) | ||
|
|
||
|
|
||
| @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 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) | ||
| delayed, regular = _build_two_linears(True), _build_two_linears(False) | ||
| regular.load_state_dict(delayed.state_dict()) | ||
| 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()) | ||
| 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()) | ||
| 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()) | ||
| 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 |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Oops, something went wrong.
Oops, something went wrong.
Add this suggestion to a batch that can be applied as a single commit.
This suggestion is invalid because no changes were made to the code.
Suggestions cannot be applied while the pull request is closed.
Suggestions cannot be applied while viewing a subset of changes.
Only one suggestion per line can be applied in a batch.
Add this suggestion to a batch that can be applied as a single commit.
Applying suggestions on deleted lines is not supported.
You must change the existing code in this line in order to create a valid suggestion.
Outdated suggestions cannot be applied.
This suggestion has been applied or marked resolved.
Suggestions cannot be applied from pending reviews.
Suggestions cannot be applied on multi-line comments.
Suggestions cannot be applied while the pull request is queued to merge.
Suggestion cannot be applied right now. Please check back later.
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Is it safe? There's also
zero_grad(set_to_none=False)There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Yes.
set_to_none=Falseleaves a zeroed tensor whereset_to_none=TrueleavesNone, so the first contribution of the step takes the add path and lands on the same value; nothing in the branch needsgradto beNoneto start a step. Said so in the comment (7ed69a4).It also turned out not to be pinned: the test ran both settings but only asserted
second != first, which cannot show pollution, because a first-step value carried into the second still differs from the first. It now compares the second step's grad against a fresh model on the same weights and input, bit-exactly.