diff --git a/tests/pytorch/test_custom_recipe.py b/tests/pytorch/test_custom_recipe.py index a1a6ef7b37..981385f14f 100644 --- a/tests/pytorch/test_custom_recipe.py +++ b/tests/pytorch/test_custom_recipe.py @@ -129,7 +129,7 @@ def quantizer_factory(role): assert inp.grad is not None -def test_custom_recipe_grouped_linear_sanity(): +def test_custom_recipe_grouped_linear_matches_current_scaling(): available, reason = te.is_fp8_available(return_reason=True) if not torch.cuda.is_available() or not available: pytest.skip(f"FP8 unsupported on this device: {reason}") @@ -144,8 +144,25 @@ def test_custom_recipe_grouped_linear_sanity(): m_splits = [16] * num_gemms batch = sum(m_splits) - model = GroupedLinear(num_gemms, in_features, out_features, params_dtype=torch.bfloat16).cuda() - inp = torch.randn(batch, in_features, device="cuda", dtype=torch.bfloat16, requires_grad=True) + model_ref = GroupedLinear( + num_gemms, in_features, out_features, params_dtype=torch.bfloat16 + ).cuda() + model_custom = GroupedLinear( + num_gemms, in_features, out_features, params_dtype=torch.bfloat16 + ).cuda() + model_custom.load_state_dict(model_ref.state_dict()) + + base_inp = torch.randn(batch, in_features, device="cuda", dtype=torch.bfloat16) + inp_ref = base_inp.clone().detach().requires_grad_(True) + inp_custom = base_inp.clone().detach().requires_grad_(True) + + with autocast(enabled=True, recipe=recipe.Float8CurrentScaling()): + out_ref = model_ref(inp_ref, m_splits) + + scale = torch.ones(out_features, device="cuda", dtype=torch.float32) + scale[0] = 1e8 + scale[1] = 1e-8 + (out_ref.float() * scale.view(1, -1)).sum().backward() def quantizer_factory(role): if role is None: @@ -157,11 +174,21 @@ def quantizer_factory(role): custom_recipe = recipe.CustomRecipe(qfactory=quantizer_factory) with autocast(enabled=True, recipe=custom_recipe): - out = model(inp, m_splits) - loss = out.float().sum() - loss.backward() + out_custom = model_custom(inp_custom, m_splits) + (out_custom.float() * scale.view(1, -1)).sum().backward() - assert inp.grad is not None + torch.testing.assert_close(out_ref, out_custom, rtol=0, atol=0) + + assert inp_ref.grad is not None and inp_custom.grad is not None + torch.testing.assert_close(inp_ref.grad, inp_custom.grad, rtol=0, atol=0) + + ref_params = dict(model_ref.named_parameters()) + custom_params = dict(model_custom.named_parameters()) + assert ref_params.keys() == custom_params.keys() + for name, ref_param in ref_params.items(): + custom_param = custom_params[name] + assert ref_param.grad is not None and custom_param.grad is not None + torch.testing.assert_close(ref_param.grad, custom_param.grad, rtol=0, atol=0) def test_custom_recipe_matches_current_scaling(): diff --git a/tests/pytorch/test_hybrid_quantization.py b/tests/pytorch/test_hybrid_quantization.py index 74ec0a05ec..abf311e503 100644 --- a/tests/pytorch/test_hybrid_quantization.py +++ b/tests/pytorch/test_hybrid_quantization.py @@ -10,6 +10,7 @@ import torch import transformer_engine.pytorch as te +import transformer_engine.pytorch.module._grouped_quantization as grouped_quantization import transformer_engine_torch as tex from hybrid_quantization_utils import ( @@ -769,30 +770,63 @@ def test_linear_builtin_delayed_scaling_rejects_save_original_input(self): with autocast(enabled=True, recipe=recipe.DelayedScaling()): model(inp) - def test_grouped_linear_classifies_requantization_safety_once_per_generation(self): + def test_grouped_linear_unsafe_custom_input_disables_save_original_input(self): counters = [{"calls": 0}, {"calls": 0}] - input_quantizers = [_CountingUnsafeIdentityQuantizer(counter) for counter in counters] - generation = [] - for input_quantizer in input_quantizers: - generation.extend((input_quantizer, IdentityQuantizer(), IdentityQuantizer())) + input_index = 0 + + def qfactory(role): + nonlocal input_index + if ( + role is not None + and role.module_type == "grouped_linear" + and role.tensor_type == "input" + ): + counter = counters[input_index % len(counters)] + input_index += 1 + return _CountingUnsafeIdentityQuantizer(counter) + return IdentityQuantizer() module = GroupedLinear( 2, - 16, - 16, + 128, + 128, bias=False, - device="meta", - ) - module.quantizers["scaling_fwd"] = generation + params_dtype=torch.bfloat16, + save_original_input=True, + ).cuda() + reference = GroupedLinear( + 2, + 128, + 128, + bias=False, + params_dtype=torch.bfloat16, + ).cuda() + reference.load_state_dict(module.state_dict()) + inp = torch.randn(128, 128, dtype=torch.bfloat16, device="cuda", requires_grad=True) + reference_inp = inp.detach().clone().requires_grad_() + custom_recipe = recipe.CustomRecipe(qfactory=qfactory) - module._validate_quantizer_generation(fwd=True) - assert module._unsafe_requantization_input_quantizer is input_quantizers[0] - assert [counter.get("safety_calls", 0) for counter in counters] == [1, 0] + with pytest.warns(UserWarning, match="Ignoring save_original_input=True"): + with autocast(enabled=True, recipe=custom_recipe): + out = module(inp, [64, 64]) + reference_out = reference(reference_inp, [64, 64]) + + calls_after_forward = [counter["calls"] for counter in counters] + out.float().sum().backward() + reference_out.float().sum().backward() - # The generation list is stable between forwards, so the O(1) identity - # guard must avoid re-running capability checks. - module._validate_quantizer_generation(fwd=True) + # Group compatibility is established when the CustomRecipe generation + # is constructed. Runtime policy reads only expert 0. assert [counter.get("safety_calls", 0) for counter in counters] == [1, 0] + assert [counter["calls"] for counter in counters] == calls_after_forward + reference_parameters = dict(reference.named_parameters()) + for name, parameter in module.named_parameters(): + torch.testing.assert_close( + parameter.grad, + reference_parameters[name].grad, + rtol=0.0, + atol=0.0, + ) @staticmethod def _counting_identity_hybrid_recipe( @@ -4200,8 +4234,8 @@ def test_transformer_layer(self): class TestHybridGroupedLinearValidation: """GroupedLinear generation-validation and split-dispatch coverage. - Structural compatibility is validated once per real quantizer generation. - Steady-state dispatch reads the first expert after that uniformity check.""" + CustomRecipe compatibility is validated once per real quantizer generation. + Built-in recipes skip that validation and steady-state generation processing.""" @pytest.mark.parametrize( "quantizers", @@ -4215,17 +4249,14 @@ class TestHybridGroupedLinearValidation: ], ) def test_uniform_lists_validate(self, quantizers): - from transformer_engine.pytorch.module.grouped_linear import ( - _validate_grouped_quantizer_list, + grouped_quantization.validate_grouped_quantizer_list( + quantizers, + operand_name="input", ) - _validate_grouped_quantizer_list(quantizers, operand_name="input") - def test_plain_custom_quantizer_uses_python_split_fallback(self, monkeypatch): - import transformer_engine.pytorch.module.grouped_linear as grouped_linear - monkeypatch.setattr( - grouped_linear.tex, + grouped_quantization.tex, "split_quantize", lambda *args, **kwargs: pytest.fail("entered native split_quantize"), ) @@ -4233,30 +4264,156 @@ def test_plain_custom_quantizer_uses_python_split_fallback(self, monkeypatch): tensor = torch.randn((4, 8)) quantizers = [_CountingPythonQuantizer(calls) for _ in range(2)] - out = grouped_linear._split_quantize_non_hybrid( + out, dbiases = grouped_quantization._split_quantize( tensor, [2, 2], quantizers, tensor.dtype, + compute_dbias=True, ) + expected_parts = torch.split(tensor, [2, 2]) + assert dbiases is not None + for actual, expected in zip(dbiases, expected_parts): + torch.testing.assert_close(actual, expected.sum(dim=0), rtol=0.0, atol=0.0) assert len(calls) == 2 - for actual, expected in zip(out, torch.split(tensor, [2, 2])): + for actual, expected in zip(out, expected_parts): torch.testing.assert_close(actual.dequantize(), expected, rtol=0.0, atol=0.0) + def test_native_bgrad_dispatch_is_quantizer_driven(self, monkeypatch): + tensor = torch.randn((4, 8)) + split_sizes = [2, 2] + quantizers = [_make_fp8_quantizer() for _ in split_sizes] + calls = [] + + def fake_bgrad_quantize(tensor_part, quantizer): + calls.append((tensor_part, quantizer)) + return tensor_part.sum(dim=0), tensor_part.clone() + + monkeypatch.setattr( + grouped_quantization.tex, + "bgrad_quantize", + fake_bgrad_quantize, + ) + monkeypatch.setattr( + grouped_quantization.tex, + "split_quantize", + lambda *args, **kwargs: pytest.fail("entered split_quantize instead of bgrad_quantize"), + ) + + outputs, dbiases = grouped_quantization._split_quantize( + tensor, + split_sizes, + quantizers, + tensor.dtype, + compute_dbias=True, + ) + + expected_parts = torch.split(tensor, split_sizes) + assert [quantizer for _, quantizer in calls] == quantizers + assert dbiases is not None + for output, dbias, expected in zip(outputs, dbiases, expected_parts): + torch.testing.assert_close(output, expected, rtol=0.0, atol=0.0) + torch.testing.assert_close(dbias, expected.sum(dim=0), rtol=0.0, atol=0.0) + + def test_native_split_without_dbias_uses_bulk_path(self, monkeypatch): + tensor = torch.randn((4, 8)) + split_sizes = [2, 2] + quantizers = [_make_fp8_quantizer() for _ in split_sizes] + calls = [] + + def fake_split_quantize( + tensor_arg, + split_sizes_arg, + quantizers_arg, + *, + disable_bulk_allocation=False, + ): + calls.append((tensor_arg, split_sizes_arg, quantizers_arg, disable_bulk_allocation)) + return torch.split(tensor_arg, split_sizes_arg) + + monkeypatch.setattr( + grouped_quantization.tex, + "bgrad_quantize", + lambda *args, **kwargs: pytest.fail("entered bgrad_quantize without a dbias request"), + ) + monkeypatch.setattr( + grouped_quantization.tex, + "split_quantize", + fake_split_quantize, + ) + + outputs, dbiases = grouped_quantization._split_quantize( + tensor, + split_sizes, + quantizers, + tensor.dtype, + disable_bulk_allocation=True, + ) + + assert dbiases is None + assert len(calls) == 1 + actual_tensor, actual_splits, actual_quantizers, actual_disable_bulk = calls[0] + assert actual_tensor is tensor + assert actual_splits is split_sizes + assert actual_quantizers is quantizers + assert actual_disable_bulk is True + for output, expected in zip(outputs, torch.split(tensor, split_sizes)): + torch.testing.assert_close(output, expected, rtol=0.0, atol=0.0) + + def test_native_bgrad_quantize_matches_individual_quantization_exactly(self): + tensor = torch.randn(32, 128, dtype=torch.bfloat16, device="cuda") + split_sizes = [16, 16] + quantizers = [_make_fp8_quantizer() for _ in split_sizes] + for quantizer in quantizers: + quantizer.internal = True + quantizer.set_usage(rowwise=True, columnwise=False) + reference_quantizers = [quantizer.copy() for quantizer in quantizers] + + outputs, dbiases = grouped_quantization._split_quantize( + tensor, + split_sizes, + quantizers, + tensor.dtype, + compute_dbias=True, + ) + + tensor_parts = torch.split(tensor, split_sizes) + reference_outputs = [ + quantizer.quantize(tensor_part) + for quantizer, tensor_part in zip(reference_quantizers, tensor_parts) + ] + for index, (output, reference_output) in enumerate(zip(outputs, reference_outputs)): + _assert_storage_data_exact( + output, + reference_output, + context=f"native grouped split {index}", + ) + assert dbiases is not None + for dbias, tensor_part in zip(dbiases, tensor_parts): + torch.testing.assert_close( + dbias, + tensor_part.sum(dim=0), + rtol=0.0, + atol=0.0, + ) + @pytest.mark.parametrize("direction", ("rowwise", "columnwise")) def test_hybrid_custom_child_uses_python_split_fallback( self, monkeypatch, direction, ): - import transformer_engine.pytorch.module.grouped_linear as grouped_linear - monkeypatch.setattr( - grouped_linear.tex, + grouped_quantization.tex, "split_quantize", lambda *args, **kwargs: pytest.fail("entered native split_quantize"), ) + monkeypatch.setattr( + grouped_quantization.tex, + "bgrad_quantize", + lambda *args, **kwargs: pytest.fail("hybrid quantization entered bgrad_quantize"), + ) calls = [] quantizers = [] for _ in range(2): @@ -4272,63 +4429,79 @@ def test_hybrid_custom_child_uses_python_split_fallback( ) quantizers.append(quantizer) - out = grouped_linear._split_quantize_hybrid( - torch.randn((4, 8)), + tensor = torch.randn((4, 8)) + out, dbiases = grouped_quantization._split_quantize( + tensor, [2, 2], quantizers, + tensor.dtype, + compute_dbias=True, ) + expected_parts = torch.split(tensor, [2, 2]) + assert dbiases is not None + for actual, expected in zip(dbiases, expected_parts): + torch.testing.assert_close(actual, expected.sum(dim=0), rtol=0.0, atol=0.0) assert len(calls) == 2 assert all(result.rowwise_sub_storage is not None for result in out) assert all(result.columnwise_sub_storage is not None for result in out) + for result, expected in zip(out, expected_parts): + torch.testing.assert_close( + result.rowwise_sub_storage.dequantize(), + expected, + rtol=0.0, + atol=0.0, + ) + torch.testing.assert_close( + result.columnwise_sub_storage.dequantize(), + expected, + rtol=0.0, + atol=0.0, + ) def test_mixed_hybrid_and_plain_raises(self): - from transformer_engine.pytorch.module.grouped_linear import ( - _validate_grouped_quantizer_list, - ) - quantizers = [ _make_hybrid_quantizer_fp8_row_fp4_col(), _make_fp8_quantizer(), _make_hybrid_quantizer_fp8_row_fp4_col(), ] with pytest.raises(ValueError, match="mix HybridQuantizer and non-hybrid"): - _validate_grouped_quantizer_list(quantizers, operand_name="input") + grouped_quantization.validate_grouped_quantizer_list( + quantizers, + operand_name="input", + ) def test_none_plus_hybrid_raises(self): - from transformer_engine.pytorch.module.grouped_linear import ( - _validate_grouped_quantizer_list, - ) - quantizers = [ _make_hybrid_quantizer_fp8_row_fp4_col(), None, _make_hybrid_quantizer_fp8_row_fp4_col(), ] with pytest.raises(ValueError, match="mix None and concrete quantizers"): - _validate_grouped_quantizer_list(quantizers, operand_name="input") + grouped_quantization.validate_grouped_quantizer_list( + quantizers, + operand_name="input", + ) def test_mixed_identity_dtype_raises(self): - from transformer_engine.pytorch.module.grouped_linear import ( - _validate_grouped_quantizer_list, - ) - quantizers = [ IdentityQuantizer(dtype=torch.bfloat16), IdentityQuantizer(dtype=torch.float16), ] with pytest.raises(ValueError, match="incompatible plain backend configurations"): - _validate_grouped_quantizer_list(quantizers, operand_name="input") + grouped_quantization.validate_grouped_quantizer_list( + quantizers, + operand_name="input", + ) def test_distinct_delayed_scaling_state_is_allowed(self): - from transformer_engine.pytorch.module.grouped_linear import ( - _validate_grouped_quantizer_list, - ) - quantizers = [_make_delayed_quantizer(), _make_delayed_quantizer()] quantizers[1].scale.fill_(2.0) quantizers[1].amax.fill_(3.0) - _validate_grouped_quantizer_list(quantizers, operand_name="input") + grouped_quantization.validate_grouped_quantizer_list( + quantizers, + operand_name="input", + ) @pytest.mark.skipif(not fp8_available, reason=reason_for_no_fp8) @pytest.mark.parametrize( @@ -4350,10 +4523,6 @@ def test_distinct_delayed_scaling_state_is_allowed(self): ], ) def test_hybrid_split_quantize_respects_parent_usage_flags(self, usage, expected): - from transformer_engine.pytorch.module.grouped_linear import ( - _split_quantize_hybrid, - ) - tensor = torch.randn(32, 128, dtype=torch.bfloat16, device="cuda") quantizers = [ HybridQuantizer( @@ -4365,15 +4534,34 @@ def test_hybrid_split_quantize_respects_parent_usage_flags(self, usage, expected for quantizer in quantizers: quantizer.set_usage(rowwise=usage[0], columnwise=usage[1]) - out = _split_quantize_hybrid(tensor, [16, 16], quantizers) + out, dbiases = grouped_quantization._split_quantize( + tensor, + [16, 16], + quantizers, + tensor.dtype, + ) + assert dbiases is None assert [storage.get_usages() for storage in out] == [expected, expected] + for index, (storage, tensor_part, quantizer) in enumerate( + zip(out, torch.split(tensor, [16, 16]), quantizers) + ): + if usage[0]: + _assert_storage_data_exact( + storage.rowwise_sub_storage, + quantizer.rowwise_quantizer.quantize(tensor_part), + context=f"grouped split {index} rowwise", + ) + if usage[1]: + _assert_storage_data_exact( + storage.columnwise_sub_storage, + quantizer.columnwise_quantizer.quantize(tensor_part), + context=f"grouped split {index} columnwise", + ) @_XFAIL_HOPPER_COLUMNWISE_PER_TENSOR_FP8 def test_columnwise_only_rowwise_dequantized_uses_transient_grouped_row(self, monkeypatch): - import transformer_engine.pytorch.module.grouped_linear as grouped_linear - - real_split_quantize = grouped_linear.tex.split_quantize + real_split_quantize = grouped_quantization.tex.split_quantize calls = [] def tracked_split_quantize(tensor, m_splits, quantizers, **kwargs): @@ -4381,7 +4569,7 @@ def tracked_split_quantize(tensor, m_splits, quantizers, **kwargs): calls.append((tensor, result)) return result - monkeypatch.setattr(grouped_linear.tex, "split_quantize", tracked_split_quantize) + monkeypatch.setattr(grouped_quantization.tex, "split_quantize", tracked_split_quantize) tensor = torch.randn(32, 128, dtype=torch.bfloat16, device="cuda") quantizers = [ HybridQuantizer( @@ -4394,8 +4582,14 @@ def tracked_split_quantize(tensor, m_splits, quantizers, **kwargs): for quantizer in quantizers: quantizer.set_usage(rowwise=False, columnwise=True) - out = grouped_linear._split_quantize_hybrid(tensor, [16, 16], quantizers) + out, dbiases = grouped_quantization._split_quantize( + tensor, + [16, 16], + quantizers, + tensor.dtype, + ) + assert dbiases is None assert len(calls) == 2 expected_columnwise_source = torch.cat( [result.dequantize(dtype=tensor.dtype) for result in calls[0][1]], dim=0 @@ -4404,20 +4598,25 @@ def tracked_split_quantize(tensor, m_splits, quantizers, **kwargs): assert calls[1][0] is not tensor assert all(storage.rowwise_sub_storage is None for storage in out) assert all(storage.columnwise_sub_storage is not None for storage in out) + for index, (storage, reference_columnwise) in enumerate(zip(out, calls[1][1])): + _assert_storage_data_exact( + storage.columnwise_sub_storage, + reference_columnwise, + context=f"transient-row grouped split {index} columnwise", + ) @requires_nvfp4 @pytest.mark.parametrize("m_splits", ([128, 128], [0, 128])) def test_nvfp4_rowwise_dequantized_preserves_source_dtype(self, monkeypatch, m_splits): - import transformer_engine.pytorch.module.grouped_linear as grouped_linear - - real_split_quantize = grouped_linear.tex.split_quantize - input_dtypes = [] + real_split_quantize = grouped_quantization.tex.split_quantize + calls = [] def tracked_split_quantize(tensor, splits, quantizers, **kwargs): - input_dtypes.append(tensor.dtype) - return real_split_quantize(tensor, splits, quantizers, **kwargs) + result = real_split_quantize(tensor, splits, quantizers, **kwargs) + calls.append((tensor, result)) + return result - monkeypatch.setattr(grouped_linear.tex, "split_quantize", tracked_split_quantize) + monkeypatch.setattr(grouped_quantization.tex, "split_quantize", tracked_split_quantize) tensor = torch.randn( sum(m_splits), 128, @@ -4433,22 +4632,49 @@ def tracked_split_quantize(tensor, splits, quantizers, **kwargs): for _ in m_splits ] - out = grouped_linear._split_quantize_hybrid(tensor, m_splits, quantizers) + out, dbiases = grouped_quantization._split_quantize( + tensor, + m_splits, + quantizers, + tensor.dtype, + ) - assert input_dtypes == [tensor.dtype, tensor.dtype] + assert dbiases is None + assert [call[0].dtype for call in calls] == [tensor.dtype, tensor.dtype] assert len(out) == len(m_splits) + expected_columnwise_source = torch.cat( + [result.dequantize(dtype=tensor.dtype) for result in calls[0][1]], dim=0 + ) + torch.testing.assert_close( + calls[1][0], + expected_columnwise_source, + rtol=0.0, + atol=0.0, + ) + for index, (storage, reference_rowwise, reference_columnwise) in enumerate( + zip(out, calls[0][1], calls[1][1]) + ): + _assert_storage_data_exact( + storage.rowwise_sub_storage, + reference_rowwise, + context=f"NVFP4 grouped split {index} rowwise", + ) + _assert_storage_data_exact( + storage.columnwise_sub_storage, + reference_columnwise, + context=f"NVFP4 grouped split {index} columnwise", + ) def test_rowwise_only_skips_columnwise_quantization(self, monkeypatch): - import transformer_engine.pytorch.module.grouped_linear as grouped_linear - - real_split_quantize = grouped_linear.tex.split_quantize + real_split_quantize = grouped_quantization.tex.split_quantize calls = [] def tracked_split_quantize(tensor, m_splits, quantizers, **kwargs): - calls.append(tensor) - return real_split_quantize(tensor, m_splits, quantizers, **kwargs) + result = real_split_quantize(tensor, m_splits, quantizers, **kwargs) + calls.append((tensor, result)) + return result - monkeypatch.setattr(grouped_linear.tex, "split_quantize", tracked_split_quantize) + monkeypatch.setattr(grouped_quantization.tex, "split_quantize", tracked_split_quantize) tensor = torch.randn(32, 128, dtype=torch.bfloat16, device="cuda") quantizers = [ HybridQuantizer( @@ -4461,24 +4687,36 @@ def tracked_split_quantize(tensor, m_splits, quantizers, **kwargs): for quantizer in quantizers: quantizer.set_usage(rowwise=True, columnwise=False) - out = grouped_linear._split_quantize_hybrid(tensor, [16, 16], quantizers) + out, dbiases = grouped_quantization._split_quantize( + tensor, + [16, 16], + quantizers, + tensor.dtype, + ) - assert calls == [tensor] + assert dbiases is None + assert len(calls) == 1 + assert calls[0][0] is tensor assert all(storage.rowwise_sub_storage is not None for storage in out) assert all(storage.columnwise_sub_storage is None for storage in out) + for index, (storage, reference_rowwise) in enumerate(zip(out, calls[0][1])): + _assert_storage_data_exact( + storage.rowwise_sub_storage, + reference_rowwise, + context=f"rowwise-only grouped split {index}", + ) @_XFAIL_HOPPER_COLUMNWISE_PER_TENSOR_FP8 def test_original_source_preserves_two_bulk_call_fast_path(self, monkeypatch): - import transformer_engine.pytorch.module.grouped_linear as grouped_linear - - real_split_quantize = grouped_linear.tex.split_quantize + real_split_quantize = grouped_quantization.tex.split_quantize calls = [] def tracked_split_quantize(tensor, m_splits, quantizers, **kwargs): - calls.append(tensor) - return real_split_quantize(tensor, m_splits, quantizers, **kwargs) + result = real_split_quantize(tensor, m_splits, quantizers, **kwargs) + calls.append((tensor, result)) + return result - monkeypatch.setattr(grouped_linear.tex, "split_quantize", tracked_split_quantize) + monkeypatch.setattr(grouped_quantization.tex, "split_quantize", tracked_split_quantize) tensor = torch.randn(32, 128, dtype=torch.bfloat16, device="cuda") quantizers = [ HybridQuantizer( @@ -4489,17 +4727,34 @@ def tracked_split_quantize(tensor, m_splits, quantizers, **kwargs): for _ in range(2) ] - out = grouped_linear._split_quantize_hybrid(tensor, [16, 16], quantizers) + out, dbiases = grouped_quantization._split_quantize( + tensor, + [16, 16], + quantizers, + tensor.dtype, + ) - assert calls == [tensor, tensor] + assert dbiases is None + assert len(calls) == 2 + assert calls[0][0] is tensor + assert calls[1][0] is tensor assert all(storage.rowwise_sub_storage is not None for storage in out) assert all(storage.columnwise_sub_storage is not None for storage in out) + for index, (storage, reference_rowwise, reference_columnwise) in enumerate( + zip(out, calls[0][1], calls[1][1]) + ): + _assert_storage_data_exact( + storage.rowwise_sub_storage, + reference_rowwise, + context=f"original-source grouped split {index} rowwise", + ) + _assert_storage_data_exact( + storage.columnwise_sub_storage, + reference_columnwise, + context=f"original-source grouped split {index} columnwise", + ) def test_validation_rejects_mixed_columnwise_source_policies(self): - from transformer_engine.pytorch.module.grouped_linear import ( - _validate_grouped_quantizer_list, - ) - quantizers = [ HybridQuantizer( rowwise_quantizer=_make_fp8_quantizer(), @@ -4510,13 +4765,12 @@ def test_validation_rejects_mixed_columnwise_source_policies(self): ] with pytest.raises(ValueError, match="mixed columnwise source policies"): - _validate_grouped_quantizer_list(quantizers, operand_name="input") + grouped_quantization.validate_grouped_quantizer_list( + quantizers, + operand_name="input", + ) def test_validation_rejects_same_family_config_mismatch(self): - from transformer_engine.pytorch.module.grouped_linear import ( - _validate_grouped_quantizer_list, - ) - quantizers = [_make_fp8_quantizer(), _make_fp8_quantizer()] quantizers[1].force_pow_2_scales = True @@ -4524,11 +4778,45 @@ def test_validation_rejects_same_family_config_mismatch(self): ValueError, match="incompatible plain backend configurations", ): - _validate_grouped_quantizer_list(quantizers, operand_name="input") + grouped_quantization.validate_grouped_quantizer_list( + quantizers, + operand_name="input", + ) - def test_validation_runs_only_with_quantizer_generation(self, monkeypatch): - import transformer_engine.pytorch.module.grouped_linear as grouped_linear + def test_builtin_recipe_skips_custom_grouped_validation(self, monkeypatch): + def unexpected_validation(*_args, **_kwargs): + pytest.fail("built-in recipes must not run custom grouped-quantizer validation") + + monkeypatch.setattr( + grouped_quantization, + "validate_grouped_quantizer_list", + unexpected_validation, + ) + model = GroupedLinear( + 2, + 128, + 128, + bias=False, + params_dtype=torch.bfloat16, + ).cuda() + tensor = torch.randn(128, 128, dtype=torch.bfloat16, device="cuda") + m_splits = torch.tensor([64, 64], dtype=torch.int64) + + with torch.no_grad(), autocast(enabled=True, recipe=recipe.DelayedScaling()): + model(tensor, m_splits) + assert model._custom_quantizer_cache == {} + + def unexpected_custom_validation(*_args, **_kwargs): + pytest.fail("built-in recipes must not validate custom quantizers per forward") + + monkeypatch.setattr( + model, + "_validate_custom_recipe_quantizers", + unexpected_custom_validation, + ) + model(tensor, m_splits) + def test_validation_runs_only_with_quantizer_generation(self, monkeypatch): def make_qfactory(columnwise_source): def qfactory(_role): return HybridQuantizer( @@ -4544,7 +4832,7 @@ def qfactory(_role): m_splits = torch.tensor([64, 64], dtype=torch.int64) original_recipe = recipe.CustomRecipe(qfactory=make_qfactory("original")) - real_validate = grouped_linear._validate_grouped_quantizer_list + real_validate = grouped_quantization.validate_grouped_quantizer_list validation_calls = [] def tracked_validate(quantizers, *, operand_name="operand"): @@ -4552,26 +4840,26 @@ def tracked_validate(quantizers, *, operand_name="operand"): return real_validate(quantizers, operand_name=operand_name) monkeypatch.setattr( - grouped_linear, - "_validate_grouped_quantizer_list", + grouped_quantization, + "validate_grouped_quantizer_list", tracked_validate, ) with torch.no_grad(), autocast(enabled=True, recipe=original_recipe): model(tensor, m_splits) first_call_count = len(validation_calls) - first_generation = model._validated_quantizer_generations["scaling_fwd"] + first_generation = model._custom_quantizer_cache["scaling_fwd"] assert first_call_count > 0 with torch.no_grad(), autocast(enabled=True, recipe=original_recipe): model(tensor, m_splits) assert len(validation_calls) == first_call_count - assert model._validated_quantizer_generations["scaling_fwd"] is first_generation + assert model._custom_quantizer_cache["scaling_fwd"] is first_generation rebuilt_recipe = recipe.CustomRecipe(qfactory=make_qfactory("rowwise_dequantized")) with torch.no_grad(), autocast(enabled=True, recipe=rebuilt_recipe): model(tensor, m_splits) - rebuilt_generation = model._validated_quantizer_generations["scaling_fwd"] + rebuilt_generation = model._custom_quantizer_cache["scaling_fwd"] assert len(validation_calls) > first_call_count assert rebuilt_generation is not first_generation assert rebuilt_generation[0].columnwise_source == "rowwise_dequantized" @@ -4597,7 +4885,7 @@ def mixed_source_qfactory(role): with pytest.raises(ValueError, match="mixed columnwise source policies"): with torch.no_grad(), autocast(enabled=True, recipe=mixed_recipe): model(tensor, m_splits) - assert model._validated_quantizer_generations["scaling_fwd"] is rebuilt_generation + assert model._custom_quantizer_cache["scaling_fwd"] is rebuilt_generation # Stale invalid recipe metadata must not affect the non-quantized path. with torch.no_grad(): @@ -4606,10 +4894,6 @@ def mixed_source_qfactory(role): @requires_fp8_and_nvfp4 def test_hybrid_split_quantize_honors_rowwise_dequantized_source(self): """NVFP4 column data must derive from the actual grouped row result.""" - from transformer_engine.pytorch.module.grouped_linear import ( - _split_quantize_hybrid, - ) - torch.manual_seed(3598) # NVFP4 grouped split-quantize requires each M split to be a multiple # of 64. @@ -4623,12 +4907,14 @@ def make_quantizer(): ) quantizers = [make_quantizer(), make_quantizer()] - actual = _split_quantize_hybrid( + actual, dbiases = grouped_quantization._split_quantize( tensor, [64, 64], quantizers, + tensor.dtype, ) + assert dbiases is None for index, (actual_part, quantizer) in enumerate(zip(actual, quantizers)): expected_columnwise = quantizer.columnwise_quantizer.quantize( actual_part.rowwise_sub_storage.dequantize() @@ -7132,12 +7418,12 @@ def _build_and_run(use_checkpoint): def _run_grouped_linear(self, recipe_obj, *, checkpoint_fn=None): """Build a GroupedLinear, run forward+backward with optional activation checkpointing around the module. Exercises the - ``_split_quantize_hybrid`` code path under recompute. + hybrid ``_split_quantize`` code path under recompute. GroupedLinear is the MoE token-dispatch kernel: a single batch is split along dim-0 into ``num_gemms`` chunks and each chunk goes through its own weight matrix. Under hybrid quantization, - ``_split_quantize_hybrid`` (``module/grouped_linear.py``) runs + ``_split_quantize`` (``module/_grouped_quantization.py``) runs ``tex.split_quantize`` twice (once per sub-quantizer direction) and zips the results into a list of ``HybridQuantizedTensor`` chunks — save-for-backward then receives a *list* of hybrid @@ -7168,7 +7454,7 @@ def _run_grouped_linear(self, recipe_obj, *, checkpoint_fn=None): @_XFAIL_HOPPER_COLUMNWISE_PER_TENSOR_FP8 def test_te_checkpoint_reentrant_grouped_linear_fp8_bitwise(self): """GroupedLinear + te.checkpoint(reentrant) under same-format FP8 - hybrid. Exercises the MoE ``_split_quantize_hybrid`` + list-of- + hybrid. Exercises the MoE ``_split_quantize`` + list-of- hybrid-tensors save-for-backward path under recompute.""" import transformer_engine.pytorch as te_pytorch diff --git a/tests/pytorch/test_identity_quantizer.py b/tests/pytorch/test_identity_quantizer.py index cb0785e3e0..9c1508ee2a 100644 --- a/tests/pytorch/test_identity_quantizer.py +++ b/tests/pytorch/test_identity_quantizer.py @@ -25,6 +25,7 @@ MXFP8Quantizer, NVFP4Quantizer, ) +from transformer_engine.pytorch.module import _grouped_quantization from transformer_engine.pytorch.tensor.identity_tensor import IdentityTensor from transformer_engine.pytorch.tensor.storage.identity_tensor_storage import ( IdentityTensorStorage, @@ -257,28 +258,38 @@ def test_make_empty_internal_returns_storage(self): assert out.dequantize().dtype == torch.bfloat16 def test_grouped_split_all_identity_uses_plain_tensor_views(self): - from transformer_engine.pytorch.module.grouped_linear import ( - _split_quantize_non_hybrid, - ) - x = torch.randn(8, 16, device="cuda", dtype=torch.bfloat16) m_splits = [3, 5] quantizers = [IdentityQuantizer(), IdentityQuantizer()] - out = _split_quantize_non_hybrid(x, m_splits, quantizers, activation_dtype=torch.bfloat16) + out, dbiases = _grouped_quantization._split_quantize( + x, + m_splits, + quantizers, + activation_dtype=torch.bfloat16, + compute_dbias=True, + ) assert all(isinstance(t, torch.Tensor) for t in out) assert not any(isinstance(t, IdentityTensorStorage) for t in out) - for actual, expected in zip(out, torch.split(x, m_splits)): + expected_splits = torch.split(x, m_splits) + for actual, expected in zip(out, expected_splits): torch.testing.assert_close(actual, expected, rtol=0.0, atol=0.0) + assert dbiases is not None + for actual, expected in zip(dbiases, expected_splits): + torch.testing.assert_close(actual, expected.sum(dim=0), rtol=0.0, atol=0.0) cast_quantizers = [ IdentityQuantizer(dtype=torch.float32), IdentityQuantizer(dtype=torch.float32), ] - cast_out = _split_quantize_non_hybrid( - x, m_splits, cast_quantizers, activation_dtype=torch.bfloat16 + cast_out, cast_dbiases = _grouped_quantization._split_quantize( + x, + m_splits, + cast_quantizers, + activation_dtype=torch.bfloat16, ) + assert cast_dbiases is None assert all(isinstance(t, IdentityTensorStorage) for t in cast_out) for actual, expected in zip(cast_out, torch.split(x, m_splits)): dequantized = actual.dequantize() @@ -291,10 +302,6 @@ def test_grouped_split_all_identity_uses_plain_tensor_views(self): ) def test_grouped_split_rejects_mixed_identity_and_quantized_operands(self): - from transformer_engine.pytorch.module.grouped_linear import ( - _validate_grouped_quantizer_list, - ) - cases = [ [IdentityQuantizer(), _mxfp8(tex.DType.kFloat8E4M3)], [ @@ -311,12 +318,11 @@ def test_grouped_split_rejects_mixed_identity_and_quantized_operands(self): for quantizers in cases: with pytest.raises(ValueError, match="mix Identity-backed and quantized"): - _validate_grouped_quantizer_list(quantizers, operand_name="input") + _grouped_quantization.validate_grouped_quantizer_list( + quantizers, operand_name="input" + ) def test_hybrid_split_forwards_disable_bulk_allocation_to_both_directions(self, monkeypatch): - import transformer_engine.pytorch.module.grouped_linear as grouped_linear - from transformer_engine.pytorch.module.grouped_linear import _split_quantize_hybrid - calls = [] def fake_split_quantize(tensor, m_splits, quantizers, *, disable_bulk_allocation=False): @@ -326,9 +332,9 @@ def fake_split_quantize(tensor, m_splits, quantizers, *, disable_bulk_allocation for tensor_part, quantizer in zip(torch.split(tensor, m_splits), quantizers) ] - monkeypatch.setattr(grouped_linear.tex, "split_quantize", fake_split_quantize) + monkeypatch.setattr(_grouped_quantization.tex, "split_quantize", fake_split_quantize) monkeypatch.setattr( - grouped_linear, + _grouped_quantization, "_supports_native_split_quantize", lambda quantizer: True, ) @@ -342,15 +348,17 @@ def fake_split_quantize(tensor, m_splits, quantizers, *, disable_bulk_allocation for _ in m_splits ] - out = _split_quantize_hybrid( + out, dbiases = _grouped_quantization._split_quantize( x, m_splits, quantizers, + activation_dtype=torch.bfloat16, disable_bulk_allocation=True, ) assert calls == [True, True] assert len(out) == len(m_splits) + assert dbiases is None @pytest.mark.skipif(not fp8_available, reason=reason_for_no_fp8) def test_grouped_linear_cpu_offload_disables_bulk_allocation_for_hybrid_input( @@ -371,15 +379,25 @@ def qfactory(role): calls = [] - def fake_split_quantize_hybrid( - tensor, m_splits, quantizers, *, disable_bulk_allocation=False, **kwargs + def fake_split_quantize( + tensor, + m_splits, + quantizers, + activation_dtype, + *, + disable_bulk_allocation=False, + **kwargs, ): - del tensor, m_splits, quantizers, kwargs + del tensor, m_splits, quantizers, activation_dtype, kwargs calls.append(disable_bulk_allocation) raise StopAfterFlagCapture("captured hybrid split kwargs") monkeypatch.setattr(grouped_linear, "is_cpu_offload_enabled", lambda: True) - monkeypatch.setattr(grouped_linear, "_split_quantize_hybrid", fake_split_quantize_hybrid) + monkeypatch.setattr( + _grouped_quantization, + "_split_quantize", + fake_split_quantize, + ) model = te.GroupedLinear(2, 64, 64, params_dtype=torch.bfloat16).cuda() x = torch.randn(64, 64, device="cuda", dtype=torch.bfloat16) diff --git a/transformer_engine/pytorch/module/_grouped_quantization.py b/transformer_engine/pytorch/module/_grouped_quantization.py new file mode 100644 index 0000000000..9bda846493 --- /dev/null +++ b/transformer_engine/pytorch/module/_grouped_quantization.py @@ -0,0 +1,388 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Grouped split-quantization helpers used by :mod:`GroupedLinear`.""" + +from typing import List, Optional, Sequence, Tuple, Union, cast + +import torch + +import transformer_engine_torch as tex + +from ..quantized_tensor import QuantizedTensorStorage, Quantizer +from ..tensor import ( + Float8BlockQuantizer, + Float8CurrentScalingQuantizer, + Float8Quantizer, + HybridQuantizer, + IdentityQuantizer, + MXFP8Quantizer, + NVFP4Quantizer, +) +from ..tensor.storage.hybrid_tensor_storage import HybridQuantizedTensorStorage +from ..utils import cast_if_needed +from ...debug.pytorch.debug_quantization import DebugQuantizer + +_NATIVE_SPLIT_QUANTIZER_TYPES = frozenset( + { + Float8Quantizer, + Float8CurrentScalingQuantizer, + Float8BlockQuantizer, + MXFP8Quantizer, + NVFP4Quantizer, + } +) + +_NATIVE_BGRAD_QUANTIZER_TYPES = frozenset( + { + Float8Quantizer, + Float8CurrentScalingQuantizer, + MXFP8Quantizer, + } +) + +_DYNAMIC_QUANTIZER_FIELDS = frozenset( + { + "rowwise_usage", + "columnwise_usage", + "internal", + "optimize_for_gemm", + } +) + + +def _supports_native_split_quantize(quantizer: Quantizer) -> bool: + """Whether ``tex.split_quantize`` has an exact converter for this quantizer.""" + return type(quantizer) in _NATIVE_SPLIT_QUANTIZER_TYPES + + +def _prefers_native_bgrad_quantize(quantizer: Quantizer) -> bool: + """Whether per-split ``bgrad_quantize`` is preferred over bulk quantization.""" + return type(quantizer) in _NATIVE_BGRAD_QUANTIZER_TYPES + + +def _uses_identity_quantizer(quantizer: Optional[Quantizer]) -> bool: + """Whether a quantizer, including a hybrid sub-quantizer, is Identity-backed.""" + if quantizer is None: + return False + if isinstance(quantizer, IdentityQuantizer): + return True + if isinstance(quantizer, HybridQuantizer): + return _uses_identity_quantizer(quantizer.rowwise_quantizer) or _uses_identity_quantizer( + quantizer.columnwise_quantizer + ) + return False + + +def _identity_quantizer_signature(quantizer: Optional[Quantizer]) -> Tuple[bool, bool]: + """Identity usage per GEMM direction: ``(rowwise, columnwise)``.""" + if isinstance(quantizer, HybridQuantizer): + return ( + _uses_identity_quantizer(quantizer.rowwise_quantizer), + _uses_identity_quantizer(quantizer.columnwise_quantizer), + ) + identity = isinstance(quantizer, IdentityQuantizer) + return (identity, identity) + + +def _backend_quantizer_signature(quantizer: Optional[Quantizer]): + """Return backend configuration that grouped kernels require to be uniform.""" + if quantizer is None: + return None + + # Identity is not registered as a torch.compile value quantizer, but its + # dtype changes the grouped GEMM input type and therefore must be uniform. + if isinstance(quantizer, IdentityQuantizer): + return (type(quantizer), (("dtype", quantizer.dtype),)) + + fields = quantizer._value_fields() + if fields is None: + # Delayed-scaling Float8Quantizer carries per-expert scale/amax tensors, + # which are intentionally different, but its emitted FP8 dtype is a + # group-wide backend choice. Other unregistered/custom quantizers retain + # the conservative exact-family behavior until they expose value fields. + fields = ("dtype",) if isinstance(quantizer, Float8Quantizer) else () + + config = [] + for name in fields: + if name in _DYNAMIC_QUANTIZER_FIELDS: + continue + value = getattr(quantizer, name) + if name == "dtype": + value = int(value) + config.append((name, value)) + return (type(quantizer), tuple(config)) + + +def _validate_backend_match( + reference: Quantizer, + quantizer: Quantizer, + operand_name: str, + direction: str, + expert_index: int, +) -> None: + """Validate one expert against the group's reference backend.""" + if type(quantizer) is not type(reference): + raise ValueError( + f"GroupedLinear {operand_name} quantizers use incompatible {direction} backend" + f" families across experts: expert 0 uses {type(reference).__name__}, but expert" + f" {expert_index} uses {type(quantizer).__name__}. Grouped operands require one" + " quantizer family per direction." + ) + reference_signature = _backend_quantizer_signature(reference) + quantizer_signature = _backend_quantizer_signature(quantizer) + if quantizer_signature != reference_signature: + raise ValueError( + f"GroupedLinear {operand_name} quantizers use incompatible {direction} backend" + f" configurations across experts: expert 0 uses {reference_signature}, but expert" + f" {expert_index} uses {quantizer_signature}. Grouped operands require the same" + " backend-relevant configuration per direction." + ) + + +def validate_grouped_quantizer_list( + quantizers: Sequence[Optional[Quantizer]], + *, + operand_name: str = "operand", +) -> None: + """Validate that one grouped operand has compatible expert quantizers.""" + if not quantizers: + return + + reference = quantizers[0] + reference_is_hybrid = isinstance(reference, HybridQuantizer) + reference_identity = _identity_quantizer_signature(reference) + + for expert_index, quantizer in enumerate(quantizers[1:], start=1): + if (quantizer is None) != (reference is None): + raise ValueError( + f"GroupedLinear {operand_name} quantizers mix None and concrete quantizers" + f" across experts: expert 0 is {type(reference).__name__}, but expert" + f" {expert_index} is {type(quantizer).__name__}." + ) + if reference is None: + continue + + quantizer_is_hybrid = isinstance(quantizer, HybridQuantizer) + if quantizer_is_hybrid != reference_is_hybrid: + raise ValueError( + f"GroupedLinear {operand_name} quantizers mix HybridQuantizer and non-hybrid" + f" quantizers across experts: expert 0 is {type(reference).__name__}, but expert" + f" {expert_index} is {type(quantizer).__name__}." + ) + + identity = _identity_quantizer_signature(quantizer) + if identity != reference_identity: + raise ValueError( + f"GroupedLinear {operand_name} quantizers mix Identity-backed and quantized" + f" directions across experts: expert 0 uses {reference_identity}, but expert" + f" {expert_index} uses {identity}." + ) + + if reference_is_hybrid: + _validate_backend_match( + reference.rowwise_quantizer, + quantizer.rowwise_quantizer, + operand_name, + "rowwise", + expert_index, + ) + _validate_backend_match( + reference.columnwise_quantizer, + quantizer.columnwise_quantizer, + operand_name, + "columnwise", + expert_index, + ) + if quantizer.columnwise_source != reference.columnwise_source: + raise ValueError( + f"GroupedLinear {operand_name} HybridQuantizer list has mixed columnwise" + " source policies across experts: expert 0 uses" + f" {reference.columnwise_source!r}, but expert {expert_index} uses" + f" {quantizer.columnwise_source!r}." + ) + else: + _validate_backend_match( + reference, + quantizer, + operand_name, + "plain", + expert_index, + ) + + +def _split_quantize_non_hybrid( + tensor: torch.Tensor, + split_sizes: Sequence[int], + quantizers: Sequence[Quantizer], + dtype: torch.dtype, + *, + disable_bulk_allocation: bool = False, + allow_identity_views: bool = True, +) -> Sequence[Union[torch.Tensor, QuantizedTensorStorage]]: + """Split and quantize one homogeneous, non-Hybrid quantizer list.""" + reference = quantizers[0] + if _supports_native_split_quantize(reference): + return tex.split_quantize( + tensor, + split_sizes, + quantizers, + disable_bulk_allocation=disable_bulk_allocation, + ) + + tensor = cast_if_needed(tensor, dtype) + if ( + allow_identity_views + # Only the base IdentityQuantizer can bypass quantization; subclasses + # may override its behavior and must go through their normal call path. + and type(reference) is IdentityQuantizer # pylint: disable=unidiomatic-typecheck + and (reference.dtype is None or reference.dtype == dtype) + ): + return torch.split(tensor, split_sizes) + + return [ + quantizer(tensor_part) + for tensor_part, quantizer in zip(torch.split(tensor, split_sizes), quantizers) + ] + + +def _split_quantize_hybrid( + tensor: torch.Tensor, + split_sizes: Sequence[int], + quantizers: Sequence[HybridQuantizer], + *, + disable_bulk_allocation: bool = False, +) -> Sequence[HybridQuantizedTensorStorage]: + """Split and quantize an all-hybrid, generation-validated operand.""" + reference = quantizers[0] + rowwise_enabled = reference.rowwise_usage + columnwise_enabled = reference.columnwise_usage + columnwise_source = reference.columnwise_source + rowwise_quantizers = [quantizer.rowwise_quantizer for quantizer in quantizers] + columnwise_quantizers = [quantizer.columnwise_quantizer for quantizer in quantizers] + + needs_rowwise_result = rowwise_enabled or ( + columnwise_enabled and columnwise_source == "rowwise_dequantized" + ) + row_results = ( + _split_quantize_non_hybrid( + tensor, + split_sizes, + rowwise_quantizers, + tensor.dtype, + disable_bulk_allocation=disable_bulk_allocation, + allow_identity_views=False, + ) + if needs_rowwise_result + else [None] * len(quantizers) + ) + + columnwise_src = tensor + if columnwise_enabled and columnwise_source == "rowwise_dequantized": + # Assemble the exact grouped row results in split order. NVFP4 padding + # and scale layout can differ from independently quantizing each split. + columnwise_src = torch.cat( + [result.dequantize(dtype=tensor.dtype) for result in row_results], + dim=0, + ) + col_results = ( + _split_quantize_non_hybrid( + columnwise_src, + split_sizes, + columnwise_quantizers, + tensor.dtype, + disable_bulk_allocation=disable_bulk_allocation, + allow_identity_views=False, + ) + if columnwise_enabled + else [None] * len(quantizers) + ) + + return [ + HybridQuantizedTensorStorage( + rowwise_storage=row if rowwise_enabled else None, + columnwise_storage=col, + quantizer=quantizer, + fake_dtype=tensor.dtype, + ) + for row, col, quantizer in zip(row_results, col_results, quantizers) + ] + + +def _split_quantize( + tensor: torch.Tensor, + split_sizes: Sequence[int], + quantizers: Optional[Sequence[Optional[Quantizer]]], + activation_dtype: torch.dtype, + *, + compute_dbias: bool = False, + disable_bulk_allocation: bool = False, +) -> Tuple[ + Sequence[Union[torch.Tensor, QuantizedTensorStorage]], + Optional[List[torch.Tensor]], +]: + """Split a grouped operand, quantizing when quantizers are provided. + + Native, hybrid, Identity, debug, and Python fallback dispatch are internal + implementation choices. ``dbiases`` is ``None`` when ``compute_dbias`` is + false and otherwise contains one reduction result per split. Quantizer lists + must be homogeneous; dispatch intentionally uses expert 0 as the reference. + """ + if quantizers is not None and len(quantizers) != len(split_sizes): + raise ValueError( + "Grouped split quantizer count does not match the number of tensor splits " + f"({len(quantizers)} != {len(split_sizes)})" + ) + + reference = quantizers[0] if quantizers else None + if reference is None: + outputs = torch.split(cast_if_needed(tensor, activation_dtype), split_sizes) + dbiases = ( + [tensor_part.sum(dim=0) for tensor_part in torch.split(tensor, split_sizes)] + if compute_dbias + else None + ) + return outputs, dbiases + + concrete_quantizers = cast(Sequence[Quantizer], quantizers) + if compute_dbias and _prefers_native_bgrad_quantize(reference): + outputs = [] + dbiases = [] + for tensor_part, quantizer in zip(torch.split(tensor, split_sizes), concrete_quantizers): + dbias, output = tex.bgrad_quantize(tensor_part, quantizer) + dbiases.append(dbias) + outputs.append(output) + return outputs, dbiases + + dbiases = ( + [tensor_part.sum(dim=0) for tensor_part in torch.split(tensor, split_sizes)] + if compute_dbias + else None + ) + if isinstance(reference, DebugQuantizer): + outputs = DebugQuantizer.multi_tensor_quantize( + tensor, + concrete_quantizers, + split_sizes, + activation_dtype, + ) + elif isinstance(reference, HybridQuantizer): + outputs = _split_quantize_hybrid( + tensor, + split_sizes, + cast(Sequence[HybridQuantizer], concrete_quantizers), + disable_bulk_allocation=disable_bulk_allocation, + ) + else: + outputs = _split_quantize_non_hybrid( + tensor, + split_sizes, + concrete_quantizers, + activation_dtype, + disable_bulk_allocation=disable_bulk_allocation, + ) + return outputs, dbiases + + +__all__ = ["_split_quantize", "validate_grouped_quantizer_list"] diff --git a/transformer_engine/pytorch/module/grouped_linear.py b/transformer_engine/pytorch/module/grouped_linear.py index 9860d48237..e04eb96c06 100644 --- a/transformer_engine/pytorch/module/grouped_linear.py +++ b/transformer_engine/pytorch/module/grouped_linear.py @@ -4,7 +4,7 @@ """GroupedLinear API""" -from typing import Union, Optional, Callable, Tuple, List, Sequence +from typing import Union, Optional, Callable, Tuple, List from itertools import chain import os import warnings @@ -32,6 +32,7 @@ _get_high_precision_init_val, ) from ._common import can_reconstruct_wgrad_input_from_original, WeightGradStore +from . import _grouped_quantization from ..quantization import FP8GlobalStateManager, QuantizerRole from ..utils import ( divide, @@ -82,370 +83,6 @@ from ...debug.pytorch.debug_quantization import DebugQuantizer from ...debug.pytorch.debug_state import TEDebugState - -_NATIVE_SPLIT_QUANTIZER_TYPES = frozenset( - { - Float8Quantizer, - Float8CurrentScalingQuantizer, - Float8BlockQuantizer, - MXFP8Quantizer, - NVFP4Quantizer, - } -) - - -def _supports_native_split_quantize(quantizer): - """Whether ``tex.split_quantize`` has an exact converter for this quantizer.""" - return type(quantizer) in _NATIVE_SPLIT_QUANTIZER_TYPES - - -def _uses_identity_quantizer(quantizer): - """Whether a quantizer, including a hybrid sub-quantizer, is Identity-backed.""" - if quantizer is None: - return False - if isinstance(quantizer, IdentityQuantizer): - return True - if isinstance(quantizer, HybridQuantizer): - return _uses_identity_quantizer(quantizer.rowwise_quantizer) or _uses_identity_quantizer( - quantizer.columnwise_quantizer - ) - return False - - -def _identity_quantizer_signature(quantizer): - """Identity usage per GEMM direction: (rowwise, columnwise).""" - if isinstance(quantizer, HybridQuantizer): - return ( - _uses_identity_quantizer(quantizer.rowwise_quantizer), - _uses_identity_quantizer(quantizer.columnwise_quantizer), - ) - identity = isinstance(quantizer, IdentityQuantizer) - return (identity, identity) - - -_DYNAMIC_QUANTIZER_SIGNATURE_FIELDS = frozenset( - { - "rowwise_usage", - "columnwise_usage", - "internal", - "optimize_for_gemm", - } -) - - -def _backend_quantizer_signature(quantizer): - """Return backend configuration that grouped kernels require to be uniform.""" - if quantizer is None: - return None - - # Identity is not registered as a torch.compile value quantizer, but its - # dtype changes the grouped GEMM input type and therefore must be uniform. - if isinstance(quantizer, IdentityQuantizer): - return (type(quantizer), (("dtype", quantizer.dtype),)) - - fields = quantizer._value_fields() - if fields is None: - # Delayed-scaling Float8Quantizer carries per-expert scale/amax tensors, - # which are intentionally different, but its emitted FP8 dtype is a - # group-wide backend choice. Other unregistered/custom quantizers retain - # the conservative exact-family behavior until they expose value fields. - fields = ("dtype",) if isinstance(quantizer, Float8Quantizer) else () - - config = [] - for name in fields: - if name in _DYNAMIC_QUANTIZER_SIGNATURE_FIELDS: - continue - value = getattr(quantizer, name) - if name == "dtype": - value = int(value) - config.append((name, value)) - return (type(quantizer), tuple(config)) - - -def _validate_backend_match(reference, quantizer, operand_name, direction, expert_index): - """Validate one expert against the group's reference backend.""" - if type(quantizer) is not type(reference): - raise ValueError( - f"GroupedLinear {operand_name} quantizers use incompatible {direction} backend" - f" families across experts: expert 0 uses {type(reference).__name__}, but expert" - f" {expert_index} uses {type(quantizer).__name__}. Grouped operands require one" - " quantizer family per direction." - ) - reference_signature = _backend_quantizer_signature(reference) - quantizer_signature = _backend_quantizer_signature(quantizer) - if quantizer_signature != reference_signature: - raise ValueError( - f"GroupedLinear {operand_name} quantizers use incompatible {direction} backend" - f" configurations across experts: expert 0 uses {reference_signature}, but expert" - f" {expert_index} uses {quantizer_signature}. Grouped operands require the same" - " backend-relevant configuration per direction." - ) - - -def _validate_grouped_quantizer_list(quantizers, *, operand_name="operand") -> None: - """Validate one grouped operand once when its quantizer generation changes.""" - if not quantizers: - return - - reference = quantizers[0] - reference_is_hybrid = isinstance(reference, HybridQuantizer) - reference_identity = _identity_quantizer_signature(reference) - - for expert_index, quantizer in enumerate(quantizers[1:], start=1): - if (quantizer is None) != (reference is None): - raise ValueError( - f"GroupedLinear {operand_name} quantizers mix None and concrete quantizers" - f" across experts: expert 0 is {type(reference).__name__}, but expert" - f" {expert_index} is {type(quantizer).__name__}." - ) - if reference is None: - continue - - quantizer_is_hybrid = isinstance(quantizer, HybridQuantizer) - if quantizer_is_hybrid != reference_is_hybrid: - raise ValueError( - f"GroupedLinear {operand_name} quantizers mix HybridQuantizer and non-hybrid" - f" quantizers across experts: expert 0 is {type(reference).__name__}, but expert" - f" {expert_index} is {type(quantizer).__name__}." - ) - - identity = _identity_quantizer_signature(quantizer) - if identity != reference_identity: - raise ValueError( - f"GroupedLinear {operand_name} quantizers mix Identity-backed and quantized" - f" directions across experts: expert 0 uses {reference_identity}, but expert" - f" {expert_index} uses {identity}." - ) - - if reference_is_hybrid: - _validate_backend_match( - reference.rowwise_quantizer, - quantizer.rowwise_quantizer, - operand_name, - "rowwise", - expert_index, - ) - _validate_backend_match( - reference.columnwise_quantizer, - quantizer.columnwise_quantizer, - operand_name, - "columnwise", - expert_index, - ) - if quantizer.columnwise_source != reference.columnwise_source: - raise ValueError( - f"GroupedLinear {operand_name} HybridQuantizer list has mixed columnwise" - " source policies across experts: expert 0 uses" - f" {reference.columnwise_source!r}, but expert {expert_index} uses" - f" {quantizer.columnwise_source!r}." - ) - else: - _validate_backend_match( - reference, - quantizer, - operand_name, - "plain", - expert_index, - ) - - -def _split_quantize_non_hybrid( - tensor, - m_splits, - quantizers, - activation_dtype, - *, - disable_bulk_allocation=False, - allow_identity_views=True, -): - """Split and quantize one homogeneous, non-Hybrid quantizer list.""" - reference = quantizers[0] - if _supports_native_split_quantize(reference): - return tex.split_quantize( - tensor, - m_splits, - quantizers, - disable_bulk_allocation=disable_bulk_allocation, - ) - - tensor = cast_if_needed(tensor, activation_dtype) - if ( - allow_identity_views - # Only the base IdentityQuantizer can bypass quantization; subclasses - # may override its behavior and must go through their normal call path. - and type(reference) is IdentityQuantizer # pylint: disable=unidiomatic-typecheck - and (reference.dtype is None or reference.dtype == activation_dtype) - ): - return torch.split(tensor, m_splits) - - return [ - quantizer(tensor_part) if quantizer is not None else tensor_part - for tensor_part, quantizer in zip(torch.split(tensor, m_splits), quantizers) - ] - - -def _split_quantize_hybrid( - tensor, - m_splits, - quantizers, - *, - disable_bulk_allocation=False, -): - """Grouped split+quantize for an all-hybrid, generation-validated operand.""" - from ..tensor.storage.hybrid_tensor_storage import HybridQuantizedTensorStorage as HybridStorage - - reference = quantizers[0] - rowwise_enabled = reference.rowwise_usage - columnwise_enabled = reference.columnwise_usage - columnwise_source = reference.columnwise_source - rowwise_quantizers = [quantizer.rowwise_quantizer for quantizer in quantizers] - columnwise_quantizers = [quantizer.columnwise_quantizer for quantizer in quantizers] - - needs_rowwise_result = rowwise_enabled or ( - columnwise_enabled and columnwise_source == "rowwise_dequantized" - ) - row_results = ( - _split_quantize_non_hybrid( - tensor, - m_splits, - rowwise_quantizers, - tensor.dtype, - disable_bulk_allocation=disable_bulk_allocation, - allow_identity_views=False, - ) - if needs_rowwise_result - else [None] * len(quantizers) - ) - - columnwise_src = tensor - if columnwise_enabled and columnwise_source == "rowwise_dequantized": - # Assemble the exact grouped row results in split order. NVFP4 padding - # and scale layout can differ from independently quantizing each split. - columnwise_src = torch.cat( - [result.dequantize(dtype=tensor.dtype) for result in row_results], - dim=0, - ) - col_results = ( - _split_quantize_non_hybrid( - columnwise_src, - m_splits, - columnwise_quantizers, - tensor.dtype, - disable_bulk_allocation=disable_bulk_allocation, - allow_identity_views=False, - ) - if columnwise_enabled - else [None] * len(quantizers) - ) - - return [ - HybridStorage( - rowwise_storage=row if rowwise_enabled else None, - columnwise_storage=col, - quantizer=q, - fake_dtype=tensor.dtype, - ) - for row, col, q in zip( - row_results, - col_results, - quantizers, - ) - ] - - -def _split_quantize( - tensor: torch.Tensor, - split_sizes: List[int], - with_quantized_output: bool, - quantizers: Optional[List[Quantizer]], - dtype: torch.dtype, - with_debug_quantizers: bool, - disable_bulk_allocation: bool, -) -> Sequence[Union[torch.Tensor, QuantizedTensorStorage]]: - """Split a tensor and quantize each part if needed.""" - if not with_quantized_output: - return torch.split(cast_if_needed(tensor, dtype), split_sizes) - - if quantizers is None or quantizers[0] is None: - raise ValueError("Quantizers are required for quantized split output") - - if with_debug_quantizers: - return DebugQuantizer.multi_tensor_quantize(tensor, quantizers, split_sizes, dtype) - - reference = quantizers[0] - if isinstance(reference, HybridQuantizer): - return _split_quantize_hybrid( - tensor, - split_sizes, - quantizers, - disable_bulk_allocation=disable_bulk_allocation, - ) - - return _split_quantize_non_hybrid( - tensor, - split_sizes, - quantizers, - dtype, - disable_bulk_allocation=disable_bulk_allocation, - ) - - -def _split_quantize_and_bias( - tensor: torch.Tensor, - split_sizes: List[int], - *, - fp8: bool, - debug: bool, - quantizers: Optional[List[Quantizer]], - dtype: torch.dtype, - use_bias: bool, - recipe: Recipe, - disable_bulk_allocation: bool, -) -> Tuple[ - Sequence[Union[torch.Tensor, QuantizedTensorStorage]], - List[Optional[torch.Tensor]], -]: - """Split grad output, quantize if needed, and compute unfused bias gradients.""" - num_splits = len(split_sizes) - grad_biases = [None] * num_splits - reference = quantizers[0] - identity = _uses_identity_quantizer(reference) - hybrid = isinstance(reference, HybridQuantizer) and not identity - - use_native_bgrad_quantize = ( - fp8 - and not debug - and not hybrid - and use_bias - and not identity - and (recipe.delayed() or recipe.float8_current_scaling() or recipe.mxfp8()) - ) - if use_native_bgrad_quantize: - outputs = [None] * num_splits - for i, tensor_part in enumerate(torch.split(tensor, split_sizes)): - grad_biases[i], outputs[i] = tex.bgrad_quantize(tensor_part, quantizers[i]) - return outputs, grad_biases - - with_quantized_output = fp8 or debug - if with_quantized_output and (use_bias or debug): - for i, tensor_part in enumerate(torch.split(tensor, split_sizes)): - grad_biases[i] = tensor_part.sum(dim=0) - - # Preserve the existing CPU-offload policy: only Hybrid split-quantize - # disables bulk allocation in backward. - disable_bulk_allocation = disable_bulk_allocation if hybrid else False - outputs = _split_quantize( - tensor, - split_sizes, - with_quantized_output=with_quantized_output, - quantizers=quantizers, - dtype=dtype, - with_debug_quantizers=debug, - disable_bulk_allocation=disable_bulk_allocation, - ) - return outputs, grad_biases - - __all__ = ["GroupedLinear"] @@ -900,14 +537,10 @@ def forward( cache_weight, skip_fp8_weight_update, save_original_input, - delayed_scaling_input_quantizer, - unsafe_requantization_input_quantizer, debug, ) = non_tensor_args - if fp8: - backward_override = FP8GlobalStateManager.get_fp8_recipe().backward_override - else: - backward_override = None + recipe = FP8GlobalStateManager.get_fp8_recipe() if fp8 else None + backward_override = recipe.backward_override if recipe is not None else None if backward_override == "high_precision": save_original_input = True elif backward_override == "dequantized": @@ -926,30 +559,30 @@ def forward( backward_needs_input = is_grad_enabled and weight_requires_grad if backward_override is None and save_original_input and backward_needs_input: - if delayed_scaling_input_quantizer is not None: - if FP8GlobalStateManager.get_fp8_recipe().custom(): + if recipe is not None and recipe.delayed(): + raise ValueError("DelayedScaling recipe is not supported with save_original_input") + + # Megatron-Core may enable this automatically to reuse an activation + # already retained by an upstream operation. Built-in recipes guarantee + # group homogeneity, while CustomRecipe generations are validated once, + # so runtime safety can be determined from expert 0. + if recipe is not None: + input_quantizer = input_quantizers[0] + if isinstance(input_quantizer, Float8Quantizer): warnings.warn( "save_original_input is incompatible with delayed-scaling quantizers " "(Float8Quantizer). Disabling save_original_input for this module.", stacklevel=2, ) save_original_input = False - else: - raise ValueError( - "DelayedScaling recipe is not supported with save_original_input" + elif not can_reconstruct_wgrad_input_from_original(input_quantizer): + warnings.warn( + "Ignoring save_original_input=True because the input quantizer cannot " + "safely reconstruct the backward operand from the original input " + f"({input_quantizer}).", + stacklevel=2, ) - - # Megatron-Core may enable this automatically to reuse an activation - # already retained by an upstream operation. The resolved quantizer - # generation is classified once in ``_validate_quantizer_generation``. - if save_original_input and unsafe_requantization_input_quantizer is not None: - warnings.warn( - "Ignoring save_original_input=True because the input quantizer cannot " - "safely reconstruct the backward operand from the original input " - f"({unsafe_requantization_input_quantizer}).", - stacklevel=2, - ) - save_original_input = False + save_original_input = False # Configure quantizers if input_quantizers[0] is not None: @@ -1036,13 +669,11 @@ def forward( # Disable bulk allocation when CPU offloading is active: offloading skips small # tensors (like scales), but bulk allocation shares storage across all tensors, # so if scales can't be offloaded, nothing in the group can be offloaded. - inputmats = _split_quantize( + inputmats, _ = _grouped_quantization._split_quantize( inp_view, m_splits, - with_quantized_output=fp8 or debug, - quantizers=input_quantizers, - dtype=activation_dtype, - with_debug_quantizers=debug, + input_quantizers, + activation_dtype, disable_bulk_allocation=cpu_offloading, ) @@ -1472,17 +1103,18 @@ def backward( rowwise=ctx.requires_dgrad, columnwise=ctx.weights_requires_grad, ) - grad_output, grad_biases = _split_quantize_and_bias( + grad_output, grad_biases = _grouped_quantization._split_quantize( grad_output_view, ctx.m_splits, - fp8=ctx.fp8, - debug=ctx.debug, - quantizers=ctx.grad_output_quantizers, - dtype=ctx.activation_dtype, - use_bias=ctx.use_bias, - recipe=ctx.fp8_recipe, - disable_bulk_allocation=ctx.cpu_offloading, + ctx.grad_output_quantizers, + ctx.activation_dtype, + compute_dbias=(ctx.fp8 or ctx.debug) and ctx.use_bias, + disable_bulk_allocation=( + ctx.cpu_offloading and isinstance(grad_output_reference, HybridQuantizer) + ), ) + if grad_biases is None: + grad_biases = [None] * N if is_dist_weight: accumulate_wgrad_into_param_main_grad = False @@ -1583,13 +1215,11 @@ def backward( input_quantizer.set_usage(rowwise=True, columnwise=True) else: input_quantizer.set_usage(rowwise=False, columnwise=True) - inputmats = _split_quantize( + inputmats, _ = _grouped_quantization._split_quantize( inp_view, ctx.m_splits, - with_quantized_output=ctx.fp8 or ctx.debug, - quantizers=ctx.input_quantizers, - dtype=ctx.activation_dtype, - with_debug_quantizers=ctx.debug, + ctx.input_quantizers, + ctx.activation_dtype, disable_bulk_allocation=ctx.cpu_offloading, ) elif ctx.backward_override == "dequantized": @@ -1824,9 +1454,8 @@ def __init__( "fwd": 3, "bwd": 2, } - self._validated_quantizer_generations = {} - self._delayed_scaling_input_quantizer = None - self._unsafe_requantization_input_quantizer = None + self._custom_quantizer_cache = {} + self._uses_custom_recipe = False if tp_group is None: self.tp_size = tp_size @@ -1915,19 +1544,22 @@ def set_meta_tensor(self, fwd: bool, recipe: Recipe) -> None: if recipe.float8_current_scaling(): self._customize_quantizers_float8_current_scaling(fwd, recipe) - self._validate_quantizer_generation(fwd) + self._uses_custom_recipe = recipe.custom() + self._validate_custom_recipe_quantizers(fwd, recipe) - def _validate_quantizer_generation(self, fwd: bool) -> None: - """Validate grouped-kernel invariants once per quantizer generation.""" - # Recipe state replaces this list object only when it constructs a new - # quantizer generation. The O(1) identity guard keeps validation off the - # steady-state forward path. Record a generation only after all of its - # operand roles pass, so a failed recipe transition is retried. + def _validate_custom_recipe_quantizers(self, fwd: bool, recipe: Recipe) -> None: + """Validate one CustomRecipe quantizer generation.""" + if not recipe.custom(): + return + + # A CustomRecipe factory may return a different quantizer for every expert, + # while grouped execution selects its implementation from expert 0. Validate + # every newly constructed list once and keep the steady-state path O(1). meta_key = "scaling_fwd" if fwd else "scaling_bwd" generation = self.quantizers.get(meta_key) if generation is None: return - if self._validated_quantizer_generations.get(meta_key) is generation: + if self._custom_quantizer_cache.get(meta_key) is generation: return if fwd: @@ -1938,33 +1570,24 @@ def _validate_quantizer_generation(self, fwd: bool) -> None: weight_quantizers = tuple( generation[self._offsets["weight"] + i * stride] for i in range(self.num_gemms) ) - _validate_grouped_quantizer_list(input_quantizers, operand_name="input") - _validate_grouped_quantizer_list(weight_quantizers, operand_name="weight") - delayed_scaling_input_quantizer = next( - (q for q in input_quantizers if isinstance(q, Float8Quantizer)), - None, + _grouped_quantization.validate_grouped_quantizer_list( + input_quantizers, operand_name="input" ) - unsafe_requantization_input_quantizer = next( - ( - q - for q in input_quantizers - if q is not None and not can_reconstruct_wgrad_input_from_original(q) - ), - None, + _grouped_quantization.validate_grouped_quantizer_list( + weight_quantizers, operand_name="weight" ) - self._delayed_scaling_input_quantizer = delayed_scaling_input_quantizer - self._unsafe_requantization_input_quantizer = unsafe_requantization_input_quantizer else: stride = self._num_fp8_tensors_per_gemm["bwd"] grad_output_quantizers = tuple( generation[self._offsets["grad_output"] + i * stride] for i in range(self.num_gemms) ) - _validate_grouped_quantizer_list( + _grouped_quantization.validate_grouped_quantizer_list( grad_output_quantizers, operand_name="grad_output", ) - self._validated_quantizer_generations[meta_key] = generation + # Cache only a fully validated generation so failures are retried. + self._custom_quantizer_cache[meta_key] = generation def get_quantizer_roles( self, @@ -2407,8 +2030,6 @@ def forward( cache_weight, skip_fp8_weight_update, self.save_original_input, - self._delayed_scaling_input_quantizer, - self._unsafe_requantization_input_quantizer, debug, ) out, new_workspaces = linear_fn( @@ -2556,13 +2177,13 @@ def _get_weight_quantizers(self) -> List[Quantizer]: return weight_quantizers def _get_quantizers(self): - if self.fp8: - # Normally validated while installing recipe metadata. Keep this - # O(1) generation guard so failed transitions cannot reuse stale - # validation state if base metadata takes an early return on retry. - self._validate_quantizer_generation(True) + if self.fp8 and self._uses_custom_recipe: + # Validation normally runs while installing metadata. Retry here because + # a failed generation remains installed when the caller catches the error. + recipe = FP8GlobalStateManager.get_fp8_recipe() + self._validate_custom_recipe_quantizers(True, recipe) if torch.is_grad_enabled(): - self._validate_quantizer_generation(False) + self._validate_custom_recipe_quantizers(False, recipe) weight_quantizers = self._get_weight_quantizers() input_quantizers, output_quantizers = (