From e53027faa5636d7bb0e09b7348b7ce0294a9304d Mon Sep 17 00:00:00 2001 From: Varun Thumbe Date: Mon, 20 Jul 2026 22:35:49 +0000 Subject: [PATCH 01/17] initial implementation for act grouped linear fusion Signed-off-by: Varun Thumbe --- tests/pytorch/test_grouped_mlp.py | 53 ++- transformer_engine/pytorch/csrc/extensions.h | 39 ++ .../pytorch/csrc/extensions/activation.cpp | 147 +++++++ .../pytorch/csrc/extensions/pybind.cpp | 37 ++ .../pytorch/ops/basic/grouped_linear.py | 129 ++++-- .../pytorch/ops/fused/__init__.py | 7 + .../ops/fused/grouped_linear_activation.py | 381 ++++++++++++++++++ 7 files changed, 758 insertions(+), 35 deletions(-) create mode 100644 transformer_engine/pytorch/ops/fused/grouped_linear_activation.py diff --git a/tests/pytorch/test_grouped_mlp.py b/tests/pytorch/test_grouped_mlp.py index 76550db5f8..dbf89245fc 100644 --- a/tests/pytorch/test_grouped_mlp.py +++ b/tests/pytorch/test_grouped_mlp.py @@ -23,6 +23,10 @@ OUTPUT_BUFFER_KEY, GRAD_INPUT_BUFFER_KEY, ) +from transformer_engine.pytorch.ops.fused.grouped_linear_activation import ( + BackwardGroupedLinearScaledActivation, + ForwardScaledActivationGroupedLinear, +) from transformer_engine.pytorch import ( QuantizedTensor, Float8CurrentScalingQuantizer, @@ -72,6 +76,15 @@ if nvfp4_available: _grouped_mlp_quantization_list.append("nvfp4_rht") +# Quantization recipes for ScaledActivation + GroupedLinear fusion +_grouped_mlp_act_grouped_linear_quantization_list: list[str] = [] +if fp8_available: + _grouped_mlp_act_grouped_linear_quantization_list.append("fp8_current_scaling") +if mxfp8_available: + _grouped_mlp_act_grouped_linear_quantization_list.append("mxfp8") +if nvfp4_available: + _grouped_mlp_act_grouped_linear_quantization_list.append("nvfp4_rht") + @pytest.fixture(autouse=True, scope="function") def _reset_rng_states_per_test(): @@ -999,14 +1012,15 @@ def _make_module(): and glu_interleave_size is None ) ) + forward_ops = module._module_groups[0]._forward_ops + backward_ops = module._module_groups[0]._backward_ops + full_grouped_mlp_fusion = False if expected_grouped_mlp_fusion: if activation_is_glu: fused_cls = te.ops.fused.GroupedMLP_CuTeGEMMGLU else: fused_cls = te.ops.fused.GroupedMLP_CuTeGEMMUnary if fused_cls.is_supported(): - forward_ops = module._module_groups[0]._forward_ops - backward_ops = module._module_groups[0]._backward_ops assert len(forward_ops) == 1 assert len(backward_ops) == 1 assert isinstance( @@ -1014,6 +1028,22 @@ def _make_module(): fused_cls, ) assert backward_ops[0][0] is forward_ops[0][0] + full_grouped_mlp_fusion = True + + # When the full FC1 + activation + FC2 fusion is unavailable, verify + # that ScaledActivation + GroupedLinear fusions cover both boundaries + # whenever grouped quantized compute is supported. + act_grouped_linear_fusion_expected = ( + not full_grouped_mlp_fusion + and te.ops.fused.act_grouped_linear_fusion_supported(fc2, module[1], recipe) + and te.ops.fused.act_grouped_linear_fusion_supported(fc1, module[1], recipe) + ) + assert any( + isinstance(op, ForwardScaledActivationGroupedLinear) for op, _ in forward_ops + ) == act_grouped_linear_fusion_expected + assert any( + isinstance(op, BackwardGroupedLinearScaledActivation) for op, _ in backward_ops + ) == act_grouped_linear_fusion_expected # Loose tols for sanity checking tols = {"rtol": 0.125, "atol": 0.25} @@ -1076,6 +1106,25 @@ def _make_module(): assert_close(fc1.weight.grad, fc1_w_ref_grad, **tols) assert_close(fc2.weight.grad, fc2_w_ref_grad, **tols) + @pytest.mark.parametrize("quantization", _grouped_mlp_act_grouped_linear_quantization_list) + def test_grouped_mlp_act_grouped_linear_fusion( + self, + *, + quantization: str, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + """Exercise ScaledActivation + GroupedLinear fusion without the full CuTe DSL MLP fusion.""" + monkeypatch.setenv("NVTE_CUTEDSL_FUSED_GROUPED_MLP", "0") + self.test_grouped_mlp( + group_size=2, + bias=False, + hidden_size=128, + quantization=quantization, + single_grouped_weight=False, + split_alignment=128, + activation="scaled_swiglu", + ) + @pytest.mark.parametrize("bias", (False, True)) @pytest.mark.parametrize("quantization", _grouped_mlp_quantization_list) @pytest.mark.parametrize( diff --git a/transformer_engine/pytorch/csrc/extensions.h b/transformer_engine/pytorch/csrc/extensions.h index 6edfbdc00e..30866f75f2 100644 --- a/transformer_engine/pytorch/csrc/extensions.h +++ b/transformer_engine/pytorch/csrc/extensions.h @@ -290,6 +290,45 @@ py::object clamped_swiglu(const at::Tensor &input, py::handle quantizer, float l py::object clamped_dswiglu(const at::Tensor &grad, const at::Tensor &input, py::handle quantizer, float limit, float alpha, float glu_linear_offset); + +/* Scaled activation + grouped quantize */ +py::object grouped_scaled_swiglu(const at::Tensor &input, const at::Tensor &act_scales, + py::handle quantizer, const size_t num_tensors, + std::optional first_dims, + std::optional tensor_offsets, + int64_t glu_interleave_size); + +py::object grouped_scaled_clamped_swiglu(const at::Tensor &input, const at::Tensor &act_scales, + py::handle quantizer, const size_t num_tensors, + std::optional first_dims, + std::optional tensor_offsets, float limit, + float alpha, float glu_linear_offset, + int64_t glu_interleave_size); + +py::object grouped_scaled_srelu(const at::Tensor &input, const at::Tensor &act_scales, + py::handle quantizer, const size_t num_tensors, + std::optional first_dims, + std::optional tensor_offsets); + +py::tuple grouped_scaled_dswiglu(const at::Tensor &grad, const at::Tensor &input, + const at::Tensor &act_scales, py::handle quantizer, + const size_t num_tensors, std::optional first_dims, + std::optional tensor_offsets, + int64_t glu_interleave_size, bool compute_scale_grad); + +py::tuple grouped_scaled_clamped_dswiglu(const at::Tensor &grad, const at::Tensor &input, + const at::Tensor &act_scales, py::handle quantizer, + const size_t num_tensors, + std::optional first_dims, + std::optional tensor_offsets, float limit, + float alpha, float glu_linear_offset, + int64_t glu_interleave_size, bool compute_scale_grad); + +py::tuple grouped_scaled_dsrelu(const at::Tensor &grad, const at::Tensor &input, + const at::Tensor &act_scales, py::handle quantizer, + const size_t num_tensors, std::optional first_dims, + std::optional tensor_offsets, + bool compute_scale_grad); /*************************************************************************************************** * LayerNorm **************************************************************************************************/ diff --git a/transformer_engine/pytorch/csrc/extensions/activation.cpp b/transformer_engine/pytorch/csrc/extensions/activation.cpp index 58a8f84f85..b8d7363e1b 100644 --- a/transformer_engine/pytorch/csrc/extensions/activation.cpp +++ b/transformer_engine/pytorch/csrc/extensions/activation.cpp @@ -342,5 +342,152 @@ py::object clamped_dswiglu(const at::Tensor& grad, const at::Tensor& input, py:: glu_linear_offset); } +/* Scaled activation + grouped quantize helpers (mirrors activation_helper / dactivation_helper). */ + +template +py::object grouped_scaled_activation_helper(const at::Tensor& input, const at::Tensor& act_scales, + py::handle quantizer, const size_t num_tensors, + std::optional first_dims, + std::optional tensor_offsets, + int shape_divisor, Args&&... args) { + init_extension(); + NVTE_CHECK(input.dim() == 2, "grouped scaled activation input must be 2D"); + NVTE_CHECK(act_scales.numel() == input.size(0), + "grouped scaled activation expects one scale per input row"); + NVTE_CHECK(shape_divisor > 0 && input.size(1) % shape_divisor == 0, + "grouped scaled activation input width is not compatible with activation"); + + auto input_tensor = input.contiguous(); + auto scales_tensor = act_scales.contiguous(); + const TensorWrapper& input_nvte = makeTransformerEngineTensor(input_tensor); + const TensorWrapper& scales_nvte = makeTransformerEngineTensor(scales_tensor); + + auto output = + at::empty({input.size(0), input.size(1) / shape_divisor}, input_tensor.options()); + const TensorWrapper& output_nvte = makeTransformerEngineTensor(output); + + auto stream = at::cuda::getCurrentCUDAStream(); + NVTE_SCOPED_GIL_RELEASE({ + act_func(input_nvte.data(), scales_nvte.data(), output_nvte.data(), + std::forward(args)..., stream); + }); + + if (quantizer.is_none()) { + return py::cast(output); + } + return group_quantize(output, quantizer, num_tensors, first_dims, std::nullopt, tensor_offsets, + std::nullopt); +} + +template +py::tuple grouped_scaled_dactivation_helper(const at::Tensor& grad, const at::Tensor& input, + const at::Tensor& act_scales, py::handle quantizer, + const size_t num_tensors, + std::optional first_dims, + std::optional tensor_offsets, + bool compute_scale_grad, Args&&... args) { + init_extension(); + NVTE_CHECK(input.dim() == 2 && grad.dim() == 2, + "grouped scaled dactivation input and grad must be 2D"); + NVTE_CHECK(act_scales.numel() == input.size(0), + "grouped scaled dactivation expects one scale per input row"); + + auto grad_tensor = grad.contiguous(); + auto input_tensor = input.contiguous(); + auto scales_tensor = act_scales.contiguous(); + auto grad_input = at::empty_like(input_tensor); + auto grad_scales = compute_scale_grad ? at::empty_like(scales_tensor) : at::Tensor(); + + const TensorWrapper& grad_nvte = makeTransformerEngineTensor(grad_tensor); + const TensorWrapper& input_nvte = makeTransformerEngineTensor(input_tensor); + const TensorWrapper& scales_nvte = makeTransformerEngineTensor(scales_tensor); + const TensorWrapper& grad_input_nvte = makeTransformerEngineTensor(grad_input); + std::optional grad_scales_nvte; + if (compute_scale_grad) { + grad_scales_nvte.emplace(makeTransformerEngineTensor(grad_scales)); + } + + auto stream = at::cuda::getCurrentCUDAStream(); + NVTE_SCOPED_GIL_RELEASE({ + dact_func(grad_nvte.data(), input_nvte.data(), scales_nvte.data(), grad_input_nvte.data(), + compute_scale_grad ? grad_scales_nvte->data() : nullptr, std::forward(args)..., + stream); + }); + + // Return both the (optionally) grouped-quantized grad input for the next + // grouped GEMM and the dense high-precision grad input so callers can reuse + // it (e.g. bias gradient) without a lossy dequantize. + py::object grad_input_out = py::cast(grad_input); + if (!quantizer.is_none()) { + grad_input_out = + group_quantize(grad_input, quantizer, num_tensors, first_dims, std::nullopt, tensor_offsets, + std::nullopt); + } + return py::make_tuple(grad_input_out, py::cast(grad_input), + compute_scale_grad ? py::cast(grad_scales) : py::none()); +} + +py::object grouped_scaled_swiglu(const at::Tensor& input, const at::Tensor& act_scales, + py::handle quantizer, const size_t num_tensors, + std::optional first_dims, + std::optional tensor_offsets, + int64_t glu_interleave_size) { + return grouped_scaled_activation_helper( + input, act_scales, quantizer, num_tensors, first_dims, tensor_offsets, /*shape_divisor=*/2, + glu_interleave_size); +} + +py::object grouped_scaled_clamped_swiglu(const at::Tensor& input, const at::Tensor& act_scales, + py::handle quantizer, const size_t num_tensors, + std::optional first_dims, + std::optional tensor_offsets, float limit, + float alpha, float glu_linear_offset, + int64_t glu_interleave_size) { + return grouped_scaled_activation_helper( + input, act_scales, quantizer, num_tensors, first_dims, tensor_offsets, /*shape_divisor=*/2, + limit, alpha, glu_linear_offset, glu_interleave_size); +} + +py::object grouped_scaled_srelu(const at::Tensor& input, const at::Tensor& act_scales, + py::handle quantizer, const size_t num_tensors, + std::optional first_dims, + std::optional tensor_offsets) { + return grouped_scaled_activation_helper( + input, act_scales, quantizer, num_tensors, first_dims, tensor_offsets, + /*shape_divisor=*/1); +} + +py::tuple grouped_scaled_dswiglu(const at::Tensor& grad, const at::Tensor& input, + const at::Tensor& act_scales, py::handle quantizer, + const size_t num_tensors, std::optional first_dims, + std::optional tensor_offsets, + int64_t glu_interleave_size, bool compute_scale_grad) { + return grouped_scaled_dactivation_helper( + grad, input, act_scales, quantizer, num_tensors, first_dims, tensor_offsets, + compute_scale_grad, glu_interleave_size); +} + +py::tuple grouped_scaled_clamped_dswiglu(const at::Tensor& grad, const at::Tensor& input, + const at::Tensor& act_scales, py::handle quantizer, + const size_t num_tensors, + std::optional first_dims, + std::optional tensor_offsets, float limit, + float alpha, float glu_linear_offset, + int64_t glu_interleave_size, bool compute_scale_grad) { + return grouped_scaled_dactivation_helper( + grad, input, act_scales, quantizer, num_tensors, first_dims, tensor_offsets, + compute_scale_grad, limit, alpha, glu_linear_offset, glu_interleave_size); +} + +py::tuple grouped_scaled_dsrelu(const at::Tensor& grad, const at::Tensor& input, + const at::Tensor& act_scales, py::handle quantizer, + const size_t num_tensors, std::optional first_dims, + std::optional tensor_offsets, + bool compute_scale_grad) { + return grouped_scaled_dactivation_helper( + grad, input, act_scales, quantizer, num_tensors, first_dims, tensor_offsets, + compute_scale_grad); +} + } // namespace pytorch } // namespace transformer_engine diff --git a/transformer_engine/pytorch/csrc/extensions/pybind.cpp b/transformer_engine/pytorch/csrc/extensions/pybind.cpp index 7e9d114be8..003c737a75 100644 --- a/transformer_engine/pytorch/csrc/extensions/pybind.cpp +++ b/transformer_engine/pytorch/csrc/extensions/pybind.cpp @@ -284,6 +284,43 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { "Backward of SwiGLU used in GPT OSS", py::arg("grad"), py::arg("fwd_input"), py::arg("quantizer"), py::arg("limit") = 7.0f, py::arg("alpha") = 1.702f, py::arg("glu_linear_offset") = 1.0f); + /* Scaled activation + grouped quantize */ + m.def("grouped_scaled_swiglu", transformer_engine::pytorch::grouped_scaled_swiglu, + "Scaled SwiGLU + grouped quantize", py::arg("input"), py::arg("act_scales"), + py::arg("quantizer"), py::arg("num_tensors"), py::arg("first_dims"), + py::arg("tensor_offsets") = py::none(), + py::arg("glu_interleave_size") = 0); + m.def("grouped_scaled_clamped_swiglu", + transformer_engine::pytorch::grouped_scaled_clamped_swiglu, + "Scaled clamped SwiGLU + grouped quantize", py::arg("input"), py::arg("act_scales"), + py::arg("quantizer"), py::arg("num_tensors"), py::arg("first_dims"), + py::arg("tensor_offsets") = py::none(), + py::arg("limit") = 7.0f, py::arg("alpha") = 1.702f, + py::arg("glu_linear_offset") = 1.0f, py::arg("glu_interleave_size") = 0); + m.def("grouped_scaled_srelu", transformer_engine::pytorch::grouped_scaled_srelu, + "Scaled SReLU + grouped quantize", py::arg("input"), py::arg("act_scales"), + py::arg("quantizer"), py::arg("num_tensors"), py::arg("first_dims"), + py::arg("tensor_offsets") = py::none()); + m.def("grouped_scaled_dswiglu", transformer_engine::pytorch::grouped_scaled_dswiglu, + "Scaled SwiGLU backward + optional grouped quantize", py::arg("grad"), + py::arg("fwd_input"), py::arg("act_scales"), py::arg("quantizer"), + py::arg("num_tensors"), py::arg("first_dims"), + py::arg("tensor_offsets") = py::none(), py::arg("glu_interleave_size") = 0, + py::arg("compute_scale_grad") = true); + m.def("grouped_scaled_clamped_dswiglu", + transformer_engine::pytorch::grouped_scaled_clamped_dswiglu, + "Scaled clamped SwiGLU backward + optional grouped quantize", py::arg("grad"), + py::arg("fwd_input"), py::arg("act_scales"), py::arg("quantizer"), + py::arg("num_tensors"), py::arg("first_dims"), + py::arg("tensor_offsets") = py::none(), py::arg("limit") = 7.0f, + py::arg("alpha") = 1.702f, py::arg("glu_linear_offset") = 1.0f, + py::arg("glu_interleave_size") = 0, py::arg("compute_scale_grad") = true); + m.def("grouped_scaled_dsrelu", transformer_engine::pytorch::grouped_scaled_dsrelu, + "Scaled SReLU backward + optional grouped quantize", py::arg("grad"), + py::arg("fwd_input"), py::arg("act_scales"), py::arg("quantizer"), + py::arg("num_tensors"), py::arg("first_dims"), + py::arg("tensor_offsets") = py::none(), + py::arg("compute_scale_grad") = true); /* DBias + DAct fusions*/ m.def("dbias_dgelu", transformer_engine::pytorch::dbias_dgelu, "DGeLU + DBias + Quantize", py::arg("grad"), py::arg("fwd_input"), py::arg("quantizer")); diff --git a/transformer_engine/pytorch/ops/basic/grouped_linear.py b/transformer_engine/pytorch/ops/basic/grouped_linear.py index 5ef0fa4339..b096962747 100644 --- a/transformer_engine/pytorch/ops/basic/grouped_linear.py +++ b/transformer_engine/pytorch/ops/basic/grouped_linear.py @@ -1244,7 +1244,7 @@ def _fuser_forward_split_quantize( def _fuser_forward_grouped_tensor( self, *, - input_: torch.Tensor, + input_: Optional[torch.Tensor] = None, split_sizes: torch.Tensor, scales: Optional[torch.Tensor], with_quantized_compute: bool, @@ -1255,40 +1255,84 @@ def _fuser_forward_grouped_tensor( weight_requires_grad: bool, device: torch.device, out_buffer: Optional[torch.Tensor] = None, + grouped_input: Optional[GroupedTensorStorage] = None, + grouped_tensor_offsets: Optional[ + tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor] + ] = None, ) -> tuple[torch.Tensor, tuple[Optional[torch.Tensor], ...]]: """Graph-safe GroupedTensor forward path (pure compute). Returns ``(output, tensors_to_save)``. ``split_sizes``, ``base_split_offsets`` and ``split_points`` are returned so that ``fuser_forward_save_ctx`` can call ``save_for_backward`` on them. + + Provide either a dense ``input_`` or a pre-built ``grouped_input``. + When both are given, ``grouped_input`` is used after validation. """ num_groups = self.num_groups has_bias = self.has_bias - base_split_offsets = tex.splits_to_offsets(split_sizes, 1) - split_points = base_split_offsets[1:].to(dtype=torch.int) - - # Flatten to 2D so the first dim is the total token count. - original_shape = list(input_.size()) - x = maybe_dequantize(input_, dtype).reshape(-1, self.in_features) - total_tokens = x.size(0) - - # Build the input GroupedTensor. - if with_quantized_compute: - input_quantizer = input_quantizers[0] - input_quantizer.set_usage(rowwise=True, columnwise=weight_requires_grad) - input_quantizer.optimize_for_gemm = True - grouped_x = tex.group_quantize(x, input_quantizer, num_groups, split_sizes) - else: - # No quantize: wrap the contiguous high-precision buffer. - grouped_x = GroupedTensorStorage( - shape=(total_tokens, self.in_features), - dtype=dtype, - num_tensors=num_groups, - quantizer=None, - data=x.reshape(-1), - first_dims=split_sizes, - tensor_offsets=base_split_offsets * self.in_features, + if grouped_tensor_offsets is None: + split_sizes, grouped_tensor_offsets = tex.splits_to_offsets_multi( + split_sizes, + device, + strides=[1, 1, self.in_features, self.out_features], + include_leading_zero=[False, True, True, True], + dtypes=[torch.int32, torch.int64, torch.int64, torch.int64], + bulk_allocate=True, ) + ( + split_points, + base_split_offsets, + input_tensor_offsets, + output_tensor_offsets, + ) = grouped_tensor_offsets + + # Build the input GroupedTensor unless a preceding fused operation + # already produced it. + if grouped_input is not None: + total_tokens, in_features = grouped_input.logical_shape + expected_quantizer = input_quantizers[0] if with_quantized_compute else None + if ( + grouped_input.quantizer is not expected_quantizer + or in_features != self.in_features + ): + raise ValueError( + "GroupedLinear received an incompatible grouped input " + f"(quantizer={grouped_input.quantizer}, " + f"logical_shape={grouped_input.logical_shape}; " + f"expected quantizer={expected_quantizer}, " + f"in_features={self.in_features})" + ) + grouped_x = grouped_input + out_shape = [total_tokens, self.out_features] + else: + # Flatten to 2D so the first dim is the total token count. + original_shape = list(input_.size()) + total_tokens = math.prod(original_shape[:-1]) + out_shape = original_shape[:-1] + [self.out_features] + x = maybe_dequantize(input_, dtype).reshape(-1, self.in_features) + if with_quantized_compute: + input_quantizer = input_quantizers[0] + input_quantizer.set_usage(rowwise=True, columnwise=weight_requires_grad) + input_quantizer.optimize_for_gemm = True + grouped_x = tex.group_quantize( + x, + input_quantizer, + num_groups, + split_sizes, + tensor_offsets=input_tensor_offsets, + ) + else: + # No quantize: wrap the contiguous high-precision buffer. + grouped_x = GroupedTensorStorage( + shape=(total_tokens, self.in_features), + dtype=dtype, + num_tensors=num_groups, + quantizer=None, + data=x.reshape(-1), + first_dims=split_sizes, + tensor_offsets=input_tensor_offsets, + ) if is_cpu_offload_enabled() and grouped_x is not None: start_offload(grouped_x) @@ -1314,7 +1358,6 @@ def _fuser_forward_grouped_tensor( ) # Allocate output buffer and wrap as a GroupedTensor view. - out_shape = original_shape[:-1] + [self.out_features] out = validate_or_alloc_output(out_buffer, out_shape, dtype, device) grouped_out = GroupedTensorStorage( shape=(total_tokens, self.out_features), @@ -1323,7 +1366,7 @@ def _fuser_forward_grouped_tensor( quantizer=None, data=out.reshape(-1), first_dims=split_sizes, - tensor_offsets=base_split_offsets * self.out_features, + tensor_offsets=output_tensor_offsets, ) # Bias: hand off to the grouped GEMM (graph-safe, fused). Plain bias @@ -1585,6 +1628,7 @@ def _fuser_backward_grouped_tensor( *, ctx: OperationContext, grad_output: torch.Tensor, + grouped_grad_output: Optional[GroupedTensorStorage] = None, ) -> tuple[ torch.Tensor, Iterable[Iterable[Optional[torch.Tensor]]], @@ -1617,15 +1661,35 @@ def _fuser_backward_grouped_tensor( else: ws, saved_tensors = saved_tensors[:num_groups], saved_tensors[num_groups:] - # Flatten grad_output to 2D (total_tokens, out_features) - # to figure out total tokens. + + # Keep the dense high-precision grad for bias-gradient computation while + # optionally using a pre-quantized grouped view for the grouped GEMMs. dy_2d = grad_output.reshape(-1, self.out_features) total_tokens = dy_2d.size(0) + grad_input_shape = list(grad_output.size())[:-1] + [self.in_features] - # Build the grad_output GroupedTensor. - # Optionally get dbias is fusion available with bgrad_group_quantize + # Build the grad_output GroupedTensor unless a preceding fused + # operation already produced it. + # Optionally get dbias if fusion is available with bgrad_group_quantize. dbias_packed = None - if with_quantized_compute: + if grouped_grad_output is not None: + expected_quantizer = ( + ctx.grad_output_quantizers[0] if with_quantized_compute else None + ) + if ( + grouped_grad_output.quantizer is not expected_quantizer + or tuple(grouped_grad_output.logical_shape) + != (total_tokens, self.out_features) + ): + raise ValueError( + "GroupedLinear received an incompatible grouped grad_output " + f"(quantizer={grouped_grad_output.quantizer}, " + f"logical_shape={grouped_grad_output.logical_shape}; " + f"expected quantizer={expected_quantizer}, " + f"logical_shape={(total_tokens, self.out_features)})" + ) + grouped_dy = grouped_grad_output + elif with_quantized_compute: grad_output_quantizer = ctx.grad_output_quantizers[0] grad_output_quantizer.set_usage( rowwise=ctx.input_requires_grad, columnwise=ctx.weight_requires_grad @@ -1681,7 +1745,6 @@ def _fuser_backward_grouped_tensor( # ---- dgrad GEMM ---------------------------------------------------- grad_input = None if ctx.input_requires_grad: - grad_input_shape = list(grad_output.size())[:-1] + [self.in_features] grad_input = validate_or_alloc_output( getattr(ctx, "dgrad_out", None), grad_input_shape, dtype, device ) diff --git a/transformer_engine/pytorch/ops/fused/__init__.py b/transformer_engine/pytorch/ops/fused/__init__.py index dc9dcd6dc3..4c2832f65d 100644 --- a/transformer_engine/pytorch/ops/fused/__init__.py +++ b/transformer_engine/pytorch/ops/fused/__init__.py @@ -12,6 +12,11 @@ from .forward_linear_bias_activation import ForwardLinearBiasActivation from .forward_linear_bias_add import ForwardLinearBiasAdd from .forward_linear_scale_add import ForwardLinearScaleAdd +from .grouped_linear_activation import ( + BackwardGroupedLinearScaledActivation, + ForwardScaledActivationGroupedLinear, + act_grouped_linear_fusion_supported, +) from .userbuffers_backward_linear import UserbuffersBackwardLinear from .userbuffers_forward_linear import UserbuffersForwardLinear @@ -21,6 +26,7 @@ register_forward_fusion(ForwardLinearBiasAdd.fuse_forward_ops) register_forward_fusion(ForwardLinearBiasActivation.fuse_forward_ops) register_forward_fusion(ForwardLinearScaleAdd.fuse_forward_ops) +register_forward_fusion(ForwardScaledActivationGroupedLinear.fuse_forward_ops) # Register backward fusions register_backward_fusion(UserbuffersBackwardLinear.fuse_backward_ops) @@ -28,6 +34,7 @@ register_backward_fusion(BackwardLinearScale.fuse_backward_ops) register_backward_fusion(BackwardActivationBias.fuse_backward_ops) register_backward_fusion(BackwardAddRMSNorm.fuse_backward_ops) +register_backward_fusion(BackwardGroupedLinearScaledActivation.fuse_backward_ops) # Import experimental fusions # Note: Registration logic is non-trivial, so submodule handles it internally. diff --git a/transformer_engine/pytorch/ops/fused/grouped_linear_activation.py b/transformer_engine/pytorch/ops/fused/grouped_linear_activation.py new file mode 100644 index 0000000000..c2b387f3fa --- /dev/null +++ b/transformer_engine/pytorch/ops/fused/grouped_linear_activation.py @@ -0,0 +1,381 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Pair fusions between scaled activations and grouped linear operations.""" + +from __future__ import annotations + +from collections.abc import Iterable +from typing import Any, Optional + +import torch + +import transformer_engine_torch as tex +from ...cpu_offload import is_cpu_offload_enabled, mark_activation_offload +from ...quantization import Recipe +from ...tensor import Quantizer +from ...utils import clear_tensor_data +from .._common import maybe_dequantize +from ..basic import GroupedLinear, ScaledClampedQGeGLU, ScaledSReLU, ScaledSwiGLU +from ..op import FusedOperation, FusibleOperation, OperationContext + + +_ScaledActivation = ScaledSwiGLU | ScaledClampedQGeGLU | ScaledSReLU +_SCALED_ACTIVATION_TYPES = (ScaledSwiGLU, ScaledClampedQGeGLU, ScaledSReLU) + + +def _grouped_scaled_activation( + activation: _ScaledActivation, + input_: torch.Tensor, + scales: torch.Tensor, + quantizer: Quantizer, + num_groups: int, + split_sizes: torch.Tensor, + tensor_offsets: torch.Tensor, +) -> torch.Tensor: + """Dispatch to the matching tex.grouped_scaled_* forward API.""" + x = input_.reshape(-1, input_.size(-1)) + s = scales.reshape(-1) + if isinstance(activation, ScaledSwiGLU): + return tex.grouped_scaled_swiglu( + x, + s, + quantizer, + num_groups, + split_sizes, + tensor_offsets, + int(activation.glu_interleave_size or 0), + ) + if isinstance(activation, ScaledClampedQGeGLU): + clamped = activation._clamped + return tex.grouped_scaled_clamped_swiglu( + x, + s, + quantizer, + num_groups, + split_sizes, + tensor_offsets, + clamped.limit, + clamped.alpha, + clamped.glu_linear_offset, + int(activation.glu_interleave_size or 0), + ) + return tex.grouped_scaled_srelu( + x, s, quantizer, num_groups, split_sizes, tensor_offsets + ) + + +def _grouped_scaled_dactivation( + activation: _ScaledActivation, + grad_output: torch.Tensor, + input_: torch.Tensor, + scales: torch.Tensor, + *, + quantizer: Quantizer, + num_groups: int, + split_sizes: torch.Tensor, + tensor_offsets: torch.Tensor, + compute_scale_grad: bool, +) -> tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]: + """Dispatch to the matching tex.grouped_scaled_d* API. + + Returns ``(grouped_dx, dense_dx, dscales)`` where ``grouped_dx`` is the + grouped-quantized grad input for the next grouped GEMM and ``dense_dx`` is + the dense high-precision grad input (reused for the bias gradient so we do + not have to dequantize ``grouped_dx``). + """ + dy = grad_output.reshape(-1, grad_output.size(-1)) + x = input_.reshape(-1, input_.size(-1)) + s = scales.reshape(-1) + if isinstance(activation, ScaledSwiGLU): + dx, dense_dx, dscales = tex.grouped_scaled_dswiglu( + dy, + x, + s, + quantizer, + num_groups, + split_sizes, + tensor_offsets, + int(activation.glu_interleave_size or 0), + compute_scale_grad, + ) + elif isinstance(activation, ScaledClampedQGeGLU): + clamped = activation._clamped + dx, dense_dx, dscales = tex.grouped_scaled_clamped_dswiglu( + dy, + x, + s, + quantizer, + num_groups, + split_sizes, + tensor_offsets, + clamped.limit, + clamped.alpha, + clamped.glu_linear_offset, + int(activation.glu_interleave_size or 0), + compute_scale_grad, + ) + else: + dx, dense_dx, dscales = tex.grouped_scaled_dsrelu( + dy, + x, + s, + quantizer, + num_groups, + split_sizes, + tensor_offsets, + compute_scale_grad, + ) + return dx, dense_dx, dscales + + +def act_grouped_linear_fusion_supported( + linear: GroupedLinear, + activation: _ScaledActivation, + recipe: Optional[Recipe], +) -> bool: + """Whether ScaledActivation + GroupedLinear can use grouped quantized compute.""" + if recipe is None or activation.activation_recompute_in_mlp: + return False + input_quantizers = [ + linear.get_quantizer("forward", 2 * group_idx) for group_idx in range(linear.num_groups) + ] + weight = linear.weight if linear.single_grouped_weight else linear.weight0 + dtype = torch.get_autocast_dtype("cuda") if torch.is_autocast_enabled() else weight.dtype + return linear._is_graph_safe_path_supported( + with_quantized_compute=True, + input_quantizers=input_quantizers, + dtype=dtype, + single_grouped_weight=linear.single_grouped_weight, + ) + + +class ForwardScaledActivationGroupedLinear(FusedOperation): + """Scaled activation + grouped quantize + grouped linear forward.""" + + def __init__(self, *, activation: _ScaledActivation, linear: GroupedLinear) -> None: + super().__init__((activation, linear)) + + def fuser_forward( + self, + basic_op_ctxs: list[OperationContext], + input_: torch.Tensor, + *, + basic_op_extra_inputs: list[tuple[torch.Tensor, ...]], + prev_op_grad_output_quantizer: Optional[Quantizer], + next_op_input_quantizer: Optional[Quantizer], + basic_op_kwargs: list[dict[str, Any]], + ) -> tuple[torch.Tensor, Iterable[Iterable[torch.Tensor]]]: + activation = self.basic_ops[0] + linear = self.basic_ops[1] + activation_ctx, linear_ctx = basic_op_ctxs + if basic_op_kwargs[0] or basic_op_kwargs[1]: + raise ValueError("Scaled activation and GroupedLinear do not expect keyword arguments") + + weight = linear.weight if linear.single_grouped_weight else linear.weight0 + device = weight.device + dtype = torch.get_autocast_dtype("cuda") if torch.is_autocast_enabled() else weight.dtype + input_ = maybe_dequantize(input_, dtype) + scales = maybe_dequantize(basic_op_extra_inputs[0][0], dtype) + + split_sizes = basic_op_extra_inputs[1][0] + if int(split_sizes.numel()) != linear.num_groups: + raise ValueError( + f"Expected {linear.num_groups} splits, but got {int(split_sizes.numel())}." + ) + split_sizes = split_sizes.to(device=device, dtype=torch.int64) + linear_scales = basic_op_extra_inputs[1][1] if linear._scale_bias else None + split_sizes, grouped_tensor_offsets = tex.splits_to_offsets_multi( + split_sizes, + device, + strides=[1, 1, linear.in_features, linear.out_features], + include_leading_zero=[False, True, True, True], + dtypes=[torch.int32, torch.int64, torch.int64, torch.int64], + bulk_allocate=True, + ) + + input_quantizers = [ + linear.get_quantizer("forward", 2 * group_idx) + for group_idx in range(linear.num_groups) + ] + weight_quantizers = [ + linear.get_quantizer("forward", 2 * group_idx + 1) + for group_idx in range(linear.num_groups) + ] + input_quantizer = input_quantizers[0] + weight_requires_grad = linear_ctx.requires_grad and weight.requires_grad + input_quantizer.set_usage(rowwise=True, columnwise=weight_requires_grad) + input_quantizer.optimize_for_gemm = True + + grouped_x = _grouped_scaled_activation( + activation, + input_, + scales, + input_quantizer, + linear.num_groups, + split_sizes, + grouped_tensor_offsets[2], + ) + + if activation_ctx.requires_grad: + if is_cpu_offload_enabled(): + mark_activation_offload(input_) + activation_ctx.input_requires_grad = True + activation_ctx.extra_input_requires_grad = basic_op_extra_inputs[0][0].requires_grad + activation_ctx.dtype = dtype + activation_ctx.save_for_backward(input_, scales) + + out, tensors_to_save = linear._fuser_forward_grouped_tensor( + split_sizes=split_sizes, + scales=linear_scales, + with_quantized_compute=True, + input_quantizers=input_quantizers, + weight_quantizers=weight_quantizers, + dtype=dtype, + input_requires_grad=linear_ctx.requires_grad, + weight_requires_grad=weight_requires_grad, + device=device, + grouped_input=grouped_x, + grouped_tensor_offsets=grouped_tensor_offsets, + ) + linear.fuser_forward_save_ctx( + basic_op_ctxs=[linear_ctx], + input_=input_, + tensors_to_save=[tensors_to_save], + requires_grad=[linear_ctx.requires_grad], + basic_op_extra_inputs=[basic_op_extra_inputs[1]], + prev_op_grad_output_quantizer=prev_op_grad_output_quantizer, + next_op_input_quantizer=next_op_input_quantizer, + basic_op_kwargs=[basic_op_kwargs[1]], + use_grouped_tensor_path=True, + ) + return out, [(), ()] + + @staticmethod + def fuse_forward_ops( + ops: list[FusibleOperation], + *, + recipe: Optional[Recipe] = None, + **unused, + ) -> list[FusibleOperation]: + """Fuse each supported ScaledActivation + GroupedLinear pair.""" + out: list[FusibleOperation] = [] + idx = 0 + while idx < len(ops): + if ( + idx + 1 < len(ops) + and isinstance(ops[idx], _SCALED_ACTIVATION_TYPES) + and isinstance(ops[idx + 1], GroupedLinear) + and act_grouped_linear_fusion_supported(ops[idx + 1], ops[idx], recipe) + ): + out.append( + ForwardScaledActivationGroupedLinear( + activation=ops[idx], + linear=ops[idx + 1], + ) + ) + idx += 2 + else: + out.append(ops[idx]) + idx += 1 + return out + + +class BackwardGroupedLinearScaledActivation(FusedOperation): + """Scaled activation backward + grouped quantize + grouped linear backward.""" + + def __init__(self, *, linear: GroupedLinear, activation: _ScaledActivation) -> None: + super().__init__((linear, activation)) + + def fuser_backward( + self, + basic_op_ctxs: list[OperationContext], + grad_output: torch.Tensor, + *, + basic_op_grad_extra_outputs: list[tuple[torch.Tensor, ...]], + ) -> tuple[ + torch.Tensor, + Iterable[Iterable[Optional[torch.Tensor]]], + Iterable[Iterable[Optional[torch.Tensor]]], + ]: + del basic_op_grad_extra_outputs + linear = self.basic_ops[0] + activation = self.basic_ops[1] + linear_ctx, activation_ctx = basic_op_ctxs + input_, scales = activation_ctx.saved_tensors + input_ = maybe_dequantize(input_, activation_ctx.dtype) + scales = maybe_dequantize(scales, activation_ctx.dtype) + grad_output = maybe_dequantize(grad_output, activation_ctx.dtype) + + split_sizes = linear_ctx.saved_tensors[0] + split_sizes, (grad_output_tensor_offsets,) = tex.splits_to_offsets_multi( + split_sizes, + input_.device, + strides=[linear.out_features], + include_leading_zero=[True], + dtypes=[torch.int64], + bulk_allocate=False, + ) + grad_output_quantizer = linear_ctx.grad_output_quantizers[0] + grad_output_quantizer.set_usage( + rowwise=linear_ctx.input_requires_grad, + columnwise=linear_ctx.weight_requires_grad, + ) + grad_output_quantizer.optimize_for_gemm = True + grouped_dy, dense_dy, grad_scales = _grouped_scaled_dactivation( + activation, + grad_output, + input_, + scales, + quantizer=grad_output_quantizer, + num_groups=linear.num_groups, + split_sizes=split_sizes, + tensor_offsets=grad_output_tensor_offsets, + compute_scale_grad=activation_ctx.extra_input_requires_grad, + ) + + # Feed the grouped-quantized grad into the grouped GEMM while also + # passing the dense high-precision grad so the bias gradient avoids a + # lossy dequantize of ``grouped_dy``. + grad_input, grad_params, grad_extra_inputs = linear._fuser_backward_grouped_tensor( + ctx=linear_ctx, + grad_output=dense_dy, + grouped_grad_output=grouped_dy, + ) + + clear_tensor_data(activation_ctx.saved_tensors[0]) + return ( + grad_input, + [grad_params[0], ()], + [grad_extra_inputs[0], (grad_scales,)], + ) + + @staticmethod + def fuse_backward_ops( + ops: list[FusibleOperation], + *, + recipe: Optional[Recipe] = None, + **unused, + ) -> list[FusibleOperation]: + """Fuse each supported GroupedLinear + ScaledActivation pair.""" + out: list[FusibleOperation] = [] + idx = 0 + while idx < len(ops): + if ( + idx + 1 < len(ops) + and isinstance(ops[idx], GroupedLinear) + and isinstance(ops[idx + 1], _SCALED_ACTIVATION_TYPES) + and act_grouped_linear_fusion_supported(ops[idx], ops[idx + 1], recipe) + ): + out.append( + BackwardGroupedLinearScaledActivation( + linear=ops[idx], + activation=ops[idx + 1], + ) + ) + idx += 2 + else: + out.append(ops[idx]) + idx += 1 + return out From c176fe2ad6a1a2f2c93f91caf219f8967837e098 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Wed, 22 Jul 2026 19:00:15 +0000 Subject: [PATCH 02/17] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- tests/pytorch/test_grouped_mlp.py | 14 +++++----- transformer_engine/pytorch/csrc/extensions.h | 3 +-- .../pytorch/csrc/extensions/activation.cpp | 21 +++++++-------- .../pytorch/csrc/extensions/pybind.cpp | 27 +++++++------------ .../pytorch/ops/basic/grouped_linear.py | 18 ++++--------- .../ops/fused/grouped_linear_activation.py | 7 ++--- 6 files changed, 35 insertions(+), 55 deletions(-) diff --git a/tests/pytorch/test_grouped_mlp.py b/tests/pytorch/test_grouped_mlp.py index dbf89245fc..fdc0c4868c 100644 --- a/tests/pytorch/test_grouped_mlp.py +++ b/tests/pytorch/test_grouped_mlp.py @@ -1038,12 +1038,14 @@ def _make_module(): and te.ops.fused.act_grouped_linear_fusion_supported(fc2, module[1], recipe) and te.ops.fused.act_grouped_linear_fusion_supported(fc1, module[1], recipe) ) - assert any( - isinstance(op, ForwardScaledActivationGroupedLinear) for op, _ in forward_ops - ) == act_grouped_linear_fusion_expected - assert any( - isinstance(op, BackwardGroupedLinearScaledActivation) for op, _ in backward_ops - ) == act_grouped_linear_fusion_expected + assert ( + any(isinstance(op, ForwardScaledActivationGroupedLinear) for op, _ in forward_ops) + == act_grouped_linear_fusion_expected + ) + assert ( + any(isinstance(op, BackwardGroupedLinearScaledActivation) for op, _ in backward_ops) + == act_grouped_linear_fusion_expected + ) # Loose tols for sanity checking tols = {"rtol": 0.125, "atol": 0.25} diff --git a/transformer_engine/pytorch/csrc/extensions.h b/transformer_engine/pytorch/csrc/extensions.h index 30866f75f2..f479d75c13 100644 --- a/transformer_engine/pytorch/csrc/extensions.h +++ b/transformer_engine/pytorch/csrc/extensions.h @@ -327,8 +327,7 @@ py::tuple grouped_scaled_clamped_dswiglu(const at::Tensor &grad, const at::Tenso py::tuple grouped_scaled_dsrelu(const at::Tensor &grad, const at::Tensor &input, const at::Tensor &act_scales, py::handle quantizer, const size_t num_tensors, std::optional first_dims, - std::optional tensor_offsets, - bool compute_scale_grad); + std::optional tensor_offsets, bool compute_scale_grad); /*************************************************************************************************** * LayerNorm **************************************************************************************************/ diff --git a/transformer_engine/pytorch/csrc/extensions/activation.cpp b/transformer_engine/pytorch/csrc/extensions/activation.cpp index b8d7363e1b..ad130078b6 100644 --- a/transformer_engine/pytorch/csrc/extensions/activation.cpp +++ b/transformer_engine/pytorch/csrc/extensions/activation.cpp @@ -362,14 +362,13 @@ py::object grouped_scaled_activation_helper(const at::Tensor& input, const at::T const TensorWrapper& input_nvte = makeTransformerEngineTensor(input_tensor); const TensorWrapper& scales_nvte = makeTransformerEngineTensor(scales_tensor); - auto output = - at::empty({input.size(0), input.size(1) / shape_divisor}, input_tensor.options()); + auto output = at::empty({input.size(0), input.size(1) / shape_divisor}, input_tensor.options()); const TensorWrapper& output_nvte = makeTransformerEngineTensor(output); auto stream = at::cuda::getCurrentCUDAStream(); NVTE_SCOPED_GIL_RELEASE({ - act_func(input_nvte.data(), scales_nvte.data(), output_nvte.data(), - std::forward(args)..., stream); + act_func(input_nvte.data(), scales_nvte.data(), output_nvte.data(), std::forward(args)..., + stream); }); if (quantizer.is_none()) { @@ -419,9 +418,8 @@ py::tuple grouped_scaled_dactivation_helper(const at::Tensor& grad, const at::Te // it (e.g. bias gradient) without a lossy dequantize. py::object grad_input_out = py::cast(grad_input); if (!quantizer.is_none()) { - grad_input_out = - group_quantize(grad_input, quantizer, num_tensors, first_dims, std::nullopt, tensor_offsets, - std::nullopt); + grad_input_out = group_quantize(grad_input, quantizer, num_tensors, first_dims, std::nullopt, + tensor_offsets, std::nullopt); } return py::make_tuple(grad_input_out, py::cast(grad_input), compute_scale_grad ? py::cast(grad_scales) : py::none()); @@ -482,11 +480,10 @@ py::tuple grouped_scaled_clamped_dswiglu(const at::Tensor& grad, const at::Tenso py::tuple grouped_scaled_dsrelu(const at::Tensor& grad, const at::Tensor& input, const at::Tensor& act_scales, py::handle quantizer, const size_t num_tensors, std::optional first_dims, - std::optional tensor_offsets, - bool compute_scale_grad) { - return grouped_scaled_dactivation_helper( - grad, input, act_scales, quantizer, num_tensors, first_dims, tensor_offsets, - compute_scale_grad); + std::optional tensor_offsets, bool compute_scale_grad) { + return grouped_scaled_dactivation_helper(grad, input, act_scales, quantizer, + num_tensors, first_dims, + tensor_offsets, compute_scale_grad); } } // namespace pytorch diff --git a/transformer_engine/pytorch/csrc/extensions/pybind.cpp b/transformer_engine/pytorch/csrc/extensions/pybind.cpp index 003c737a75..944152f10b 100644 --- a/transformer_engine/pytorch/csrc/extensions/pybind.cpp +++ b/transformer_engine/pytorch/csrc/extensions/pybind.cpp @@ -288,39 +288,32 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("grouped_scaled_swiglu", transformer_engine::pytorch::grouped_scaled_swiglu, "Scaled SwiGLU + grouped quantize", py::arg("input"), py::arg("act_scales"), py::arg("quantizer"), py::arg("num_tensors"), py::arg("first_dims"), - py::arg("tensor_offsets") = py::none(), - py::arg("glu_interleave_size") = 0); - m.def("grouped_scaled_clamped_swiglu", - transformer_engine::pytorch::grouped_scaled_clamped_swiglu, + py::arg("tensor_offsets") = py::none(), py::arg("glu_interleave_size") = 0); + m.def("grouped_scaled_clamped_swiglu", transformer_engine::pytorch::grouped_scaled_clamped_swiglu, "Scaled clamped SwiGLU + grouped quantize", py::arg("input"), py::arg("act_scales"), py::arg("quantizer"), py::arg("num_tensors"), py::arg("first_dims"), - py::arg("tensor_offsets") = py::none(), - py::arg("limit") = 7.0f, py::arg("alpha") = 1.702f, + py::arg("tensor_offsets") = py::none(), py::arg("limit") = 7.0f, py::arg("alpha") = 1.702f, py::arg("glu_linear_offset") = 1.0f, py::arg("glu_interleave_size") = 0); m.def("grouped_scaled_srelu", transformer_engine::pytorch::grouped_scaled_srelu, "Scaled SReLU + grouped quantize", py::arg("input"), py::arg("act_scales"), py::arg("quantizer"), py::arg("num_tensors"), py::arg("first_dims"), py::arg("tensor_offsets") = py::none()); m.def("grouped_scaled_dswiglu", transformer_engine::pytorch::grouped_scaled_dswiglu, - "Scaled SwiGLU backward + optional grouped quantize", py::arg("grad"), - py::arg("fwd_input"), py::arg("act_scales"), py::arg("quantizer"), - py::arg("num_tensors"), py::arg("first_dims"), + "Scaled SwiGLU backward + optional grouped quantize", py::arg("grad"), py::arg("fwd_input"), + py::arg("act_scales"), py::arg("quantizer"), py::arg("num_tensors"), py::arg("first_dims"), py::arg("tensor_offsets") = py::none(), py::arg("glu_interleave_size") = 0, py::arg("compute_scale_grad") = true); m.def("grouped_scaled_clamped_dswiglu", transformer_engine::pytorch::grouped_scaled_clamped_dswiglu, "Scaled clamped SwiGLU backward + optional grouped quantize", py::arg("grad"), - py::arg("fwd_input"), py::arg("act_scales"), py::arg("quantizer"), - py::arg("num_tensors"), py::arg("first_dims"), - py::arg("tensor_offsets") = py::none(), py::arg("limit") = 7.0f, + py::arg("fwd_input"), py::arg("act_scales"), py::arg("quantizer"), py::arg("num_tensors"), + py::arg("first_dims"), py::arg("tensor_offsets") = py::none(), py::arg("limit") = 7.0f, py::arg("alpha") = 1.702f, py::arg("glu_linear_offset") = 1.0f, py::arg("glu_interleave_size") = 0, py::arg("compute_scale_grad") = true); m.def("grouped_scaled_dsrelu", transformer_engine::pytorch::grouped_scaled_dsrelu, - "Scaled SReLU backward + optional grouped quantize", py::arg("grad"), - py::arg("fwd_input"), py::arg("act_scales"), py::arg("quantizer"), - py::arg("num_tensors"), py::arg("first_dims"), - py::arg("tensor_offsets") = py::none(), - py::arg("compute_scale_grad") = true); + "Scaled SReLU backward + optional grouped quantize", py::arg("grad"), py::arg("fwd_input"), + py::arg("act_scales"), py::arg("quantizer"), py::arg("num_tensors"), py::arg("first_dims"), + py::arg("tensor_offsets") = py::none(), py::arg("compute_scale_grad") = true); /* DBias + DAct fusions*/ m.def("dbias_dgelu", transformer_engine::pytorch::dbias_dgelu, "DGeLU + DBias + Quantize", py::arg("grad"), py::arg("fwd_input"), py::arg("quantizer")); diff --git a/transformer_engine/pytorch/ops/basic/grouped_linear.py b/transformer_engine/pytorch/ops/basic/grouped_linear.py index b096962747..88a6276a4e 100644 --- a/transformer_engine/pytorch/ops/basic/grouped_linear.py +++ b/transformer_engine/pytorch/ops/basic/grouped_linear.py @@ -1292,10 +1292,7 @@ def _fuser_forward_grouped_tensor( if grouped_input is not None: total_tokens, in_features = grouped_input.logical_shape expected_quantizer = input_quantizers[0] if with_quantized_compute else None - if ( - grouped_input.quantizer is not expected_quantizer - or in_features != self.in_features - ): + if grouped_input.quantizer is not expected_quantizer or in_features != self.in_features: raise ValueError( "GroupedLinear received an incompatible grouped input " f"(quantizer={grouped_input.quantizer}, " @@ -1661,7 +1658,6 @@ def _fuser_backward_grouped_tensor( else: ws, saved_tensors = saved_tensors[:num_groups], saved_tensors[num_groups:] - # Keep the dense high-precision grad for bias-gradient computation while # optionally using a pre-quantized grouped view for the grouped GEMMs. dy_2d = grad_output.reshape(-1, self.out_features) @@ -1673,14 +1669,10 @@ def _fuser_backward_grouped_tensor( # Optionally get dbias if fusion is available with bgrad_group_quantize. dbias_packed = None if grouped_grad_output is not None: - expected_quantizer = ( - ctx.grad_output_quantizers[0] if with_quantized_compute else None - ) - if ( - grouped_grad_output.quantizer is not expected_quantizer - or tuple(grouped_grad_output.logical_shape) - != (total_tokens, self.out_features) - ): + expected_quantizer = ctx.grad_output_quantizers[0] if with_quantized_compute else None + if grouped_grad_output.quantizer is not expected_quantizer or tuple( + grouped_grad_output.logical_shape + ) != (total_tokens, self.out_features): raise ValueError( "GroupedLinear received an incompatible grouped grad_output " f"(quantizer={grouped_grad_output.quantizer}, " diff --git a/transformer_engine/pytorch/ops/fused/grouped_linear_activation.py b/transformer_engine/pytorch/ops/fused/grouped_linear_activation.py index c2b387f3fa..3fbd7a8dd7 100644 --- a/transformer_engine/pytorch/ops/fused/grouped_linear_activation.py +++ b/transformer_engine/pytorch/ops/fused/grouped_linear_activation.py @@ -61,9 +61,7 @@ def _grouped_scaled_activation( clamped.glu_linear_offset, int(activation.glu_interleave_size or 0), ) - return tex.grouped_scaled_srelu( - x, s, quantizer, num_groups, split_sizes, tensor_offsets - ) + return tex.grouped_scaled_srelu(x, s, quantizer, num_groups, split_sizes, tensor_offsets) def _grouped_scaled_dactivation( @@ -196,8 +194,7 @@ def fuser_forward( ) input_quantizers = [ - linear.get_quantizer("forward", 2 * group_idx) - for group_idx in range(linear.num_groups) + linear.get_quantizer("forward", 2 * group_idx) for group_idx in range(linear.num_groups) ] weight_quantizers = [ linear.get_quantizer("forward", 2 * group_idx + 1) From af2fa133842a8b3520cb0cb37f13ef055d98f189 Mon Sep 17 00:00:00 2001 From: Varun Thumbe Date: Thu, 23 Jul 2026 07:27:10 +0000 Subject: [PATCH 03/17] address review comments + cleanup Signed-off-by: Varun Thumbe --- tests/pytorch/test_grouped_mlp.py | 27 +- .../pytorch/csrc/extensions/activation.cpp | 2 + .../pytorch/ops/basic/activation.py | 60 ++-- .../pytorch/ops/basic/grouped_linear.py | 289 ++++++++++-------- .../pytorch/ops/fused/__init__.py | 12 +- .../backward_activation_grouped_linear.py | 179 +++++++++++ ...y => forward_activation_grouped_linear.py} | 178 +---------- 7 files changed, 432 insertions(+), 315 deletions(-) create mode 100644 transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py rename transformer_engine/pytorch/ops/fused/{grouped_linear_activation.py => forward_activation_grouped_linear.py} (56%) diff --git a/tests/pytorch/test_grouped_mlp.py b/tests/pytorch/test_grouped_mlp.py index fdc0c4868c..f2421ee4dd 100644 --- a/tests/pytorch/test_grouped_mlp.py +++ b/tests/pytorch/test_grouped_mlp.py @@ -23,8 +23,10 @@ OUTPUT_BUFFER_KEY, GRAD_INPUT_BUFFER_KEY, ) -from transformer_engine.pytorch.ops.fused.grouped_linear_activation import ( - BackwardGroupedLinearScaledActivation, +from transformer_engine.pytorch.ops.fused.backward_activation_grouped_linear import ( + BackwardScaledActivationGroupedLinear, +) +from transformer_engine.pytorch.ops.fused.forward_activation_grouped_linear import ( ForwardScaledActivationGroupedLinear, ) from transformer_engine.pytorch import ( @@ -737,6 +739,13 @@ def test_grouped_mlp( maybe_skip_quantization(quantization, dims=in_shape, device=device, dtype=dtype) if dtype == torch.bfloat16 and not is_bf16_available(): pytest.skip("BF16 requires SM 8.0+") + if os.environ.get("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "0") == "0" and ( + single_grouped_weight or single_grouped_bias + ): + pytest.skip( + "single_grouped_weight/single_grouped_bias requires" + " NVTE_GROUPED_LINEAR_SINGLE_PARAM=1" + ) if single_grouped_weight and quantization != "mxfp8": pytest.skip("single_grouped_weight is only supported for MXFP8 quantization") if single_grouped_bias and not bias: @@ -1043,7 +1052,7 @@ def _make_module(): == act_grouped_linear_fusion_expected ) assert ( - any(isinstance(op, BackwardGroupedLinearScaledActivation) for op, _ in backward_ops) + any(isinstance(op, BackwardScaledActivationGroupedLinear) for op, _ in backward_ops) == act_grouped_linear_fusion_expected ) @@ -1199,6 +1208,10 @@ def test_grouped_mlp_single_weight_numerics( ) -> None: """single_grouped_weight=True/False should match exactly for fused MXFP8 grouped MLP.""" + if os.environ.get("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "0") == "0": + pytest.skip( + "single_grouped_weight requires NVTE_GROUPED_LINEAR_SINGLE_PARAM=1" + ) if not te.ops.fused.GroupedMLP_CuTeGEMMGLU.is_supported(): pytest.skip("MXFP8 fused grouped MLP is not supported on this system") @@ -1517,6 +1530,10 @@ def test_grouped_mlp_overwrite_main_grad( that read ``.grad`` don't see stale bytes from the cached dummy). """ + if os.environ.get("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "0") == "0" and single_grouped_weight: + pytest.skip( + "single_grouped_weight requires NVTE_GROUPED_LINEAR_SINGLE_PARAM=1" + ) if not te.ops.fused.GroupedMLP_CuTeGEMMGLU.is_supported(): pytest.skip("MXFP8 fused grouped MLP is not supported on this system") @@ -1648,6 +1665,10 @@ def test_grouped_mlp_cuda_graph_safe_mxfp8( ) -> None: """Grouped MLP forward+backward should be CUDA graph capturable (MXFP8).""" + if os.environ.get("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "0") == "0" and single_grouped_weight: + pytest.skip( + "single_grouped_weight requires NVTE_GROUPED_LINEAR_SINGLE_PARAM=1" + ) if not te.ops.fused.GroupedMLP_CuTeGEMMGLU.is_supported(): pytest.skip("MXFP8 fused grouped MLP is not supported on this system") if dtype not in (torch.bfloat16, torch.float16): diff --git a/transformer_engine/pytorch/csrc/extensions/activation.cpp b/transformer_engine/pytorch/csrc/extensions/activation.cpp index ad130078b6..bce92fa3ab 100644 --- a/transformer_engine/pytorch/csrc/extensions/activation.cpp +++ b/transformer_engine/pytorch/csrc/extensions/activation.cpp @@ -362,6 +362,8 @@ py::object grouped_scaled_activation_helper(const at::Tensor& input, const at::T const TensorWrapper& input_nvte = makeTransformerEngineTensor(input_tensor); const TensorWrapper& scales_nvte = makeTransformerEngineTensor(scales_tensor); + // Keep the dense activation in the input dtype. It is only a transient buffer when a + // quantizer is provided; the quantized return path never exposes it to the caller. auto output = at::empty({input.size(0), input.size(1) / shape_divisor}, input_tensor.options()); const TensorWrapper& output_nvte = makeTransformerEngineTensor(output); diff --git a/transformer_engine/pytorch/ops/basic/activation.py b/transformer_engine/pytorch/ops/basic/activation.py index f4beffe90c..42a6ef2243 100644 --- a/transformer_engine/pytorch/ops/basic/activation.py +++ b/transformer_engine/pytorch/ops/basic/activation.py @@ -348,18 +348,8 @@ def _activation_backward_impl(self, *args, **kwargs) -> torch.Tensor: return tex.dsrelu(*args, **kwargs) -class ScaledSReLU(BasicOperation): - r"""Squared ReLU with per-row post-scaling. - - If the SReLU output has shape ``(d_1, ..., d_n)``, it is multiplied - with an extra input tensor of shape ``(d_1, ..., d_{n-1})``. - - Parameters - ---------- - activation_recompute_in_mlp : bool, default = ``False`` - Enable fused grouped MLP kernels to recompute activation outputs - during backward when supported instead of saving them. - """ +class _ScaledUnary(BasicOperation, metaclass=abc.ABCMeta): + """Unary activation with per-row scales (fused grouped MLP middle op).""" num_extra_inputs: int = 1 @@ -367,6 +357,18 @@ def __init__(self, *, activation_recompute_in_mlp: bool = False) -> None: super().__init__() self.activation_recompute_in_mlp: bool = activation_recompute_in_mlp + @abc.abstractmethod + def _unary_forward(self, input_: torch.Tensor) -> torch.Tensor: + """Apply the unary activation.""" + + @abc.abstractmethod + def _unary_backward( + self, + grad_output: torch.Tensor, + input_: torch.Tensor, + ) -> torch.Tensor: + """Apply the unary activation backward pass.""" + def op_forward(self, *args, **kwargs) -> None: raise RuntimeError( f"{self.__class__.__name__} operation has " @@ -410,7 +412,7 @@ def fuser_forward( x = maybe_dequantize(input_.contiguous(), dtype) scales = maybe_dequantize(extra_input, dtype) - y = tex.srelu(x, None) * scales.unsqueeze(-1) + y = self._unary_forward(x) * scales.unsqueeze(-1) ctx = basic_op_ctxs[0] if ctx.requires_grad: @@ -450,19 +452,43 @@ def fuser_backward( grad_input = None if ctx.input_requires_grad: - grad_srelu_out = grad_output * scales.unsqueeze(-1) - grad_input = tex.dsrelu(grad_srelu_out, x, None) + grad_unary_out = grad_output * scales.unsqueeze(-1) + grad_input = self._unary_backward(grad_unary_out, x) grad_extra_input = None if ctx.extra_input_requires_grad: - srelu_out = tex.srelu(x, None) - grad_extra_input = torch.linalg.vecdot(srelu_out, grad_output) + unary_out = self._unary_forward(x) + grad_extra_input = torch.linalg.vecdot(unary_out, grad_output) clear_tensor_data(ctx.saved_tensors[0]) return grad_input, [()], [(grad_extra_input,)] +class ScaledSReLU(_ScaledUnary): + r"""Squared ReLU with per-row post-scaling. + + If the SReLU output has shape ``(d_1, ..., d_n)``, it is multiplied + with an extra input tensor of shape ``(d_1, ..., d_{n-1})``. + + Parameters + ---------- + activation_recompute_in_mlp : bool, default = ``False`` + Enable fused grouped MLP kernels to recompute activation outputs + during backward when supported instead of saving them. + """ + + def _unary_forward(self, input_: torch.Tensor) -> torch.Tensor: + return tex.srelu(input_, None) + + def _unary_backward( + self, + grad_output: torch.Tensor, + input_: torch.Tensor, + ) -> torch.Tensor: + return tex.dsrelu(grad_output, input_, None) + + class SReGLU(_ActivationOperation): r"""Squared Rectified Gated Linear Unit diff --git a/transformer_engine/pytorch/ops/basic/grouped_linear.py b/transformer_engine/pytorch/ops/basic/grouped_linear.py index 88a6276a4e..fffa4d0610 100644 --- a/transformer_engine/pytorch/ops/basic/grouped_linear.py +++ b/transformer_engine/pytorch/ops/basic/grouped_linear.py @@ -1009,7 +1009,7 @@ def fuser_forward( ) if use_grouped_tensor_path: - out, tensors_to_save = self._fuser_forward_grouped_tensor( + out, tensors_to_save = self._fuser_forward_graph_safe( input_=input_, split_sizes=split_sizes, scales=scales, @@ -1241,10 +1241,10 @@ def _fuser_forward_split_quantize( saved.extend(ws) return out, tuple(saved) - def _fuser_forward_grouped_tensor( + def _fuser_forward_graph_safe( self, *, - input_: Optional[torch.Tensor] = None, + input_: torch.Tensor, split_sizes: torch.Tensor, scales: Optional[torch.Tensor], with_quantized_compute: bool, @@ -1255,81 +1255,93 @@ def _fuser_forward_grouped_tensor( weight_requires_grad: bool, device: torch.device, out_buffer: Optional[torch.Tensor] = None, - grouped_input: Optional[GroupedTensorStorage] = None, - grouped_tensor_offsets: Optional[ - tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor] - ] = None, ) -> tuple[torch.Tensor, tuple[Optional[torch.Tensor], ...]]: - """Graph-safe GroupedTensor forward path (pure compute). - Returns ``(output, tensors_to_save)``. ``split_sizes``, - ``base_split_offsets`` and ``split_points`` are returned so that - ``fuser_forward_save_ctx`` can call ``save_for_backward`` on them. - - Provide either a dense ``input_`` or a pre-built ``grouped_input``. - When both are given, ``grouped_input`` is used after validation. - """ + """Build graph-safe grouped input storage and run grouped GEMM.""" num_groups = self.num_groups - has_bias = self.has_bias - - if grouped_tensor_offsets is None: - split_sizes, grouped_tensor_offsets = tex.splits_to_offsets_multi( + split_sizes, grouped_tensor_offsets = tex.splits_to_offsets_multi( + split_sizes, + device, + strides=[1, 1, self.in_features, self.out_features], + include_leading_zero=[False, True, True, True], + dtypes=[torch.int32, torch.int64, torch.int64, torch.int64], + bulk_allocate=True, + ) + input_tensor_offsets = grouped_tensor_offsets[2] + original_shape = list(input_.size()) + total_tokens = math.prod(original_shape[:-1]) + x = maybe_dequantize(input_, dtype).reshape(-1, self.in_features) + if with_quantized_compute: + input_quantizer = input_quantizers[0] + input_quantizer.set_usage(rowwise=True, columnwise=weight_requires_grad) + input_quantizer.optimize_for_gemm = True + grouped_x = tex.group_quantize( + x, + input_quantizer, + num_groups, split_sizes, - device, - strides=[1, 1, self.in_features, self.out_features], - include_leading_zero=[False, True, True, True], - dtypes=[torch.int32, torch.int64, torch.int64, torch.int64], - bulk_allocate=True, + tensor_offsets=input_tensor_offsets, ) - ( - split_points, - base_split_offsets, - input_tensor_offsets, - output_tensor_offsets, - ) = grouped_tensor_offsets - - # Build the input GroupedTensor unless a preceding fused operation - # already produced it. - if grouped_input is not None: - total_tokens, in_features = grouped_input.logical_shape - expected_quantizer = input_quantizers[0] if with_quantized_compute else None - if grouped_input.quantizer is not expected_quantizer or in_features != self.in_features: - raise ValueError( - "GroupedLinear received an incompatible grouped input " - f"(quantizer={grouped_input.quantizer}, " - f"logical_shape={grouped_input.logical_shape}; " - f"expected quantizer={expected_quantizer}, " - f"in_features={self.in_features})" - ) - grouped_x = grouped_input - out_shape = [total_tokens, self.out_features] else: - # Flatten to 2D so the first dim is the total token count. - original_shape = list(input_.size()) - total_tokens = math.prod(original_shape[:-1]) - out_shape = original_shape[:-1] + [self.out_features] - x = maybe_dequantize(input_, dtype).reshape(-1, self.in_features) - if with_quantized_compute: - input_quantizer = input_quantizers[0] - input_quantizer.set_usage(rowwise=True, columnwise=weight_requires_grad) - input_quantizer.optimize_for_gemm = True - grouped_x = tex.group_quantize( - x, - input_quantizer, - num_groups, - split_sizes, - tensor_offsets=input_tensor_offsets, - ) - else: - # No quantize: wrap the contiguous high-precision buffer. - grouped_x = GroupedTensorStorage( - shape=(total_tokens, self.in_features), - dtype=dtype, - num_tensors=num_groups, - quantizer=None, - data=x.reshape(-1), - first_dims=split_sizes, - tensor_offsets=input_tensor_offsets, - ) + grouped_x = GroupedTensorStorage( + shape=(total_tokens, self.in_features), + dtype=dtype, + num_tensors=num_groups, + quantizer=None, + data=x.reshape(-1), + first_dims=split_sizes, + tensor_offsets=input_tensor_offsets, + ) + return self._fuser_forward_grouped_tensor( + grouped_input=grouped_x, + split_sizes=split_sizes, + scales=scales, + with_quantized_compute=with_quantized_compute, + input_quantizers=input_quantizers, + weight_quantizers=weight_quantizers, + dtype=dtype, + input_requires_grad=input_requires_grad, + weight_requires_grad=weight_requires_grad, + device=device, + split_points=grouped_tensor_offsets[0], + base_split_offsets=grouped_tensor_offsets[1], + output_tensor_offsets=grouped_tensor_offsets[3], + out_buffer=out_buffer, + out_shape=original_shape[:-1] + [self.out_features], + ) + + def _fuser_forward_grouped_tensor( + self, + *, + grouped_input: GroupedTensorStorage, + split_sizes: torch.Tensor, + scales: Optional[torch.Tensor], + with_quantized_compute: bool, + input_quantizers: list[Optional[Quantizer]], + weight_quantizers: list[Optional[Quantizer]], + dtype: torch.dtype, + input_requires_grad: bool, + weight_requires_grad: bool, + device: torch.device, + split_points: torch.Tensor, + base_split_offsets: torch.Tensor, + output_tensor_offsets: torch.Tensor, + out_buffer: Optional[torch.Tensor] = None, + out_shape: list[int], + ) -> tuple[torch.Tensor, tuple[Optional[torch.Tensor], ...]]: + """Run grouped GEMM with a pre-built grouped input.""" + num_groups = self.num_groups + has_bias = self.has_bias + total_tokens, in_features = grouped_input.logical_shape + expected_quantizer = input_quantizers[0] if with_quantized_compute else None + if grouped_input.quantizer is not expected_quantizer or in_features != self.in_features: + raise ValueError( + "GroupedLinear received an incompatible grouped input " + f"(quantizer={grouped_input.quantizer}, " + f"logical_shape={grouped_input.logical_shape}; " + f"expected quantizer={expected_quantizer}, " + f"in_features={self.in_features})" + ) + grouped_x = grouped_input if is_cpu_offload_enabled() and grouped_x is not None: start_offload(grouped_x) @@ -1428,7 +1440,7 @@ def fuser_backward( ctx = basic_op_ctxs[0] # Dispatch to the path used in forward (saved as ``ctx.use_grouped_tensor_path``). if getattr(ctx, "use_grouped_tensor_path", False): - return self._fuser_backward_grouped_tensor( + return self._fuser_backward_graph_safe( ctx=ctx, grad_output=grad_output, ) @@ -1617,15 +1629,78 @@ def _fuser_backward_split_quantize( grad_extra = (None, grad_scales) if self._scale_bias else (None,) return grad_input, [grad_params], [grad_extra] - # ================================================================== - # Graph-safe backward: counterpart of `_fuser_forward_grouped_tensor`. - # ================================================================== + def _fuser_backward_graph_safe( + self, + *, + ctx: OperationContext, + grad_output: torch.Tensor, + ) -> tuple[ + torch.Tensor, + Iterable[Iterable[Optional[torch.Tensor]]], + Iterable[Iterable[Optional[torch.Tensor]]], + ]: + """Build graph-safe grouped grad-output storage and run grouped GEMMs.""" + num_groups = self.num_groups + has_bias = self.has_bias + dtype = ctx.dtype + with_quantized_compute = bool(getattr(ctx, "with_quantized_compute", False)) + split_sizes = ctx.saved_tensors[0] + base_split_offsets = ctx.saved_tensors[1] + dy_2d = grad_output.reshape(-1, self.out_features) + total_tokens = dy_2d.size(0) + + dbias_packed = None + if with_quantized_compute: + grad_output_quantizer = ctx.grad_output_quantizers[0] + grad_output_quantizer.set_usage( + rowwise=ctx.input_requires_grad, + columnwise=ctx.weight_requires_grad, + ) + grad_output_quantizer.optimize_for_gemm = True + if ( + has_bias + and not self._scale_bias + and isinstance(grad_output_quantizer, MXFP8Quantizer) + ): + grouped_dy, dbias_packed = tex.bgrad_group_quantize( + dy_2d, + grad_output_quantizer, + num_groups, + split_sizes, + ) + else: + grouped_dy = tex.group_quantize( + dy_2d, + grad_output_quantizer, + num_groups, + split_sizes, + ) + else: + dy_2d = maybe_dequantize(dy_2d, dtype) + grouped_dy = GroupedTensorStorage( + shape=(total_tokens, self.out_features), + dtype=dtype, + num_tensors=num_groups, + quantizer=None, + data=dy_2d.reshape(-1), + first_dims=split_sizes, + tensor_offsets=base_split_offsets * self.out_features, + ) + + return self._fuser_backward_grouped_tensor( + ctx=ctx, + grad_output=grad_output, + grouped_grad_output=grouped_dy, + dbias_packed=dbias_packed, + ) + def _fuser_backward_grouped_tensor( self, *, ctx: OperationContext, grad_output: torch.Tensor, - grouped_grad_output: Optional[GroupedTensorStorage] = None, + grouped_grad_output: GroupedTensorStorage, + dbias_packed: Optional[torch.Tensor] = None, ) -> tuple[ torch.Tensor, Iterable[Iterable[Optional[torch.Tensor]]], @@ -1664,54 +1739,18 @@ def _fuser_backward_grouped_tensor( total_tokens = dy_2d.size(0) grad_input_shape = list(grad_output.size())[:-1] + [self.in_features] - # Build the grad_output GroupedTensor unless a preceding fused - # operation already produced it. - # Optionally get dbias if fusion is available with bgrad_group_quantize. - dbias_packed = None - if grouped_grad_output is not None: - expected_quantizer = ctx.grad_output_quantizers[0] if with_quantized_compute else None - if grouped_grad_output.quantizer is not expected_quantizer or tuple( - grouped_grad_output.logical_shape - ) != (total_tokens, self.out_features): - raise ValueError( - "GroupedLinear received an incompatible grouped grad_output " - f"(quantizer={grouped_grad_output.quantizer}, " - f"logical_shape={grouped_grad_output.logical_shape}; " - f"expected quantizer={expected_quantizer}, " - f"logical_shape={(total_tokens, self.out_features)})" - ) - grouped_dy = grouped_grad_output - elif with_quantized_compute: - grad_output_quantizer = ctx.grad_output_quantizers[0] - grad_output_quantizer.set_usage( - rowwise=ctx.input_requires_grad, columnwise=ctx.weight_requires_grad - ) - grad_output_quantizer.optimize_for_gemm = True - - if ( - has_bias - and not self._scale_bias - and isinstance(grad_output_quantizer, MXFP8Quantizer) - ): - grouped_dy, dbias_packed = tex.bgrad_group_quantize( - dy_2d, grad_output_quantizer, num_groups, split_sizes - ) - else: - grouped_dy = tex.group_quantize( - dy_2d, grad_output_quantizer, num_groups, split_sizes - ) - else: - dy_2d = maybe_dequantize(dy_2d, dtype) - # Wrap BF16/FP16 buffer as a GroupedTensor for grouped gemm - grouped_dy = GroupedTensorStorage( - shape=(total_tokens, self.out_features), - dtype=dtype, - num_tensors=num_groups, - quantizer=None, - data=dy_2d.reshape(-1), - first_dims=split_sizes, - tensor_offsets=base_split_offsets * self.out_features, + expected_quantizer = ctx.grad_output_quantizers[0] if with_quantized_compute else None + if grouped_grad_output.quantizer is not expected_quantizer or tuple( + grouped_grad_output.logical_shape + ) != (total_tokens, self.out_features): + raise ValueError( + "GroupedLinear received an incompatible grouped grad_output " + f"(quantizer={grouped_grad_output.quantizer}, " + f"logical_shape={grouped_grad_output.logical_shape}; " + f"expected quantizer={expected_quantizer}, " + f"logical_shape={(total_tokens, self.out_features)})" ) + grouped_dy = grouped_grad_output # Bias Grads compute if not already computed in bgrad_group_quantize final_bias_grads: Optional[torch.Tensor] = None diff --git a/transformer_engine/pytorch/ops/fused/__init__.py b/transformer_engine/pytorch/ops/fused/__init__.py index 4c2832f65d..85191a9807 100644 --- a/transformer_engine/pytorch/ops/fused/__init__.py +++ b/transformer_engine/pytorch/ops/fused/__init__.py @@ -6,17 +6,17 @@ from ..fuser import register_backward_fusion, register_forward_fusion from .backward_activation_bias import BackwardActivationBias +from .backward_activation_grouped_linear import BackwardScaledActivationGroupedLinear from .backward_add_rmsnorm import BackwardAddRMSNorm from .backward_linear_add import BackwardLinearAdd from .backward_linear_scale import BackwardLinearScale -from .forward_linear_bias_activation import ForwardLinearBiasActivation -from .forward_linear_bias_add import ForwardLinearBiasAdd -from .forward_linear_scale_add import ForwardLinearScaleAdd -from .grouped_linear_activation import ( - BackwardGroupedLinearScaledActivation, +from .forward_activation_grouped_linear import ( ForwardScaledActivationGroupedLinear, act_grouped_linear_fusion_supported, ) +from .forward_linear_bias_activation import ForwardLinearBiasActivation +from .forward_linear_bias_add import ForwardLinearBiasAdd +from .forward_linear_scale_add import ForwardLinearScaleAdd from .userbuffers_backward_linear import UserbuffersBackwardLinear from .userbuffers_forward_linear import UserbuffersForwardLinear @@ -34,7 +34,7 @@ register_backward_fusion(BackwardLinearScale.fuse_backward_ops) register_backward_fusion(BackwardActivationBias.fuse_backward_ops) register_backward_fusion(BackwardAddRMSNorm.fuse_backward_ops) -register_backward_fusion(BackwardGroupedLinearScaledActivation.fuse_backward_ops) +register_backward_fusion(BackwardScaledActivationGroupedLinear.fuse_backward_ops) # Import experimental fusions # Note: Registration logic is non-trivial, so submodule handles it internally. diff --git a/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py b/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py new file mode 100644 index 0000000000..ad8b4b5a7a --- /dev/null +++ b/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py @@ -0,0 +1,179 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Fused scaled activation + grouped linear backward.""" + +from __future__ import annotations + +from collections.abc import Iterable +from typing import Optional + +import torch + +import transformer_engine_torch as tex +from ...quantization import Recipe +from ...tensor import Quantizer +from ...utils import clear_tensor_data +from .._common import maybe_dequantize +from ..basic import GroupedLinear, ScaledClampedQGeGLU, ScaledSReLU, ScaledSwiGLU +from ..op import FusedOperation, FusibleOperation, OperationContext +from .forward_activation_grouped_linear import ( + _SCALED_ACTIVATION_TYPES, + _ScaledActivation, + act_grouped_linear_fusion_supported, +) + + +def _grouped_scaled_dactivation( + activation: _ScaledActivation, + grad_output: torch.Tensor, + input_: torch.Tensor, + scales: torch.Tensor, + *, + quantizer: Quantizer, + num_groups: int, + split_sizes: torch.Tensor, + tensor_offsets: torch.Tensor, + compute_scale_grad: bool, +) -> tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]: + """Dispatch a grouped scaled activation backward pass.""" + dy = grad_output.reshape(-1, grad_output.size(-1)) + x = input_.reshape(-1, input_.size(-1)) + s = scales.reshape(-1) + if isinstance(activation, ScaledSwiGLU): + return tex.grouped_scaled_dswiglu( + dy, + x, + s, + quantizer, + num_groups, + split_sizes, + tensor_offsets, + int(activation.glu_interleave_size or 0), + compute_scale_grad, + ) + if isinstance(activation, ScaledClampedQGeGLU): + clamped = activation._clamped + return tex.grouped_scaled_clamped_dswiglu( + dy, + x, + s, + quantizer, + num_groups, + split_sizes, + tensor_offsets, + clamped.limit, + clamped.alpha, + clamped.glu_linear_offset, + int(activation.glu_interleave_size or 0), + compute_scale_grad, + ) + if isinstance(activation, ScaledSReLU): + return tex.grouped_scaled_dsrelu( + dy, + x, + s, + quantizer, + num_groups, + split_sizes, + tensor_offsets, + compute_scale_grad, + ) + raise TypeError(f"Unsupported scaled activation type ({type(activation).__name__})") + + +class BackwardScaledActivationGroupedLinear(FusedOperation): + """Scaled activation backward + grouped quantize + grouped linear backward.""" + + def __init__(self, *, linear: GroupedLinear, activation: _ScaledActivation) -> None: + super().__init__((linear, activation)) + + def fuser_backward( + self, + basic_op_ctxs: list[OperationContext], + grad_output: torch.Tensor, + *, + basic_op_grad_extra_outputs: list[tuple[torch.Tensor, ...]], + ) -> tuple[ + torch.Tensor, + Iterable[Iterable[Optional[torch.Tensor]]], + Iterable[Iterable[Optional[torch.Tensor]]], + ]: + del basic_op_grad_extra_outputs + linear = self.basic_ops[0] + activation = self.basic_ops[1] + linear_ctx, activation_ctx = basic_op_ctxs + input_, scales = activation_ctx.saved_tensors + input_ = maybe_dequantize(input_, activation_ctx.dtype) + scales = maybe_dequantize(scales, activation_ctx.dtype) + grad_output = maybe_dequantize(grad_output, activation_ctx.dtype) + + split_sizes = linear_ctx.saved_tensors[0] + split_sizes, (grad_output_tensor_offsets,) = tex.splits_to_offsets_multi( + split_sizes, + input_.device, + strides=[linear.out_features], + include_leading_zero=[True], + dtypes=[torch.int64], + bulk_allocate=False, + ) + grad_output_quantizer = linear_ctx.grad_output_quantizers[0] + grad_output_quantizer.set_usage( + rowwise=linear_ctx.input_requires_grad, + columnwise=linear_ctx.weight_requires_grad, + ) + grad_output_quantizer.optimize_for_gemm = True + grouped_dy, dense_dy, grad_scales = _grouped_scaled_dactivation( + activation, + grad_output, + input_, + scales, + quantizer=grad_output_quantizer, + num_groups=linear.num_groups, + split_sizes=split_sizes, + tensor_offsets=grad_output_tensor_offsets, + compute_scale_grad=activation_ctx.extra_input_requires_grad, + ) + + grad_input, grad_params, grad_extra_inputs = linear._fuser_backward_grouped_tensor( + ctx=linear_ctx, + grad_output=dense_dy, + grouped_grad_output=grouped_dy, + ) + + clear_tensor_data(activation_ctx.saved_tensors[0]) + return ( + grad_input, + [grad_params[0], ()], + [grad_extra_inputs[0], (grad_scales,)], + ) + + @staticmethod + def fuse_backward_ops( + ops: list[FusibleOperation], + *, + recipe: Optional[Recipe] = None, + **unused, + ) -> list[FusibleOperation]: + """Fuse each supported GroupedLinear + ScaledActivation pair.""" + out: list[FusibleOperation] = [] + idx = 0 + while idx < len(ops): + if ( + idx + 1 < len(ops) + and isinstance(ops[idx], GroupedLinear) + and isinstance(ops[idx + 1], _SCALED_ACTIVATION_TYPES) + and act_grouped_linear_fusion_supported(ops[idx], ops[idx + 1], recipe) + ): + out.append( + BackwardScaledActivationGroupedLinear( + linear=ops[idx], + activation=ops[idx + 1], + ) + ) + idx += 2 + else: + out.append(ops[idx]) + idx += 1 + return out diff --git a/transformer_engine/pytorch/ops/fused/grouped_linear_activation.py b/transformer_engine/pytorch/ops/fused/forward_activation_grouped_linear.py similarity index 56% rename from transformer_engine/pytorch/ops/fused/grouped_linear_activation.py rename to transformer_engine/pytorch/ops/fused/forward_activation_grouped_linear.py index 3fbd7a8dd7..f004eb943f 100644 --- a/transformer_engine/pytorch/ops/fused/grouped_linear_activation.py +++ b/transformer_engine/pytorch/ops/fused/forward_activation_grouped_linear.py @@ -2,7 +2,7 @@ # # See LICENSE for license information. -"""Pair fusions between scaled activations and grouped linear operations.""" +"""Fused scaled activation + grouped linear forward.""" from __future__ import annotations @@ -15,14 +15,15 @@ from ...cpu_offload import is_cpu_offload_enabled, mark_activation_offload from ...quantization import Recipe from ...tensor import Quantizer -from ...utils import clear_tensor_data from .._common import maybe_dequantize from ..basic import GroupedLinear, ScaledClampedQGeGLU, ScaledSReLU, ScaledSwiGLU +from ..basic.activation import _ScaledUnary +from ..basic.swiglu import _ScaledGLU from ..op import FusedOperation, FusibleOperation, OperationContext -_ScaledActivation = ScaledSwiGLU | ScaledClampedQGeGLU | ScaledSReLU -_SCALED_ACTIVATION_TYPES = (ScaledSwiGLU, ScaledClampedQGeGLU, ScaledSReLU) +_ScaledActivation = _ScaledGLU | _ScaledUnary +_SCALED_ACTIVATION_TYPES = (_ScaledGLU, _ScaledUnary) def _grouped_scaled_activation( @@ -34,7 +35,7 @@ def _grouped_scaled_activation( split_sizes: torch.Tensor, tensor_offsets: torch.Tensor, ) -> torch.Tensor: - """Dispatch to the matching tex.grouped_scaled_* forward API.""" + """Dispatch a grouped scaled activation.""" x = input_.reshape(-1, input_.size(-1)) s = scales.reshape(-1) if isinstance(activation, ScaledSwiGLU): @@ -61,71 +62,16 @@ def _grouped_scaled_activation( clamped.glu_linear_offset, int(activation.glu_interleave_size or 0), ) - return tex.grouped_scaled_srelu(x, s, quantizer, num_groups, split_sizes, tensor_offsets) - - -def _grouped_scaled_dactivation( - activation: _ScaledActivation, - grad_output: torch.Tensor, - input_: torch.Tensor, - scales: torch.Tensor, - *, - quantizer: Quantizer, - num_groups: int, - split_sizes: torch.Tensor, - tensor_offsets: torch.Tensor, - compute_scale_grad: bool, -) -> tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]: - """Dispatch to the matching tex.grouped_scaled_d* API. - - Returns ``(grouped_dx, dense_dx, dscales)`` where ``grouped_dx`` is the - grouped-quantized grad input for the next grouped GEMM and ``dense_dx`` is - the dense high-precision grad input (reused for the bias gradient so we do - not have to dequantize ``grouped_dx``). - """ - dy = grad_output.reshape(-1, grad_output.size(-1)) - x = input_.reshape(-1, input_.size(-1)) - s = scales.reshape(-1) - if isinstance(activation, ScaledSwiGLU): - dx, dense_dx, dscales = tex.grouped_scaled_dswiglu( - dy, - x, - s, - quantizer, - num_groups, - split_sizes, - tensor_offsets, - int(activation.glu_interleave_size or 0), - compute_scale_grad, - ) - elif isinstance(activation, ScaledClampedQGeGLU): - clamped = activation._clamped - dx, dense_dx, dscales = tex.grouped_scaled_clamped_dswiglu( - dy, + if isinstance(activation, ScaledSReLU): + return tex.grouped_scaled_srelu( x, s, quantizer, num_groups, split_sizes, tensor_offsets, - clamped.limit, - clamped.alpha, - clamped.glu_linear_offset, - int(activation.glu_interleave_size or 0), - compute_scale_grad, ) - else: - dx, dense_dx, dscales = tex.grouped_scaled_dsrelu( - dy, - x, - s, - quantizer, - num_groups, - split_sizes, - tensor_offsets, - compute_scale_grad, - ) - return dx, dense_dx, dscales + raise TypeError(f"Unsupported scaled activation type ({type(activation).__name__})") def act_grouped_linear_fusion_supported( @@ -224,6 +170,7 @@ def fuser_forward( activation_ctx.save_for_backward(input_, scales) out, tensors_to_save = linear._fuser_forward_grouped_tensor( + grouped_input=grouped_x, split_sizes=split_sizes, scales=linear_scales, with_quantized_compute=True, @@ -233,8 +180,10 @@ def fuser_forward( input_requires_grad=linear_ctx.requires_grad, weight_requires_grad=weight_requires_grad, device=device, - grouped_input=grouped_x, - grouped_tensor_offsets=grouped_tensor_offsets, + split_points=grouped_tensor_offsets[0], + base_split_offsets=grouped_tensor_offsets[1], + output_tensor_offsets=grouped_tensor_offsets[3], + out_shape=list(input_.size())[:-1] + [linear.out_features], ) linear.fuser_forward_save_ctx( basic_op_ctxs=[linear_ctx], @@ -277,102 +226,3 @@ def fuse_forward_ops( out.append(ops[idx]) idx += 1 return out - - -class BackwardGroupedLinearScaledActivation(FusedOperation): - """Scaled activation backward + grouped quantize + grouped linear backward.""" - - def __init__(self, *, linear: GroupedLinear, activation: _ScaledActivation) -> None: - super().__init__((linear, activation)) - - def fuser_backward( - self, - basic_op_ctxs: list[OperationContext], - grad_output: torch.Tensor, - *, - basic_op_grad_extra_outputs: list[tuple[torch.Tensor, ...]], - ) -> tuple[ - torch.Tensor, - Iterable[Iterable[Optional[torch.Tensor]]], - Iterable[Iterable[Optional[torch.Tensor]]], - ]: - del basic_op_grad_extra_outputs - linear = self.basic_ops[0] - activation = self.basic_ops[1] - linear_ctx, activation_ctx = basic_op_ctxs - input_, scales = activation_ctx.saved_tensors - input_ = maybe_dequantize(input_, activation_ctx.dtype) - scales = maybe_dequantize(scales, activation_ctx.dtype) - grad_output = maybe_dequantize(grad_output, activation_ctx.dtype) - - split_sizes = linear_ctx.saved_tensors[0] - split_sizes, (grad_output_tensor_offsets,) = tex.splits_to_offsets_multi( - split_sizes, - input_.device, - strides=[linear.out_features], - include_leading_zero=[True], - dtypes=[torch.int64], - bulk_allocate=False, - ) - grad_output_quantizer = linear_ctx.grad_output_quantizers[0] - grad_output_quantizer.set_usage( - rowwise=linear_ctx.input_requires_grad, - columnwise=linear_ctx.weight_requires_grad, - ) - grad_output_quantizer.optimize_for_gemm = True - grouped_dy, dense_dy, grad_scales = _grouped_scaled_dactivation( - activation, - grad_output, - input_, - scales, - quantizer=grad_output_quantizer, - num_groups=linear.num_groups, - split_sizes=split_sizes, - tensor_offsets=grad_output_tensor_offsets, - compute_scale_grad=activation_ctx.extra_input_requires_grad, - ) - - # Feed the grouped-quantized grad into the grouped GEMM while also - # passing the dense high-precision grad so the bias gradient avoids a - # lossy dequantize of ``grouped_dy``. - grad_input, grad_params, grad_extra_inputs = linear._fuser_backward_grouped_tensor( - ctx=linear_ctx, - grad_output=dense_dy, - grouped_grad_output=grouped_dy, - ) - - clear_tensor_data(activation_ctx.saved_tensors[0]) - return ( - grad_input, - [grad_params[0], ()], - [grad_extra_inputs[0], (grad_scales,)], - ) - - @staticmethod - def fuse_backward_ops( - ops: list[FusibleOperation], - *, - recipe: Optional[Recipe] = None, - **unused, - ) -> list[FusibleOperation]: - """Fuse each supported GroupedLinear + ScaledActivation pair.""" - out: list[FusibleOperation] = [] - idx = 0 - while idx < len(ops): - if ( - idx + 1 < len(ops) - and isinstance(ops[idx], GroupedLinear) - and isinstance(ops[idx + 1], _SCALED_ACTIVATION_TYPES) - and act_grouped_linear_fusion_supported(ops[idx], ops[idx + 1], recipe) - ): - out.append( - BackwardGroupedLinearScaledActivation( - linear=ops[idx], - activation=ops[idx + 1], - ) - ) - idx += 2 - else: - out.append(ops[idx]) - idx += 1 - return out From 28046ae91c329272a4baecc9f3b4ae43bda1b34a Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Thu, 23 Jul 2026 07:35:52 +0000 Subject: [PATCH 04/17] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- tests/pytorch/test_grouped_mlp.py | 12 +++--------- 1 file changed, 3 insertions(+), 9 deletions(-) diff --git a/tests/pytorch/test_grouped_mlp.py b/tests/pytorch/test_grouped_mlp.py index f2421ee4dd..5a7f280962 100644 --- a/tests/pytorch/test_grouped_mlp.py +++ b/tests/pytorch/test_grouped_mlp.py @@ -1209,9 +1209,7 @@ def test_grouped_mlp_single_weight_numerics( """single_grouped_weight=True/False should match exactly for fused MXFP8 grouped MLP.""" if os.environ.get("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "0") == "0": - pytest.skip( - "single_grouped_weight requires NVTE_GROUPED_LINEAR_SINGLE_PARAM=1" - ) + pytest.skip("single_grouped_weight requires NVTE_GROUPED_LINEAR_SINGLE_PARAM=1") if not te.ops.fused.GroupedMLP_CuTeGEMMGLU.is_supported(): pytest.skip("MXFP8 fused grouped MLP is not supported on this system") @@ -1531,9 +1529,7 @@ def test_grouped_mlp_overwrite_main_grad( """ if os.environ.get("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "0") == "0" and single_grouped_weight: - pytest.skip( - "single_grouped_weight requires NVTE_GROUPED_LINEAR_SINGLE_PARAM=1" - ) + pytest.skip("single_grouped_weight requires NVTE_GROUPED_LINEAR_SINGLE_PARAM=1") if not te.ops.fused.GroupedMLP_CuTeGEMMGLU.is_supported(): pytest.skip("MXFP8 fused grouped MLP is not supported on this system") @@ -1666,9 +1662,7 @@ def test_grouped_mlp_cuda_graph_safe_mxfp8( """Grouped MLP forward+backward should be CUDA graph capturable (MXFP8).""" if os.environ.get("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "0") == "0" and single_grouped_weight: - pytest.skip( - "single_grouped_weight requires NVTE_GROUPED_LINEAR_SINGLE_PARAM=1" - ) + pytest.skip("single_grouped_weight requires NVTE_GROUPED_LINEAR_SINGLE_PARAM=1") if not te.ops.fused.GroupedMLP_CuTeGEMMGLU.is_supported(): pytest.skip("MXFP8 fused grouped MLP is not supported on this system") if dtype not in (torch.bfloat16, torch.float16): From ff77e39e406066618a4686689ab1e1209afb56e9 Mon Sep 17 00:00:00 2001 From: Varun Thumbe Date: Thu, 23 Jul 2026 21:58:49 +0000 Subject: [PATCH 05/17] address review comment + cleanup Signed-off-by: Varun Thumbe --- transformer_engine/pytorch/csrc/extensions.h | 24 +++ .../pytorch/csrc/extensions/activation.cpp | 187 ++++++++++++++---- .../pytorch/csrc/extensions/pybind.cpp | 21 ++ .../pytorch/ops/basic/activation.py | 60 ++++-- .../pytorch/ops/basic/swiglu.py | 153 +++++++------- .../backward_activation_grouped_linear.py | 13 +- 6 files changed, 321 insertions(+), 137 deletions(-) diff --git a/transformer_engine/pytorch/csrc/extensions.h b/transformer_engine/pytorch/csrc/extensions.h index f479d75c13..e51c4ecc7c 100644 --- a/transformer_engine/pytorch/csrc/extensions.h +++ b/transformer_engine/pytorch/csrc/extensions.h @@ -291,6 +291,30 @@ py::object clamped_swiglu(const at::Tensor &input, py::handle quantizer, float l py::object clamped_dswiglu(const at::Tensor &grad, const at::Tensor &input, py::handle quantizer, float limit, float alpha, float glu_linear_offset); +/* Scaled activation */ +py::object scaled_swiglu(const at::Tensor &input, const at::Tensor &act_scales, + py::handle quantizer, int64_t glu_interleave_size); + +py::object scaled_clamped_swiglu(const at::Tensor &input, const at::Tensor &act_scales, + py::handle quantizer, float limit, float alpha, + float glu_linear_offset, int64_t glu_interleave_size); + +py::object scaled_srelu(const at::Tensor &input, const at::Tensor &act_scales, + py::handle quantizer); + +py::tuple scaled_dswiglu(const at::Tensor &grad, const at::Tensor &input, + const at::Tensor &act_scales, py::handle quantizer, + int64_t glu_interleave_size, bool compute_scale_grad); + +py::tuple scaled_clamped_dswiglu(const at::Tensor &grad, const at::Tensor &input, + const at::Tensor &act_scales, py::handle quantizer, float limit, + float alpha, float glu_linear_offset, int64_t glu_interleave_size, + bool compute_scale_grad); + +py::tuple scaled_dsrelu(const at::Tensor &grad, const at::Tensor &input, + const at::Tensor &act_scales, py::handle quantizer, + bool compute_scale_grad); + /* Scaled activation + grouped quantize */ py::object grouped_scaled_swiglu(const at::Tensor &input, const at::Tensor &act_scales, py::handle quantizer, const size_t num_tensors, diff --git a/transformer_engine/pytorch/csrc/extensions/activation.cpp b/transformer_engine/pytorch/csrc/extensions/activation.cpp index bce92fa3ab..4c4bf1a32a 100644 --- a/transformer_engine/pytorch/csrc/extensions/activation.cpp +++ b/transformer_engine/pytorch/csrc/extensions/activation.cpp @@ -342,29 +342,31 @@ py::object clamped_dswiglu(const at::Tensor& grad, const at::Tensor& input, py:: glu_linear_offset); } -/* Scaled activation + grouped quantize helpers (mirrors activation_helper / dactivation_helper). */ +/* Scaled activation helpers (activation + per-row scale via nvte_scaled_*). + * + * Grouped variants reuse the dense compute helpers, then optionally apply + * group_quantize. Keep the nvte_scaled_* launch path in one place. + */ template -py::object grouped_scaled_activation_helper(const at::Tensor& input, const at::Tensor& act_scales, - py::handle quantizer, const size_t num_tensors, - std::optional first_dims, - std::optional tensor_offsets, - int shape_divisor, Args&&... args) { +at::Tensor scaled_activation_compute(const at::Tensor& input, const at::Tensor& act_scales, + int shape_divisor, Args&&... args) { init_extension(); - NVTE_CHECK(input.dim() == 2, "grouped scaled activation input must be 2D"); - NVTE_CHECK(act_scales.numel() == input.size(0), - "grouped scaled activation expects one scale per input row"); - NVTE_CHECK(shape_divisor > 0 && input.size(1) % shape_divisor == 0, - "grouped scaled activation input width is not compatible with activation"); + NVTE_CHECK(input.dim() >= 1, "scaled activation input must have at least 1 dimension"); + NVTE_CHECK(shape_divisor > 0 && input.size(-1) % shape_divisor == 0, + "scaled activation input width is not compatible with activation"); auto input_tensor = input.contiguous(); - auto scales_tensor = act_scales.contiguous(); + auto scales_tensor = act_scales.contiguous().reshape({-1}); + const int64_t rows = input_tensor.numel() / input_tensor.size(-1); + NVTE_CHECK(scales_tensor.numel() == rows, "scaled activation expects one scale per input row"); + + std::vector output_sizes(input_tensor.sizes().begin(), input_tensor.sizes().end()); + output_sizes.back() /= shape_divisor; + auto output = at::empty(output_sizes, input_tensor.options()); + const TensorWrapper& input_nvte = makeTransformerEngineTensor(input_tensor); const TensorWrapper& scales_nvte = makeTransformerEngineTensor(scales_tensor); - - // Keep the dense activation in the input dtype. It is only a transient buffer when a - // quantizer is provided; the quantized return path never exposes it to the caller. - auto output = at::empty({input.size(0), input.size(1) / shape_divisor}, input_tensor.options()); const TensorWrapper& output_nvte = makeTransformerEngineTensor(output); auto stream = at::cuda::getCurrentCUDAStream(); @@ -372,40 +374,37 @@ py::object grouped_scaled_activation_helper(const at::Tensor& input, const at::T act_func(input_nvte.data(), scales_nvte.data(), output_nvte.data(), std::forward(args)..., stream); }); - - if (quantizer.is_none()) { - return py::cast(output); - } - return group_quantize(output, quantizer, num_tensors, first_dims, std::nullopt, tensor_offsets, - std::nullopt); + return output; } template -py::tuple grouped_scaled_dactivation_helper(const at::Tensor& grad, const at::Tensor& input, - const at::Tensor& act_scales, py::handle quantizer, - const size_t num_tensors, - std::optional first_dims, - std::optional tensor_offsets, - bool compute_scale_grad, Args&&... args) { +std::tuple scaled_dactivation_compute(const at::Tensor& grad, + const at::Tensor& input, + const at::Tensor& act_scales, + bool compute_scale_grad, + Args&&... args) { init_extension(); - NVTE_CHECK(input.dim() == 2 && grad.dim() == 2, - "grouped scaled dactivation input and grad must be 2D"); - NVTE_CHECK(act_scales.numel() == input.size(0), - "grouped scaled dactivation expects one scale per input row"); + NVTE_CHECK(input.dim() >= 1 && grad.dim() >= 1, + "scaled dactivation input and grad must have at least 1 dimension"); auto grad_tensor = grad.contiguous(); auto input_tensor = input.contiguous(); auto scales_tensor = act_scales.contiguous(); + const int64_t rows = input_tensor.numel() / input_tensor.size(-1); + NVTE_CHECK(scales_tensor.numel() == rows, "scaled dactivation expects one scale per input row"); + + auto scales_flat = scales_tensor.reshape({-1}); auto grad_input = at::empty_like(input_tensor); auto grad_scales = compute_scale_grad ? at::empty_like(scales_tensor) : at::Tensor(); + auto grad_scales_flat = compute_scale_grad ? grad_scales.reshape({-1}) : at::Tensor(); const TensorWrapper& grad_nvte = makeTransformerEngineTensor(grad_tensor); const TensorWrapper& input_nvte = makeTransformerEngineTensor(input_tensor); - const TensorWrapper& scales_nvte = makeTransformerEngineTensor(scales_tensor); + const TensorWrapper& scales_nvte = makeTransformerEngineTensor(scales_flat); const TensorWrapper& grad_input_nvte = makeTransformerEngineTensor(grad_input); std::optional grad_scales_nvte; if (compute_scale_grad) { - grad_scales_nvte.emplace(makeTransformerEngineTensor(grad_scales)); + grad_scales_nvte.emplace(makeTransformerEngineTensor(grad_scales_flat)); } auto stream = at::cuda::getCurrentCUDAStream(); @@ -414,17 +413,123 @@ py::tuple grouped_scaled_dactivation_helper(const at::Tensor& grad, const at::Te compute_scale_grad ? grad_scales_nvte->data() : nullptr, std::forward(args)..., stream); }); + return {grad_input, grad_scales}; +} + +py::object maybe_quantize(const at::Tensor& tensor, py::handle quantizer) { + if (quantizer.is_none()) { + return py::cast(tensor); + } + auto quantizer_cpp = convert_quantizer(quantizer); + const TensorWrapper& tensor_nvte = makeTransformerEngineTensor(tensor); + const auto shape_te = tensor_nvte.shape(); + const std::vector shape(shape_te.data, shape_te.data + shape_te.ndim); + auto fake_dtype = GetTransformerEngineDType(tensor.scalar_type()); + auto [out_nvte, out_py] = quantizer_cpp->create_tensor(shape, fake_dtype); + quantizer_cpp->quantize(tensor_nvte, out_nvte); + return out_py; +} + +py::object maybe_group_quantize(const at::Tensor& tensor, py::handle quantizer, + const size_t num_tensors, std::optional first_dims, + std::optional tensor_offsets) { + if (quantizer.is_none()) { + return py::cast(tensor); + } + return group_quantize(tensor, quantizer, num_tensors, first_dims, std::nullopt, tensor_offsets, + std::nullopt); +} + +template +py::object scaled_activation_helper(const at::Tensor& input, const at::Tensor& act_scales, + py::handle quantizer, int shape_divisor, Args&&... args) { + auto output = scaled_activation_compute(input, act_scales, shape_divisor, + std::forward(args)...); + return maybe_quantize(output, quantizer); +} + +template +py::tuple scaled_dactivation_helper(const at::Tensor& grad, const at::Tensor& input, + const at::Tensor& act_scales, py::handle quantizer, + bool compute_scale_grad, Args&&... args) { + auto [grad_input, grad_scales] = scaled_dactivation_compute( + grad, input, act_scales, compute_scale_grad, std::forward(args)...); + return py::make_tuple(maybe_quantize(grad_input, quantizer), + compute_scale_grad ? py::cast(grad_scales) : py::none()); +} + +template +py::object grouped_scaled_activation_helper(const at::Tensor& input, const at::Tensor& act_scales, + py::handle quantizer, const size_t num_tensors, + std::optional first_dims, + std::optional tensor_offsets, + int shape_divisor, Args&&... args) { + NVTE_CHECK(input.dim() == 2, "grouped scaled activation input must be 2D"); + auto output = scaled_activation_compute(input, act_scales, shape_divisor, + std::forward(args)...); + return maybe_group_quantize(output, quantizer, num_tensors, first_dims, tensor_offsets); +} +template +py::tuple grouped_scaled_dactivation_helper(const at::Tensor& grad, const at::Tensor& input, + const at::Tensor& act_scales, py::handle quantizer, + const size_t num_tensors, + std::optional first_dims, + std::optional tensor_offsets, + bool compute_scale_grad, Args&&... args) { + NVTE_CHECK(input.dim() == 2 && grad.dim() == 2, + "grouped scaled dactivation input and grad must be 2D"); + auto [grad_input, grad_scales] = scaled_dactivation_compute( + grad, input, act_scales, compute_scale_grad, std::forward(args)...); // Return both the (optionally) grouped-quantized grad input for the next // grouped GEMM and the dense high-precision grad input so callers can reuse // it (e.g. bias gradient) without a lossy dequantize. - py::object grad_input_out = py::cast(grad_input); - if (!quantizer.is_none()) { - grad_input_out = group_quantize(grad_input, quantizer, num_tensors, first_dims, std::nullopt, - tensor_offsets, std::nullopt); - } - return py::make_tuple(grad_input_out, py::cast(grad_input), - compute_scale_grad ? py::cast(grad_scales) : py::none()); + return py::make_tuple( + maybe_group_quantize(grad_input, quantizer, num_tensors, first_dims, tensor_offsets), + py::cast(grad_input), compute_scale_grad ? py::cast(grad_scales) : py::none()); +} + +py::object scaled_swiglu(const at::Tensor& input, const at::Tensor& act_scales, + py::handle quantizer, int64_t glu_interleave_size) { + return scaled_activation_helper(input, act_scales, quantizer, + /*shape_divisor=*/2, glu_interleave_size); +} + +py::object scaled_clamped_swiglu(const at::Tensor& input, const at::Tensor& act_scales, + py::handle quantizer, float limit, float alpha, + float glu_linear_offset, int64_t glu_interleave_size) { + return scaled_activation_helper( + input, act_scales, quantizer, /*shape_divisor=*/2, limit, alpha, glu_linear_offset, + glu_interleave_size); +} + +py::object scaled_srelu(const at::Tensor& input, const at::Tensor& act_scales, + py::handle quantizer) { + return scaled_activation_helper(input, act_scales, quantizer, + /*shape_divisor=*/1); +} + +py::tuple scaled_dswiglu(const at::Tensor& grad, const at::Tensor& input, + const at::Tensor& act_scales, py::handle quantizer, + int64_t glu_interleave_size, bool compute_scale_grad) { + return scaled_dactivation_helper( + grad, input, act_scales, quantizer, compute_scale_grad, glu_interleave_size); +} + +py::tuple scaled_clamped_dswiglu(const at::Tensor& grad, const at::Tensor& input, + const at::Tensor& act_scales, py::handle quantizer, float limit, + float alpha, float glu_linear_offset, int64_t glu_interleave_size, + bool compute_scale_grad) { + return scaled_dactivation_helper( + grad, input, act_scales, quantizer, compute_scale_grad, limit, alpha, glu_linear_offset, + glu_interleave_size); +} + +py::tuple scaled_dsrelu(const at::Tensor& grad, const at::Tensor& input, + const at::Tensor& act_scales, py::handle quantizer, + bool compute_scale_grad) { + return scaled_dactivation_helper(grad, input, act_scales, quantizer, + compute_scale_grad); } py::object grouped_scaled_swiglu(const at::Tensor& input, const at::Tensor& act_scales, diff --git a/transformer_engine/pytorch/csrc/extensions/pybind.cpp b/transformer_engine/pytorch/csrc/extensions/pybind.cpp index 944152f10b..fbba54fcb5 100644 --- a/transformer_engine/pytorch/csrc/extensions/pybind.cpp +++ b/transformer_engine/pytorch/csrc/extensions/pybind.cpp @@ -284,6 +284,27 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { "Backward of SwiGLU used in GPT OSS", py::arg("grad"), py::arg("fwd_input"), py::arg("quantizer"), py::arg("limit") = 7.0f, py::arg("alpha") = 1.702f, py::arg("glu_linear_offset") = 1.0f); + /* Scaled activation */ + m.def("scaled_swiglu", transformer_engine::pytorch::scaled_swiglu, "Scaled SwiGLU activation", + py::arg("input"), py::arg("act_scales"), py::arg("quantizer"), + py::arg("glu_interleave_size") = 0); + m.def("scaled_clamped_swiglu", transformer_engine::pytorch::scaled_clamped_swiglu, + "Scaled clamped SwiGLU activation", py::arg("input"), py::arg("act_scales"), + py::arg("quantizer"), py::arg("limit") = 7.0f, py::arg("alpha") = 1.702f, + py::arg("glu_linear_offset") = 1.0f, py::arg("glu_interleave_size") = 0); + m.def("scaled_srelu", transformer_engine::pytorch::scaled_srelu, "Scaled SReLU activation", + py::arg("input"), py::arg("act_scales"), py::arg("quantizer")); + m.def("scaled_dswiglu", transformer_engine::pytorch::scaled_dswiglu, "Scaled SwiGLU backward", + py::arg("grad"), py::arg("fwd_input"), py::arg("act_scales"), py::arg("quantizer"), + py::arg("glu_interleave_size") = 0, py::arg("compute_scale_grad") = true); + m.def("scaled_clamped_dswiglu", transformer_engine::pytorch::scaled_clamped_dswiglu, + "Scaled clamped SwiGLU backward", py::arg("grad"), py::arg("fwd_input"), + py::arg("act_scales"), py::arg("quantizer"), py::arg("limit") = 7.0f, + py::arg("alpha") = 1.702f, py::arg("glu_linear_offset") = 1.0f, + py::arg("glu_interleave_size") = 0, py::arg("compute_scale_grad") = true); + m.def("scaled_dsrelu", transformer_engine::pytorch::scaled_dsrelu, "Scaled SReLU backward", + py::arg("grad"), py::arg("fwd_input"), py::arg("act_scales"), py::arg("quantizer"), + py::arg("compute_scale_grad") = true); /* Scaled activation + grouped quantize */ m.def("grouped_scaled_swiglu", transformer_engine::pytorch::grouped_scaled_swiglu, "Scaled SwiGLU + grouped quantize", py::arg("input"), py::arg("act_scales"), diff --git a/transformer_engine/pytorch/ops/basic/activation.py b/transformer_engine/pytorch/ops/basic/activation.py index 42a6ef2243..3102bd6c04 100644 --- a/transformer_engine/pytorch/ops/basic/activation.py +++ b/transformer_engine/pytorch/ops/basic/activation.py @@ -358,16 +358,23 @@ def __init__(self, *, activation_recompute_in_mlp: bool = False) -> None: self.activation_recompute_in_mlp: bool = activation_recompute_in_mlp @abc.abstractmethod - def _unary_forward(self, input_: torch.Tensor) -> torch.Tensor: - """Apply the unary activation.""" + def _scaled_unary_forward( + self, + input_: torch.Tensor, + scales: torch.Tensor, + ) -> torch.Tensor: + """Apply the scaled unary activation.""" @abc.abstractmethod - def _unary_backward( + def _scaled_unary_backward( self, grad_output: torch.Tensor, input_: torch.Tensor, - ) -> torch.Tensor: - """Apply the unary activation backward pass.""" + scales: torch.Tensor, + *, + compute_scale_grad: bool, + ) -> tuple[torch.Tensor, Optional[torch.Tensor]]: + """Apply the scaled unary activation backward pass.""" def op_forward(self, *args, **kwargs) -> None: raise RuntimeError( @@ -412,7 +419,7 @@ def fuser_forward( x = maybe_dequantize(input_.contiguous(), dtype) scales = maybe_dequantize(extra_input, dtype) - y = self._unary_forward(x) * scales.unsqueeze(-1) + y = self._scaled_unary_forward(x, scales) ctx = basic_op_ctxs[0] if ctx.requires_grad: @@ -450,15 +457,14 @@ def fuser_backward( scales = maybe_dequantize(scales, ctx.dtype) grad_output = maybe_dequantize(grad_output.contiguous(), ctx.dtype) - grad_input = None - if ctx.input_requires_grad: - grad_unary_out = grad_output * scales.unsqueeze(-1) - grad_input = self._unary_backward(grad_unary_out, x) - - grad_extra_input = None - if ctx.extra_input_requires_grad: - unary_out = self._unary_forward(x) - grad_extra_input = torch.linalg.vecdot(unary_out, grad_output) + grad_input, grad_extra_input = self._scaled_unary_backward( + grad_output, + x, + scales, + compute_scale_grad=ctx.extra_input_requires_grad, + ) + if not ctx.input_requires_grad: + grad_input = None clear_tensor_data(ctx.saved_tensors[0]) @@ -478,16 +484,28 @@ class ScaledSReLU(_ScaledUnary): during backward when supported instead of saving them. """ - def _unary_forward(self, input_: torch.Tensor) -> torch.Tensor: - return tex.srelu(input_, None) - - def _unary_backward( + def _scaled_unary_forward( self, - grad_output: torch.Tensor, input_: torch.Tensor, + scales: torch.Tensor, ) -> torch.Tensor: - return tex.dsrelu(grad_output, input_, None) + return tex.scaled_srelu(input_, scales, None) + def _scaled_unary_backward( + self, + grad_output: torch.Tensor, + input_: torch.Tensor, + scales: torch.Tensor, + *, + compute_scale_grad: bool, + ) -> tuple[torch.Tensor, Optional[torch.Tensor]]: + return tex.scaled_dsrelu( + grad_output, + input_, + scales, + None, + compute_scale_grad, + ) class SReGLU(_ActivationOperation): r"""Squared Rectified Gated Linear Unit diff --git a/transformer_engine/pytorch/ops/basic/swiglu.py b/transformer_engine/pytorch/ops/basic/swiglu.py index 02f330ede3..9c3511fa04 100644 --- a/transformer_engine/pytorch/ops/basic/swiglu.py +++ b/transformer_engine/pytorch/ops/basic/swiglu.py @@ -387,14 +387,21 @@ def __init__( self.glu_interleave_size: Optional[int] = glu_interleave_size self.activation_recompute_in_mlp: bool = activation_recompute_in_mlp - def _glu_forward(self, swiglu_in: torch.Tensor) -> torch.Tensor: + def _scaled_glu_forward( + self, + input_: torch.Tensor, + scales: torch.Tensor, + ) -> torch.Tensor: raise NotImplementedError - def _glu_backward( + def _scaled_glu_backward( self, - grad_swiglu_out: torch.Tensor, - swiglu_in: torch.Tensor, - ) -> torch.Tensor: + grad_output: torch.Tensor, + input_: torch.Tensor, + scales: torch.Tensor, + *, + compute_scale_grad: bool, + ) -> tuple[torch.Tensor, Optional[torch.Tensor]]: raise NotImplementedError def op_forward(self, *args, **kwargs) -> None: @@ -442,22 +449,7 @@ def fuser_forward( # Make sure inputs are in correct dtype input_ = maybe_dequantize(input_, dtype) scales = maybe_dequantize(extra_input, dtype) - - # Remove gate interleaving if needed - swiglu_in = input_ - if self.glu_interleave_size is not None: - shape = swiglu_in.size() - swiglu_in = swiglu_in.reshape( - -1, - shape[-1] // (2 * self.glu_interleave_size), - 2, - self.glu_interleave_size, - ) - swiglu_in = swiglu_in.transpose(1, 2).contiguous() - swiglu_in = swiglu_in.view(shape) - - swiglu_out = self._glu_forward(swiglu_in) - out = swiglu_out * scales.unsqueeze(-1) + out = self._scaled_glu_forward(input_, scales) # Save state for backward pass ctx = basic_op_ctxs[0] @@ -469,7 +461,7 @@ def fuser_forward( ctx.dtype = dtype ctx.save_for_backward( input_, - scales if ctx.input_requires_grad else None, + scales if ctx.input_requires_grad or ctx.extra_input_requires_grad else None, ) return out, [()] @@ -498,41 +490,14 @@ def fuser_backward( scales = maybe_dequantize(scales, ctx.dtype) grad_output = maybe_dequantize(grad_output, ctx.dtype) - # Remove gate interleaving if needed - swiglu_in = input_ - if self.glu_interleave_size is not None: - shape = swiglu_in.size() - swiglu_in = swiglu_in.reshape( - -1, - shape[-1] // (2 * self.glu_interleave_size), - 2, - self.glu_interleave_size, - ) - swiglu_in = swiglu_in.transpose(1, 2).contiguous() - swiglu_in = swiglu_in.view(shape) - - # Compute input grad - grad_input = None - if ctx.input_requires_grad: - grad_swiglu_out = grad_output * scales.unsqueeze(-1) - grad_swiglu_in = self._glu_backward(grad_swiglu_out, swiglu_in) - grad_input = grad_swiglu_in - if self.glu_interleave_size is not None: - shape = grad_input.size() - grad_input = grad_input.reshape( - -1, - 2, - shape[-1] // (2 * self.glu_interleave_size), - self.glu_interleave_size, - ) - grad_input = grad_input.transpose(1, 2).contiguous() - grad_input = grad_input.view(shape) - - # Compute scales grad by recomputing GLU - grad_extra_input = None - if ctx.extra_input_requires_grad: - swiglu_out = self._glu_forward(swiglu_in) - grad_extra_input = torch.linalg.vecdot(swiglu_out, grad_output) + grad_input, grad_extra_input = self._scaled_glu_backward( + grad_output, + input_, + scales, + compute_scale_grad=ctx.extra_input_requires_grad, + ) + if not ctx.input_requires_grad: + grad_input = None # Clear input tensor if possible clear_tensor_data(ctx.saved_tensors[0]) # input_ @@ -558,15 +523,34 @@ class ScaledSwiGLU(_ScaledGLU): """ - def _glu_forward(self, swiglu_in: torch.Tensor) -> torch.Tensor: - return tex.swiglu(swiglu_in, None) - - def _glu_backward( + def _scaled_glu_forward( self, - grad_swiglu_out: torch.Tensor, - swiglu_in: torch.Tensor, + input_: torch.Tensor, + scales: torch.Tensor, ) -> torch.Tensor: - return tex.dswiglu(grad_swiglu_out, swiglu_in, None) + return tex.scaled_swiglu( + input_, + scales, + None, + int(self.glu_interleave_size or 0), + ) + + def _scaled_glu_backward( + self, + grad_output: torch.Tensor, + input_: torch.Tensor, + scales: torch.Tensor, + *, + compute_scale_grad: bool, + ) -> tuple[torch.Tensor, Optional[torch.Tensor]]: + return tex.scaled_dswiglu( + grad_output, + input_, + scales, + None, + int(self.glu_interleave_size or 0), + compute_scale_grad, + ) class ScaledClampedQGeGLU(_ScaledGLU): @@ -614,16 +598,39 @@ def __init__( glu_linear_offset=glu_linear_offset, ) - def _glu_forward(self, swiglu_in: torch.Tensor) -> torch.Tensor: - return self._clamped._tex_clamped_swiglu_forward(swiglu_in, None) - - def _glu_backward( + def _scaled_glu_forward( self, - grad_swiglu_out: torch.Tensor, - swiglu_in: torch.Tensor, + input_: torch.Tensor, + scales: torch.Tensor, ) -> torch.Tensor: - return self._clamped._tex_clamped_dswiglu( - grad_swiglu_out, - swiglu_in, + clamped = self._clamped + return tex.scaled_clamped_swiglu( + input_, + scales, None, + clamped.limit, + clamped.alpha, + clamped.glu_linear_offset, + int(self.glu_interleave_size or 0), ) + + def _scaled_glu_backward( + self, + grad_output: torch.Tensor, + input_: torch.Tensor, + scales: torch.Tensor, + *, + compute_scale_grad: bool, + ) -> tuple[torch.Tensor, Optional[torch.Tensor]]: + clamped = self._clamped + return tex.scaled_clamped_dswiglu( + grad_output, + input_, + scales, + None, + clamped.limit, + clamped.alpha, + clamped.glu_linear_offset, + int(self.glu_interleave_size or 0), + compute_scale_grad, + ) \ No newline at end of file diff --git a/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py b/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py index ad8b4b5a7a..579900aeed 100644 --- a/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py +++ b/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py @@ -96,14 +96,23 @@ def fuser_backward( *, basic_op_grad_extra_outputs: list[tuple[torch.Tensor, ...]], ) -> tuple[ - torch.Tensor, + Optional[torch.Tensor], Iterable[Iterable[Optional[torch.Tensor]]], Iterable[Iterable[Optional[torch.Tensor]]], ]: - del basic_op_grad_extra_outputs linear = self.basic_ops[0] activation = self.basic_ops[1] linear_ctx, activation_ctx = basic_op_ctxs + + if not linear_ctx.requires_grad: + _, _, act_grad_extra_inputs = activation.fuser_backward( + [activation_ctx], + grad_output, + basic_op_grad_extra_outputs=[basic_op_grad_extra_outputs[1]], + ) + return None, [(), ()], [(), act_grad_extra_inputs[0]] + + del basic_op_grad_extra_outputs input_, scales = activation_ctx.saved_tensors input_ = maybe_dequantize(input_, activation_ctx.dtype) scales = maybe_dequantize(scales, activation_ctx.dtype) From 59a070b774da4e3e271d9fff5f3850d660d94922 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Thu, 23 Jul 2026 21:59:50 +0000 Subject: [PATCH 06/17] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- transformer_engine/pytorch/csrc/extensions/activation.cpp | 4 ++-- transformer_engine/pytorch/ops/basic/activation.py | 1 + transformer_engine/pytorch/ops/basic/swiglu.py | 2 +- 3 files changed, 4 insertions(+), 3 deletions(-) diff --git a/transformer_engine/pytorch/csrc/extensions/activation.cpp b/transformer_engine/pytorch/csrc/extensions/activation.cpp index 4c4bf1a32a..7c486a522e 100644 --- a/transformer_engine/pytorch/csrc/extensions/activation.cpp +++ b/transformer_engine/pytorch/csrc/extensions/activation.cpp @@ -512,8 +512,8 @@ py::object scaled_srelu(const at::Tensor& input, const at::Tensor& act_scales, py::tuple scaled_dswiglu(const at::Tensor& grad, const at::Tensor& input, const at::Tensor& act_scales, py::handle quantizer, int64_t glu_interleave_size, bool compute_scale_grad) { - return scaled_dactivation_helper( - grad, input, act_scales, quantizer, compute_scale_grad, glu_interleave_size); + return scaled_dactivation_helper(grad, input, act_scales, quantizer, + compute_scale_grad, glu_interleave_size); } py::tuple scaled_clamped_dswiglu(const at::Tensor& grad, const at::Tensor& input, diff --git a/transformer_engine/pytorch/ops/basic/activation.py b/transformer_engine/pytorch/ops/basic/activation.py index 3102bd6c04..8137051322 100644 --- a/transformer_engine/pytorch/ops/basic/activation.py +++ b/transformer_engine/pytorch/ops/basic/activation.py @@ -507,6 +507,7 @@ def _scaled_unary_backward( compute_scale_grad, ) + class SReGLU(_ActivationOperation): r"""Squared Rectified Gated Linear Unit diff --git a/transformer_engine/pytorch/ops/basic/swiglu.py b/transformer_engine/pytorch/ops/basic/swiglu.py index 9c3511fa04..fb663c0480 100644 --- a/transformer_engine/pytorch/ops/basic/swiglu.py +++ b/transformer_engine/pytorch/ops/basic/swiglu.py @@ -633,4 +633,4 @@ def _scaled_glu_backward( clamped.glu_linear_offset, int(self.glu_interleave_size or 0), compute_scale_grad, - ) \ No newline at end of file + ) From 4f84ce0e085cbdec39315b82d3dd74f685ff807e Mon Sep 17 00:00:00 2001 From: Varun Thumbe Date: Fri, 24 Jul 2026 23:35:35 +0000 Subject: [PATCH 07/17] fix CI + handle activation requires_grad=false and fusion optimizations for precomputed tensor offsets in backward Signed-off-by: Varun Thumbe --- qa/L0_pytorch_unittest/test.sh | 1 + tests/pytorch/test_grouped_mlp.py | 34 ++-------- .../pytorch/ops/basic/grouped_linear.py | 68 ++++++++++++------- .../backward_activation_grouped_linear.py | 17 ++--- .../forward_activation_grouped_linear.py | 1 + .../pytorch/ops/fused/grouped_mlp.py | 49 ++++++++++--- 6 files changed, 101 insertions(+), 69 deletions(-) diff --git a/qa/L0_pytorch_unittest/test.sh b/qa/L0_pytorch_unittest/test.sh index 5d767ba4d1..726027aea6 100644 --- a/qa/L0_pytorch_unittest/test.sh +++ b/qa/L0_pytorch_unittest/test.sh @@ -69,6 +69,7 @@ NVTE_DISABLE_TRITON_AUTOTUNING=1 NVIDIA_TF32_OVERRIDE=0 python3 -m pytest --tb=a PYTORCH_JIT=0 NVTE_TORCH_COMPILE=0 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_grouped_linear.xml $TE_PATH/tests/pytorch/test_grouped_linear.py || test_fail "test_grouped_linear.py" PYTORCH_JIT=0 NVTE_TORCH_COMPILE=0 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_ops_grouped_linear_distributed_weight.xml $TE_PATH/tests/pytorch/test_ops_grouped_linear_distributed_weight.py || test_fail "test_ops_grouped_linear_distributed_weight.py" NVTE_GROUPED_LINEAR_SINGLE_PARAM=1 NVTE_CUTEDSL_FUSED_GROUPED_MLP=1 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_grouped_mlp.xml $TE_PATH/tests/pytorch/test_grouped_mlp.py || test_fail "test_grouped_mlp.py" +NVTE_GROUPED_LINEAR_SINGLE_PARAM=1 NVTE_CUTEDSL_FUSED_GROUPED_MLP=0 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_grouped_mlp_without_cutedsl_fusion.xml $TE_PATH/tests/pytorch/test_grouped_mlp.py || test_fail "test_grouped_mlp.py without CuTe DSL fusion" if [ "$RET" -ne 0 ]; then echo "Error in the following test cases:$FAILED_CASES" diff --git a/tests/pytorch/test_grouped_mlp.py b/tests/pytorch/test_grouped_mlp.py index 5a7f280962..07c9b8a673 100644 --- a/tests/pytorch/test_grouped_mlp.py +++ b/tests/pytorch/test_grouped_mlp.py @@ -78,14 +78,9 @@ if nvfp4_available: _grouped_mlp_quantization_list.append("nvfp4_rht") -# Quantization recipes for ScaledActivation + GroupedLinear fusion -_grouped_mlp_act_grouped_linear_quantization_list: list[str] = [] +# Quantization recipes supported by ScaledActivation + GroupedLinear fusion if fp8_available: - _grouped_mlp_act_grouped_linear_quantization_list.append("fp8_current_scaling") -if mxfp8_available: - _grouped_mlp_act_grouped_linear_quantization_list.append("mxfp8") -if nvfp4_available: - _grouped_mlp_act_grouped_linear_quantization_list.append("nvfp4_rht") + _grouped_mlp_quantization_list.append("fp8_current_scaling") @pytest.fixture(autouse=True, scope="function") @@ -1017,8 +1012,10 @@ def _make_module(): or ( quantization == "nvfp4_rht" and dtype == torch.bfloat16 - and activation == "scaled_srelu" - and glu_interleave_size is None + and ( + ( not activation_is_glu and glu_interleave_size is None) + or (activation_is_glu and glu_interleave_size == 32) + ) ) ) forward_ops = module._module_groups[0]._forward_ops @@ -1117,25 +1114,6 @@ def _make_module(): assert_close(fc1.weight.grad, fc1_w_ref_grad, **tols) assert_close(fc2.weight.grad, fc2_w_ref_grad, **tols) - @pytest.mark.parametrize("quantization", _grouped_mlp_act_grouped_linear_quantization_list) - def test_grouped_mlp_act_grouped_linear_fusion( - self, - *, - quantization: str, - monkeypatch: pytest.MonkeyPatch, - ) -> None: - """Exercise ScaledActivation + GroupedLinear fusion without the full CuTe DSL MLP fusion.""" - monkeypatch.setenv("NVTE_CUTEDSL_FUSED_GROUPED_MLP", "0") - self.test_grouped_mlp( - group_size=2, - bias=False, - hidden_size=128, - quantization=quantization, - single_grouped_weight=False, - split_alignment=128, - activation="scaled_swiglu", - ) - @pytest.mark.parametrize("bias", (False, True)) @pytest.mark.parametrize("quantization", _grouped_mlp_quantization_list) @pytest.mark.parametrize( diff --git a/transformer_engine/pytorch/ops/basic/grouped_linear.py b/transformer_engine/pytorch/ops/basic/grouped_linear.py index 741b229e1e..e633bade3f 100644 --- a/transformer_engine/pytorch/ops/basic/grouped_linear.py +++ b/transformer_engine/pytorch/ops/basic/grouped_linear.py @@ -1115,14 +1115,17 @@ def fuser_forward_save_ctx( # temporary workspaces freshly created in each forward pass. if is_cpu_offload_enabled(): saved = tensors_to_save[0] - offset = 4 if self._scale_bias else 3 + # Metadata prefix: + # [split_sizes, base_split_offsets, split_points, + # input_tensor_offsets, output_tensor_offsets, (scales?)] + offset = 6 if self._scale_bias else 5 if use_grouped_tensor_path: - # Layout: [split_sizes, base_split_offsets, split_points, (scales?), grouped_x, *weights] + # Layout: [..., grouped_x, *weights] grouped_x = saved[offset] if grouped_x is not None: mark_activation_offload(grouped_x) else: - # Layout: [split_sizes, None, None, (scales?), *xs, *ws] + # Layout: [..., *xs, *ws] live_xs = [t for t in saved[offset : offset + self.num_groups] if t is not None] if live_xs: mark_activation_offload(*live_xs) @@ -1148,9 +1151,9 @@ def fuser_forward_save_ctx( ctx.weight_quantizers = weight_quantizers ctx.grad_output_quantizers = grad_output_quantizers ctx.grad_input_quantizers = None - # ``split_sizes`` and ``base_split_offsets`` are routed through - # ``save_for_backward`` (see ``_fuser_forward_split_quantize`` and - # ``_fuser_forward_grouped_tensor`` for the saved-tensor layout). + # ``split_sizes``, offset metadata, and related tensors are routed + # through ``save_for_backward`` (see ``_fuser_forward_split_quantize`` + # and ``_fuser_forward_grouped_tensor`` for the saved-tensor layout). if torch.is_autocast_enabled(): ctx.dtype = torch.get_autocast_dtype("cuda") else: @@ -1264,12 +1267,13 @@ def _fuser_forward_split_quantize( # Build the tuple of tensors to save for backward. Layout: # [split_sizes, base_split_offsets, split_points, + # input_tensor_offsets, output_tensor_offsets, # (scales if scale_bias), *xs, *ws] - # ``base_split_offsets`` and ``split_points`` are unused on the - # split-quantize backward path but are included as ``None`` so the - # saved-tensor layout matches the graph-safe - # ``_fuser_forward_grouped_tensor`` path (and the fused MLP forward). - saved: list[Optional[torch.Tensor]] = [split_sizes, None, None] + # Offset metadata slots are unused on the split-quantize backward path + # but are included as ``None`` so the saved-tensor layout matches the + # graph-safe ``_fuser_forward_grouped_tensor`` path (and fused forwards + # that share this GroupedLinear context contract). + saved: list[Optional[torch.Tensor]] = [split_sizes, None, None, None, None] if self._scale_bias: saved.append(scales) saved.extend(xs) @@ -1339,6 +1343,7 @@ def _fuser_forward_graph_safe( device=device, split_points=grouped_tensor_offsets[0], base_split_offsets=grouped_tensor_offsets[1], + input_tensor_offsets=grouped_tensor_offsets[2], output_tensor_offsets=grouped_tensor_offsets[3], out_buffer=out_buffer, out_shape=original_shape[:-1] + [self.out_features], @@ -1359,6 +1364,7 @@ def _fuser_forward_grouped_tensor( device: torch.device, split_points: torch.Tensor, base_split_offsets: torch.Tensor, + input_tensor_offsets: torch.Tensor, output_tensor_offsets: torch.Tensor, out_buffer: Optional[torch.Tensor] = None, out_shape: list[int], @@ -1445,14 +1451,24 @@ def _fuser_forward_grouped_tensor( # Build the tuple of tensors to save for backward. Layout: # [split_sizes, base_split_offsets, split_points, + # input_tensor_offsets, output_tensor_offsets, # (scales if _scale_bias), grouped_x, *weights] + # ``output_tensor_offsets`` matches the linear output row layout and is + # reused as ``grad_output`` offsets in backward (including fused + # activation + grouped linear backward). if grouped_x is not None: # (For FP8 per tensor current scaling on Hopper --> Free Rowwise Data # in backward pass) if with_quantized_compute and grouped_x.columnwise_data is not None: grouped_x.rowwise_data = None grouped_x.scale_inv = None - saved: list[Optional[torch.Tensor]] = [split_sizes, base_split_offsets, split_points] + saved: list[Optional[torch.Tensor]] = [ + split_sizes, + base_split_offsets, + split_points, + input_tensor_offsets, + output_tensor_offsets, + ] if self._scale_bias: saved.append(scales) saved.append(grouped_x) @@ -1502,13 +1518,14 @@ def _fuser_backward_split_quantize( # Saved tensors from forward pass. Layout: # [split_sizes, base_split_offsets, split_points, + # input_tensor_offsets, output_tensor_offsets, # (scales if _scale_bias), *xs, *ws] - # ``base_split_offsets`` and ``split_points`` are unused on this path - # but are present so the saved-tensor layout matches the graph-safe - # path (and the fused MLP forward). + # Offset metadata beyond ``split_sizes`` is unused on this path but is + # present so the saved-tensor layout matches the graph-safe path (and + # fused forwards that share this GroupedLinear context contract). saved_tensors = ctx.saved_tensors split_sizes = saved_tensors[0] - saved_tensors = saved_tensors[3:] + saved_tensors = saved_tensors[5:] scales = None if self._scale_bias: scales, saved_tensors = saved_tensors[0], saved_tensors[1:] @@ -1687,7 +1704,7 @@ def _fuser_backward_graph_safe( dtype = ctx.dtype with_quantized_compute = bool(getattr(ctx, "with_quantized_compute", False)) split_sizes = ctx.saved_tensors[0] - base_split_offsets = ctx.saved_tensors[1] + output_tensor_offsets = ctx.saved_tensors[4] dy_2d = grad_output.reshape(-1, self.out_features) total_tokens = dy_2d.size(0) @@ -1710,6 +1727,7 @@ def _fuser_backward_graph_safe( grad_output_quantizer, num_groups, split_sizes, + tensor_offsets=output_tensor_offsets, ) else: grouped_dy = tex.group_quantize( @@ -1717,6 +1735,7 @@ def _fuser_backward_graph_safe( grad_output_quantizer, num_groups, split_sizes, + tensor_offsets=output_tensor_offsets, ) else: dy_2d = maybe_dequantize(dy_2d, dtype) @@ -1727,7 +1746,7 @@ def _fuser_backward_graph_safe( quantizer=None, data=dy_2d.reshape(-1), first_dims=split_sizes, - tensor_offsets=base_split_offsets * self.out_features, + tensor_offsets=output_tensor_offsets, ) return self._fuser_backward_grouped_tensor( @@ -1759,14 +1778,17 @@ def _fuser_backward_grouped_tensor( # Saved tensors from forward pass # Layout: [split_sizes, base_split_offsets, split_points, + # input_tensor_offsets, output_tensor_offsets, # (scales if _scale_bias), grouped_x, *weights] - # ``split_points`` is unused on this path but is present so the - # saved-tensor layout matches the fused MLP forward (which needs it - # for the cuDNN grouped GEMM kernel). + # ``split_points`` / ``output_tensor_offsets`` are unused on this path + # but are present so the saved-tensor layout matches the fused MLP / + # activation-fusion forwards that share this GroupedLinear context + # contract. saved_tensors = ctx.saved_tensors split_sizes = saved_tensors[0] base_split_offsets = saved_tensors[1] - saved_tensors = saved_tensors[3:] + input_tensor_offsets = saved_tensors[3] + saved_tensors = saved_tensors[5:] scales = None if self._scale_bias: scales, saved_tensors = saved_tensors[0], saved_tensors[1:] @@ -1829,7 +1851,7 @@ def _fuser_backward_grouped_tensor( quantizer=None, data=grad_input.reshape(-1), first_dims=split_sizes, - tensor_offsets=base_split_offsets * self.in_features, + tensor_offsets=input_tensor_offsets, ) general_grouped_gemm_for_grouped_tensor( dist_dgrad_weights if is_dist_weight else ws, diff --git a/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py b/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py index 579900aeed..34bbb34d30 100644 --- a/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py +++ b/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py @@ -112,6 +112,14 @@ def fuser_backward( ) return None, [(), ()], [(), act_grad_extra_inputs[0]] + if not activation_ctx.requires_grad: + grad_input, grad_params, grad_extra_inputs = linear.fuser_backward( + [linear_ctx], + grad_output, + basic_op_grad_extra_outputs=[basic_op_grad_extra_outputs[0]], + ) + return grad_input, [grad_params[0], ()], [grad_extra_inputs[0], (None,)] + del basic_op_grad_extra_outputs input_, scales = activation_ctx.saved_tensors input_ = maybe_dequantize(input_, activation_ctx.dtype) @@ -119,14 +127,7 @@ def fuser_backward( grad_output = maybe_dequantize(grad_output, activation_ctx.dtype) split_sizes = linear_ctx.saved_tensors[0] - split_sizes, (grad_output_tensor_offsets,) = tex.splits_to_offsets_multi( - split_sizes, - input_.device, - strides=[linear.out_features], - include_leading_zero=[True], - dtypes=[torch.int64], - bulk_allocate=False, - ) + grad_output_tensor_offsets = linear_ctx.saved_tensors[4] grad_output_quantizer = linear_ctx.grad_output_quantizers[0] grad_output_quantizer.set_usage( rowwise=linear_ctx.input_requires_grad, diff --git a/transformer_engine/pytorch/ops/fused/forward_activation_grouped_linear.py b/transformer_engine/pytorch/ops/fused/forward_activation_grouped_linear.py index f004eb943f..58bfb0d43e 100644 --- a/transformer_engine/pytorch/ops/fused/forward_activation_grouped_linear.py +++ b/transformer_engine/pytorch/ops/fused/forward_activation_grouped_linear.py @@ -182,6 +182,7 @@ def fuser_forward( device=device, split_points=grouped_tensor_offsets[0], base_split_offsets=grouped_tensor_offsets[1], + input_tensor_offsets=grouped_tensor_offsets[2], output_tensor_offsets=grouped_tensor_offsets[3], out_shape=list(input_.size())[:-1] + [linear.out_features], ) diff --git a/transformer_engine/pytorch/ops/fused/grouped_mlp.py b/transformer_engine/pytorch/ops/fused/grouped_mlp.py index 83954a9b3d..b4f6505654 100644 --- a/transformer_engine/pytorch/ops/fused/grouped_mlp.py +++ b/transformer_engine/pytorch/ops/fused/grouped_mlp.py @@ -980,14 +980,29 @@ def fuser_forward( split_points, base_split_offsets, fc1_x_tensor_offsets, + fc1_out_tensor_offsets, fc2_x_tensor_offsets, fc2_out_tensor_offsets, ) = tex.splits_to_offsets_multi( split_sizes, device, - strides=[1, 1, fc1_weight_shape[1], fc2_weight_shape[1], fc2_weight_shape[0]], - include_leading_zero=[False, True, True, True, True], - dtypes=[torch.int32, torch.int64, torch.int64, torch.int64, torch.int64], + strides=[ + 1, + 1, + fc1_weight_shape[1], + fc1_weight_shape[0], + fc2_weight_shape[1], + fc2_weight_shape[0], + ], + include_leading_zero=[False, True, True, True, True, True], + dtypes=[ + torch.int32, + torch.int64, + torch.int64, + torch.int64, + torch.int64, + torch.int64, + ], bulk_allocate=True, ) @@ -1549,6 +1564,10 @@ def fuser_forward( split_sizes, base_split_offsets, split_points, + fc1_x_tensor_offsets, + fc1_out_tensor_offsets, + fc2_x_tensor_offsets, + fc2_out_tensor_offsets, grouped_fc1_x, *fc1_weight_tensors, activation_in, @@ -1601,12 +1620,22 @@ def fuser_backward( # Saved tensors from the joint forward. # Layout: [split_sizes, base_split_offsets, split_points, + # fc1_x_tensor_offsets, fc1_out_tensor_offsets, + # fc2_x_tensor_offsets, fc2_out_tensor_offsets, # grouped_fc1_x, *fc1_weights, # activation_in, scales, # grouped_fc2_x, *fc2_weights] saved_tensors = fc1_ctx.saved_tensors - split_sizes, base_split_offsets, split_points = saved_tensors[:3] - saved_tensors = saved_tensors[3:] + ( + split_sizes, + base_split_offsets, + split_points, + fc1_x_tensor_offsets, + fc1_out_tensor_offsets, + fc2_x_tensor_offsets, + fc2_out_tensor_offsets, + ) = saved_tensors[:7] + saved_tensors = saved_tensors[7:] grouped_fc1_x, saved_tensors = saved_tensors[0], saved_tensors[1:] if fc1_op.single_grouped_weight: grouped_fc1_weight, saved_tensors = saved_tensors[0], saved_tensors[1:] @@ -1673,6 +1702,7 @@ def fuser_backward( fc2_grad_output_quantizer, num_groups, split_sizes, + tensor_offsets=fc2_out_tensor_offsets, ) else: grouped_fc2_dy = _group_quantize_for_grouped_mlp( @@ -1680,7 +1710,7 @@ def fuser_backward( fc2_grad_output_quantizer, num_groups, split_sizes, - tensor_offsets=base_split_offsets * fc2_weight_shape[0], + tensor_offsets=fc2_out_tensor_offsets, ) use_nvfp4 = ( @@ -1901,7 +1931,7 @@ def fuser_backward( fc2_input_quantizer, num_groups, split_sizes, - tensor_offsets=base_split_offsets * fc2_weight_shape[1], + tensor_offsets=fc2_x_tensor_offsets, ) else: sfd_col_d_srelu_tensor = fc2_dgrad_kernel_out.get("sfd_col_d_srelu_tensor") @@ -1923,7 +1953,7 @@ def fuser_backward( scale_inv=None, columnwise_scale_inv=fc2_x_col_scale.reshape(-1), first_dims=split_sizes, - tensor_offsets=base_split_offsets * fc2_weight_shape[1], + tensor_offsets=fc2_x_tensor_offsets, with_gemm_swizzled_scales=True, ) @@ -1966,7 +1996,7 @@ def fuser_backward( fc1_bias_grads = [dbias_2d[group_idx] for group_idx in range(num_groups)] # FC1 grad output for dgrad and wgrad GEMMs - fc1_dy_tensor_offsets = base_split_offsets * fc1_weight_shape[0] + fc1_dy_tensor_offsets = fc1_out_tensor_offsets fc1_grad_output_quantizer = fc1_ctx.grad_output_quantizers[0] if use_nvfp4: fc1_grad_output_quantizer.set_usage( @@ -2057,7 +2087,6 @@ def fuser_backward( layout="NN", ) else: - fc1_x_tensor_offsets = base_split_offsets * fc1_weight_shape[1] grouped_grad_input = GroupedTensor( shape=(out_shape[0], fc1_weight_shape[1]), dtype=dtype, From a51d67810035c917e80a522b885b3251a6686f1d Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Fri, 24 Jul 2026 23:36:36 +0000 Subject: [PATCH 08/17] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- tests/pytorch/test_grouped_mlp.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/pytorch/test_grouped_mlp.py b/tests/pytorch/test_grouped_mlp.py index 07c9b8a673..899abc36c0 100644 --- a/tests/pytorch/test_grouped_mlp.py +++ b/tests/pytorch/test_grouped_mlp.py @@ -1013,7 +1013,7 @@ def _make_module(): quantization == "nvfp4_rht" and dtype == torch.bfloat16 and ( - ( not activation_is_glu and glu_interleave_size is None) + (not activation_is_glu and glu_interleave_size is None) or (activation_is_glu and glu_interleave_size == 32) ) ) From dfa295610a77170745ec36f6df4930be26fdc300 Mon Sep 17 00:00:00 2001 From: Varun Thumbe Date: Fri, 24 Jul 2026 23:43:48 +0000 Subject: [PATCH 09/17] precomputed tensor offsets in grouped linear as well Signed-off-by: Varun Thumbe --- .../pytorch/module/grouped_linear.py | 40 ++++++++++++++----- 1 file changed, 31 insertions(+), 9 deletions(-) diff --git a/transformer_engine/pytorch/module/grouped_linear.py b/transformer_engine/pytorch/module/grouped_linear.py index 32963c07ca..7d7c87e2b9 100644 --- a/transformer_engine/pytorch/module/grouped_linear.py +++ b/transformer_engine/pytorch/module/grouped_linear.py @@ -192,7 +192,7 @@ def _make_grouped_tensor( *, num_gemms: int, split_sizes: torch.Tensor, - base_split_offsets: torch.Tensor, + tensor_offsets: torch.Tensor, last_dim: int, dtype: torch.dtype, ) -> GroupedTensorStorage: @@ -204,7 +204,7 @@ def _make_grouped_tensor( quantizer=None, data=data.reshape(-1), first_dims=split_sizes, - tensor_offsets=base_split_offsets * last_dim, + tensor_offsets=tensor_offsets, ) @staticmethod @@ -335,8 +335,18 @@ def _forward_grouped_tensor( out_features = weights[0].size(0) weight_requires_grad = weights[0].requires_grad - split_sizes = m_splits.to(device=device) - base_split_offsets = tex.splits_to_offsets(split_sizes, 1) + split_sizes, ( + base_split_offsets, + input_tensor_offsets, + output_tensor_offsets, + ) = tex.splits_to_offsets_multi( + m_splits, + device, + strides=[1, in_features, out_features], + include_leading_zero=[True, True, True], + dtypes=[torch.int64, torch.int64, torch.int64], + bulk_allocate=True, + ) inp_view = inp.reshape(-1, in_features) x = cast_if_needed(inp_view, activation_dtype) @@ -347,13 +357,19 @@ def _forward_grouped_tensor( columnwise=is_grad_enabled and weight_requires_grad, ) input_quantizer.optimize_for_gemm = True - grouped_x = tex.group_quantize(x, input_quantizer, num_gemms, split_sizes) + grouped_x = tex.group_quantize( + x, + input_quantizer, + num_gemms, + split_sizes, + tensor_offsets=input_tensor_offsets, + ) else: grouped_x = _GroupedLinear._make_grouped_tensor( x, num_gemms=num_gemms, split_sizes=split_sizes, - base_split_offsets=base_split_offsets, + tensor_offsets=input_tensor_offsets, last_dim=in_features, dtype=activation_dtype, ) @@ -382,7 +398,7 @@ def _forward_grouped_tensor( out, num_gemms=num_gemms, split_sizes=split_sizes, - base_split_offsets=base_split_offsets, + tensor_offsets=output_tensor_offsets, last_dim=out_features, dtype=activation_dtype, ) @@ -427,6 +443,8 @@ def _forward_grouped_tensor( *weights_to_save, split_sizes, base_split_offsets, + input_tensor_offsets, + output_tensor_offsets, ) ctx.save_for_backward(*tensors_to_save) ctx.tensor_objects = tensor_objects @@ -837,6 +855,8 @@ def _backward_grouped_tensor( weights = saved_tensors[1 : 1 + N] split_sizes = saved_tensors[1 + N] base_split_offsets = saved_tensors[2 + N] + input_tensor_offsets = saved_tensors[3 + N] + output_tensor_offsets = saved_tensors[4 + N] origin_weights = [None] * N main_grads = [None] * N @@ -873,6 +893,7 @@ def _backward_grouped_tensor( grad_output_quantizer, N, split_sizes, + tensor_offsets=output_tensor_offsets, ) else: grouped_dy = tex.group_quantize( @@ -880,13 +901,14 @@ def _backward_grouped_tensor( grad_output_quantizer, N, split_sizes, + tensor_offsets=output_tensor_offsets, ) else: grouped_dy = _GroupedLinear._make_grouped_tensor( dy_2d, num_gemms=N, split_sizes=split_sizes, - base_split_offsets=base_split_offsets, + tensor_offsets=output_tensor_offsets, last_dim=ctx.weights_shape_0, dtype=ctx.activation_dtype, ) @@ -918,7 +940,7 @@ def _backward_grouped_tensor( dgrad, num_gemms=N, split_sizes=split_sizes, - base_split_offsets=base_split_offsets, + tensor_offsets=input_tensor_offsets, last_dim=ctx.weights_shape_1, dtype=ctx.activation_dtype, ) From d6cccfa714ece364bba29dd5a352c8ac06e9c064 Mon Sep 17 00:00:00 2001 From: Varun Thumbe Date: Fri, 24 Jul 2026 23:49:17 +0000 Subject: [PATCH 10/17] fix lint Signed-off-by: Varun Thumbe --- .../pytorch/ops/fused/backward_activation_grouped_linear.py | 2 +- .../pytorch/ops/fused/forward_activation_grouped_linear.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py b/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py index 34bbb34d30..f50803b740 100644 --- a/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py +++ b/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py @@ -164,7 +164,7 @@ def fuse_backward_ops( ops: list[FusibleOperation], *, recipe: Optional[Recipe] = None, - **unused, + **unused, # pylint: disable=unused-argument ) -> list[FusibleOperation]: """Fuse each supported GroupedLinear + ScaledActivation pair.""" out: list[FusibleOperation] = [] diff --git a/transformer_engine/pytorch/ops/fused/forward_activation_grouped_linear.py b/transformer_engine/pytorch/ops/fused/forward_activation_grouped_linear.py index 58bfb0d43e..4f0b0d627c 100644 --- a/transformer_engine/pytorch/ops/fused/forward_activation_grouped_linear.py +++ b/transformer_engine/pytorch/ops/fused/forward_activation_grouped_linear.py @@ -204,7 +204,7 @@ def fuse_forward_ops( ops: list[FusibleOperation], *, recipe: Optional[Recipe] = None, - **unused, + **unused, # pylint: disable=unused-argument ) -> list[FusibleOperation]: """Fuse each supported ScaledActivation + GroupedLinear pair.""" out: list[FusibleOperation] = [] From 2a2f6ca324f5e4555043873671e7200b9bee85be Mon Sep 17 00:00:00 2001 From: Varun Thumbe Date: Wed, 29 Jul 2026 17:58:54 +0000 Subject: [PATCH 11/17] unecsaary based on op infra Signed-off-by: Varun Thumbe --- .../ops/fused/backward_activation_grouped_linear.py | 8 -------- 1 file changed, 8 deletions(-) diff --git a/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py b/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py index f50803b740..d31ed19b63 100644 --- a/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py +++ b/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py @@ -112,14 +112,6 @@ def fuser_backward( ) return None, [(), ()], [(), act_grad_extra_inputs[0]] - if not activation_ctx.requires_grad: - grad_input, grad_params, grad_extra_inputs = linear.fuser_backward( - [linear_ctx], - grad_output, - basic_op_grad_extra_outputs=[basic_op_grad_extra_outputs[0]], - ) - return grad_input, [grad_params[0], ()], [grad_extra_inputs[0], (None,)] - del basic_op_grad_extra_outputs input_, scales = activation_ctx.saved_tensors input_ = maybe_dequantize(input_, activation_ctx.dtype) From 63e6b404cf597035b3960331a39329337cbe13af Mon Sep 17 00:00:00 2001 From: Varun Thumbe Date: Wed, 29 Jul 2026 19:40:34 +0000 Subject: [PATCH 12/17] better name Signed-off-by: Varun Thumbe --- .../pytorch/ops/fused/backward_activation_grouped_linear.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py b/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py index d31ed19b63..0cb17fe0d1 100644 --- a/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py +++ b/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py @@ -119,7 +119,7 @@ def fuser_backward( grad_output = maybe_dequantize(grad_output, activation_ctx.dtype) split_sizes = linear_ctx.saved_tensors[0] - grad_output_tensor_offsets = linear_ctx.saved_tensors[4] + activation_input_tensor_offsets = linear_ctx.saved_tensors[4] grad_output_quantizer = linear_ctx.grad_output_quantizers[0] grad_output_quantizer.set_usage( rowwise=linear_ctx.input_requires_grad, @@ -134,7 +134,7 @@ def fuser_backward( quantizer=grad_output_quantizer, num_groups=linear.num_groups, split_sizes=split_sizes, - tensor_offsets=grad_output_tensor_offsets, + tensor_offsets=activation_input_tensor_offsets, compute_scale_grad=activation_ctx.extra_input_requires_grad, ) From 802fb46df8c628005c50afd1c914859ecee73aef Mon Sep 17 00:00:00 2001 From: Varun Thumbe Date: Mon, 3 Aug 2026 02:36:53 +0000 Subject: [PATCH 13/17] activation+group_quantize fusion via ops infra Signed-off-by: Varun Thumbe [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci ugly solution for dbias fusion Signed-off-by: Varun Thumbe [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- tests/pytorch/test_fusible_ops.py | 93 ++++++- tests/pytorch/test_grouped_mlp.py | 23 -- tests/pytorch/utils.py | 9 + .../pytorch/csrc/extensions/activation.cpp | 41 +++- transformer_engine/pytorch/ops/_common.py | 23 +- .../pytorch/ops/basic/activation.py | 139 +++++++++-- .../pytorch/ops/basic/grouped_linear.py | 202 +++++++++++---- .../pytorch/ops/basic/swiglu.py | 173 ++++++++++++- .../pytorch/ops/fused/__init__.py | 7 - .../backward_activation_grouped_linear.py | 181 -------------- .../forward_activation_grouped_linear.py | 229 ------------------ 11 files changed, 586 insertions(+), 534 deletions(-) delete mode 100644 transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py delete mode 100644 transformer_engine/pytorch/ops/fused/forward_activation_grouped_linear.py diff --git a/tests/pytorch/test_fusible_ops.py b/tests/pytorch/test_fusible_ops.py index 66857d8125..802fb2ecd6 100644 --- a/tests/pytorch/test_fusible_ops.py +++ b/tests/pytorch/test_fusible_ops.py @@ -44,6 +44,7 @@ QuantizerRole, is_bf16_available, ) +from transformer_engine.pytorch.tensor import GroupedTensorStorage # Import utility functions from utils import ( @@ -2223,6 +2224,84 @@ def test_grouped_linear( if bias: assert_close_grads(getattr(op, f"bias{group_idx}"), bs_ref[group_idx], **tols) + def test_grouped_linear_normalizes_grouped_storage(self, device: torch.device = "cuda") -> None: + """Validate grouped-input normalization and the legacy-path rejection.""" + num_groups, in_features = 2, 16 + split_sizes = torch.tensor([8, 12], dtype=torch.int64, device=device) + tensor_offsets = GroupedTensorStorage.make_tensor_offsets(split_sizes, in_features) + data = torch.randn(int(split_sizes.sum()), in_features, dtype=torch.bfloat16, device=device) + grouped_input = GroupedTensorStorage( + shape=tuple(data.shape), + dtype=data.dtype, + num_tensors=num_groups, + quantizer=None, + data=data.reshape(-1), + first_dims=split_sizes, + tensor_offsets=tensor_offsets, + ) + op = te_ops.GroupedLinear( + num_groups, in_features, in_features, bias=False, dtype=data.dtype, device=device + ) + + normalized = op._prepare_grouped_input( + grouped_input, + expected_features=in_features, + split_sizes=split_sizes, + tensor_offsets=tensor_offsets, + dtype=data.dtype, + quantizer=None, + ) + assert normalized is not grouped_input + assert normalized.rowwise_data.data_ptr() == data.data_ptr() + assert normalized.first_dims.data_ptr() == split_sizes.data_ptr() + assert normalized.tensor_offsets.data_ptr() == tensor_offsets.data_ptr() + + with pytest.raises(ValueError, match="legacy split-quantize path"): + op._fuser_forward_split_quantize( + input_=grouped_input, + split_sizes=split_sizes, + scales=None, + with_quantized_compute=False, + input_quantizers=[None] * num_groups, + weight_quantizers=[None] * num_groups, + dtype=data.dtype, + input_requires_grad=False, + weight_requires_grad=False, + device=torch.device(device), + ) + + active_quantizer = object() + prequantized = GroupedTensorStorage( + shape=tuple(data.shape), + dtype=data.dtype, + num_tensors=num_groups, + quantizer=active_quantizer, + data=data.reshape(-1), + first_dims=split_sizes, + tensor_offsets=tensor_offsets, + ) + assert ( + op._prepare_grouped_input( + prequantized, + expected_features=in_features, + split_sizes=split_sizes, + tensor_offsets=tensor_offsets, + dtype=data.dtype, + quantizer=active_quantizer, + ) + is prequantized + ) + + with pytest.raises(ValueError, match="incompatible grouped input"): + op._prepare_grouped_input( + prequantized, + expected_features=in_features, + split_sizes=split_sizes, + tensor_offsets=tensor_offsets, + dtype=data.dtype, + quantizer=object(), + ) + def test_grouped_linear_caller_buffers( self, *, @@ -2281,9 +2360,17 @@ def build() -> te_ops.Sequential: }, ) - # Forward output aliases the last op's output buffer with no copy. - assert y.data_ptr() == out_buf.data_ptr() - torch.testing.assert_close(y, y_ref, rtol=0, atol=0) + # The graph-safe path returns grouped storage. Its rowwise backing + # buffer aliases the caller-provided output with no copy. + assert isinstance(y, GroupedTensorStorage) + assert y.rowwise_data.data_ptr() == out_buf.data_ptr() + assert isinstance(y_ref, GroupedTensorStorage) + torch.testing.assert_close( + y.rowwise_data.reshape(y.logical_shape), + y_ref.rowwise_data.reshape(y_ref.logical_shape), + rtol=0, + atol=0, + ) y.backward(dy) diff --git a/tests/pytorch/test_grouped_mlp.py b/tests/pytorch/test_grouped_mlp.py index 245bc355cd..df483cabff 100644 --- a/tests/pytorch/test_grouped_mlp.py +++ b/tests/pytorch/test_grouped_mlp.py @@ -24,12 +24,6 @@ OUTPUT_BUFFER_KEY, GRAD_INPUT_BUFFER_KEY, ) -from transformer_engine.pytorch.ops.fused.backward_activation_grouped_linear import ( - BackwardScaledActivationGroupedLinear, -) -from transformer_engine.pytorch.ops.fused.forward_activation_grouped_linear import ( - ForwardScaledActivationGroupedLinear, -) from transformer_engine.pytorch import ( QuantizedTensor, Float8CurrentScalingQuantizer, @@ -1067,23 +1061,6 @@ def _make_module(): assert backward_ops[0][0] is forward_ops[0][0] full_grouped_mlp_fusion = True - # When the full FC1 + activation + FC2 fusion is unavailable, verify - # that ScaledActivation + GroupedLinear fusions cover both boundaries - # whenever grouped quantized compute is supported. - act_grouped_linear_fusion_expected = ( - not full_grouped_mlp_fusion - and te.ops.fused.act_grouped_linear_fusion_supported(fc2, module[1], recipe) - and te.ops.fused.act_grouped_linear_fusion_supported(fc1, module[1], recipe) - ) - assert ( - any(isinstance(op, ForwardScaledActivationGroupedLinear) for op, _ in forward_ops) - == act_grouped_linear_fusion_expected - ) - assert ( - any(isinstance(op, BackwardScaledActivationGroupedLinear) for op, _ in backward_ops) - == act_grouped_linear_fusion_expected - ) - # Loose tols for sanity checking tols = {"rtol": 0.125, "atol": 0.25} if quantization in ("nvfp4", "nvfp4_row_scaled", "nvfp4_4over6", "nvfp4_rht"): diff --git a/tests/pytorch/utils.py b/tests/pytorch/utils.py index 21601d8cdd..83fd4b47bb 100644 --- a/tests/pytorch/utils.py +++ b/tests/pytorch/utils.py @@ -21,6 +21,7 @@ from transformer_engine.pytorch import InferenceParams, QuantizedTensor from transformer_engine.pytorch import DType from transformer_engine.pytorch.attention.dot_product_attention import _attention_backends +from transformer_engine.pytorch.tensor import GroupedTensorStorage from transformer_engine.pytorch.attention.dot_product_attention.utils import ( get_attention_backend, AttentionParams, @@ -464,6 +465,14 @@ def assert_close( it can handle quantized tensors. """ + if isinstance(actual, GroupedTensorStorage) and actual.quantizer is None: + if actual.rowwise_data is None: + raise ValueError("Cannot compare a GroupedTensor without rowwise data") + actual = actual.rowwise_data.reshape(actual.logical_shape) + if isinstance(expected, GroupedTensorStorage) and expected.quantizer is None: + if expected.rowwise_data is None: + raise ValueError("Cannot compare a GroupedTensor without rowwise data") + expected = expected.rowwise_data.reshape(expected.logical_shape) if isinstance(actual, QuantizedTensor): actual = actual.dequantize() if isinstance(expected, QuantizedTensor): diff --git a/transformer_engine/pytorch/csrc/extensions/activation.cpp b/transformer_engine/pytorch/csrc/extensions/activation.cpp index 7c486a522e..3e0542be2e 100644 --- a/transformer_engine/pytorch/csrc/extensions/activation.cpp +++ b/transformer_engine/pytorch/csrc/extensions/activation.cpp @@ -430,14 +430,23 @@ py::object maybe_quantize(const at::Tensor& tensor, py::handle quantizer) { return out_py; } -py::object maybe_group_quantize(const at::Tensor& tensor, py::handle quantizer, - const size_t num_tensors, std::optional first_dims, - std::optional tensor_offsets) { - if (quantizer.is_none()) { - return py::cast(tensor); - } - return group_quantize(tensor, quantizer, num_tensors, first_dims, std::nullopt, tensor_offsets, - std::nullopt); +py::object make_grouped_tensor_storage(const at::Tensor& tensor, const size_t num_tensors, + std::optional first_dims, + std::optional tensor_offsets) { + py::handle grouped_tensor_storage_class( + reinterpret_cast(GroupedTensorStoragePythonClass)); + py::dict kwargs; + kwargs["shape"] = py::make_tuple(tensor.size(0), tensor.size(1)); + kwargs["dtype"] = py::cast(GetATenDType(GetTransformerEngineDType(tensor.scalar_type()))); + kwargs["num_tensors"] = py::cast(num_tensors); + kwargs["quantizer"] = py::none(); + kwargs["data"] = py::cast(tensor.reshape({-1})); + kwargs["first_dims"] = first_dims.has_value() ? py::cast(*first_dims) : py::none(); + kwargs["tensor_offsets"] = tensor_offsets.has_value() ? py::cast(*tensor_offsets) : py::none(); + PyObject* result = + PyObject_Call(grouped_tensor_storage_class.ptr(), py::tuple().ptr(), kwargs.ptr()); + NVTE_CHECK(result != nullptr, "Failed to construct GroupedTensorStorage"); + return py::reinterpret_steal(result); } template @@ -467,7 +476,11 @@ py::object grouped_scaled_activation_helper(const at::Tensor& input, const at::T NVTE_CHECK(input.dim() == 2, "grouped scaled activation input must be 2D"); auto output = scaled_activation_compute(input, act_scales, shape_divisor, std::forward(args)...); - return maybe_group_quantize(output, quantizer, num_tensors, first_dims, tensor_offsets); + if (quantizer.is_none()) { + return make_grouped_tensor_storage(output, num_tensors, first_dims, tensor_offsets); + } + return group_quantize(output, quantizer, num_tensors, first_dims, std::nullopt, tensor_offsets, + std::nullopt); } template @@ -484,9 +497,13 @@ py::tuple grouped_scaled_dactivation_helper(const at::Tensor& grad, const at::Te // Return both the (optionally) grouped-quantized grad input for the next // grouped GEMM and the dense high-precision grad input so callers can reuse // it (e.g. bias gradient) without a lossy dequantize. - return py::make_tuple( - maybe_group_quantize(grad_input, quantizer, num_tensors, first_dims, tensor_offsets), - py::cast(grad_input), compute_scale_grad ? py::cast(grad_scales) : py::none()); + auto grouped_grad_input = + quantizer.is_none() + ? make_grouped_tensor_storage(grad_input, num_tensors, first_dims, tensor_offsets) + : group_quantize(grad_input, quantizer, num_tensors, first_dims, std::nullopt, + tensor_offsets, std::nullopt); + return py::make_tuple(grouped_grad_input, py::cast(grad_input), + compute_scale_grad ? py::cast(grad_scales) : py::none()); } py::object scaled_swiglu(const at::Tensor& input, const at::Tensor& act_scales, diff --git a/transformer_engine/pytorch/ops/_common.py b/transformer_engine/pytorch/ops/_common.py index 607346ce30..219bbb506d 100644 --- a/transformer_engine/pytorch/ops/_common.py +++ b/transformer_engine/pytorch/ops/_common.py @@ -10,10 +10,13 @@ import torch +import transformer_engine_torch as tex from transformer_engine_torch import FP8TensorMeta +from ..constants import TE_DType from ..torch_version import torch_version from ..quantization import FP8GlobalStateManager from ..tensor.float8_tensor import Float8Tensor +from ..tensor.storage.grouped_tensor_storage import GroupedTensorStorage from ..quantized_tensor import QuantizedTensorStorage from ..utils import canonicalize_dtype @@ -55,9 +58,25 @@ def is_quantized_tensor(tensor: torch.Tensor | QuantizedTensorStorage) -> bool: def maybe_dequantize( tensor: torch.Tensor | QuantizedTensorStorage, dtype: torch.dtype | None = None -) -> torch.Tensor: +) -> torch.Tensor | GroupedTensorStorage: """Dequantize tensor to given dtype or just convert if not a quantized tensor""" - if is_quantized_tensor(tensor): + if isinstance(tensor, GroupedTensorStorage): + if tensor.quantizer is not None: + output_dtype = dtype if dtype is not None else tensor.fake_dtype + return tex.group_dequantize(tensor, TE_DType[output_dtype]) + if dtype is None or tensor.rowwise_data.dtype == dtype: + return tensor + return GroupedTensorStorage( + tensor.logical_shape, + dtype, + num_tensors=tensor.num_tensors, + quantizer=None, + data=tensor.rowwise_data.to(dtype=dtype), + first_dims=tensor.first_dims, + last_dims=tensor.last_dims, + tensor_offsets=tensor.tensor_offsets, + ) + elif is_quantized_tensor(tensor): return tensor.dequantize(dtype=dtype) if dtype is not None and tensor.dtype != dtype: tensor = tensor.to(dtype) diff --git a/transformer_engine/pytorch/ops/basic/activation.py b/transformer_engine/pytorch/ops/basic/activation.py index 8137051322..93341ea1be 100644 --- a/transformer_engine/pytorch/ops/basic/activation.py +++ b/transformer_engine/pytorch/ops/basic/activation.py @@ -15,6 +15,7 @@ from ...constants import DType from ...cpu_offload import is_cpu_offload_enabled, mark_activation_offload from ...tensor.float8_tensor import Float8CurrentScalingQuantizer, Quantizer +from ...tensor.storage.grouped_tensor_storage import GroupedTensorStorage from ...utils import clear_tensor_data from ..op import BasicOperation, OperationContext from .._common import maybe_dequantize @@ -362,20 +363,47 @@ def _scaled_unary_forward( self, input_: torch.Tensor, scales: torch.Tensor, + quantizer: Optional[Quantizer], ) -> torch.Tensor: """Apply the scaled unary activation.""" + @abc.abstractmethod + def _grouped_scaled_unary_forward( + self, + input_: torch.Tensor, + scales: torch.Tensor, + quantizer: Optional[Quantizer], + grouped_input: GroupedTensorStorage, + ) -> torch.Tensor: + """Apply the scaled unary activation to grouped input.""" + @abc.abstractmethod def _scaled_unary_backward( self, grad_output: torch.Tensor, input_: torch.Tensor, scales: torch.Tensor, + quantizer: Optional[Quantizer], *, compute_scale_grad: bool, ) -> tuple[torch.Tensor, Optional[torch.Tensor]]: """Apply the scaled unary activation backward pass.""" + @abc.abstractmethod + def _grouped_scaled_unary_backward( + self, + grad_output: torch.Tensor, + input_: torch.Tensor, + scales: torch.Tensor, + quantizer: Optional[Quantizer], + *, + num_groups: int, + first_dims: torch.Tensor, + tensor_offsets: Optional[torch.Tensor], + compute_scale_grad: bool, + ) -> tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]: + """Apply the scaled unary activation backward pass to grouped input.""" + def op_forward(self, *args, **kwargs) -> None: raise RuntimeError( f"{self.__class__.__name__} operation has " @@ -398,8 +426,8 @@ def fuser_forward( input_: torch.Tensor, *, basic_op_extra_inputs: list[tuple[torch.Tensor, ...]], - prev_op_grad_output_quantizer: Optional[Quantizer], # pylint: disable=unused-argument - next_op_input_quantizer: Optional[Quantizer], # pylint: disable=unused-argument + prev_op_grad_output_quantizer: Optional[Quantizer], + next_op_input_quantizer: Optional[Quantizer], basic_op_kwargs: list[dict[str, Any]], # pylint: disable=unused-argument ) -> tuple[torch.Tensor, Iterable[Iterable[torch.Tensor]]]: if self.activation_recompute_in_mlp: @@ -417,9 +445,17 @@ def fuser_forward( else: dtype = extra_input.dtype - x = maybe_dequantize(input_.contiguous(), dtype) + grouped_input = input_ if isinstance(input_, GroupedTensorStorage) else None + x = maybe_dequantize(input_, dtype) + if isinstance(x, GroupedTensorStorage): + x = x.rowwise_data.reshape(x.logical_shape) scales = maybe_dequantize(extra_input, dtype) - y = self._scaled_unary_forward(x, scales) + if grouped_input is None: + y = self._scaled_unary_forward(x, scales, next_op_input_quantizer) + else: + y = self._grouped_scaled_unary_forward( + x, scales, next_op_input_quantizer, grouped_input + ) ctx = basic_op_ctxs[0] if ctx.requires_grad: @@ -428,7 +464,11 @@ def fuser_forward( ctx.input_requires_grad = True ctx.extra_input_requires_grad = extra_input.requires_grad ctx.dtype = dtype - ctx.save_for_backward(x, scales) + ctx.save_for_backward( + grouped_input if grouped_input is not None else x, + scales, + ) + ctx.prev_op_grad_output_quantizer = prev_op_grad_output_quantizer return y, [()] @@ -452,17 +492,41 @@ def fuser_backward( ) ctx = basic_op_ctxs[0] - x, scales = ctx.saved_tensors - x = maybe_dequantize(x.contiguous(), ctx.dtype) + input_, scales = ctx.saved_tensors + grouped_input = input_ if isinstance(input_, GroupedTensorStorage) else None + first_dims = grouped_input.first_dims if grouped_input is not None else None + tensor_offsets = grouped_input.tensor_offsets if grouped_input is not None else None + x = maybe_dequantize(input_, ctx.dtype) + if isinstance(x, GroupedTensorStorage): + x = x.rowwise_data.reshape(x.logical_shape) scales = maybe_dequantize(scales, ctx.dtype) - grad_output = maybe_dequantize(grad_output.contiguous(), ctx.dtype) - - grad_input, grad_extra_input = self._scaled_unary_backward( - grad_output, - x, - scales, - compute_scale_grad=ctx.extra_input_requires_grad, - ) + grad_output = maybe_dequantize(grad_output, ctx.dtype) + if isinstance(grad_output, GroupedTensorStorage): + grad_output = grad_output.rowwise_data.reshape(grad_output.logical_shape) + + if first_dims is None: + grad_input, grad_extra_input = self._scaled_unary_backward( + grad_output, + x, + scales, + ctx.prev_op_grad_output_quantizer, + compute_scale_grad=ctx.extra_input_requires_grad, + ) + else: + grad_input, dense_grad_input, grad_extra_input = self._grouped_scaled_unary_backward( + grad_output, + x, + scales, + ctx.prev_op_grad_output_quantizer, + num_groups=int(first_dims.numel()), + first_dims=first_dims, + tensor_offsets=tensor_offsets, + compute_scale_grad=ctx.extra_input_requires_grad, + ) + # Preserve the pre-quantize result for the preceding + # GroupedLinear's dbias/dscale reduction. ``grad_input`` remains + # quantized for its dgrad and wgrad GEMMs. + grad_input._dense_for_dbias = dense_grad_input if not ctx.input_requires_grad: grad_input = None @@ -488,14 +552,32 @@ def _scaled_unary_forward( self, input_: torch.Tensor, scales: torch.Tensor, + quantizer: Optional[Quantizer], + ) -> torch.Tensor: + return tex.scaled_srelu(input_, scales, quantizer) + + def _grouped_scaled_unary_forward( + self, + input_: torch.Tensor, + scales: torch.Tensor, + quantizer: Optional[Quantizer], + grouped_input: GroupedTensorStorage, ) -> torch.Tensor: - return tex.scaled_srelu(input_, scales, None) + return tex.grouped_scaled_srelu( + input_, + scales.reshape(-1), + quantizer, + grouped_input.num_tensors, + grouped_input.first_dims, + None, + ) def _scaled_unary_backward( self, grad_output: torch.Tensor, input_: torch.Tensor, scales: torch.Tensor, + quantizer: Optional[Quantizer], *, compute_scale_grad: bool, ) -> tuple[torch.Tensor, Optional[torch.Tensor]]: @@ -503,7 +585,30 @@ def _scaled_unary_backward( grad_output, input_, scales, - None, + quantizer, + compute_scale_grad, + ) + + def _grouped_scaled_unary_backward( + self, + grad_output: torch.Tensor, + input_: torch.Tensor, + scales: torch.Tensor, + quantizer: Optional[Quantizer], + *, + num_groups: int, + first_dims: torch.Tensor, + tensor_offsets: Optional[torch.Tensor], + compute_scale_grad: bool, + ) -> tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]: + return tex.grouped_scaled_dsrelu( + grad_output, + input_, + scales.reshape(-1), + quantizer, + num_groups, + first_dims, + tensor_offsets, compute_scale_grad, ) diff --git a/transformer_engine/pytorch/ops/basic/grouped_linear.py b/transformer_engine/pytorch/ops/basic/grouped_linear.py index bc8dbed6c9..9ca15b4646 100644 --- a/transformer_engine/pytorch/ops/basic/grouped_linear.py +++ b/transformer_engine/pytorch/ops/basic/grouped_linear.py @@ -695,14 +695,25 @@ def pre_fuser_forward(self, *, requires_grad: bool) -> None: ) weight_requires_grad = requires_grad and weight_requires_grad + input_quantizers = [ + self.get_quantizer("forward", 2 * group_idx) for group_idx in range(self.num_groups) + ] + # Configure quantizer usages for group_idx in range(self.num_groups): - input_quantizer = self.get_quantizer("forward", 2 * group_idx) + input_quantizer = input_quantizers[group_idx] weight_quantizer = self.get_quantizer("forward", 2 * group_idx + 1) grad_output_quantizer = self.get_quantizer("backward", group_idx) input_quantizer.set_usage(rowwise=True, columnwise=weight_requires_grad) + # Activation fusion quantizes with this quantizer before this + # GroupedLinear's fuser_forward runs. All quantization paths + # support GEMM-optimized scale fusion, so configure it here. + input_quantizer.optimize_for_gemm = True weight_quantizer.set_usage(rowwise=True, columnwise=requires_grad) grad_output_quantizer.set_usage(rowwise=True, columnwise=weight_requires_grad) + # Activation backward quantizes with this quantizer before + # this GroupedLinear's backward path runs. + grad_output_quantizer.optimize_for_gemm = True def reset_recipe_state(self, *, recipe: Optional[Recipe]) -> None: super().reset_recipe_state(recipe=recipe) @@ -997,7 +1008,7 @@ def _get_grouped_bias_for_gemm( def fuser_forward( self, basic_op_ctxs: list[OperationContext], - input_: torch.Tensor, + input_: torch.Tensor | GroupedTensorStorage, *, basic_op_extra_inputs: list[tuple[torch.Tensor, ...]], prev_op_grad_output_quantizer: Optional[Quantizer], @@ -1191,7 +1202,7 @@ def fuser_forward_save_ctx( def _fuser_forward_split_quantize( self, *, - input_: torch.Tensor, + input_: torch.Tensor | GroupedTensorStorage, split_sizes: torch.Tensor, scales: Optional[torch.Tensor], with_quantized_compute: bool, @@ -1204,6 +1215,11 @@ def _fuser_forward_split_quantize( out_buffer: Optional[torch.Tensor] = None, ) -> tuple[torch.Tensor, tuple[Optional[torch.Tensor], ...]]: """Legacy ``tex.split_quantize`` + ``general_grouped_gemm`` flow.""" + if isinstance(input_, GroupedTensorStorage): + raise ValueError( + "GroupedLinear cannot use GroupedTensorStorage with the legacy " + "split-quantize path; use the grouped-tensor path instead." + ) num_groups = self.num_groups has_bias = self.has_bias @@ -1327,30 +1343,23 @@ def _fuser_forward_graph_safe( bulk_allocate=True, ) input_tensor_offsets = grouped_tensor_offsets[2] - original_shape = list(input_.size()) - total_tokens = math.prod(original_shape[:-1]) - x = maybe_dequantize(input_, dtype).reshape(-1, self.in_features) - if with_quantized_compute: - input_quantizer = input_quantizers[0] + original_shape = ( + list(input_.logical_shape) + if isinstance(input_, GroupedTensorStorage) + else list(input_.size()) + ) + input_quantizer = input_quantizers[0] if with_quantized_compute else None + if input_quantizer is not None: input_quantizer.set_usage(rowwise=True, columnwise=weight_requires_grad) input_quantizer.optimize_for_gemm = True - grouped_x = tex.group_quantize( - x, - input_quantizer, - num_groups, - split_sizes, - tensor_offsets=input_tensor_offsets, - ) - else: - grouped_x = GroupedTensorStorage( - shape=(total_tokens, self.in_features), - dtype=dtype, - num_tensors=num_groups, - quantizer=None, - data=x.reshape(-1), - first_dims=split_sizes, - tensor_offsets=input_tensor_offsets, - ) + grouped_x = self._prepare_grouped_input( + input_, + expected_features=self.in_features, + split_sizes=split_sizes, + tensor_offsets=input_tensor_offsets, + dtype=dtype, + quantizer=input_quantizer, + ) return self._fuser_forward_grouped_tensor( grouped_input=grouped_x, split_sizes=split_sizes, @@ -1370,6 +1379,76 @@ def _fuser_forward_graph_safe( out_shape=original_shape[:-1] + [self.out_features], ) + def _prepare_grouped_input( + self, + input_: torch.Tensor | GroupedTensorStorage, + *, + expected_features: int, + split_sizes: torch.Tensor, + tensor_offsets: torch.Tensor, + dtype: torch.dtype, + quantizer: Optional[Quantizer], + ) -> GroupedTensorStorage: + """Normalize an input to grouped storage for grouped GEMM. + + A dense input is never treated as an already-quantized grouped value. + Grouped inputs retain their quantization only when it is compatible with + the active quantizer; otherwise their rowwise data is dequantized and + converted to the requested high-precision dtype before being rewrapped. + """ + if isinstance(input_, GroupedTensorStorage): + if input_.num_tensors != self.num_groups: + raise ValueError( + f"GroupedLinear expected {self.num_groups} input groups, " + f"got {input_.num_tensors}." + ) + total_tokens, in_features = input_.logical_shape + if in_features != expected_features: + raise ValueError( + f"GroupedLinear expected input features={expected_features}, got {in_features}." + ) + input_first_dims = input_.first_dims if input_.first_dims is not None else split_sizes + input_offsets = ( + input_.tensor_offsets if input_.tensor_offsets is not None else tensor_offsets + ) + if quantizer is not None and input_.quantizer is not None: + if input_.quantizer is not quantizer: + raise ValueError( + "GroupedLinear received an incompatible grouped input " + f"(quantizer={input_.quantizer}; expected quantizer={quantizer})." + ) + return input_ + source = maybe_dequantize(input_, dtype) + assert isinstance(source, GroupedTensorStorage) + if source.rowwise_data is None: + raise RuntimeError("GroupedLinear requires rowwise data for grouped input.") + rowwise_data = source.rowwise_data.reshape(source.logical_shape) + first_dims = source.first_dims if source.first_dims is not None else input_first_dims + offsets = source.tensor_offsets if source.tensor_offsets is not None else input_offsets + else: + rowwise_data = maybe_dequantize(input_, dtype).reshape(-1, expected_features) + total_tokens = rowwise_data.size(0) + first_dims = split_sizes + offsets = tensor_offsets + + if quantizer is not None: + return tex.group_quantize( + rowwise_data, + quantizer, + self.num_groups, + first_dims, + tensor_offsets=offsets, + ) + return GroupedTensorStorage( + shape=(total_tokens, expected_features), + dtype=dtype, + num_tensors=self.num_groups, + quantizer=None, + data=rowwise_data.reshape(-1), + first_dims=first_dims, + tensor_offsets=offsets, + ) + def _fuser_forward_grouped_tensor( self, *, @@ -1428,9 +1507,11 @@ def _fuser_forward_grouped_tensor( dtype=dtype, ) - # Allocate output buffer and wrap as a GroupedTensor view. + # Allocate output buffer and wrap it in a GroupedTensor. This remains + # compatible with the fuser/autograd boundary while carrying grouped + # layout metadata to a subsequent grouped operation. out = validate_or_alloc_output(out_buffer, out_shape, dtype, device) - grouped_out = GroupedTensorStorage( + grouped_out = GroupedTensor( shape=(total_tokens, self.out_features), dtype=dtype, num_tensors=num_groups, @@ -1497,12 +1578,12 @@ def _fuser_forward_grouped_tensor( saved.append(grouped_weights) else: saved.extend(grouped_weights) - return out, tuple(saved) + return grouped_out, tuple(saved) def fuser_backward( self, basic_op_ctxs: list[OperationContext], - grad_output: torch.Tensor, + grad_output: torch.Tensor | GroupedTensorStorage, *, basic_op_grad_extra_outputs: list[tuple[torch.Tensor, ...]], ) -> tuple[ @@ -1713,7 +1794,7 @@ def _fuser_backward_graph_safe( self, *, ctx: OperationContext, - grad_output: torch.Tensor, + grad_output: torch.Tensor | GroupedTensorStorage, ) -> tuple[ torch.Tensor, Iterable[Iterable[Optional[torch.Tensor]]], @@ -1726,8 +1807,6 @@ def _fuser_backward_graph_safe( with_quantized_compute = bool(getattr(ctx, "with_quantized_compute", False)) split_sizes = ctx.saved_tensors[0] output_tensor_offsets = ctx.saved_tensors[4] - dy_2d = grad_output.reshape(-1, self.out_features) - total_tokens = dy_2d.size(0) dbias_packed = None if with_quantized_compute: @@ -1742,7 +1821,13 @@ def _fuser_backward_graph_safe( fuse_bgrad = isinstance(grad_output_quantizer, MXFP8Quantizer) or ( isinstance(grad_output_quantizer, Float8BlockQuantizer) and ctx.input_requires_grad ) - if has_bias and not self._scale_bias and fuse_bgrad: + if ( + not isinstance(grad_output, GroupedTensorStorage) + and has_bias + and not self._scale_bias + and fuse_bgrad + ): + dy_2d = grad_output.reshape(-1, self.out_features) grouped_dy, dbias_packed = tex.bgrad_group_quantize( dy_2d, grad_output_quantizer, @@ -1751,23 +1836,22 @@ def _fuser_backward_graph_safe( tensor_offsets=output_tensor_offsets, ) else: - grouped_dy = tex.group_quantize( - dy_2d, - grad_output_quantizer, - num_groups, - split_sizes, + grouped_dy = self._prepare_grouped_input( + grad_output, + expected_features=self.out_features, + split_sizes=split_sizes, tensor_offsets=output_tensor_offsets, + dtype=dtype, + quantizer=grad_output_quantizer, ) else: - dy_2d = maybe_dequantize(dy_2d, dtype) - grouped_dy = GroupedTensorStorage( - shape=(total_tokens, self.out_features), + grouped_dy = self._prepare_grouped_input( + grad_output, + expected_features=self.out_features, + split_sizes=split_sizes, + tensor_offsets=output_tensor_offsets, dtype=dtype, - num_tensors=num_groups, quantizer=None, - data=dy_2d.reshape(-1), - first_dims=split_sizes, - tensor_offsets=output_tensor_offsets, ) return self._fuser_backward_grouped_tensor( @@ -1819,11 +1903,21 @@ def _fuser_backward_grouped_tensor( else: ws, saved_tensors = saved_tensors[:num_groups], saved_tensors[num_groups:] - # Keep the dense high-precision grad for bias-gradient computation while - # optionally using a pre-quantized grouped view for the grouped GEMMs. - dy_2d = grad_output.reshape(-1, self.out_features) - total_tokens = dy_2d.size(0) - grad_input_shape = list(grad_output.size())[:-1] + [self.in_features] + # Keep grouped grad output quantized for GEMM. Materialize a dense + # high-precision view only when bias-gradient computation needs it. + dy_2d = None + if isinstance(grad_output, GroupedTensorStorage): + total_tokens, out_features = grad_output.logical_shape + if out_features != self.out_features: + raise ValueError( + f"GroupedLinear expected grad-output features={self.out_features}, " + f"got {out_features}." + ) + grad_input_shape = [total_tokens, self.in_features] + else: + dy_2d = grad_output.reshape(-1, self.out_features) + total_tokens = dy_2d.size(0) + grad_input_shape = list(grad_output.size())[:-1] + [self.in_features] expected_quantizer = ctx.grad_output_quantizers[0] if with_quantized_compute else None if grouped_grad_output.quantizer is not expected_quantizer or tuple( @@ -1842,6 +1936,12 @@ def _fuser_backward_grouped_tensor( final_bias_grads: Optional[torch.Tensor] = None grad_scales: Optional[torch.Tensor] = None if has_bias: + if dy_2d is None: + dy_2d = getattr(grad_output, "_dense_for_dbias", None) + if dy_2d is None: + dy_2d = maybe_dequantize(grad_output, dtype) + if isinstance(dy_2d, GroupedTensorStorage): + dy_2d = dy_2d.rowwise_data.reshape(dy_2d.logical_shape) if self._scale_bias: bias_packed = torch.stack(self._get_bias_tensors(dtype)) scales_f32 = scales.to(dtype=torch.float32) @@ -1858,6 +1958,8 @@ def _fuser_backward_grouped_tensor( final_bias_grads = [dbias_packed.to(dtype=dtype)] else: final_bias_grads = [dbias_packed[idx].to(dtype=dtype) for idx in range(num_groups)] + if hasattr(grad_output, "_dense_for_dbias"): + grad_output._dense_for_dbias = None # ---- dgrad GEMM ---------------------------------------------------- grad_input = None diff --git a/transformer_engine/pytorch/ops/basic/swiglu.py b/transformer_engine/pytorch/ops/basic/swiglu.py index fb663c0480..f1fe591496 100644 --- a/transformer_engine/pytorch/ops/basic/swiglu.py +++ b/transformer_engine/pytorch/ops/basic/swiglu.py @@ -14,6 +14,7 @@ from ...constants import DType from ...cpu_offload import is_cpu_offload_enabled, mark_activation_offload from ...tensor import Float8CurrentScalingQuantizer, Quantizer +from ...tensor.storage.grouped_tensor_storage import GroupedTensorStorage from ...utils import clear_tensor_data from ..op import BasicOperation, OperationContext from .._common import maybe_dequantize @@ -391,6 +392,16 @@ def _scaled_glu_forward( self, input_: torch.Tensor, scales: torch.Tensor, + quantizer: Optional[Quantizer], + ) -> torch.Tensor: + raise NotImplementedError + + def _grouped_scaled_glu_forward( + self, + input_: torch.Tensor, + scales: torch.Tensor, + quantizer: Optional[Quantizer], + grouped_input: GroupedTensorStorage, ) -> torch.Tensor: raise NotImplementedError @@ -399,11 +410,26 @@ def _scaled_glu_backward( grad_output: torch.Tensor, input_: torch.Tensor, scales: torch.Tensor, + quantizer: Optional[Quantizer], *, compute_scale_grad: bool, ) -> tuple[torch.Tensor, Optional[torch.Tensor]]: raise NotImplementedError + def _grouped_scaled_glu_backward( + self, + grad_output: torch.Tensor, + input_: torch.Tensor, + scales: torch.Tensor, + quantizer: Optional[Quantizer], + *, + num_groups: int, + first_dims: torch.Tensor, + tensor_offsets: Optional[torch.Tensor], + compute_scale_grad: bool, + ) -> tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]: + raise NotImplementedError + def op_forward(self, *args, **kwargs) -> None: raise RuntimeError( f"{self.__class__.__name__} operation has " @@ -447,9 +473,17 @@ def fuser_forward( dtype = extra_input.dtype # Make sure inputs are in correct dtype + grouped_input = input_ if isinstance(input_, GroupedTensorStorage) else None input_ = maybe_dequantize(input_, dtype) + if isinstance(input_, GroupedTensorStorage): + input_ = input_.rowwise_data.reshape(input_.logical_shape) scales = maybe_dequantize(extra_input, dtype) - out = self._scaled_glu_forward(input_, scales) + if grouped_input is None: + out = self._scaled_glu_forward(input_, scales, next_op_input_quantizer) + else: + out = self._grouped_scaled_glu_forward( + input_, scales, next_op_input_quantizer, grouped_input + ) # Save state for backward pass ctx = basic_op_ctxs[0] @@ -460,9 +494,10 @@ def fuser_forward( ctx.extra_input_requires_grad = extra_input.requires_grad ctx.dtype = dtype ctx.save_for_backward( - input_, + grouped_input if grouped_input is not None else input_, scales if ctx.input_requires_grad or ctx.extra_input_requires_grad else None, ) + ctx.prev_op_grad_output_quantizer = prev_op_grad_output_quantizer return out, [()] @@ -485,17 +520,41 @@ def fuser_backward( ctx = basic_op_ctxs[0] input_, scales = ctx.saved_tensors + grouped_input = input_ if isinstance(input_, GroupedTensorStorage) else None + first_dims = grouped_input.first_dims if grouped_input is not None else None + tensor_offsets = grouped_input.tensor_offsets if grouped_input is not None else None input_ = maybe_dequantize(input_, ctx.dtype) + if isinstance(input_, GroupedTensorStorage): + input_ = input_.rowwise_data.reshape(input_.logical_shape) if scales is not None: scales = maybe_dequantize(scales, ctx.dtype) grad_output = maybe_dequantize(grad_output, ctx.dtype) + if isinstance(grad_output, GroupedTensorStorage): + grad_output = grad_output.rowwise_data.reshape(grad_output.logical_shape) - grad_input, grad_extra_input = self._scaled_glu_backward( - grad_output, - input_, - scales, - compute_scale_grad=ctx.extra_input_requires_grad, - ) + if grouped_input is None: + grad_input, grad_extra_input = self._scaled_glu_backward( + grad_output, + input_, + scales, + ctx.prev_op_grad_output_quantizer, + compute_scale_grad=ctx.extra_input_requires_grad, + ) + else: + grad_input, dense_grad_input, grad_extra_input = self._grouped_scaled_glu_backward( + grad_output, + input_, + scales, + ctx.prev_op_grad_output_quantizer, + num_groups=int(first_dims.numel()), + first_dims=first_dims, + tensor_offsets=tensor_offsets, + compute_scale_grad=ctx.extra_input_requires_grad, + ) + # Preserve the pre-quantize result for the preceding + # GroupedLinear's dbias/dscale reduction. ``grad_input`` remains + # quantized for its dgrad and wgrad GEMMs. + grad_input._dense_for_dbias = dense_grad_input if not ctx.input_requires_grad: grad_input = None @@ -527,10 +586,28 @@ def _scaled_glu_forward( self, input_: torch.Tensor, scales: torch.Tensor, + quantizer: Optional[Quantizer], ) -> torch.Tensor: return tex.scaled_swiglu( input_, scales, + quantizer, + int(self.glu_interleave_size or 0), + ) + + def _grouped_scaled_glu_forward( + self, + input_: torch.Tensor, + scales: torch.Tensor, + quantizer: Optional[Quantizer], + grouped_input: GroupedTensorStorage, + ) -> torch.Tensor: + return tex.grouped_scaled_swiglu( + input_, + scales.reshape(-1), + quantizer, + grouped_input.num_tensors, + grouped_input.first_dims, None, int(self.glu_interleave_size or 0), ) @@ -540,6 +617,7 @@ def _scaled_glu_backward( grad_output: torch.Tensor, input_: torch.Tensor, scales: torch.Tensor, + quantizer: Optional[Quantizer], *, compute_scale_grad: bool, ) -> tuple[torch.Tensor, Optional[torch.Tensor]]: @@ -547,7 +625,31 @@ def _scaled_glu_backward( grad_output, input_, scales, - None, + quantizer, + int(self.glu_interleave_size or 0), + compute_scale_grad, + ) + + def _grouped_scaled_glu_backward( + self, + grad_output: torch.Tensor, + input_: torch.Tensor, + scales: torch.Tensor, + quantizer: Optional[Quantizer], + *, + num_groups: int, + first_dims: torch.Tensor, + tensor_offsets: Optional[torch.Tensor], + compute_scale_grad: bool, + ) -> tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]: + return tex.grouped_scaled_dswiglu( + grad_output, + input_, + scales.reshape(-1), + quantizer, + num_groups, + first_dims, + tensor_offsets, int(self.glu_interleave_size or 0), compute_scale_grad, ) @@ -602,11 +704,33 @@ def _scaled_glu_forward( self, input_: torch.Tensor, scales: torch.Tensor, + quantizer: Optional[Quantizer], ) -> torch.Tensor: clamped = self._clamped return tex.scaled_clamped_swiglu( input_, scales, + quantizer, + clamped.limit, + clamped.alpha, + clamped.glu_linear_offset, + int(self.glu_interleave_size or 0), + ) + + def _grouped_scaled_glu_forward( + self, + input_: torch.Tensor, + scales: torch.Tensor, + quantizer: Optional[Quantizer], + grouped_input: GroupedTensorStorage, + ) -> torch.Tensor: + clamped = self._clamped + return tex.grouped_scaled_clamped_swiglu( + input_, + scales.reshape(-1), + quantizer, + grouped_input.num_tensors, + grouped_input.first_dims, None, clamped.limit, clamped.alpha, @@ -619,6 +743,7 @@ def _scaled_glu_backward( grad_output: torch.Tensor, input_: torch.Tensor, scales: torch.Tensor, + quantizer: Optional[Quantizer], *, compute_scale_grad: bool, ) -> tuple[torch.Tensor, Optional[torch.Tensor]]: @@ -627,7 +752,35 @@ def _scaled_glu_backward( grad_output, input_, scales, - None, + quantizer, + clamped.limit, + clamped.alpha, + clamped.glu_linear_offset, + int(self.glu_interleave_size or 0), + compute_scale_grad, + ) + + def _grouped_scaled_glu_backward( + self, + grad_output: torch.Tensor, + input_: torch.Tensor, + scales: torch.Tensor, + quantizer: Optional[Quantizer], + *, + num_groups: int, + first_dims: torch.Tensor, + tensor_offsets: Optional[torch.Tensor], + compute_scale_grad: bool, + ) -> tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]: + clamped = self._clamped + return tex.grouped_scaled_clamped_dswiglu( + grad_output, + input_, + scales.reshape(-1), + quantizer, + num_groups, + first_dims, + tensor_offsets, clamped.limit, clamped.alpha, clamped.glu_linear_offset, diff --git a/transformer_engine/pytorch/ops/fused/__init__.py b/transformer_engine/pytorch/ops/fused/__init__.py index 85191a9807..dc9dcd6dc3 100644 --- a/transformer_engine/pytorch/ops/fused/__init__.py +++ b/transformer_engine/pytorch/ops/fused/__init__.py @@ -6,14 +6,9 @@ from ..fuser import register_backward_fusion, register_forward_fusion from .backward_activation_bias import BackwardActivationBias -from .backward_activation_grouped_linear import BackwardScaledActivationGroupedLinear from .backward_add_rmsnorm import BackwardAddRMSNorm from .backward_linear_add import BackwardLinearAdd from .backward_linear_scale import BackwardLinearScale -from .forward_activation_grouped_linear import ( - ForwardScaledActivationGroupedLinear, - act_grouped_linear_fusion_supported, -) from .forward_linear_bias_activation import ForwardLinearBiasActivation from .forward_linear_bias_add import ForwardLinearBiasAdd from .forward_linear_scale_add import ForwardLinearScaleAdd @@ -26,7 +21,6 @@ register_forward_fusion(ForwardLinearBiasAdd.fuse_forward_ops) register_forward_fusion(ForwardLinearBiasActivation.fuse_forward_ops) register_forward_fusion(ForwardLinearScaleAdd.fuse_forward_ops) -register_forward_fusion(ForwardScaledActivationGroupedLinear.fuse_forward_ops) # Register backward fusions register_backward_fusion(UserbuffersBackwardLinear.fuse_backward_ops) @@ -34,7 +28,6 @@ register_backward_fusion(BackwardLinearScale.fuse_backward_ops) register_backward_fusion(BackwardActivationBias.fuse_backward_ops) register_backward_fusion(BackwardAddRMSNorm.fuse_backward_ops) -register_backward_fusion(BackwardScaledActivationGroupedLinear.fuse_backward_ops) # Import experimental fusions # Note: Registration logic is non-trivial, so submodule handles it internally. diff --git a/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py b/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py deleted file mode 100644 index 0cb17fe0d1..0000000000 --- a/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py +++ /dev/null @@ -1,181 +0,0 @@ -# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# -# See LICENSE for license information. - -"""Fused scaled activation + grouped linear backward.""" - -from __future__ import annotations - -from collections.abc import Iterable -from typing import Optional - -import torch - -import transformer_engine_torch as tex -from ...quantization import Recipe -from ...tensor import Quantizer -from ...utils import clear_tensor_data -from .._common import maybe_dequantize -from ..basic import GroupedLinear, ScaledClampedQGeGLU, ScaledSReLU, ScaledSwiGLU -from ..op import FusedOperation, FusibleOperation, OperationContext -from .forward_activation_grouped_linear import ( - _SCALED_ACTIVATION_TYPES, - _ScaledActivation, - act_grouped_linear_fusion_supported, -) - - -def _grouped_scaled_dactivation( - activation: _ScaledActivation, - grad_output: torch.Tensor, - input_: torch.Tensor, - scales: torch.Tensor, - *, - quantizer: Quantizer, - num_groups: int, - split_sizes: torch.Tensor, - tensor_offsets: torch.Tensor, - compute_scale_grad: bool, -) -> tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]: - """Dispatch a grouped scaled activation backward pass.""" - dy = grad_output.reshape(-1, grad_output.size(-1)) - x = input_.reshape(-1, input_.size(-1)) - s = scales.reshape(-1) - if isinstance(activation, ScaledSwiGLU): - return tex.grouped_scaled_dswiglu( - dy, - x, - s, - quantizer, - num_groups, - split_sizes, - tensor_offsets, - int(activation.glu_interleave_size or 0), - compute_scale_grad, - ) - if isinstance(activation, ScaledClampedQGeGLU): - clamped = activation._clamped - return tex.grouped_scaled_clamped_dswiglu( - dy, - x, - s, - quantizer, - num_groups, - split_sizes, - tensor_offsets, - clamped.limit, - clamped.alpha, - clamped.glu_linear_offset, - int(activation.glu_interleave_size or 0), - compute_scale_grad, - ) - if isinstance(activation, ScaledSReLU): - return tex.grouped_scaled_dsrelu( - dy, - x, - s, - quantizer, - num_groups, - split_sizes, - tensor_offsets, - compute_scale_grad, - ) - raise TypeError(f"Unsupported scaled activation type ({type(activation).__name__})") - - -class BackwardScaledActivationGroupedLinear(FusedOperation): - """Scaled activation backward + grouped quantize + grouped linear backward.""" - - def __init__(self, *, linear: GroupedLinear, activation: _ScaledActivation) -> None: - super().__init__((linear, activation)) - - def fuser_backward( - self, - basic_op_ctxs: list[OperationContext], - grad_output: torch.Tensor, - *, - basic_op_grad_extra_outputs: list[tuple[torch.Tensor, ...]], - ) -> tuple[ - Optional[torch.Tensor], - Iterable[Iterable[Optional[torch.Tensor]]], - Iterable[Iterable[Optional[torch.Tensor]]], - ]: - linear = self.basic_ops[0] - activation = self.basic_ops[1] - linear_ctx, activation_ctx = basic_op_ctxs - - if not linear_ctx.requires_grad: - _, _, act_grad_extra_inputs = activation.fuser_backward( - [activation_ctx], - grad_output, - basic_op_grad_extra_outputs=[basic_op_grad_extra_outputs[1]], - ) - return None, [(), ()], [(), act_grad_extra_inputs[0]] - - del basic_op_grad_extra_outputs - input_, scales = activation_ctx.saved_tensors - input_ = maybe_dequantize(input_, activation_ctx.dtype) - scales = maybe_dequantize(scales, activation_ctx.dtype) - grad_output = maybe_dequantize(grad_output, activation_ctx.dtype) - - split_sizes = linear_ctx.saved_tensors[0] - activation_input_tensor_offsets = linear_ctx.saved_tensors[4] - grad_output_quantizer = linear_ctx.grad_output_quantizers[0] - grad_output_quantizer.set_usage( - rowwise=linear_ctx.input_requires_grad, - columnwise=linear_ctx.weight_requires_grad, - ) - grad_output_quantizer.optimize_for_gemm = True - grouped_dy, dense_dy, grad_scales = _grouped_scaled_dactivation( - activation, - grad_output, - input_, - scales, - quantizer=grad_output_quantizer, - num_groups=linear.num_groups, - split_sizes=split_sizes, - tensor_offsets=activation_input_tensor_offsets, - compute_scale_grad=activation_ctx.extra_input_requires_grad, - ) - - grad_input, grad_params, grad_extra_inputs = linear._fuser_backward_grouped_tensor( - ctx=linear_ctx, - grad_output=dense_dy, - grouped_grad_output=grouped_dy, - ) - - clear_tensor_data(activation_ctx.saved_tensors[0]) - return ( - grad_input, - [grad_params[0], ()], - [grad_extra_inputs[0], (grad_scales,)], - ) - - @staticmethod - def fuse_backward_ops( - ops: list[FusibleOperation], - *, - recipe: Optional[Recipe] = None, - **unused, # pylint: disable=unused-argument - ) -> list[FusibleOperation]: - """Fuse each supported GroupedLinear + ScaledActivation pair.""" - out: list[FusibleOperation] = [] - idx = 0 - while idx < len(ops): - if ( - idx + 1 < len(ops) - and isinstance(ops[idx], GroupedLinear) - and isinstance(ops[idx + 1], _SCALED_ACTIVATION_TYPES) - and act_grouped_linear_fusion_supported(ops[idx], ops[idx + 1], recipe) - ): - out.append( - BackwardScaledActivationGroupedLinear( - linear=ops[idx], - activation=ops[idx + 1], - ) - ) - idx += 2 - else: - out.append(ops[idx]) - idx += 1 - return out diff --git a/transformer_engine/pytorch/ops/fused/forward_activation_grouped_linear.py b/transformer_engine/pytorch/ops/fused/forward_activation_grouped_linear.py deleted file mode 100644 index 4f0b0d627c..0000000000 --- a/transformer_engine/pytorch/ops/fused/forward_activation_grouped_linear.py +++ /dev/null @@ -1,229 +0,0 @@ -# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# -# See LICENSE for license information. - -"""Fused scaled activation + grouped linear forward.""" - -from __future__ import annotations - -from collections.abc import Iterable -from typing import Any, Optional - -import torch - -import transformer_engine_torch as tex -from ...cpu_offload import is_cpu_offload_enabled, mark_activation_offload -from ...quantization import Recipe -from ...tensor import Quantizer -from .._common import maybe_dequantize -from ..basic import GroupedLinear, ScaledClampedQGeGLU, ScaledSReLU, ScaledSwiGLU -from ..basic.activation import _ScaledUnary -from ..basic.swiglu import _ScaledGLU -from ..op import FusedOperation, FusibleOperation, OperationContext - - -_ScaledActivation = _ScaledGLU | _ScaledUnary -_SCALED_ACTIVATION_TYPES = (_ScaledGLU, _ScaledUnary) - - -def _grouped_scaled_activation( - activation: _ScaledActivation, - input_: torch.Tensor, - scales: torch.Tensor, - quantizer: Quantizer, - num_groups: int, - split_sizes: torch.Tensor, - tensor_offsets: torch.Tensor, -) -> torch.Tensor: - """Dispatch a grouped scaled activation.""" - x = input_.reshape(-1, input_.size(-1)) - s = scales.reshape(-1) - if isinstance(activation, ScaledSwiGLU): - return tex.grouped_scaled_swiglu( - x, - s, - quantizer, - num_groups, - split_sizes, - tensor_offsets, - int(activation.glu_interleave_size or 0), - ) - if isinstance(activation, ScaledClampedQGeGLU): - clamped = activation._clamped - return tex.grouped_scaled_clamped_swiglu( - x, - s, - quantizer, - num_groups, - split_sizes, - tensor_offsets, - clamped.limit, - clamped.alpha, - clamped.glu_linear_offset, - int(activation.glu_interleave_size or 0), - ) - if isinstance(activation, ScaledSReLU): - return tex.grouped_scaled_srelu( - x, - s, - quantizer, - num_groups, - split_sizes, - tensor_offsets, - ) - raise TypeError(f"Unsupported scaled activation type ({type(activation).__name__})") - - -def act_grouped_linear_fusion_supported( - linear: GroupedLinear, - activation: _ScaledActivation, - recipe: Optional[Recipe], -) -> bool: - """Whether ScaledActivation + GroupedLinear can use grouped quantized compute.""" - if recipe is None or activation.activation_recompute_in_mlp: - return False - input_quantizers = [ - linear.get_quantizer("forward", 2 * group_idx) for group_idx in range(linear.num_groups) - ] - weight = linear.weight if linear.single_grouped_weight else linear.weight0 - dtype = torch.get_autocast_dtype("cuda") if torch.is_autocast_enabled() else weight.dtype - return linear._is_graph_safe_path_supported( - with_quantized_compute=True, - input_quantizers=input_quantizers, - dtype=dtype, - single_grouped_weight=linear.single_grouped_weight, - ) - - -class ForwardScaledActivationGroupedLinear(FusedOperation): - """Scaled activation + grouped quantize + grouped linear forward.""" - - def __init__(self, *, activation: _ScaledActivation, linear: GroupedLinear) -> None: - super().__init__((activation, linear)) - - def fuser_forward( - self, - basic_op_ctxs: list[OperationContext], - input_: torch.Tensor, - *, - basic_op_extra_inputs: list[tuple[torch.Tensor, ...]], - prev_op_grad_output_quantizer: Optional[Quantizer], - next_op_input_quantizer: Optional[Quantizer], - basic_op_kwargs: list[dict[str, Any]], - ) -> tuple[torch.Tensor, Iterable[Iterable[torch.Tensor]]]: - activation = self.basic_ops[0] - linear = self.basic_ops[1] - activation_ctx, linear_ctx = basic_op_ctxs - if basic_op_kwargs[0] or basic_op_kwargs[1]: - raise ValueError("Scaled activation and GroupedLinear do not expect keyword arguments") - - weight = linear.weight if linear.single_grouped_weight else linear.weight0 - device = weight.device - dtype = torch.get_autocast_dtype("cuda") if torch.is_autocast_enabled() else weight.dtype - input_ = maybe_dequantize(input_, dtype) - scales = maybe_dequantize(basic_op_extra_inputs[0][0], dtype) - - split_sizes = basic_op_extra_inputs[1][0] - if int(split_sizes.numel()) != linear.num_groups: - raise ValueError( - f"Expected {linear.num_groups} splits, but got {int(split_sizes.numel())}." - ) - split_sizes = split_sizes.to(device=device, dtype=torch.int64) - linear_scales = basic_op_extra_inputs[1][1] if linear._scale_bias else None - split_sizes, grouped_tensor_offsets = tex.splits_to_offsets_multi( - split_sizes, - device, - strides=[1, 1, linear.in_features, linear.out_features], - include_leading_zero=[False, True, True, True], - dtypes=[torch.int32, torch.int64, torch.int64, torch.int64], - bulk_allocate=True, - ) - - input_quantizers = [ - linear.get_quantizer("forward", 2 * group_idx) for group_idx in range(linear.num_groups) - ] - weight_quantizers = [ - linear.get_quantizer("forward", 2 * group_idx + 1) - for group_idx in range(linear.num_groups) - ] - input_quantizer = input_quantizers[0] - weight_requires_grad = linear_ctx.requires_grad and weight.requires_grad - input_quantizer.set_usage(rowwise=True, columnwise=weight_requires_grad) - input_quantizer.optimize_for_gemm = True - - grouped_x = _grouped_scaled_activation( - activation, - input_, - scales, - input_quantizer, - linear.num_groups, - split_sizes, - grouped_tensor_offsets[2], - ) - - if activation_ctx.requires_grad: - if is_cpu_offload_enabled(): - mark_activation_offload(input_) - activation_ctx.input_requires_grad = True - activation_ctx.extra_input_requires_grad = basic_op_extra_inputs[0][0].requires_grad - activation_ctx.dtype = dtype - activation_ctx.save_for_backward(input_, scales) - - out, tensors_to_save = linear._fuser_forward_grouped_tensor( - grouped_input=grouped_x, - split_sizes=split_sizes, - scales=linear_scales, - with_quantized_compute=True, - input_quantizers=input_quantizers, - weight_quantizers=weight_quantizers, - dtype=dtype, - input_requires_grad=linear_ctx.requires_grad, - weight_requires_grad=weight_requires_grad, - device=device, - split_points=grouped_tensor_offsets[0], - base_split_offsets=grouped_tensor_offsets[1], - input_tensor_offsets=grouped_tensor_offsets[2], - output_tensor_offsets=grouped_tensor_offsets[3], - out_shape=list(input_.size())[:-1] + [linear.out_features], - ) - linear.fuser_forward_save_ctx( - basic_op_ctxs=[linear_ctx], - input_=input_, - tensors_to_save=[tensors_to_save], - requires_grad=[linear_ctx.requires_grad], - basic_op_extra_inputs=[basic_op_extra_inputs[1]], - prev_op_grad_output_quantizer=prev_op_grad_output_quantizer, - next_op_input_quantizer=next_op_input_quantizer, - basic_op_kwargs=[basic_op_kwargs[1]], - use_grouped_tensor_path=True, - ) - return out, [(), ()] - - @staticmethod - def fuse_forward_ops( - ops: list[FusibleOperation], - *, - recipe: Optional[Recipe] = None, - **unused, # pylint: disable=unused-argument - ) -> list[FusibleOperation]: - """Fuse each supported ScaledActivation + GroupedLinear pair.""" - out: list[FusibleOperation] = [] - idx = 0 - while idx < len(ops): - if ( - idx + 1 < len(ops) - and isinstance(ops[idx], _SCALED_ACTIVATION_TYPES) - and isinstance(ops[idx + 1], GroupedLinear) - and act_grouped_linear_fusion_supported(ops[idx + 1], ops[idx], recipe) - ): - out.append( - ForwardScaledActivationGroupedLinear( - activation=ops[idx], - linear=ops[idx + 1], - ) - ) - idx += 2 - else: - out.append(ops[idx]) - idx += 1 - return out From 8dc7ce5a23b7efe73ce50f9555925151146db004 Mon Sep 17 00:00:00 2001 From: Varun Thumbe Date: Mon, 3 Aug 2026 06:59:02 +0000 Subject: [PATCH 14/17] revert to have grouped_linear + activation fusion for now Signed-off-by: Varun Thumbe --- tests/pytorch/test_fusible_ops.py | 93 +------ tests/pytorch/test_grouped_mlp.py | 23 ++ tests/pytorch/utils.py | 9 - .../pytorch/csrc/extensions/activation.cpp | 41 +--- transformer_engine/pytorch/ops/_common.py | 23 +- .../pytorch/ops/basic/activation.py | 139 ++--------- .../pytorch/ops/basic/grouped_linear.py | 202 ++++----------- .../pytorch/ops/basic/swiglu.py | 173 +------------ .../pytorch/ops/fused/__init__.py | 7 + .../backward_activation_grouped_linear.py | 181 ++++++++++++++ .../forward_activation_grouped_linear.py | 229 ++++++++++++++++++ 11 files changed, 534 insertions(+), 586 deletions(-) create mode 100644 transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py create mode 100644 transformer_engine/pytorch/ops/fused/forward_activation_grouped_linear.py diff --git a/tests/pytorch/test_fusible_ops.py b/tests/pytorch/test_fusible_ops.py index 802fb2ecd6..66857d8125 100644 --- a/tests/pytorch/test_fusible_ops.py +++ b/tests/pytorch/test_fusible_ops.py @@ -44,7 +44,6 @@ QuantizerRole, is_bf16_available, ) -from transformer_engine.pytorch.tensor import GroupedTensorStorage # Import utility functions from utils import ( @@ -2224,84 +2223,6 @@ def test_grouped_linear( if bias: assert_close_grads(getattr(op, f"bias{group_idx}"), bs_ref[group_idx], **tols) - def test_grouped_linear_normalizes_grouped_storage(self, device: torch.device = "cuda") -> None: - """Validate grouped-input normalization and the legacy-path rejection.""" - num_groups, in_features = 2, 16 - split_sizes = torch.tensor([8, 12], dtype=torch.int64, device=device) - tensor_offsets = GroupedTensorStorage.make_tensor_offsets(split_sizes, in_features) - data = torch.randn(int(split_sizes.sum()), in_features, dtype=torch.bfloat16, device=device) - grouped_input = GroupedTensorStorage( - shape=tuple(data.shape), - dtype=data.dtype, - num_tensors=num_groups, - quantizer=None, - data=data.reshape(-1), - first_dims=split_sizes, - tensor_offsets=tensor_offsets, - ) - op = te_ops.GroupedLinear( - num_groups, in_features, in_features, bias=False, dtype=data.dtype, device=device - ) - - normalized = op._prepare_grouped_input( - grouped_input, - expected_features=in_features, - split_sizes=split_sizes, - tensor_offsets=tensor_offsets, - dtype=data.dtype, - quantizer=None, - ) - assert normalized is not grouped_input - assert normalized.rowwise_data.data_ptr() == data.data_ptr() - assert normalized.first_dims.data_ptr() == split_sizes.data_ptr() - assert normalized.tensor_offsets.data_ptr() == tensor_offsets.data_ptr() - - with pytest.raises(ValueError, match="legacy split-quantize path"): - op._fuser_forward_split_quantize( - input_=grouped_input, - split_sizes=split_sizes, - scales=None, - with_quantized_compute=False, - input_quantizers=[None] * num_groups, - weight_quantizers=[None] * num_groups, - dtype=data.dtype, - input_requires_grad=False, - weight_requires_grad=False, - device=torch.device(device), - ) - - active_quantizer = object() - prequantized = GroupedTensorStorage( - shape=tuple(data.shape), - dtype=data.dtype, - num_tensors=num_groups, - quantizer=active_quantizer, - data=data.reshape(-1), - first_dims=split_sizes, - tensor_offsets=tensor_offsets, - ) - assert ( - op._prepare_grouped_input( - prequantized, - expected_features=in_features, - split_sizes=split_sizes, - tensor_offsets=tensor_offsets, - dtype=data.dtype, - quantizer=active_quantizer, - ) - is prequantized - ) - - with pytest.raises(ValueError, match="incompatible grouped input"): - op._prepare_grouped_input( - prequantized, - expected_features=in_features, - split_sizes=split_sizes, - tensor_offsets=tensor_offsets, - dtype=data.dtype, - quantizer=object(), - ) - def test_grouped_linear_caller_buffers( self, *, @@ -2360,17 +2281,9 @@ def build() -> te_ops.Sequential: }, ) - # The graph-safe path returns grouped storage. Its rowwise backing - # buffer aliases the caller-provided output with no copy. - assert isinstance(y, GroupedTensorStorage) - assert y.rowwise_data.data_ptr() == out_buf.data_ptr() - assert isinstance(y_ref, GroupedTensorStorage) - torch.testing.assert_close( - y.rowwise_data.reshape(y.logical_shape), - y_ref.rowwise_data.reshape(y_ref.logical_shape), - rtol=0, - atol=0, - ) + # Forward output aliases the last op's output buffer with no copy. + assert y.data_ptr() == out_buf.data_ptr() + torch.testing.assert_close(y, y_ref, rtol=0, atol=0) y.backward(dy) diff --git a/tests/pytorch/test_grouped_mlp.py b/tests/pytorch/test_grouped_mlp.py index df483cabff..245bc355cd 100644 --- a/tests/pytorch/test_grouped_mlp.py +++ b/tests/pytorch/test_grouped_mlp.py @@ -24,6 +24,12 @@ OUTPUT_BUFFER_KEY, GRAD_INPUT_BUFFER_KEY, ) +from transformer_engine.pytorch.ops.fused.backward_activation_grouped_linear import ( + BackwardScaledActivationGroupedLinear, +) +from transformer_engine.pytorch.ops.fused.forward_activation_grouped_linear import ( + ForwardScaledActivationGroupedLinear, +) from transformer_engine.pytorch import ( QuantizedTensor, Float8CurrentScalingQuantizer, @@ -1061,6 +1067,23 @@ def _make_module(): assert backward_ops[0][0] is forward_ops[0][0] full_grouped_mlp_fusion = True + # When the full FC1 + activation + FC2 fusion is unavailable, verify + # that ScaledActivation + GroupedLinear fusions cover both boundaries + # whenever grouped quantized compute is supported. + act_grouped_linear_fusion_expected = ( + not full_grouped_mlp_fusion + and te.ops.fused.act_grouped_linear_fusion_supported(fc2, module[1], recipe) + and te.ops.fused.act_grouped_linear_fusion_supported(fc1, module[1], recipe) + ) + assert ( + any(isinstance(op, ForwardScaledActivationGroupedLinear) for op, _ in forward_ops) + == act_grouped_linear_fusion_expected + ) + assert ( + any(isinstance(op, BackwardScaledActivationGroupedLinear) for op, _ in backward_ops) + == act_grouped_linear_fusion_expected + ) + # Loose tols for sanity checking tols = {"rtol": 0.125, "atol": 0.25} if quantization in ("nvfp4", "nvfp4_row_scaled", "nvfp4_4over6", "nvfp4_rht"): diff --git a/tests/pytorch/utils.py b/tests/pytorch/utils.py index 83fd4b47bb..21601d8cdd 100644 --- a/tests/pytorch/utils.py +++ b/tests/pytorch/utils.py @@ -21,7 +21,6 @@ from transformer_engine.pytorch import InferenceParams, QuantizedTensor from transformer_engine.pytorch import DType from transformer_engine.pytorch.attention.dot_product_attention import _attention_backends -from transformer_engine.pytorch.tensor import GroupedTensorStorage from transformer_engine.pytorch.attention.dot_product_attention.utils import ( get_attention_backend, AttentionParams, @@ -465,14 +464,6 @@ def assert_close( it can handle quantized tensors. """ - if isinstance(actual, GroupedTensorStorage) and actual.quantizer is None: - if actual.rowwise_data is None: - raise ValueError("Cannot compare a GroupedTensor without rowwise data") - actual = actual.rowwise_data.reshape(actual.logical_shape) - if isinstance(expected, GroupedTensorStorage) and expected.quantizer is None: - if expected.rowwise_data is None: - raise ValueError("Cannot compare a GroupedTensor without rowwise data") - expected = expected.rowwise_data.reshape(expected.logical_shape) if isinstance(actual, QuantizedTensor): actual = actual.dequantize() if isinstance(expected, QuantizedTensor): diff --git a/transformer_engine/pytorch/csrc/extensions/activation.cpp b/transformer_engine/pytorch/csrc/extensions/activation.cpp index 3e0542be2e..7c486a522e 100644 --- a/transformer_engine/pytorch/csrc/extensions/activation.cpp +++ b/transformer_engine/pytorch/csrc/extensions/activation.cpp @@ -430,23 +430,14 @@ py::object maybe_quantize(const at::Tensor& tensor, py::handle quantizer) { return out_py; } -py::object make_grouped_tensor_storage(const at::Tensor& tensor, const size_t num_tensors, - std::optional first_dims, - std::optional tensor_offsets) { - py::handle grouped_tensor_storage_class( - reinterpret_cast(GroupedTensorStoragePythonClass)); - py::dict kwargs; - kwargs["shape"] = py::make_tuple(tensor.size(0), tensor.size(1)); - kwargs["dtype"] = py::cast(GetATenDType(GetTransformerEngineDType(tensor.scalar_type()))); - kwargs["num_tensors"] = py::cast(num_tensors); - kwargs["quantizer"] = py::none(); - kwargs["data"] = py::cast(tensor.reshape({-1})); - kwargs["first_dims"] = first_dims.has_value() ? py::cast(*first_dims) : py::none(); - kwargs["tensor_offsets"] = tensor_offsets.has_value() ? py::cast(*tensor_offsets) : py::none(); - PyObject* result = - PyObject_Call(grouped_tensor_storage_class.ptr(), py::tuple().ptr(), kwargs.ptr()); - NVTE_CHECK(result != nullptr, "Failed to construct GroupedTensorStorage"); - return py::reinterpret_steal(result); +py::object maybe_group_quantize(const at::Tensor& tensor, py::handle quantizer, + const size_t num_tensors, std::optional first_dims, + std::optional tensor_offsets) { + if (quantizer.is_none()) { + return py::cast(tensor); + } + return group_quantize(tensor, quantizer, num_tensors, first_dims, std::nullopt, tensor_offsets, + std::nullopt); } template @@ -476,11 +467,7 @@ py::object grouped_scaled_activation_helper(const at::Tensor& input, const at::T NVTE_CHECK(input.dim() == 2, "grouped scaled activation input must be 2D"); auto output = scaled_activation_compute(input, act_scales, shape_divisor, std::forward(args)...); - if (quantizer.is_none()) { - return make_grouped_tensor_storage(output, num_tensors, first_dims, tensor_offsets); - } - return group_quantize(output, quantizer, num_tensors, first_dims, std::nullopt, tensor_offsets, - std::nullopt); + return maybe_group_quantize(output, quantizer, num_tensors, first_dims, tensor_offsets); } template @@ -497,13 +484,9 @@ py::tuple grouped_scaled_dactivation_helper(const at::Tensor& grad, const at::Te // Return both the (optionally) grouped-quantized grad input for the next // grouped GEMM and the dense high-precision grad input so callers can reuse // it (e.g. bias gradient) without a lossy dequantize. - auto grouped_grad_input = - quantizer.is_none() - ? make_grouped_tensor_storage(grad_input, num_tensors, first_dims, tensor_offsets) - : group_quantize(grad_input, quantizer, num_tensors, first_dims, std::nullopt, - tensor_offsets, std::nullopt); - return py::make_tuple(grouped_grad_input, py::cast(grad_input), - compute_scale_grad ? py::cast(grad_scales) : py::none()); + return py::make_tuple( + maybe_group_quantize(grad_input, quantizer, num_tensors, first_dims, tensor_offsets), + py::cast(grad_input), compute_scale_grad ? py::cast(grad_scales) : py::none()); } py::object scaled_swiglu(const at::Tensor& input, const at::Tensor& act_scales, diff --git a/transformer_engine/pytorch/ops/_common.py b/transformer_engine/pytorch/ops/_common.py index 219bbb506d..607346ce30 100644 --- a/transformer_engine/pytorch/ops/_common.py +++ b/transformer_engine/pytorch/ops/_common.py @@ -10,13 +10,10 @@ import torch -import transformer_engine_torch as tex from transformer_engine_torch import FP8TensorMeta -from ..constants import TE_DType from ..torch_version import torch_version from ..quantization import FP8GlobalStateManager from ..tensor.float8_tensor import Float8Tensor -from ..tensor.storage.grouped_tensor_storage import GroupedTensorStorage from ..quantized_tensor import QuantizedTensorStorage from ..utils import canonicalize_dtype @@ -58,25 +55,9 @@ def is_quantized_tensor(tensor: torch.Tensor | QuantizedTensorStorage) -> bool: def maybe_dequantize( tensor: torch.Tensor | QuantizedTensorStorage, dtype: torch.dtype | None = None -) -> torch.Tensor | GroupedTensorStorage: +) -> torch.Tensor: """Dequantize tensor to given dtype or just convert if not a quantized tensor""" - if isinstance(tensor, GroupedTensorStorage): - if tensor.quantizer is not None: - output_dtype = dtype if dtype is not None else tensor.fake_dtype - return tex.group_dequantize(tensor, TE_DType[output_dtype]) - if dtype is None or tensor.rowwise_data.dtype == dtype: - return tensor - return GroupedTensorStorage( - tensor.logical_shape, - dtype, - num_tensors=tensor.num_tensors, - quantizer=None, - data=tensor.rowwise_data.to(dtype=dtype), - first_dims=tensor.first_dims, - last_dims=tensor.last_dims, - tensor_offsets=tensor.tensor_offsets, - ) - elif is_quantized_tensor(tensor): + if is_quantized_tensor(tensor): return tensor.dequantize(dtype=dtype) if dtype is not None and tensor.dtype != dtype: tensor = tensor.to(dtype) diff --git a/transformer_engine/pytorch/ops/basic/activation.py b/transformer_engine/pytorch/ops/basic/activation.py index 93341ea1be..8137051322 100644 --- a/transformer_engine/pytorch/ops/basic/activation.py +++ b/transformer_engine/pytorch/ops/basic/activation.py @@ -15,7 +15,6 @@ from ...constants import DType from ...cpu_offload import is_cpu_offload_enabled, mark_activation_offload from ...tensor.float8_tensor import Float8CurrentScalingQuantizer, Quantizer -from ...tensor.storage.grouped_tensor_storage import GroupedTensorStorage from ...utils import clear_tensor_data from ..op import BasicOperation, OperationContext from .._common import maybe_dequantize @@ -363,47 +362,20 @@ def _scaled_unary_forward( self, input_: torch.Tensor, scales: torch.Tensor, - quantizer: Optional[Quantizer], ) -> torch.Tensor: """Apply the scaled unary activation.""" - @abc.abstractmethod - def _grouped_scaled_unary_forward( - self, - input_: torch.Tensor, - scales: torch.Tensor, - quantizer: Optional[Quantizer], - grouped_input: GroupedTensorStorage, - ) -> torch.Tensor: - """Apply the scaled unary activation to grouped input.""" - @abc.abstractmethod def _scaled_unary_backward( self, grad_output: torch.Tensor, input_: torch.Tensor, scales: torch.Tensor, - quantizer: Optional[Quantizer], *, compute_scale_grad: bool, ) -> tuple[torch.Tensor, Optional[torch.Tensor]]: """Apply the scaled unary activation backward pass.""" - @abc.abstractmethod - def _grouped_scaled_unary_backward( - self, - grad_output: torch.Tensor, - input_: torch.Tensor, - scales: torch.Tensor, - quantizer: Optional[Quantizer], - *, - num_groups: int, - first_dims: torch.Tensor, - tensor_offsets: Optional[torch.Tensor], - compute_scale_grad: bool, - ) -> tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]: - """Apply the scaled unary activation backward pass to grouped input.""" - def op_forward(self, *args, **kwargs) -> None: raise RuntimeError( f"{self.__class__.__name__} operation has " @@ -426,8 +398,8 @@ def fuser_forward( input_: torch.Tensor, *, basic_op_extra_inputs: list[tuple[torch.Tensor, ...]], - prev_op_grad_output_quantizer: Optional[Quantizer], - next_op_input_quantizer: Optional[Quantizer], + prev_op_grad_output_quantizer: Optional[Quantizer], # pylint: disable=unused-argument + next_op_input_quantizer: Optional[Quantizer], # pylint: disable=unused-argument basic_op_kwargs: list[dict[str, Any]], # pylint: disable=unused-argument ) -> tuple[torch.Tensor, Iterable[Iterable[torch.Tensor]]]: if self.activation_recompute_in_mlp: @@ -445,17 +417,9 @@ def fuser_forward( else: dtype = extra_input.dtype - grouped_input = input_ if isinstance(input_, GroupedTensorStorage) else None - x = maybe_dequantize(input_, dtype) - if isinstance(x, GroupedTensorStorage): - x = x.rowwise_data.reshape(x.logical_shape) + x = maybe_dequantize(input_.contiguous(), dtype) scales = maybe_dequantize(extra_input, dtype) - if grouped_input is None: - y = self._scaled_unary_forward(x, scales, next_op_input_quantizer) - else: - y = self._grouped_scaled_unary_forward( - x, scales, next_op_input_quantizer, grouped_input - ) + y = self._scaled_unary_forward(x, scales) ctx = basic_op_ctxs[0] if ctx.requires_grad: @@ -464,11 +428,7 @@ def fuser_forward( ctx.input_requires_grad = True ctx.extra_input_requires_grad = extra_input.requires_grad ctx.dtype = dtype - ctx.save_for_backward( - grouped_input if grouped_input is not None else x, - scales, - ) - ctx.prev_op_grad_output_quantizer = prev_op_grad_output_quantizer + ctx.save_for_backward(x, scales) return y, [()] @@ -492,41 +452,17 @@ def fuser_backward( ) ctx = basic_op_ctxs[0] - input_, scales = ctx.saved_tensors - grouped_input = input_ if isinstance(input_, GroupedTensorStorage) else None - first_dims = grouped_input.first_dims if grouped_input is not None else None - tensor_offsets = grouped_input.tensor_offsets if grouped_input is not None else None - x = maybe_dequantize(input_, ctx.dtype) - if isinstance(x, GroupedTensorStorage): - x = x.rowwise_data.reshape(x.logical_shape) + x, scales = ctx.saved_tensors + x = maybe_dequantize(x.contiguous(), ctx.dtype) scales = maybe_dequantize(scales, ctx.dtype) - grad_output = maybe_dequantize(grad_output, ctx.dtype) - if isinstance(grad_output, GroupedTensorStorage): - grad_output = grad_output.rowwise_data.reshape(grad_output.logical_shape) - - if first_dims is None: - grad_input, grad_extra_input = self._scaled_unary_backward( - grad_output, - x, - scales, - ctx.prev_op_grad_output_quantizer, - compute_scale_grad=ctx.extra_input_requires_grad, - ) - else: - grad_input, dense_grad_input, grad_extra_input = self._grouped_scaled_unary_backward( - grad_output, - x, - scales, - ctx.prev_op_grad_output_quantizer, - num_groups=int(first_dims.numel()), - first_dims=first_dims, - tensor_offsets=tensor_offsets, - compute_scale_grad=ctx.extra_input_requires_grad, - ) - # Preserve the pre-quantize result for the preceding - # GroupedLinear's dbias/dscale reduction. ``grad_input`` remains - # quantized for its dgrad and wgrad GEMMs. - grad_input._dense_for_dbias = dense_grad_input + grad_output = maybe_dequantize(grad_output.contiguous(), ctx.dtype) + + grad_input, grad_extra_input = self._scaled_unary_backward( + grad_output, + x, + scales, + compute_scale_grad=ctx.extra_input_requires_grad, + ) if not ctx.input_requires_grad: grad_input = None @@ -552,32 +488,14 @@ def _scaled_unary_forward( self, input_: torch.Tensor, scales: torch.Tensor, - quantizer: Optional[Quantizer], - ) -> torch.Tensor: - return tex.scaled_srelu(input_, scales, quantizer) - - def _grouped_scaled_unary_forward( - self, - input_: torch.Tensor, - scales: torch.Tensor, - quantizer: Optional[Quantizer], - grouped_input: GroupedTensorStorage, ) -> torch.Tensor: - return tex.grouped_scaled_srelu( - input_, - scales.reshape(-1), - quantizer, - grouped_input.num_tensors, - grouped_input.first_dims, - None, - ) + return tex.scaled_srelu(input_, scales, None) def _scaled_unary_backward( self, grad_output: torch.Tensor, input_: torch.Tensor, scales: torch.Tensor, - quantizer: Optional[Quantizer], *, compute_scale_grad: bool, ) -> tuple[torch.Tensor, Optional[torch.Tensor]]: @@ -585,30 +503,7 @@ def _scaled_unary_backward( grad_output, input_, scales, - quantizer, - compute_scale_grad, - ) - - def _grouped_scaled_unary_backward( - self, - grad_output: torch.Tensor, - input_: torch.Tensor, - scales: torch.Tensor, - quantizer: Optional[Quantizer], - *, - num_groups: int, - first_dims: torch.Tensor, - tensor_offsets: Optional[torch.Tensor], - compute_scale_grad: bool, - ) -> tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]: - return tex.grouped_scaled_dsrelu( - grad_output, - input_, - scales.reshape(-1), - quantizer, - num_groups, - first_dims, - tensor_offsets, + None, compute_scale_grad, ) diff --git a/transformer_engine/pytorch/ops/basic/grouped_linear.py b/transformer_engine/pytorch/ops/basic/grouped_linear.py index 9ca15b4646..bc8dbed6c9 100644 --- a/transformer_engine/pytorch/ops/basic/grouped_linear.py +++ b/transformer_engine/pytorch/ops/basic/grouped_linear.py @@ -695,25 +695,14 @@ def pre_fuser_forward(self, *, requires_grad: bool) -> None: ) weight_requires_grad = requires_grad and weight_requires_grad - input_quantizers = [ - self.get_quantizer("forward", 2 * group_idx) for group_idx in range(self.num_groups) - ] - # Configure quantizer usages for group_idx in range(self.num_groups): - input_quantizer = input_quantizers[group_idx] + input_quantizer = self.get_quantizer("forward", 2 * group_idx) weight_quantizer = self.get_quantizer("forward", 2 * group_idx + 1) grad_output_quantizer = self.get_quantizer("backward", group_idx) input_quantizer.set_usage(rowwise=True, columnwise=weight_requires_grad) - # Activation fusion quantizes with this quantizer before this - # GroupedLinear's fuser_forward runs. All quantization paths - # support GEMM-optimized scale fusion, so configure it here. - input_quantizer.optimize_for_gemm = True weight_quantizer.set_usage(rowwise=True, columnwise=requires_grad) grad_output_quantizer.set_usage(rowwise=True, columnwise=weight_requires_grad) - # Activation backward quantizes with this quantizer before - # this GroupedLinear's backward path runs. - grad_output_quantizer.optimize_for_gemm = True def reset_recipe_state(self, *, recipe: Optional[Recipe]) -> None: super().reset_recipe_state(recipe=recipe) @@ -1008,7 +997,7 @@ def _get_grouped_bias_for_gemm( def fuser_forward( self, basic_op_ctxs: list[OperationContext], - input_: torch.Tensor | GroupedTensorStorage, + input_: torch.Tensor, *, basic_op_extra_inputs: list[tuple[torch.Tensor, ...]], prev_op_grad_output_quantizer: Optional[Quantizer], @@ -1202,7 +1191,7 @@ def fuser_forward_save_ctx( def _fuser_forward_split_quantize( self, *, - input_: torch.Tensor | GroupedTensorStorage, + input_: torch.Tensor, split_sizes: torch.Tensor, scales: Optional[torch.Tensor], with_quantized_compute: bool, @@ -1215,11 +1204,6 @@ def _fuser_forward_split_quantize( out_buffer: Optional[torch.Tensor] = None, ) -> tuple[torch.Tensor, tuple[Optional[torch.Tensor], ...]]: """Legacy ``tex.split_quantize`` + ``general_grouped_gemm`` flow.""" - if isinstance(input_, GroupedTensorStorage): - raise ValueError( - "GroupedLinear cannot use GroupedTensorStorage with the legacy " - "split-quantize path; use the grouped-tensor path instead." - ) num_groups = self.num_groups has_bias = self.has_bias @@ -1343,23 +1327,30 @@ def _fuser_forward_graph_safe( bulk_allocate=True, ) input_tensor_offsets = grouped_tensor_offsets[2] - original_shape = ( - list(input_.logical_shape) - if isinstance(input_, GroupedTensorStorage) - else list(input_.size()) - ) - input_quantizer = input_quantizers[0] if with_quantized_compute else None - if input_quantizer is not None: + original_shape = list(input_.size()) + total_tokens = math.prod(original_shape[:-1]) + x = maybe_dequantize(input_, dtype).reshape(-1, self.in_features) + if with_quantized_compute: + input_quantizer = input_quantizers[0] input_quantizer.set_usage(rowwise=True, columnwise=weight_requires_grad) input_quantizer.optimize_for_gemm = True - grouped_x = self._prepare_grouped_input( - input_, - expected_features=self.in_features, - split_sizes=split_sizes, - tensor_offsets=input_tensor_offsets, - dtype=dtype, - quantizer=input_quantizer, - ) + grouped_x = tex.group_quantize( + x, + input_quantizer, + num_groups, + split_sizes, + tensor_offsets=input_tensor_offsets, + ) + else: + grouped_x = GroupedTensorStorage( + shape=(total_tokens, self.in_features), + dtype=dtype, + num_tensors=num_groups, + quantizer=None, + data=x.reshape(-1), + first_dims=split_sizes, + tensor_offsets=input_tensor_offsets, + ) return self._fuser_forward_grouped_tensor( grouped_input=grouped_x, split_sizes=split_sizes, @@ -1379,76 +1370,6 @@ def _fuser_forward_graph_safe( out_shape=original_shape[:-1] + [self.out_features], ) - def _prepare_grouped_input( - self, - input_: torch.Tensor | GroupedTensorStorage, - *, - expected_features: int, - split_sizes: torch.Tensor, - tensor_offsets: torch.Tensor, - dtype: torch.dtype, - quantizer: Optional[Quantizer], - ) -> GroupedTensorStorage: - """Normalize an input to grouped storage for grouped GEMM. - - A dense input is never treated as an already-quantized grouped value. - Grouped inputs retain their quantization only when it is compatible with - the active quantizer; otherwise their rowwise data is dequantized and - converted to the requested high-precision dtype before being rewrapped. - """ - if isinstance(input_, GroupedTensorStorage): - if input_.num_tensors != self.num_groups: - raise ValueError( - f"GroupedLinear expected {self.num_groups} input groups, " - f"got {input_.num_tensors}." - ) - total_tokens, in_features = input_.logical_shape - if in_features != expected_features: - raise ValueError( - f"GroupedLinear expected input features={expected_features}, got {in_features}." - ) - input_first_dims = input_.first_dims if input_.first_dims is not None else split_sizes - input_offsets = ( - input_.tensor_offsets if input_.tensor_offsets is not None else tensor_offsets - ) - if quantizer is not None and input_.quantizer is not None: - if input_.quantizer is not quantizer: - raise ValueError( - "GroupedLinear received an incompatible grouped input " - f"(quantizer={input_.quantizer}; expected quantizer={quantizer})." - ) - return input_ - source = maybe_dequantize(input_, dtype) - assert isinstance(source, GroupedTensorStorage) - if source.rowwise_data is None: - raise RuntimeError("GroupedLinear requires rowwise data for grouped input.") - rowwise_data = source.rowwise_data.reshape(source.logical_shape) - first_dims = source.first_dims if source.first_dims is not None else input_first_dims - offsets = source.tensor_offsets if source.tensor_offsets is not None else input_offsets - else: - rowwise_data = maybe_dequantize(input_, dtype).reshape(-1, expected_features) - total_tokens = rowwise_data.size(0) - first_dims = split_sizes - offsets = tensor_offsets - - if quantizer is not None: - return tex.group_quantize( - rowwise_data, - quantizer, - self.num_groups, - first_dims, - tensor_offsets=offsets, - ) - return GroupedTensorStorage( - shape=(total_tokens, expected_features), - dtype=dtype, - num_tensors=self.num_groups, - quantizer=None, - data=rowwise_data.reshape(-1), - first_dims=first_dims, - tensor_offsets=offsets, - ) - def _fuser_forward_grouped_tensor( self, *, @@ -1507,11 +1428,9 @@ def _fuser_forward_grouped_tensor( dtype=dtype, ) - # Allocate output buffer and wrap it in a GroupedTensor. This remains - # compatible with the fuser/autograd boundary while carrying grouped - # layout metadata to a subsequent grouped operation. + # Allocate output buffer and wrap as a GroupedTensor view. out = validate_or_alloc_output(out_buffer, out_shape, dtype, device) - grouped_out = GroupedTensor( + grouped_out = GroupedTensorStorage( shape=(total_tokens, self.out_features), dtype=dtype, num_tensors=num_groups, @@ -1578,12 +1497,12 @@ def _fuser_forward_grouped_tensor( saved.append(grouped_weights) else: saved.extend(grouped_weights) - return grouped_out, tuple(saved) + return out, tuple(saved) def fuser_backward( self, basic_op_ctxs: list[OperationContext], - grad_output: torch.Tensor | GroupedTensorStorage, + grad_output: torch.Tensor, *, basic_op_grad_extra_outputs: list[tuple[torch.Tensor, ...]], ) -> tuple[ @@ -1794,7 +1713,7 @@ def _fuser_backward_graph_safe( self, *, ctx: OperationContext, - grad_output: torch.Tensor | GroupedTensorStorage, + grad_output: torch.Tensor, ) -> tuple[ torch.Tensor, Iterable[Iterable[Optional[torch.Tensor]]], @@ -1807,6 +1726,8 @@ def _fuser_backward_graph_safe( with_quantized_compute = bool(getattr(ctx, "with_quantized_compute", False)) split_sizes = ctx.saved_tensors[0] output_tensor_offsets = ctx.saved_tensors[4] + dy_2d = grad_output.reshape(-1, self.out_features) + total_tokens = dy_2d.size(0) dbias_packed = None if with_quantized_compute: @@ -1821,13 +1742,7 @@ def _fuser_backward_graph_safe( fuse_bgrad = isinstance(grad_output_quantizer, MXFP8Quantizer) or ( isinstance(grad_output_quantizer, Float8BlockQuantizer) and ctx.input_requires_grad ) - if ( - not isinstance(grad_output, GroupedTensorStorage) - and has_bias - and not self._scale_bias - and fuse_bgrad - ): - dy_2d = grad_output.reshape(-1, self.out_features) + if has_bias and not self._scale_bias and fuse_bgrad: grouped_dy, dbias_packed = tex.bgrad_group_quantize( dy_2d, grad_output_quantizer, @@ -1836,22 +1751,23 @@ def _fuser_backward_graph_safe( tensor_offsets=output_tensor_offsets, ) else: - grouped_dy = self._prepare_grouped_input( - grad_output, - expected_features=self.out_features, - split_sizes=split_sizes, + grouped_dy = tex.group_quantize( + dy_2d, + grad_output_quantizer, + num_groups, + split_sizes, tensor_offsets=output_tensor_offsets, - dtype=dtype, - quantizer=grad_output_quantizer, ) else: - grouped_dy = self._prepare_grouped_input( - grad_output, - expected_features=self.out_features, - split_sizes=split_sizes, - tensor_offsets=output_tensor_offsets, + dy_2d = maybe_dequantize(dy_2d, dtype) + grouped_dy = GroupedTensorStorage( + shape=(total_tokens, self.out_features), dtype=dtype, + num_tensors=num_groups, quantizer=None, + data=dy_2d.reshape(-1), + first_dims=split_sizes, + tensor_offsets=output_tensor_offsets, ) return self._fuser_backward_grouped_tensor( @@ -1903,21 +1819,11 @@ def _fuser_backward_grouped_tensor( else: ws, saved_tensors = saved_tensors[:num_groups], saved_tensors[num_groups:] - # Keep grouped grad output quantized for GEMM. Materialize a dense - # high-precision view only when bias-gradient computation needs it. - dy_2d = None - if isinstance(grad_output, GroupedTensorStorage): - total_tokens, out_features = grad_output.logical_shape - if out_features != self.out_features: - raise ValueError( - f"GroupedLinear expected grad-output features={self.out_features}, " - f"got {out_features}." - ) - grad_input_shape = [total_tokens, self.in_features] - else: - dy_2d = grad_output.reshape(-1, self.out_features) - total_tokens = dy_2d.size(0) - grad_input_shape = list(grad_output.size())[:-1] + [self.in_features] + # Keep the dense high-precision grad for bias-gradient computation while + # optionally using a pre-quantized grouped view for the grouped GEMMs. + dy_2d = grad_output.reshape(-1, self.out_features) + total_tokens = dy_2d.size(0) + grad_input_shape = list(grad_output.size())[:-1] + [self.in_features] expected_quantizer = ctx.grad_output_quantizers[0] if with_quantized_compute else None if grouped_grad_output.quantizer is not expected_quantizer or tuple( @@ -1936,12 +1842,6 @@ def _fuser_backward_grouped_tensor( final_bias_grads: Optional[torch.Tensor] = None grad_scales: Optional[torch.Tensor] = None if has_bias: - if dy_2d is None: - dy_2d = getattr(grad_output, "_dense_for_dbias", None) - if dy_2d is None: - dy_2d = maybe_dequantize(grad_output, dtype) - if isinstance(dy_2d, GroupedTensorStorage): - dy_2d = dy_2d.rowwise_data.reshape(dy_2d.logical_shape) if self._scale_bias: bias_packed = torch.stack(self._get_bias_tensors(dtype)) scales_f32 = scales.to(dtype=torch.float32) @@ -1958,8 +1858,6 @@ def _fuser_backward_grouped_tensor( final_bias_grads = [dbias_packed.to(dtype=dtype)] else: final_bias_grads = [dbias_packed[idx].to(dtype=dtype) for idx in range(num_groups)] - if hasattr(grad_output, "_dense_for_dbias"): - grad_output._dense_for_dbias = None # ---- dgrad GEMM ---------------------------------------------------- grad_input = None diff --git a/transformer_engine/pytorch/ops/basic/swiglu.py b/transformer_engine/pytorch/ops/basic/swiglu.py index f1fe591496..fb663c0480 100644 --- a/transformer_engine/pytorch/ops/basic/swiglu.py +++ b/transformer_engine/pytorch/ops/basic/swiglu.py @@ -14,7 +14,6 @@ from ...constants import DType from ...cpu_offload import is_cpu_offload_enabled, mark_activation_offload from ...tensor import Float8CurrentScalingQuantizer, Quantizer -from ...tensor.storage.grouped_tensor_storage import GroupedTensorStorage from ...utils import clear_tensor_data from ..op import BasicOperation, OperationContext from .._common import maybe_dequantize @@ -392,16 +391,6 @@ def _scaled_glu_forward( self, input_: torch.Tensor, scales: torch.Tensor, - quantizer: Optional[Quantizer], - ) -> torch.Tensor: - raise NotImplementedError - - def _grouped_scaled_glu_forward( - self, - input_: torch.Tensor, - scales: torch.Tensor, - quantizer: Optional[Quantizer], - grouped_input: GroupedTensorStorage, ) -> torch.Tensor: raise NotImplementedError @@ -410,26 +399,11 @@ def _scaled_glu_backward( grad_output: torch.Tensor, input_: torch.Tensor, scales: torch.Tensor, - quantizer: Optional[Quantizer], *, compute_scale_grad: bool, ) -> tuple[torch.Tensor, Optional[torch.Tensor]]: raise NotImplementedError - def _grouped_scaled_glu_backward( - self, - grad_output: torch.Tensor, - input_: torch.Tensor, - scales: torch.Tensor, - quantizer: Optional[Quantizer], - *, - num_groups: int, - first_dims: torch.Tensor, - tensor_offsets: Optional[torch.Tensor], - compute_scale_grad: bool, - ) -> tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]: - raise NotImplementedError - def op_forward(self, *args, **kwargs) -> None: raise RuntimeError( f"{self.__class__.__name__} operation has " @@ -473,17 +447,9 @@ def fuser_forward( dtype = extra_input.dtype # Make sure inputs are in correct dtype - grouped_input = input_ if isinstance(input_, GroupedTensorStorage) else None input_ = maybe_dequantize(input_, dtype) - if isinstance(input_, GroupedTensorStorage): - input_ = input_.rowwise_data.reshape(input_.logical_shape) scales = maybe_dequantize(extra_input, dtype) - if grouped_input is None: - out = self._scaled_glu_forward(input_, scales, next_op_input_quantizer) - else: - out = self._grouped_scaled_glu_forward( - input_, scales, next_op_input_quantizer, grouped_input - ) + out = self._scaled_glu_forward(input_, scales) # Save state for backward pass ctx = basic_op_ctxs[0] @@ -494,10 +460,9 @@ def fuser_forward( ctx.extra_input_requires_grad = extra_input.requires_grad ctx.dtype = dtype ctx.save_for_backward( - grouped_input if grouped_input is not None else input_, + input_, scales if ctx.input_requires_grad or ctx.extra_input_requires_grad else None, ) - ctx.prev_op_grad_output_quantizer = prev_op_grad_output_quantizer return out, [()] @@ -520,41 +485,17 @@ def fuser_backward( ctx = basic_op_ctxs[0] input_, scales = ctx.saved_tensors - grouped_input = input_ if isinstance(input_, GroupedTensorStorage) else None - first_dims = grouped_input.first_dims if grouped_input is not None else None - tensor_offsets = grouped_input.tensor_offsets if grouped_input is not None else None input_ = maybe_dequantize(input_, ctx.dtype) - if isinstance(input_, GroupedTensorStorage): - input_ = input_.rowwise_data.reshape(input_.logical_shape) if scales is not None: scales = maybe_dequantize(scales, ctx.dtype) grad_output = maybe_dequantize(grad_output, ctx.dtype) - if isinstance(grad_output, GroupedTensorStorage): - grad_output = grad_output.rowwise_data.reshape(grad_output.logical_shape) - if grouped_input is None: - grad_input, grad_extra_input = self._scaled_glu_backward( - grad_output, - input_, - scales, - ctx.prev_op_grad_output_quantizer, - compute_scale_grad=ctx.extra_input_requires_grad, - ) - else: - grad_input, dense_grad_input, grad_extra_input = self._grouped_scaled_glu_backward( - grad_output, - input_, - scales, - ctx.prev_op_grad_output_quantizer, - num_groups=int(first_dims.numel()), - first_dims=first_dims, - tensor_offsets=tensor_offsets, - compute_scale_grad=ctx.extra_input_requires_grad, - ) - # Preserve the pre-quantize result for the preceding - # GroupedLinear's dbias/dscale reduction. ``grad_input`` remains - # quantized for its dgrad and wgrad GEMMs. - grad_input._dense_for_dbias = dense_grad_input + grad_input, grad_extra_input = self._scaled_glu_backward( + grad_output, + input_, + scales, + compute_scale_grad=ctx.extra_input_requires_grad, + ) if not ctx.input_requires_grad: grad_input = None @@ -586,28 +527,10 @@ def _scaled_glu_forward( self, input_: torch.Tensor, scales: torch.Tensor, - quantizer: Optional[Quantizer], ) -> torch.Tensor: return tex.scaled_swiglu( input_, scales, - quantizer, - int(self.glu_interleave_size or 0), - ) - - def _grouped_scaled_glu_forward( - self, - input_: torch.Tensor, - scales: torch.Tensor, - quantizer: Optional[Quantizer], - grouped_input: GroupedTensorStorage, - ) -> torch.Tensor: - return tex.grouped_scaled_swiglu( - input_, - scales.reshape(-1), - quantizer, - grouped_input.num_tensors, - grouped_input.first_dims, None, int(self.glu_interleave_size or 0), ) @@ -617,7 +540,6 @@ def _scaled_glu_backward( grad_output: torch.Tensor, input_: torch.Tensor, scales: torch.Tensor, - quantizer: Optional[Quantizer], *, compute_scale_grad: bool, ) -> tuple[torch.Tensor, Optional[torch.Tensor]]: @@ -625,31 +547,7 @@ def _scaled_glu_backward( grad_output, input_, scales, - quantizer, - int(self.glu_interleave_size or 0), - compute_scale_grad, - ) - - def _grouped_scaled_glu_backward( - self, - grad_output: torch.Tensor, - input_: torch.Tensor, - scales: torch.Tensor, - quantizer: Optional[Quantizer], - *, - num_groups: int, - first_dims: torch.Tensor, - tensor_offsets: Optional[torch.Tensor], - compute_scale_grad: bool, - ) -> tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]: - return tex.grouped_scaled_dswiglu( - grad_output, - input_, - scales.reshape(-1), - quantizer, - num_groups, - first_dims, - tensor_offsets, + None, int(self.glu_interleave_size or 0), compute_scale_grad, ) @@ -704,33 +602,11 @@ def _scaled_glu_forward( self, input_: torch.Tensor, scales: torch.Tensor, - quantizer: Optional[Quantizer], ) -> torch.Tensor: clamped = self._clamped return tex.scaled_clamped_swiglu( input_, scales, - quantizer, - clamped.limit, - clamped.alpha, - clamped.glu_linear_offset, - int(self.glu_interleave_size or 0), - ) - - def _grouped_scaled_glu_forward( - self, - input_: torch.Tensor, - scales: torch.Tensor, - quantizer: Optional[Quantizer], - grouped_input: GroupedTensorStorage, - ) -> torch.Tensor: - clamped = self._clamped - return tex.grouped_scaled_clamped_swiglu( - input_, - scales.reshape(-1), - quantizer, - grouped_input.num_tensors, - grouped_input.first_dims, None, clamped.limit, clamped.alpha, @@ -743,7 +619,6 @@ def _scaled_glu_backward( grad_output: torch.Tensor, input_: torch.Tensor, scales: torch.Tensor, - quantizer: Optional[Quantizer], *, compute_scale_grad: bool, ) -> tuple[torch.Tensor, Optional[torch.Tensor]]: @@ -752,35 +627,7 @@ def _scaled_glu_backward( grad_output, input_, scales, - quantizer, - clamped.limit, - clamped.alpha, - clamped.glu_linear_offset, - int(self.glu_interleave_size or 0), - compute_scale_grad, - ) - - def _grouped_scaled_glu_backward( - self, - grad_output: torch.Tensor, - input_: torch.Tensor, - scales: torch.Tensor, - quantizer: Optional[Quantizer], - *, - num_groups: int, - first_dims: torch.Tensor, - tensor_offsets: Optional[torch.Tensor], - compute_scale_grad: bool, - ) -> tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]: - clamped = self._clamped - return tex.grouped_scaled_clamped_dswiglu( - grad_output, - input_, - scales.reshape(-1), - quantizer, - num_groups, - first_dims, - tensor_offsets, + None, clamped.limit, clamped.alpha, clamped.glu_linear_offset, diff --git a/transformer_engine/pytorch/ops/fused/__init__.py b/transformer_engine/pytorch/ops/fused/__init__.py index dc9dcd6dc3..85191a9807 100644 --- a/transformer_engine/pytorch/ops/fused/__init__.py +++ b/transformer_engine/pytorch/ops/fused/__init__.py @@ -6,9 +6,14 @@ from ..fuser import register_backward_fusion, register_forward_fusion from .backward_activation_bias import BackwardActivationBias +from .backward_activation_grouped_linear import BackwardScaledActivationGroupedLinear from .backward_add_rmsnorm import BackwardAddRMSNorm from .backward_linear_add import BackwardLinearAdd from .backward_linear_scale import BackwardLinearScale +from .forward_activation_grouped_linear import ( + ForwardScaledActivationGroupedLinear, + act_grouped_linear_fusion_supported, +) from .forward_linear_bias_activation import ForwardLinearBiasActivation from .forward_linear_bias_add import ForwardLinearBiasAdd from .forward_linear_scale_add import ForwardLinearScaleAdd @@ -21,6 +26,7 @@ register_forward_fusion(ForwardLinearBiasAdd.fuse_forward_ops) register_forward_fusion(ForwardLinearBiasActivation.fuse_forward_ops) register_forward_fusion(ForwardLinearScaleAdd.fuse_forward_ops) +register_forward_fusion(ForwardScaledActivationGroupedLinear.fuse_forward_ops) # Register backward fusions register_backward_fusion(UserbuffersBackwardLinear.fuse_backward_ops) @@ -28,6 +34,7 @@ register_backward_fusion(BackwardLinearScale.fuse_backward_ops) register_backward_fusion(BackwardActivationBias.fuse_backward_ops) register_backward_fusion(BackwardAddRMSNorm.fuse_backward_ops) +register_backward_fusion(BackwardScaledActivationGroupedLinear.fuse_backward_ops) # Import experimental fusions # Note: Registration logic is non-trivial, so submodule handles it internally. diff --git a/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py b/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py new file mode 100644 index 0000000000..0cb17fe0d1 --- /dev/null +++ b/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py @@ -0,0 +1,181 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Fused scaled activation + grouped linear backward.""" + +from __future__ import annotations + +from collections.abc import Iterable +from typing import Optional + +import torch + +import transformer_engine_torch as tex +from ...quantization import Recipe +from ...tensor import Quantizer +from ...utils import clear_tensor_data +from .._common import maybe_dequantize +from ..basic import GroupedLinear, ScaledClampedQGeGLU, ScaledSReLU, ScaledSwiGLU +from ..op import FusedOperation, FusibleOperation, OperationContext +from .forward_activation_grouped_linear import ( + _SCALED_ACTIVATION_TYPES, + _ScaledActivation, + act_grouped_linear_fusion_supported, +) + + +def _grouped_scaled_dactivation( + activation: _ScaledActivation, + grad_output: torch.Tensor, + input_: torch.Tensor, + scales: torch.Tensor, + *, + quantizer: Quantizer, + num_groups: int, + split_sizes: torch.Tensor, + tensor_offsets: torch.Tensor, + compute_scale_grad: bool, +) -> tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]: + """Dispatch a grouped scaled activation backward pass.""" + dy = grad_output.reshape(-1, grad_output.size(-1)) + x = input_.reshape(-1, input_.size(-1)) + s = scales.reshape(-1) + if isinstance(activation, ScaledSwiGLU): + return tex.grouped_scaled_dswiglu( + dy, + x, + s, + quantizer, + num_groups, + split_sizes, + tensor_offsets, + int(activation.glu_interleave_size or 0), + compute_scale_grad, + ) + if isinstance(activation, ScaledClampedQGeGLU): + clamped = activation._clamped + return tex.grouped_scaled_clamped_dswiglu( + dy, + x, + s, + quantizer, + num_groups, + split_sizes, + tensor_offsets, + clamped.limit, + clamped.alpha, + clamped.glu_linear_offset, + int(activation.glu_interleave_size or 0), + compute_scale_grad, + ) + if isinstance(activation, ScaledSReLU): + return tex.grouped_scaled_dsrelu( + dy, + x, + s, + quantizer, + num_groups, + split_sizes, + tensor_offsets, + compute_scale_grad, + ) + raise TypeError(f"Unsupported scaled activation type ({type(activation).__name__})") + + +class BackwardScaledActivationGroupedLinear(FusedOperation): + """Scaled activation backward + grouped quantize + grouped linear backward.""" + + def __init__(self, *, linear: GroupedLinear, activation: _ScaledActivation) -> None: + super().__init__((linear, activation)) + + def fuser_backward( + self, + basic_op_ctxs: list[OperationContext], + grad_output: torch.Tensor, + *, + basic_op_grad_extra_outputs: list[tuple[torch.Tensor, ...]], + ) -> tuple[ + Optional[torch.Tensor], + Iterable[Iterable[Optional[torch.Tensor]]], + Iterable[Iterable[Optional[torch.Tensor]]], + ]: + linear = self.basic_ops[0] + activation = self.basic_ops[1] + linear_ctx, activation_ctx = basic_op_ctxs + + if not linear_ctx.requires_grad: + _, _, act_grad_extra_inputs = activation.fuser_backward( + [activation_ctx], + grad_output, + basic_op_grad_extra_outputs=[basic_op_grad_extra_outputs[1]], + ) + return None, [(), ()], [(), act_grad_extra_inputs[0]] + + del basic_op_grad_extra_outputs + input_, scales = activation_ctx.saved_tensors + input_ = maybe_dequantize(input_, activation_ctx.dtype) + scales = maybe_dequantize(scales, activation_ctx.dtype) + grad_output = maybe_dequantize(grad_output, activation_ctx.dtype) + + split_sizes = linear_ctx.saved_tensors[0] + activation_input_tensor_offsets = linear_ctx.saved_tensors[4] + grad_output_quantizer = linear_ctx.grad_output_quantizers[0] + grad_output_quantizer.set_usage( + rowwise=linear_ctx.input_requires_grad, + columnwise=linear_ctx.weight_requires_grad, + ) + grad_output_quantizer.optimize_for_gemm = True + grouped_dy, dense_dy, grad_scales = _grouped_scaled_dactivation( + activation, + grad_output, + input_, + scales, + quantizer=grad_output_quantizer, + num_groups=linear.num_groups, + split_sizes=split_sizes, + tensor_offsets=activation_input_tensor_offsets, + compute_scale_grad=activation_ctx.extra_input_requires_grad, + ) + + grad_input, grad_params, grad_extra_inputs = linear._fuser_backward_grouped_tensor( + ctx=linear_ctx, + grad_output=dense_dy, + grouped_grad_output=grouped_dy, + ) + + clear_tensor_data(activation_ctx.saved_tensors[0]) + return ( + grad_input, + [grad_params[0], ()], + [grad_extra_inputs[0], (grad_scales,)], + ) + + @staticmethod + def fuse_backward_ops( + ops: list[FusibleOperation], + *, + recipe: Optional[Recipe] = None, + **unused, # pylint: disable=unused-argument + ) -> list[FusibleOperation]: + """Fuse each supported GroupedLinear + ScaledActivation pair.""" + out: list[FusibleOperation] = [] + idx = 0 + while idx < len(ops): + if ( + idx + 1 < len(ops) + and isinstance(ops[idx], GroupedLinear) + and isinstance(ops[idx + 1], _SCALED_ACTIVATION_TYPES) + and act_grouped_linear_fusion_supported(ops[idx], ops[idx + 1], recipe) + ): + out.append( + BackwardScaledActivationGroupedLinear( + linear=ops[idx], + activation=ops[idx + 1], + ) + ) + idx += 2 + else: + out.append(ops[idx]) + idx += 1 + return out diff --git a/transformer_engine/pytorch/ops/fused/forward_activation_grouped_linear.py b/transformer_engine/pytorch/ops/fused/forward_activation_grouped_linear.py new file mode 100644 index 0000000000..4f0b0d627c --- /dev/null +++ b/transformer_engine/pytorch/ops/fused/forward_activation_grouped_linear.py @@ -0,0 +1,229 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Fused scaled activation + grouped linear forward.""" + +from __future__ import annotations + +from collections.abc import Iterable +from typing import Any, Optional + +import torch + +import transformer_engine_torch as tex +from ...cpu_offload import is_cpu_offload_enabled, mark_activation_offload +from ...quantization import Recipe +from ...tensor import Quantizer +from .._common import maybe_dequantize +from ..basic import GroupedLinear, ScaledClampedQGeGLU, ScaledSReLU, ScaledSwiGLU +from ..basic.activation import _ScaledUnary +from ..basic.swiglu import _ScaledGLU +from ..op import FusedOperation, FusibleOperation, OperationContext + + +_ScaledActivation = _ScaledGLU | _ScaledUnary +_SCALED_ACTIVATION_TYPES = (_ScaledGLU, _ScaledUnary) + + +def _grouped_scaled_activation( + activation: _ScaledActivation, + input_: torch.Tensor, + scales: torch.Tensor, + quantizer: Quantizer, + num_groups: int, + split_sizes: torch.Tensor, + tensor_offsets: torch.Tensor, +) -> torch.Tensor: + """Dispatch a grouped scaled activation.""" + x = input_.reshape(-1, input_.size(-1)) + s = scales.reshape(-1) + if isinstance(activation, ScaledSwiGLU): + return tex.grouped_scaled_swiglu( + x, + s, + quantizer, + num_groups, + split_sizes, + tensor_offsets, + int(activation.glu_interleave_size or 0), + ) + if isinstance(activation, ScaledClampedQGeGLU): + clamped = activation._clamped + return tex.grouped_scaled_clamped_swiglu( + x, + s, + quantizer, + num_groups, + split_sizes, + tensor_offsets, + clamped.limit, + clamped.alpha, + clamped.glu_linear_offset, + int(activation.glu_interleave_size or 0), + ) + if isinstance(activation, ScaledSReLU): + return tex.grouped_scaled_srelu( + x, + s, + quantizer, + num_groups, + split_sizes, + tensor_offsets, + ) + raise TypeError(f"Unsupported scaled activation type ({type(activation).__name__})") + + +def act_grouped_linear_fusion_supported( + linear: GroupedLinear, + activation: _ScaledActivation, + recipe: Optional[Recipe], +) -> bool: + """Whether ScaledActivation + GroupedLinear can use grouped quantized compute.""" + if recipe is None or activation.activation_recompute_in_mlp: + return False + input_quantizers = [ + linear.get_quantizer("forward", 2 * group_idx) for group_idx in range(linear.num_groups) + ] + weight = linear.weight if linear.single_grouped_weight else linear.weight0 + dtype = torch.get_autocast_dtype("cuda") if torch.is_autocast_enabled() else weight.dtype + return linear._is_graph_safe_path_supported( + with_quantized_compute=True, + input_quantizers=input_quantizers, + dtype=dtype, + single_grouped_weight=linear.single_grouped_weight, + ) + + +class ForwardScaledActivationGroupedLinear(FusedOperation): + """Scaled activation + grouped quantize + grouped linear forward.""" + + def __init__(self, *, activation: _ScaledActivation, linear: GroupedLinear) -> None: + super().__init__((activation, linear)) + + def fuser_forward( + self, + basic_op_ctxs: list[OperationContext], + input_: torch.Tensor, + *, + basic_op_extra_inputs: list[tuple[torch.Tensor, ...]], + prev_op_grad_output_quantizer: Optional[Quantizer], + next_op_input_quantizer: Optional[Quantizer], + basic_op_kwargs: list[dict[str, Any]], + ) -> tuple[torch.Tensor, Iterable[Iterable[torch.Tensor]]]: + activation = self.basic_ops[0] + linear = self.basic_ops[1] + activation_ctx, linear_ctx = basic_op_ctxs + if basic_op_kwargs[0] or basic_op_kwargs[1]: + raise ValueError("Scaled activation and GroupedLinear do not expect keyword arguments") + + weight = linear.weight if linear.single_grouped_weight else linear.weight0 + device = weight.device + dtype = torch.get_autocast_dtype("cuda") if torch.is_autocast_enabled() else weight.dtype + input_ = maybe_dequantize(input_, dtype) + scales = maybe_dequantize(basic_op_extra_inputs[0][0], dtype) + + split_sizes = basic_op_extra_inputs[1][0] + if int(split_sizes.numel()) != linear.num_groups: + raise ValueError( + f"Expected {linear.num_groups} splits, but got {int(split_sizes.numel())}." + ) + split_sizes = split_sizes.to(device=device, dtype=torch.int64) + linear_scales = basic_op_extra_inputs[1][1] if linear._scale_bias else None + split_sizes, grouped_tensor_offsets = tex.splits_to_offsets_multi( + split_sizes, + device, + strides=[1, 1, linear.in_features, linear.out_features], + include_leading_zero=[False, True, True, True], + dtypes=[torch.int32, torch.int64, torch.int64, torch.int64], + bulk_allocate=True, + ) + + input_quantizers = [ + linear.get_quantizer("forward", 2 * group_idx) for group_idx in range(linear.num_groups) + ] + weight_quantizers = [ + linear.get_quantizer("forward", 2 * group_idx + 1) + for group_idx in range(linear.num_groups) + ] + input_quantizer = input_quantizers[0] + weight_requires_grad = linear_ctx.requires_grad and weight.requires_grad + input_quantizer.set_usage(rowwise=True, columnwise=weight_requires_grad) + input_quantizer.optimize_for_gemm = True + + grouped_x = _grouped_scaled_activation( + activation, + input_, + scales, + input_quantizer, + linear.num_groups, + split_sizes, + grouped_tensor_offsets[2], + ) + + if activation_ctx.requires_grad: + if is_cpu_offload_enabled(): + mark_activation_offload(input_) + activation_ctx.input_requires_grad = True + activation_ctx.extra_input_requires_grad = basic_op_extra_inputs[0][0].requires_grad + activation_ctx.dtype = dtype + activation_ctx.save_for_backward(input_, scales) + + out, tensors_to_save = linear._fuser_forward_grouped_tensor( + grouped_input=grouped_x, + split_sizes=split_sizes, + scales=linear_scales, + with_quantized_compute=True, + input_quantizers=input_quantizers, + weight_quantizers=weight_quantizers, + dtype=dtype, + input_requires_grad=linear_ctx.requires_grad, + weight_requires_grad=weight_requires_grad, + device=device, + split_points=grouped_tensor_offsets[0], + base_split_offsets=grouped_tensor_offsets[1], + input_tensor_offsets=grouped_tensor_offsets[2], + output_tensor_offsets=grouped_tensor_offsets[3], + out_shape=list(input_.size())[:-1] + [linear.out_features], + ) + linear.fuser_forward_save_ctx( + basic_op_ctxs=[linear_ctx], + input_=input_, + tensors_to_save=[tensors_to_save], + requires_grad=[linear_ctx.requires_grad], + basic_op_extra_inputs=[basic_op_extra_inputs[1]], + prev_op_grad_output_quantizer=prev_op_grad_output_quantizer, + next_op_input_quantizer=next_op_input_quantizer, + basic_op_kwargs=[basic_op_kwargs[1]], + use_grouped_tensor_path=True, + ) + return out, [(), ()] + + @staticmethod + def fuse_forward_ops( + ops: list[FusibleOperation], + *, + recipe: Optional[Recipe] = None, + **unused, # pylint: disable=unused-argument + ) -> list[FusibleOperation]: + """Fuse each supported ScaledActivation + GroupedLinear pair.""" + out: list[FusibleOperation] = [] + idx = 0 + while idx < len(ops): + if ( + idx + 1 < len(ops) + and isinstance(ops[idx], _SCALED_ACTIVATION_TYPES) + and isinstance(ops[idx + 1], GroupedLinear) + and act_grouped_linear_fusion_supported(ops[idx + 1], ops[idx], recipe) + ): + out.append( + ForwardScaledActivationGroupedLinear( + activation=ops[idx], + linear=ops[idx + 1], + ) + ) + idx += 2 + else: + out.append(ops[idx]) + idx += 1 + return out From 099357b7047cbd483c3788e301d388dfe808e8da Mon Sep 17 00:00:00 2001 From: Varun Thumbe Date: Wed, 5 Aug 2026 19:08:30 +0000 Subject: [PATCH 15/17] get rid of fused ops now Signed-off-by: Varun Thumbe --- qa/L0_pytorch_unittest/test.sh | 1 - tests/pytorch/test_grouped_mlp.py | 34 +-- transformer_engine/pytorch/csrc/extensions.h | 37 --- .../pytorch/csrc/extensions/activation.cpp | 108 +-------- .../pytorch/csrc/extensions/pybind.cpp | 30 --- .../pytorch/ops/basic/grouped_linear.py | 170 +++---------- .../pytorch/ops/fused/__init__.py | 7 - .../backward_activation_grouped_linear.py | 181 -------------- .../forward_activation_grouped_linear.py | 229 ------------------ 9 files changed, 39 insertions(+), 758 deletions(-) delete mode 100644 transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py delete mode 100644 transformer_engine/pytorch/ops/fused/forward_activation_grouped_linear.py diff --git a/qa/L0_pytorch_unittest/test.sh b/qa/L0_pytorch_unittest/test.sh index 1dc714faf3..39d3e79e62 100644 --- a/qa/L0_pytorch_unittest/test.sh +++ b/qa/L0_pytorch_unittest/test.sh @@ -70,7 +70,6 @@ NVTE_ALLOW_NONDETERMINISTIC_ALGO=0 NVTE_DISABLE_TRITON_AUTOTUNING=1 NVIDIA_TF32_ PYTORCH_JIT=0 NVTE_TORCH_COMPILE=0 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_grouped_linear.xml $TE_PATH/tests/pytorch/test_grouped_linear.py || test_fail "test_grouped_linear.py" PYTORCH_JIT=0 NVTE_TORCH_COMPILE=0 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_ops_grouped_linear_distributed_weight.xml $TE_PATH/tests/pytorch/test_ops_grouped_linear_distributed_weight.py || test_fail "test_ops_grouped_linear_distributed_weight.py" NVTE_GROUPED_LINEAR_SINGLE_PARAM=1 NVTE_CUTEDSL_FUSED_GROUPED_MLP=1 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_grouped_mlp.xml $TE_PATH/tests/pytorch/test_grouped_mlp.py || test_fail "test_grouped_mlp.py" -NVTE_GROUPED_LINEAR_SINGLE_PARAM=1 NVTE_CUTEDSL_FUSED_GROUPED_MLP=0 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_grouped_mlp_without_cutedsl_fusion.xml $TE_PATH/tests/pytorch/test_grouped_mlp.py || test_fail "test_grouped_mlp.py without CuTe DSL fusion" if [ "$RET" -ne 0 ]; then echo "Error in the following test cases:$FAILED_CASES" diff --git a/tests/pytorch/test_grouped_mlp.py b/tests/pytorch/test_grouped_mlp.py index 245bc355cd..9c9a7065a2 100644 --- a/tests/pytorch/test_grouped_mlp.py +++ b/tests/pytorch/test_grouped_mlp.py @@ -24,12 +24,6 @@ OUTPUT_BUFFER_KEY, GRAD_INPUT_BUFFER_KEY, ) -from transformer_engine.pytorch.ops.fused.backward_activation_grouped_linear import ( - BackwardScaledActivationGroupedLinear, -) -from transformer_engine.pytorch.ops.fused.forward_activation_grouped_linear import ( - ForwardScaledActivationGroupedLinear, -) from transformer_engine.pytorch import ( QuantizedTensor, Float8CurrentScalingQuantizer, @@ -79,11 +73,6 @@ if nvfp4_available: _grouped_mlp_quantization_list.append("nvfp4_rht") -# Quantization recipes supported by ScaledActivation + GroupedLinear fusion -if fp8_available: - _grouped_mlp_quantization_list.append("fp8_current_scaling") - - @pytest.fixture(autouse=True, scope="function") def _reset_rng_states_per_test(): """Restore torch, CUDA, and Python ``random`` before each test in this module.""" @@ -1049,15 +1038,14 @@ def _make_module(): ) ) ) - forward_ops = module._module_groups[0]._forward_ops - backward_ops = module._module_groups[0]._backward_ops - full_grouped_mlp_fusion = False if expected_grouped_mlp_fusion: if activation_is_glu: fused_cls = te.ops.fused.GroupedMLP_CuTeGEMMGLU else: fused_cls = te.ops.fused.GroupedMLP_CuTeGEMMUnary if fused_cls.is_supported(): + forward_ops = module._module_groups[0]._forward_ops + backward_ops = module._module_groups[0]._backward_ops assert len(forward_ops) == 1 assert len(backward_ops) == 1 assert isinstance( @@ -1065,24 +1053,6 @@ def _make_module(): fused_cls, ) assert backward_ops[0][0] is forward_ops[0][0] - full_grouped_mlp_fusion = True - - # When the full FC1 + activation + FC2 fusion is unavailable, verify - # that ScaledActivation + GroupedLinear fusions cover both boundaries - # whenever grouped quantized compute is supported. - act_grouped_linear_fusion_expected = ( - not full_grouped_mlp_fusion - and te.ops.fused.act_grouped_linear_fusion_supported(fc2, module[1], recipe) - and te.ops.fused.act_grouped_linear_fusion_supported(fc1, module[1], recipe) - ) - assert ( - any(isinstance(op, ForwardScaledActivationGroupedLinear) for op, _ in forward_ops) - == act_grouped_linear_fusion_expected - ) - assert ( - any(isinstance(op, BackwardScaledActivationGroupedLinear) for op, _ in backward_ops) - == act_grouped_linear_fusion_expected - ) # Loose tols for sanity checking tols = {"rtol": 0.125, "atol": 0.25} diff --git a/transformer_engine/pytorch/csrc/extensions.h b/transformer_engine/pytorch/csrc/extensions.h index 9b51df7e09..8248a63680 100644 --- a/transformer_engine/pytorch/csrc/extensions.h +++ b/transformer_engine/pytorch/csrc/extensions.h @@ -315,43 +315,6 @@ py::tuple scaled_dsrelu(const at::Tensor &grad, const at::Tensor &input, const at::Tensor &act_scales, py::handle quantizer, bool compute_scale_grad); -/* Scaled activation + grouped quantize */ -py::object grouped_scaled_swiglu(const at::Tensor &input, const at::Tensor &act_scales, - py::handle quantizer, const size_t num_tensors, - std::optional first_dims, - std::optional tensor_offsets, - int64_t glu_interleave_size); - -py::object grouped_scaled_clamped_swiglu(const at::Tensor &input, const at::Tensor &act_scales, - py::handle quantizer, const size_t num_tensors, - std::optional first_dims, - std::optional tensor_offsets, float limit, - float alpha, float glu_linear_offset, - int64_t glu_interleave_size); - -py::object grouped_scaled_srelu(const at::Tensor &input, const at::Tensor &act_scales, - py::handle quantizer, const size_t num_tensors, - std::optional first_dims, - std::optional tensor_offsets); - -py::tuple grouped_scaled_dswiglu(const at::Tensor &grad, const at::Tensor &input, - const at::Tensor &act_scales, py::handle quantizer, - const size_t num_tensors, std::optional first_dims, - std::optional tensor_offsets, - int64_t glu_interleave_size, bool compute_scale_grad); - -py::tuple grouped_scaled_clamped_dswiglu(const at::Tensor &grad, const at::Tensor &input, - const at::Tensor &act_scales, py::handle quantizer, - const size_t num_tensors, - std::optional first_dims, - std::optional tensor_offsets, float limit, - float alpha, float glu_linear_offset, - int64_t glu_interleave_size, bool compute_scale_grad); - -py::tuple grouped_scaled_dsrelu(const at::Tensor &grad, const at::Tensor &input, - const at::Tensor &act_scales, py::handle quantizer, - const size_t num_tensors, std::optional first_dims, - std::optional tensor_offsets, bool compute_scale_grad); /*************************************************************************************************** * LayerNorm **************************************************************************************************/ diff --git a/transformer_engine/pytorch/csrc/extensions/activation.cpp b/transformer_engine/pytorch/csrc/extensions/activation.cpp index 7c486a522e..544ff92c1b 100644 --- a/transformer_engine/pytorch/csrc/extensions/activation.cpp +++ b/transformer_engine/pytorch/csrc/extensions/activation.cpp @@ -342,11 +342,7 @@ py::object clamped_dswiglu(const at::Tensor& grad, const at::Tensor& input, py:: glu_linear_offset); } -/* Scaled activation helpers (activation + per-row scale via nvte_scaled_*). - * - * Grouped variants reuse the dense compute helpers, then optionally apply - * group_quantize. Keep the nvte_scaled_* launch path in one place. - */ +/* Scaled activation helpers (activation + per-row scale via nvte_scaled_*). */ template at::Tensor scaled_activation_compute(const at::Tensor& input, const at::Tensor& act_scales, @@ -430,16 +426,6 @@ py::object maybe_quantize(const at::Tensor& tensor, py::handle quantizer) { return out_py; } -py::object maybe_group_quantize(const at::Tensor& tensor, py::handle quantizer, - const size_t num_tensors, std::optional first_dims, - std::optional tensor_offsets) { - if (quantizer.is_none()) { - return py::cast(tensor); - } - return group_quantize(tensor, quantizer, num_tensors, first_dims, std::nullopt, tensor_offsets, - std::nullopt); -} - template py::object scaled_activation_helper(const at::Tensor& input, const at::Tensor& act_scales, py::handle quantizer, int shape_divisor, Args&&... args) { @@ -458,37 +444,6 @@ py::tuple scaled_dactivation_helper(const at::Tensor& grad, const at::Tensor& in compute_scale_grad ? py::cast(grad_scales) : py::none()); } -template -py::object grouped_scaled_activation_helper(const at::Tensor& input, const at::Tensor& act_scales, - py::handle quantizer, const size_t num_tensors, - std::optional first_dims, - std::optional tensor_offsets, - int shape_divisor, Args&&... args) { - NVTE_CHECK(input.dim() == 2, "grouped scaled activation input must be 2D"); - auto output = scaled_activation_compute(input, act_scales, shape_divisor, - std::forward(args)...); - return maybe_group_quantize(output, quantizer, num_tensors, first_dims, tensor_offsets); -} - -template -py::tuple grouped_scaled_dactivation_helper(const at::Tensor& grad, const at::Tensor& input, - const at::Tensor& act_scales, py::handle quantizer, - const size_t num_tensors, - std::optional first_dims, - std::optional tensor_offsets, - bool compute_scale_grad, Args&&... args) { - NVTE_CHECK(input.dim() == 2 && grad.dim() == 2, - "grouped scaled dactivation input and grad must be 2D"); - auto [grad_input, grad_scales] = scaled_dactivation_compute( - grad, input, act_scales, compute_scale_grad, std::forward(args)...); - // Return both the (optionally) grouped-quantized grad input for the next - // grouped GEMM and the dense high-precision grad input so callers can reuse - // it (e.g. bias gradient) without a lossy dequantize. - return py::make_tuple( - maybe_group_quantize(grad_input, quantizer, num_tensors, first_dims, tensor_offsets), - py::cast(grad_input), compute_scale_grad ? py::cast(grad_scales) : py::none()); -} - py::object scaled_swiglu(const at::Tensor& input, const at::Tensor& act_scales, py::handle quantizer, int64_t glu_interleave_size) { return scaled_activation_helper(input, act_scales, quantizer, @@ -532,66 +487,5 @@ py::tuple scaled_dsrelu(const at::Tensor& grad, const at::Tensor& input, compute_scale_grad); } -py::object grouped_scaled_swiglu(const at::Tensor& input, const at::Tensor& act_scales, - py::handle quantizer, const size_t num_tensors, - std::optional first_dims, - std::optional tensor_offsets, - int64_t glu_interleave_size) { - return grouped_scaled_activation_helper( - input, act_scales, quantizer, num_tensors, first_dims, tensor_offsets, /*shape_divisor=*/2, - glu_interleave_size); -} - -py::object grouped_scaled_clamped_swiglu(const at::Tensor& input, const at::Tensor& act_scales, - py::handle quantizer, const size_t num_tensors, - std::optional first_dims, - std::optional tensor_offsets, float limit, - float alpha, float glu_linear_offset, - int64_t glu_interleave_size) { - return grouped_scaled_activation_helper( - input, act_scales, quantizer, num_tensors, first_dims, tensor_offsets, /*shape_divisor=*/2, - limit, alpha, glu_linear_offset, glu_interleave_size); -} - -py::object grouped_scaled_srelu(const at::Tensor& input, const at::Tensor& act_scales, - py::handle quantizer, const size_t num_tensors, - std::optional first_dims, - std::optional tensor_offsets) { - return grouped_scaled_activation_helper( - input, act_scales, quantizer, num_tensors, first_dims, tensor_offsets, - /*shape_divisor=*/1); -} - -py::tuple grouped_scaled_dswiglu(const at::Tensor& grad, const at::Tensor& input, - const at::Tensor& act_scales, py::handle quantizer, - const size_t num_tensors, std::optional first_dims, - std::optional tensor_offsets, - int64_t glu_interleave_size, bool compute_scale_grad) { - return grouped_scaled_dactivation_helper( - grad, input, act_scales, quantizer, num_tensors, first_dims, tensor_offsets, - compute_scale_grad, glu_interleave_size); -} - -py::tuple grouped_scaled_clamped_dswiglu(const at::Tensor& grad, const at::Tensor& input, - const at::Tensor& act_scales, py::handle quantizer, - const size_t num_tensors, - std::optional first_dims, - std::optional tensor_offsets, float limit, - float alpha, float glu_linear_offset, - int64_t glu_interleave_size, bool compute_scale_grad) { - return grouped_scaled_dactivation_helper( - grad, input, act_scales, quantizer, num_tensors, first_dims, tensor_offsets, - compute_scale_grad, limit, alpha, glu_linear_offset, glu_interleave_size); -} - -py::tuple grouped_scaled_dsrelu(const at::Tensor& grad, const at::Tensor& input, - const at::Tensor& act_scales, py::handle quantizer, - const size_t num_tensors, std::optional first_dims, - std::optional tensor_offsets, bool compute_scale_grad) { - return grouped_scaled_dactivation_helper(grad, input, act_scales, quantizer, - num_tensors, first_dims, - tensor_offsets, compute_scale_grad); -} - } // namespace pytorch } // namespace transformer_engine diff --git a/transformer_engine/pytorch/csrc/extensions/pybind.cpp b/transformer_engine/pytorch/csrc/extensions/pybind.cpp index 386125368a..e406dc3446 100644 --- a/transformer_engine/pytorch/csrc/extensions/pybind.cpp +++ b/transformer_engine/pytorch/csrc/extensions/pybind.cpp @@ -306,36 +306,6 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("scaled_dsrelu", transformer_engine::pytorch::scaled_dsrelu, "Scaled SReLU backward", py::arg("grad"), py::arg("fwd_input"), py::arg("act_scales"), py::arg("quantizer"), py::arg("compute_scale_grad") = true); - /* Scaled activation + grouped quantize */ - m.def("grouped_scaled_swiglu", transformer_engine::pytorch::grouped_scaled_swiglu, - "Scaled SwiGLU + grouped quantize", py::arg("input"), py::arg("act_scales"), - py::arg("quantizer"), py::arg("num_tensors"), py::arg("first_dims"), - py::arg("tensor_offsets") = py::none(), py::arg("glu_interleave_size") = 0); - m.def("grouped_scaled_clamped_swiglu", transformer_engine::pytorch::grouped_scaled_clamped_swiglu, - "Scaled clamped SwiGLU + grouped quantize", py::arg("input"), py::arg("act_scales"), - py::arg("quantizer"), py::arg("num_tensors"), py::arg("first_dims"), - py::arg("tensor_offsets") = py::none(), py::arg("limit") = 7.0f, py::arg("alpha") = 1.702f, - py::arg("glu_linear_offset") = 1.0f, py::arg("glu_interleave_size") = 0); - m.def("grouped_scaled_srelu", transformer_engine::pytorch::grouped_scaled_srelu, - "Scaled SReLU + grouped quantize", py::arg("input"), py::arg("act_scales"), - py::arg("quantizer"), py::arg("num_tensors"), py::arg("first_dims"), - py::arg("tensor_offsets") = py::none()); - m.def("grouped_scaled_dswiglu", transformer_engine::pytorch::grouped_scaled_dswiglu, - "Scaled SwiGLU backward + optional grouped quantize", py::arg("grad"), py::arg("fwd_input"), - py::arg("act_scales"), py::arg("quantizer"), py::arg("num_tensors"), py::arg("first_dims"), - py::arg("tensor_offsets") = py::none(), py::arg("glu_interleave_size") = 0, - py::arg("compute_scale_grad") = true); - m.def("grouped_scaled_clamped_dswiglu", - transformer_engine::pytorch::grouped_scaled_clamped_dswiglu, - "Scaled clamped SwiGLU backward + optional grouped quantize", py::arg("grad"), - py::arg("fwd_input"), py::arg("act_scales"), py::arg("quantizer"), py::arg("num_tensors"), - py::arg("first_dims"), py::arg("tensor_offsets") = py::none(), py::arg("limit") = 7.0f, - py::arg("alpha") = 1.702f, py::arg("glu_linear_offset") = 1.0f, - py::arg("glu_interleave_size") = 0, py::arg("compute_scale_grad") = true); - m.def("grouped_scaled_dsrelu", transformer_engine::pytorch::grouped_scaled_dsrelu, - "Scaled SReLU backward + optional grouped quantize", py::arg("grad"), py::arg("fwd_input"), - py::arg("act_scales"), py::arg("quantizer"), py::arg("num_tensors"), py::arg("first_dims"), - py::arg("tensor_offsets") = py::none(), py::arg("compute_scale_grad") = true); /* DBias + DAct fusions*/ m.def("dbias_dgelu", transformer_engine::pytorch::dbias_dgelu, "DGeLU + DBias + Quantize", py::arg("grad"), py::arg("fwd_input"), py::arg("quantizer")); diff --git a/transformer_engine/pytorch/ops/basic/grouped_linear.py b/transformer_engine/pytorch/ops/basic/grouped_linear.py index bc8dbed6c9..98c5a537c5 100644 --- a/transformer_engine/pytorch/ops/basic/grouped_linear.py +++ b/transformer_engine/pytorch/ops/basic/grouped_linear.py @@ -1064,7 +1064,7 @@ def fuser_forward( ) if use_grouped_tensor_path: - out, tensors_to_save = self._fuser_forward_graph_safe( + out, tensors_to_save = self._fuser_forward_grouped_tensor( input_=input_, split_sizes=split_sizes, scales=scales, @@ -1292,8 +1292,7 @@ def _fuser_forward_split_quantize( # (scales if scale_bias), *xs, *ws] # Offset metadata slots are unused on the split-quantize backward path # but are included as ``None`` so the saved-tensor layout matches the - # graph-safe ``_fuser_forward_grouped_tensor`` path (and fused forwards - # that share this GroupedLinear context contract). + # graph-safe ``_fuser_forward_grouped_tensor`` path. saved: list[Optional[torch.Tensor]] = [split_sizes, None, None, None, None] if self._scale_bias: saved.append(scales) @@ -1301,7 +1300,7 @@ def _fuser_forward_split_quantize( saved.extend(ws) return out, tuple(saved) - def _fuser_forward_graph_safe( + def _fuser_forward_grouped_tensor( self, *, input_: torch.Tensor, @@ -1326,10 +1325,13 @@ def _fuser_forward_graph_safe( dtypes=[torch.int32, torch.int64, torch.int64, torch.int64], bulk_allocate=True, ) + split_points = grouped_tensor_offsets[0] + base_split_offsets = grouped_tensor_offsets[1] input_tensor_offsets = grouped_tensor_offsets[2] + output_tensor_offsets = grouped_tensor_offsets[3] original_shape = list(input_.size()) - total_tokens = math.prod(original_shape[:-1]) x = maybe_dequantize(input_, dtype).reshape(-1, self.in_features) + total_tokens = x.size(0) if with_quantized_compute: input_quantizer = input_quantizers[0] input_quantizer.set_usage(rowwise=True, columnwise=weight_requires_grad) @@ -1351,59 +1353,7 @@ def _fuser_forward_graph_safe( first_dims=split_sizes, tensor_offsets=input_tensor_offsets, ) - return self._fuser_forward_grouped_tensor( - grouped_input=grouped_x, - split_sizes=split_sizes, - scales=scales, - with_quantized_compute=with_quantized_compute, - input_quantizers=input_quantizers, - weight_quantizers=weight_quantizers, - dtype=dtype, - input_requires_grad=input_requires_grad, - weight_requires_grad=weight_requires_grad, - device=device, - split_points=grouped_tensor_offsets[0], - base_split_offsets=grouped_tensor_offsets[1], - input_tensor_offsets=grouped_tensor_offsets[2], - output_tensor_offsets=grouped_tensor_offsets[3], - out_buffer=out_buffer, - out_shape=original_shape[:-1] + [self.out_features], - ) - - def _fuser_forward_grouped_tensor( - self, - *, - grouped_input: GroupedTensorStorage, - split_sizes: torch.Tensor, - scales: Optional[torch.Tensor], - with_quantized_compute: bool, - input_quantizers: list[Optional[Quantizer]], - weight_quantizers: list[Optional[Quantizer]], - dtype: torch.dtype, - input_requires_grad: bool, - weight_requires_grad: bool, - device: torch.device, - split_points: torch.Tensor, - base_split_offsets: torch.Tensor, - input_tensor_offsets: torch.Tensor, - output_tensor_offsets: torch.Tensor, - out_buffer: Optional[torch.Tensor] = None, - out_shape: list[int], - ) -> tuple[torch.Tensor, tuple[Optional[torch.Tensor], ...]]: - """Run grouped GEMM with a pre-built grouped input.""" - num_groups = self.num_groups has_bias = self.has_bias - total_tokens, in_features = grouped_input.logical_shape - expected_quantizer = input_quantizers[0] if with_quantized_compute else None - if grouped_input.quantizer is not expected_quantizer or in_features != self.in_features: - raise ValueError( - "GroupedLinear received an incompatible grouped input " - f"(quantizer={grouped_input.quantizer}, " - f"logical_shape={grouped_input.logical_shape}; " - f"expected quantizer={expected_quantizer}, " - f"in_features={self.in_features})" - ) - grouped_x = grouped_input if is_cpu_offload_enabled() and grouped_x is not None: start_offload(grouped_x) @@ -1429,6 +1379,7 @@ def _fuser_forward_grouped_tensor( ) # Allocate output buffer and wrap as a GroupedTensor view. + out_shape = original_shape[:-1] + [self.out_features] out = validate_or_alloc_output(out_buffer, out_shape, dtype, device) grouped_out = GroupedTensorStorage( shape=(total_tokens, self.out_features), @@ -1475,8 +1426,7 @@ def _fuser_forward_grouped_tensor( # input_tensor_offsets, output_tensor_offsets, # (scales if _scale_bias), grouped_x, *weights] # ``output_tensor_offsets`` matches the linear output row layout and is - # reused as ``grad_output`` offsets in backward (including fused - # activation + grouped linear backward). + # reused as ``grad_output`` offsets in backward. if grouped_x is not None: # (For FP8 per tensor current scaling on Hopper --> Free Rowwise Data # in backward pass) @@ -1513,7 +1463,7 @@ def fuser_backward( ctx = basic_op_ctxs[0] # Dispatch to the path used in forward (saved as ``ctx.use_grouped_tensor_path``). if getattr(ctx, "use_grouped_tensor_path", False): - return self._fuser_backward_graph_safe( + return self._fuser_backward_grouped_tensor( ctx=ctx, grad_output=grad_output, ) @@ -1542,8 +1492,7 @@ def _fuser_backward_split_quantize( # input_tensor_offsets, output_tensor_offsets, # (scales if _scale_bias), *xs, *ws] # Offset metadata beyond ``split_sizes`` is unused on this path but is - # present so the saved-tensor layout matches the graph-safe path (and - # fused forwards that share this GroupedLinear context contract). + # present so the saved-tensor layout matches the graph-safe path. saved_tensors = ctx.saved_tensors split_sizes = saved_tensors[0] saved_tensors = saved_tensors[5:] @@ -1709,7 +1658,7 @@ def _fuser_backward_split_quantize( grad_extra = (None, grad_scales) if self._scale_bias else (None,) return grad_input, [grad_params], [grad_extra] - def _fuser_backward_graph_safe( + def _fuser_backward_grouped_tensor( self, *, ctx: OperationContext, @@ -1719,15 +1668,36 @@ def _fuser_backward_graph_safe( Iterable[Iterable[Optional[torch.Tensor]]], Iterable[Iterable[Optional[torch.Tensor]]], ]: - """Build graph-safe grouped grad-output storage and run grouped GEMMs.""" + """Graph-safe GroupedTensor backward path.""" num_groups = self.num_groups has_bias = self.has_bias + weights, is_dist_weight, dist_dgrad_weights = self._backward_weight_setup() + device = weights[0].device dtype = ctx.dtype with_quantized_compute = bool(getattr(ctx, "with_quantized_compute", False)) - split_sizes = ctx.saved_tensors[0] - output_tensor_offsets = ctx.saved_tensors[4] + + # Saved tensors from forward pass. Layout: + # [split_sizes, base_split_offsets, split_points, + # input_tensor_offsets, output_tensor_offsets, + # (scales if _scale_bias), grouped_x, *weights] + saved_tensors = ctx.saved_tensors + split_sizes = saved_tensors[0] + base_split_offsets = saved_tensors[1] + input_tensor_offsets = saved_tensors[3] + output_tensor_offsets = saved_tensors[4] + saved_tensors = saved_tensors[5:] + scales = None + if self._scale_bias: + scales, saved_tensors = saved_tensors[0], saved_tensors[1:] + grouped_x, saved_tensors = saved_tensors[0], saved_tensors[1:] + if self.single_grouped_weight: + ws, saved_tensors = saved_tensors[0], saved_tensors[1:] + else: + ws, saved_tensors = saved_tensors[:num_groups], saved_tensors[num_groups:] + dy_2d = grad_output.reshape(-1, self.out_features) total_tokens = dy_2d.size(0) + grad_input_shape = list(grad_output.size())[:-1] + [self.in_features] dbias_packed = None if with_quantized_compute: @@ -1770,74 +1740,6 @@ def _fuser_backward_graph_safe( tensor_offsets=output_tensor_offsets, ) - return self._fuser_backward_grouped_tensor( - ctx=ctx, - grad_output=grad_output, - grouped_grad_output=grouped_dy, - dbias_packed=dbias_packed, - ) - - def _fuser_backward_grouped_tensor( - self, - *, - ctx: OperationContext, - grad_output: torch.Tensor, - grouped_grad_output: GroupedTensorStorage, - dbias_packed: Optional[torch.Tensor] = None, - ) -> tuple[ - torch.Tensor, - Iterable[Iterable[Optional[torch.Tensor]]], - Iterable[Iterable[Optional[torch.Tensor]]], - ]: - num_groups = self.num_groups - has_bias = self.has_bias - weights, is_dist_weight, dist_dgrad_weights = self._backward_weight_setup() - device = weights[0].device - dtype = ctx.dtype - - with_quantized_compute = bool(getattr(ctx, "with_quantized_compute", False)) - - # Saved tensors from forward pass - # Layout: [split_sizes, base_split_offsets, split_points, - # input_tensor_offsets, output_tensor_offsets, - # (scales if _scale_bias), grouped_x, *weights] - # ``split_points`` / ``output_tensor_offsets`` are unused on this path - # but are present so the saved-tensor layout matches the fused MLP / - # activation-fusion forwards that share this GroupedLinear context - # contract. - saved_tensors = ctx.saved_tensors - split_sizes = saved_tensors[0] - base_split_offsets = saved_tensors[1] - input_tensor_offsets = saved_tensors[3] - saved_tensors = saved_tensors[5:] - scales = None - if self._scale_bias: - scales, saved_tensors = saved_tensors[0], saved_tensors[1:] - grouped_x, saved_tensors = saved_tensors[0], saved_tensors[1:] - if self.single_grouped_weight: - ws, saved_tensors = saved_tensors[0], saved_tensors[1:] - else: - ws, saved_tensors = saved_tensors[:num_groups], saved_tensors[num_groups:] - - # Keep the dense high-precision grad for bias-gradient computation while - # optionally using a pre-quantized grouped view for the grouped GEMMs. - dy_2d = grad_output.reshape(-1, self.out_features) - total_tokens = dy_2d.size(0) - grad_input_shape = list(grad_output.size())[:-1] + [self.in_features] - - expected_quantizer = ctx.grad_output_quantizers[0] if with_quantized_compute else None - if grouped_grad_output.quantizer is not expected_quantizer or tuple( - grouped_grad_output.logical_shape - ) != (total_tokens, self.out_features): - raise ValueError( - "GroupedLinear received an incompatible grouped grad_output " - f"(quantizer={grouped_grad_output.quantizer}, " - f"logical_shape={grouped_grad_output.logical_shape}; " - f"expected quantizer={expected_quantizer}, " - f"logical_shape={(total_tokens, self.out_features)})" - ) - grouped_dy = grouped_grad_output - # Bias Grads compute if not already computed in bgrad_group_quantize final_bias_grads: Optional[torch.Tensor] = None grad_scales: Optional[torch.Tensor] = None diff --git a/transformer_engine/pytorch/ops/fused/__init__.py b/transformer_engine/pytorch/ops/fused/__init__.py index 85191a9807..dc9dcd6dc3 100644 --- a/transformer_engine/pytorch/ops/fused/__init__.py +++ b/transformer_engine/pytorch/ops/fused/__init__.py @@ -6,14 +6,9 @@ from ..fuser import register_backward_fusion, register_forward_fusion from .backward_activation_bias import BackwardActivationBias -from .backward_activation_grouped_linear import BackwardScaledActivationGroupedLinear from .backward_add_rmsnorm import BackwardAddRMSNorm from .backward_linear_add import BackwardLinearAdd from .backward_linear_scale import BackwardLinearScale -from .forward_activation_grouped_linear import ( - ForwardScaledActivationGroupedLinear, - act_grouped_linear_fusion_supported, -) from .forward_linear_bias_activation import ForwardLinearBiasActivation from .forward_linear_bias_add import ForwardLinearBiasAdd from .forward_linear_scale_add import ForwardLinearScaleAdd @@ -26,7 +21,6 @@ register_forward_fusion(ForwardLinearBiasAdd.fuse_forward_ops) register_forward_fusion(ForwardLinearBiasActivation.fuse_forward_ops) register_forward_fusion(ForwardLinearScaleAdd.fuse_forward_ops) -register_forward_fusion(ForwardScaledActivationGroupedLinear.fuse_forward_ops) # Register backward fusions register_backward_fusion(UserbuffersBackwardLinear.fuse_backward_ops) @@ -34,7 +28,6 @@ register_backward_fusion(BackwardLinearScale.fuse_backward_ops) register_backward_fusion(BackwardActivationBias.fuse_backward_ops) register_backward_fusion(BackwardAddRMSNorm.fuse_backward_ops) -register_backward_fusion(BackwardScaledActivationGroupedLinear.fuse_backward_ops) # Import experimental fusions # Note: Registration logic is non-trivial, so submodule handles it internally. diff --git a/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py b/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py deleted file mode 100644 index 0cb17fe0d1..0000000000 --- a/transformer_engine/pytorch/ops/fused/backward_activation_grouped_linear.py +++ /dev/null @@ -1,181 +0,0 @@ -# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# -# See LICENSE for license information. - -"""Fused scaled activation + grouped linear backward.""" - -from __future__ import annotations - -from collections.abc import Iterable -from typing import Optional - -import torch - -import transformer_engine_torch as tex -from ...quantization import Recipe -from ...tensor import Quantizer -from ...utils import clear_tensor_data -from .._common import maybe_dequantize -from ..basic import GroupedLinear, ScaledClampedQGeGLU, ScaledSReLU, ScaledSwiGLU -from ..op import FusedOperation, FusibleOperation, OperationContext -from .forward_activation_grouped_linear import ( - _SCALED_ACTIVATION_TYPES, - _ScaledActivation, - act_grouped_linear_fusion_supported, -) - - -def _grouped_scaled_dactivation( - activation: _ScaledActivation, - grad_output: torch.Tensor, - input_: torch.Tensor, - scales: torch.Tensor, - *, - quantizer: Quantizer, - num_groups: int, - split_sizes: torch.Tensor, - tensor_offsets: torch.Tensor, - compute_scale_grad: bool, -) -> tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]: - """Dispatch a grouped scaled activation backward pass.""" - dy = grad_output.reshape(-1, grad_output.size(-1)) - x = input_.reshape(-1, input_.size(-1)) - s = scales.reshape(-1) - if isinstance(activation, ScaledSwiGLU): - return tex.grouped_scaled_dswiglu( - dy, - x, - s, - quantizer, - num_groups, - split_sizes, - tensor_offsets, - int(activation.glu_interleave_size or 0), - compute_scale_grad, - ) - if isinstance(activation, ScaledClampedQGeGLU): - clamped = activation._clamped - return tex.grouped_scaled_clamped_dswiglu( - dy, - x, - s, - quantizer, - num_groups, - split_sizes, - tensor_offsets, - clamped.limit, - clamped.alpha, - clamped.glu_linear_offset, - int(activation.glu_interleave_size or 0), - compute_scale_grad, - ) - if isinstance(activation, ScaledSReLU): - return tex.grouped_scaled_dsrelu( - dy, - x, - s, - quantizer, - num_groups, - split_sizes, - tensor_offsets, - compute_scale_grad, - ) - raise TypeError(f"Unsupported scaled activation type ({type(activation).__name__})") - - -class BackwardScaledActivationGroupedLinear(FusedOperation): - """Scaled activation backward + grouped quantize + grouped linear backward.""" - - def __init__(self, *, linear: GroupedLinear, activation: _ScaledActivation) -> None: - super().__init__((linear, activation)) - - def fuser_backward( - self, - basic_op_ctxs: list[OperationContext], - grad_output: torch.Tensor, - *, - basic_op_grad_extra_outputs: list[tuple[torch.Tensor, ...]], - ) -> tuple[ - Optional[torch.Tensor], - Iterable[Iterable[Optional[torch.Tensor]]], - Iterable[Iterable[Optional[torch.Tensor]]], - ]: - linear = self.basic_ops[0] - activation = self.basic_ops[1] - linear_ctx, activation_ctx = basic_op_ctxs - - if not linear_ctx.requires_grad: - _, _, act_grad_extra_inputs = activation.fuser_backward( - [activation_ctx], - grad_output, - basic_op_grad_extra_outputs=[basic_op_grad_extra_outputs[1]], - ) - return None, [(), ()], [(), act_grad_extra_inputs[0]] - - del basic_op_grad_extra_outputs - input_, scales = activation_ctx.saved_tensors - input_ = maybe_dequantize(input_, activation_ctx.dtype) - scales = maybe_dequantize(scales, activation_ctx.dtype) - grad_output = maybe_dequantize(grad_output, activation_ctx.dtype) - - split_sizes = linear_ctx.saved_tensors[0] - activation_input_tensor_offsets = linear_ctx.saved_tensors[4] - grad_output_quantizer = linear_ctx.grad_output_quantizers[0] - grad_output_quantizer.set_usage( - rowwise=linear_ctx.input_requires_grad, - columnwise=linear_ctx.weight_requires_grad, - ) - grad_output_quantizer.optimize_for_gemm = True - grouped_dy, dense_dy, grad_scales = _grouped_scaled_dactivation( - activation, - grad_output, - input_, - scales, - quantizer=grad_output_quantizer, - num_groups=linear.num_groups, - split_sizes=split_sizes, - tensor_offsets=activation_input_tensor_offsets, - compute_scale_grad=activation_ctx.extra_input_requires_grad, - ) - - grad_input, grad_params, grad_extra_inputs = linear._fuser_backward_grouped_tensor( - ctx=linear_ctx, - grad_output=dense_dy, - grouped_grad_output=grouped_dy, - ) - - clear_tensor_data(activation_ctx.saved_tensors[0]) - return ( - grad_input, - [grad_params[0], ()], - [grad_extra_inputs[0], (grad_scales,)], - ) - - @staticmethod - def fuse_backward_ops( - ops: list[FusibleOperation], - *, - recipe: Optional[Recipe] = None, - **unused, # pylint: disable=unused-argument - ) -> list[FusibleOperation]: - """Fuse each supported GroupedLinear + ScaledActivation pair.""" - out: list[FusibleOperation] = [] - idx = 0 - while idx < len(ops): - if ( - idx + 1 < len(ops) - and isinstance(ops[idx], GroupedLinear) - and isinstance(ops[idx + 1], _SCALED_ACTIVATION_TYPES) - and act_grouped_linear_fusion_supported(ops[idx], ops[idx + 1], recipe) - ): - out.append( - BackwardScaledActivationGroupedLinear( - linear=ops[idx], - activation=ops[idx + 1], - ) - ) - idx += 2 - else: - out.append(ops[idx]) - idx += 1 - return out diff --git a/transformer_engine/pytorch/ops/fused/forward_activation_grouped_linear.py b/transformer_engine/pytorch/ops/fused/forward_activation_grouped_linear.py deleted file mode 100644 index 4f0b0d627c..0000000000 --- a/transformer_engine/pytorch/ops/fused/forward_activation_grouped_linear.py +++ /dev/null @@ -1,229 +0,0 @@ -# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# -# See LICENSE for license information. - -"""Fused scaled activation + grouped linear forward.""" - -from __future__ import annotations - -from collections.abc import Iterable -from typing import Any, Optional - -import torch - -import transformer_engine_torch as tex -from ...cpu_offload import is_cpu_offload_enabled, mark_activation_offload -from ...quantization import Recipe -from ...tensor import Quantizer -from .._common import maybe_dequantize -from ..basic import GroupedLinear, ScaledClampedQGeGLU, ScaledSReLU, ScaledSwiGLU -from ..basic.activation import _ScaledUnary -from ..basic.swiglu import _ScaledGLU -from ..op import FusedOperation, FusibleOperation, OperationContext - - -_ScaledActivation = _ScaledGLU | _ScaledUnary -_SCALED_ACTIVATION_TYPES = (_ScaledGLU, _ScaledUnary) - - -def _grouped_scaled_activation( - activation: _ScaledActivation, - input_: torch.Tensor, - scales: torch.Tensor, - quantizer: Quantizer, - num_groups: int, - split_sizes: torch.Tensor, - tensor_offsets: torch.Tensor, -) -> torch.Tensor: - """Dispatch a grouped scaled activation.""" - x = input_.reshape(-1, input_.size(-1)) - s = scales.reshape(-1) - if isinstance(activation, ScaledSwiGLU): - return tex.grouped_scaled_swiglu( - x, - s, - quantizer, - num_groups, - split_sizes, - tensor_offsets, - int(activation.glu_interleave_size or 0), - ) - if isinstance(activation, ScaledClampedQGeGLU): - clamped = activation._clamped - return tex.grouped_scaled_clamped_swiglu( - x, - s, - quantizer, - num_groups, - split_sizes, - tensor_offsets, - clamped.limit, - clamped.alpha, - clamped.glu_linear_offset, - int(activation.glu_interleave_size or 0), - ) - if isinstance(activation, ScaledSReLU): - return tex.grouped_scaled_srelu( - x, - s, - quantizer, - num_groups, - split_sizes, - tensor_offsets, - ) - raise TypeError(f"Unsupported scaled activation type ({type(activation).__name__})") - - -def act_grouped_linear_fusion_supported( - linear: GroupedLinear, - activation: _ScaledActivation, - recipe: Optional[Recipe], -) -> bool: - """Whether ScaledActivation + GroupedLinear can use grouped quantized compute.""" - if recipe is None or activation.activation_recompute_in_mlp: - return False - input_quantizers = [ - linear.get_quantizer("forward", 2 * group_idx) for group_idx in range(linear.num_groups) - ] - weight = linear.weight if linear.single_grouped_weight else linear.weight0 - dtype = torch.get_autocast_dtype("cuda") if torch.is_autocast_enabled() else weight.dtype - return linear._is_graph_safe_path_supported( - with_quantized_compute=True, - input_quantizers=input_quantizers, - dtype=dtype, - single_grouped_weight=linear.single_grouped_weight, - ) - - -class ForwardScaledActivationGroupedLinear(FusedOperation): - """Scaled activation + grouped quantize + grouped linear forward.""" - - def __init__(self, *, activation: _ScaledActivation, linear: GroupedLinear) -> None: - super().__init__((activation, linear)) - - def fuser_forward( - self, - basic_op_ctxs: list[OperationContext], - input_: torch.Tensor, - *, - basic_op_extra_inputs: list[tuple[torch.Tensor, ...]], - prev_op_grad_output_quantizer: Optional[Quantizer], - next_op_input_quantizer: Optional[Quantizer], - basic_op_kwargs: list[dict[str, Any]], - ) -> tuple[torch.Tensor, Iterable[Iterable[torch.Tensor]]]: - activation = self.basic_ops[0] - linear = self.basic_ops[1] - activation_ctx, linear_ctx = basic_op_ctxs - if basic_op_kwargs[0] or basic_op_kwargs[1]: - raise ValueError("Scaled activation and GroupedLinear do not expect keyword arguments") - - weight = linear.weight if linear.single_grouped_weight else linear.weight0 - device = weight.device - dtype = torch.get_autocast_dtype("cuda") if torch.is_autocast_enabled() else weight.dtype - input_ = maybe_dequantize(input_, dtype) - scales = maybe_dequantize(basic_op_extra_inputs[0][0], dtype) - - split_sizes = basic_op_extra_inputs[1][0] - if int(split_sizes.numel()) != linear.num_groups: - raise ValueError( - f"Expected {linear.num_groups} splits, but got {int(split_sizes.numel())}." - ) - split_sizes = split_sizes.to(device=device, dtype=torch.int64) - linear_scales = basic_op_extra_inputs[1][1] if linear._scale_bias else None - split_sizes, grouped_tensor_offsets = tex.splits_to_offsets_multi( - split_sizes, - device, - strides=[1, 1, linear.in_features, linear.out_features], - include_leading_zero=[False, True, True, True], - dtypes=[torch.int32, torch.int64, torch.int64, torch.int64], - bulk_allocate=True, - ) - - input_quantizers = [ - linear.get_quantizer("forward", 2 * group_idx) for group_idx in range(linear.num_groups) - ] - weight_quantizers = [ - linear.get_quantizer("forward", 2 * group_idx + 1) - for group_idx in range(linear.num_groups) - ] - input_quantizer = input_quantizers[0] - weight_requires_grad = linear_ctx.requires_grad and weight.requires_grad - input_quantizer.set_usage(rowwise=True, columnwise=weight_requires_grad) - input_quantizer.optimize_for_gemm = True - - grouped_x = _grouped_scaled_activation( - activation, - input_, - scales, - input_quantizer, - linear.num_groups, - split_sizes, - grouped_tensor_offsets[2], - ) - - if activation_ctx.requires_grad: - if is_cpu_offload_enabled(): - mark_activation_offload(input_) - activation_ctx.input_requires_grad = True - activation_ctx.extra_input_requires_grad = basic_op_extra_inputs[0][0].requires_grad - activation_ctx.dtype = dtype - activation_ctx.save_for_backward(input_, scales) - - out, tensors_to_save = linear._fuser_forward_grouped_tensor( - grouped_input=grouped_x, - split_sizes=split_sizes, - scales=linear_scales, - with_quantized_compute=True, - input_quantizers=input_quantizers, - weight_quantizers=weight_quantizers, - dtype=dtype, - input_requires_grad=linear_ctx.requires_grad, - weight_requires_grad=weight_requires_grad, - device=device, - split_points=grouped_tensor_offsets[0], - base_split_offsets=grouped_tensor_offsets[1], - input_tensor_offsets=grouped_tensor_offsets[2], - output_tensor_offsets=grouped_tensor_offsets[3], - out_shape=list(input_.size())[:-1] + [linear.out_features], - ) - linear.fuser_forward_save_ctx( - basic_op_ctxs=[linear_ctx], - input_=input_, - tensors_to_save=[tensors_to_save], - requires_grad=[linear_ctx.requires_grad], - basic_op_extra_inputs=[basic_op_extra_inputs[1]], - prev_op_grad_output_quantizer=prev_op_grad_output_quantizer, - next_op_input_quantizer=next_op_input_quantizer, - basic_op_kwargs=[basic_op_kwargs[1]], - use_grouped_tensor_path=True, - ) - return out, [(), ()] - - @staticmethod - def fuse_forward_ops( - ops: list[FusibleOperation], - *, - recipe: Optional[Recipe] = None, - **unused, # pylint: disable=unused-argument - ) -> list[FusibleOperation]: - """Fuse each supported ScaledActivation + GroupedLinear pair.""" - out: list[FusibleOperation] = [] - idx = 0 - while idx < len(ops): - if ( - idx + 1 < len(ops) - and isinstance(ops[idx], _SCALED_ACTIVATION_TYPES) - and isinstance(ops[idx + 1], GroupedLinear) - and act_grouped_linear_fusion_supported(ops[idx + 1], ops[idx], recipe) - ): - out.append( - ForwardScaledActivationGroupedLinear( - activation=ops[idx], - linear=ops[idx + 1], - ) - ) - idx += 2 - else: - out.append(ops[idx]) - idx += 1 - return out From 9c0c9f2e8f5ae6955623ec5ac99a5a94219ce850 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Wed, 5 Aug 2026 19:09:38 +0000 Subject: [PATCH 16/17] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- tests/pytorch/test_grouped_mlp.py | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/pytorch/test_grouped_mlp.py b/tests/pytorch/test_grouped_mlp.py index 9c9a7065a2..28f7360641 100644 --- a/tests/pytorch/test_grouped_mlp.py +++ b/tests/pytorch/test_grouped_mlp.py @@ -73,6 +73,7 @@ if nvfp4_available: _grouped_mlp_quantization_list.append("nvfp4_rht") + @pytest.fixture(autouse=True, scope="function") def _reset_rng_states_per_test(): """Restore torch, CUDA, and Python ``random`` before each test in this module.""" From 07034ab891083a6859c8b65a5b0938bfaed53a05 Mon Sep 17 00:00:00 2001 From: Varun Thumbe Date: Wed, 5 Aug 2026 21:24:24 +0000 Subject: [PATCH 17/17] improve comments Signed-off-by: Varun Thumbe --- transformer_engine/pytorch/ops/basic/grouped_linear.py | 10 ++++++++++ 1 file changed, 10 insertions(+) diff --git a/transformer_engine/pytorch/ops/basic/grouped_linear.py b/transformer_engine/pytorch/ops/basic/grouped_linear.py index 98c5a537c5..c94d78fecd 100644 --- a/transformer_engine/pytorch/ops/basic/grouped_linear.py +++ b/transformer_engine/pytorch/ops/basic/grouped_linear.py @@ -1330,8 +1330,10 @@ def _fuser_forward_grouped_tensor( input_tensor_offsets = grouped_tensor_offsets[2] output_tensor_offsets = grouped_tensor_offsets[3] original_shape = list(input_.size()) + # Flatten to 2D so the first dim is the total token count. x = maybe_dequantize(input_, dtype).reshape(-1, self.in_features) total_tokens = x.size(0) + # Build the input GroupedTensorStorage for input. if with_quantized_compute: input_quantizer = input_quantizers[0] input_quantizer.set_usage(rowwise=True, columnwise=weight_requires_grad) @@ -1344,6 +1346,7 @@ def _fuser_forward_grouped_tensor( tensor_offsets=input_tensor_offsets, ) else: + # No quantize: wrap the contiguous high-precision buffer. grouped_x = GroupedTensorStorage( shape=(total_tokens, self.in_features), dtype=dtype, @@ -1658,6 +1661,9 @@ def _fuser_backward_split_quantize( grad_extra = (None, grad_scales) if self._scale_bias else (None,) return grad_input, [grad_params], [grad_extra] + # ================================================================== + # Graph-safe backward: counterpart of `_fuser_forward_grouped_tensor`. + # ================================================================== def _fuser_backward_grouped_tensor( self, *, @@ -1695,10 +1701,14 @@ def _fuser_backward_grouped_tensor( else: ws, saved_tensors = saved_tensors[:num_groups], saved_tensors[num_groups:] + # Flatten grad_output to 2D (total_tokens, out_features) + # to figure out total tokens. dy_2d = grad_output.reshape(-1, self.out_features) total_tokens = dy_2d.size(0) grad_input_shape = list(grad_output.size())[:-1] + [self.in_features] + # Build the grad_output GroupedTensor. + # Optionally get dbias is fusion available with bgrad_group_quantize dbias_packed = None if with_quantized_compute: grad_output_quantizer = ctx.grad_output_quantizers[0]