Repository navigation
Grouped Linear CPU Overhead Reductions By Removing Unnecessary Python Logic #3613
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
base: main
Are you sure you want to change the base?
Changes from all commits
86758da
89c80c6
f551085
a7a1010
0fef1c5
a65404b
033def7
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -4,7 +4,10 @@ | |
|
|
||
| from __future__ import annotations | ||
|
|
||
|
|
||
| from collections import Counter | ||
| from collections.abc import Iterable, Sequence | ||
| from unittest.mock import Mock | ||
| import io | ||
| import os | ||
| import math | ||
|
|
@@ -15,14 +18,17 @@ | |
|
|
||
| import transformer_engine | ||
| import torch | ||
| from torch.utils.checkpoint import checkpoint | ||
|
|
||
| import transformer_engine.common.recipe | ||
| import transformer_engine.pytorch as te | ||
| import transformer_engine.pytorch.ops as te_ops | ||
| from transformer_engine.pytorch.ops.basic.grouped_linear import ( | ||
| OUTPUT_BUFFER_KEY, | ||
| GRAD_INPUT_BUFFER_KEY, | ||
| is_op_fuser_grouped_tensor_path_supported, | ||
| ) | ||
| from transformer_engine.pytorch.quantization import FP8GlobalStateManager | ||
| from transformer_engine.pytorch.ops.fuser import OperationFuser | ||
| from transformer_engine.pytorch._extra_state import UNSAFE_PICKLE_EXTRA_STATE_ENV | ||
|
|
||
|
|
@@ -5508,3 +5514,181 @@ def test_linear_inference_loop( | |
| quantization=quantization, | ||
| recipe=recipe, | ||
| ) | ||
|
|
||
|
|
||
| @pytest.mark.parametrize("num_groups", [1, 8, 32]) | ||
| @pytest.mark.parametrize("recompute", [False, True]) | ||
| @pytest.mark.parametrize("single_grouped_weight", [False, True]) | ||
| def test_basic_grouped_runtime_quantizers( | ||
| monkeypatch, num_groups, recompute, single_grouped_weight | ||
| ): | ||
| """Reuse cached references while configuring each consumed weight quantizer.""" | ||
| fp8_recipe = transformer_engine.common.recipe.MXFP8BlockScaling() | ||
| if not is_op_fuser_grouped_tensor_path_supported(fp8_recipe, torch.bfloat16): | ||
| pytest.skip("Native MXFP8 grouped GEMM is unavailable.") | ||
| FP8GlobalStateManager.reset() | ||
| monkeypatch.setenv("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "1") | ||
| op = te_ops.GroupedLinear( | ||
| num_groups, | ||
| 128, | ||
| 128, | ||
| bias=False, | ||
| dtype=torch.bfloat16, | ||
| device="cuda", | ||
| single_grouped_weight=single_grouped_weight, | ||
| ) | ||
| model = te_ops.Sequential(op) | ||
| x = torch.randn(num_groups * 128, 128, device="cuda", dtype=torch.bfloat16, requires_grad=True) | ||
| splits = torch.full((num_groups,), 128, device="cuda", dtype=torch.int64) | ||
|
|
||
| def forward(inp): | ||
| with te.autocast(enabled=True, recipe=fp8_recipe): | ||
| return model(inp, splits) | ||
|
|
||
| def step(): | ||
| out = checkpoint(forward, x, use_reentrant=True) if recompute else forward(x) | ||
| out.sum().backward() | ||
|
|
||
| step() | ||
| counts = Counter() | ||
| original = MXFP8Quantizer.set_usage | ||
|
|
||
| def set_usage(q, **kwargs): | ||
| counts[id(q)] += 1 | ||
| return original(q, **kwargs) | ||
|
|
||
| monkeypatch.setattr(MXFP8Quantizer, "set_usage", set_usage) | ||
| original_getter = op.get_quantizer | ||
| monkeypatch.setattr(op, "get_quantizer", Mock(side_effect=AssertionError("runtime getter"))) | ||
| step() | ||
| forwards = 2 if recompute else 1 | ||
| assert counts[id(op._input_quantizers[0])] == forwards | ||
| assert counts[id(op._grad_output_quantizers[0])] == forwards + 1 | ||
| for i, quantizer in enumerate(op._weight_quantizers): | ||
| expected = forwards if not single_grouped_weight or i == 0 else 0 | ||
| assert counts[id(quantizer)] == expected | ||
| assert all(counts[id(q)] == 0 for q in op._input_quantizers[1:]) | ||
| assert all(counts[id(q)] == 0 for q in op._grad_output_quantizers[1:]) | ||
| assert len(op._input_quantizers) == num_groups | ||
| assert len(op._weight_quantizers) == num_groups | ||
| assert len(op._grad_output_quantizers) == num_groups | ||
| assert op._input_quantizers[0] is not op._weight_quantizers[0] | ||
| assert op._input_quantizers[0] is not op._grad_output_quantizers[0] | ||
| # Delayed scaling keeps distinct per-expert state, and resetting invalidates cached refs. | ||
| monkeypatch.setattr(op, "get_quantizer", original_getter) | ||
| old_input = op._input_quantizers[0] | ||
| with te.autocast(enabled=True, recipe=transformer_engine.common.recipe.DelayedScaling()): | ||
| op.reset_recipe_state(recipe=transformer_engine.common.recipe.DelayedScaling()) | ||
| assert op._input_quantizers[0] is not old_input | ||
| assert len({id(q) for q in op._input_quantizers}) == num_groups | ||
|
|
||
|
|
||
| @pytest.mark.parametrize( | ||
| "fp8_recipe", | ||
| [ | ||
| transformer_engine.common.recipe.MXFP8BlockScaling(), | ||
| transformer_engine.common.recipe.Float8CurrentScaling(), | ||
| ], | ||
| ) | ||
| def test_grouped_discrete_weight_quantizers(fp8_recipe): | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I disagree with this contract. The op's weight quantizer should take precedence, and the param's quantizer should adapt to match the op's quantizer. This way you can have the same logic for configuring the weight quantizer, no matter if the weights are quantized or not. |
||
| """Discrete weights use their own quantizers and preserve already quantized tensors.""" | ||
| if not is_op_fuser_grouped_tensor_path_supported(fp8_recipe, torch.bfloat16): | ||
| pytest.skip("Native grouped GEMM is unavailable.") | ||
| FP8GlobalStateManager.reset() | ||
| op = te_ops.GroupedLinear(3, 128, 128, bias=False, dtype=torch.bfloat16, device="cuda") | ||
| with te.autocast(enabled=True, recipe=fp8_recipe): | ||
| op.reset_recipe_state(recipe=fp8_recipe) | ||
| quantizers = op._weight_quantizers | ||
| assert len({id(q) for q in quantizers}) == 3 | ||
| for q in quantizers: | ||
| q.set_usage(rowwise=True, columnwise=True) | ||
| weights = [torch.randn(128, 128, device="cuda", dtype=torch.bfloat16) for _ in range(3)] | ||
| reference = [q(w) for w, q in zip(weights, quantizers)] | ||
| weights[1] = reference[1] | ||
| quantizer_spies = [Mock(wraps=q) for q in quantizers] | ||
| outputs = op._get_discrete_weights_for_gemm( | ||
| weights, | ||
| quantizer_spies, | ||
| columnwise_usage=True, | ||
| with_quantized_compute=True, | ||
| dtype=torch.bfloat16, | ||
| ) | ||
| for i in (0, 2): | ||
| quantizer_spies[i].assert_called_once_with(weights[i]) | ||
| quantizer_spies[i].set_usage.assert_called_once_with(rowwise=True, columnwise=True) | ||
| quantizer_spies[1].assert_not_called() | ||
| quantizer_spies[1].set_usage.assert_not_called() | ||
| assert outputs[1] is weights[1] | ||
| for output, ref in zip(outputs, reference): | ||
| torch.testing.assert_close(output.dequantize(), ref.dequantize(), rtol=0, atol=0) | ||
| scales = [w._rowwise_scale_inv if fp8_recipe.mxfp8() else w._scale_inv for w in outputs] | ||
| assert len({s.data_ptr() for s in scales}) == 3 | ||
|
|
||
|
|
||
| @pytest.mark.parametrize("num_groups", [1, 8]) | ||
| @pytest.mark.parametrize("fused", [False, True]) | ||
| @pytest.mark.parametrize("delayed", [False, True]) | ||
| def test_basic_grouped_wgrad_storage_and_hooks(monkeypatch, fused, delayed, num_groups): | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. This is testing advanced features in |
||
| """Unfused wgrads own independent storage; external buffers and hooks retain their contract.""" | ||
| from transformer_engine.pytorch.ops.basic import grouped_linear as gl | ||
|
|
||
| if not is_op_fuser_grouped_tensor_path_supported(None, torch.bfloat16): | ||
| pytest.skip("Native grouped GEMM is unavailable.") | ||
| op = te_ops.GroupedLinear( | ||
| num_groups, | ||
| 128, | ||
| 128, | ||
| bias=False, | ||
| device="cuda", | ||
| dtype=torch.bfloat16, | ||
| accumulate_into_main_grad=fused, | ||
| delay_wgrad_compute=delayed, | ||
| ) | ||
| model = te_ops.Sequential(op) | ||
| flat = torch.full((num_groups, 128, 128), 0.25, device="cuda", dtype=torch.float32) | ||
| weights = [getattr(op, f"weight{i}") for i in range(num_groups)] | ||
| for i, weight in enumerate(weights): | ||
| with torch.no_grad(): | ||
| weight.fill_(1) | ||
| if fused: | ||
| weight.main_grad = flat[i] | ||
| weight.overwrite_main_grad = False | ||
| weight.grad_added_to_main_grad = False | ||
| x = torch.full( | ||
| (num_groups * 128, 128), 0.125, device="cuda", dtype=torch.bfloat16, requires_grad=True | ||
| ) | ||
| splits = torch.full((num_groups,), 128, device="cuda", dtype=torch.int64) | ||
| seen = [] | ||
| original = gl.general_grouped_gemm_for_grouped_tensor | ||
|
|
||
| def gemm(a, b, out, **kwargs): | ||
| if kwargs.get("layout") == "NT": | ||
| assert len(out) == num_groups | ||
| num_storages = len({g.untyped_storage().data_ptr() for g in out}) | ||
| assert num_storages == (1 if fused else num_groups) | ||
| assert len({g.data_ptr() for g in out}) == num_groups | ||
| if fused: | ||
| assert out[0].data_ptr() == flat.data_ptr() | ||
| else: | ||
| assert all(g._base is None for g in out) | ||
| seen.append(out) | ||
| return original(a, b, out, **kwargs) | ||
|
|
||
| monkeypatch.setattr(gl, "general_grouped_gemm_for_grouped_tensor", gemm) | ||
| hook = Mock() | ||
| op.register_wgrad_accumulation_and_reduce_hooks(hook) | ||
| y = model(x, splits) | ||
| torch.testing.assert_close(y, torch.full_like(y, 16), rtol=0, atol=0) | ||
| y.backward(torch.full_like(y, 0.5)) | ||
| if delayed: | ||
| assert not seen | ||
| op.backward_dw() | ||
| hook.assert_called_once() | ||
| assert len(seen) == 1 | ||
| for i, weight in enumerate(weights): | ||
| grad = weight.main_grad if fused else weight.grad | ||
| torch.testing.assert_close( | ||
| grad, torch.full_like(grad, 8.25 if fused else 8), rtol=0, atol=0 | ||
| ) | ||
| if fused: | ||
| assert grad.data_ptr() == flat[i].data_ptr() | ||
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.
I think this test is a liability. Caching is an internal implementation detail, but the test implies it's a contract.