diff --git a/tests/pytorch/test_grouped_mlp.py b/tests/pytorch/test_grouped_mlp.py index 76550db5f8..5a7f280962 100644 --- a/tests/pytorch/test_grouped_mlp.py +++ b/tests/pytorch/test_grouped_mlp.py @@ -23,6 +23,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, @@ -72,6 +78,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(): @@ -724,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: @@ -999,14 +1021,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 +1037,24 @@ 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} @@ -1076,6 +1117,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( @@ -1148,6 +1208,8 @@ 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") @@ -1466,6 +1528,8 @@ 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") @@ -1597,6 +1661,8 @@ 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.h b/transformer_engine/pytorch/csrc/extensions.h index 6edfbdc00e..f479d75c13 100644 --- a/transformer_engine/pytorch/csrc/extensions.h +++ b/transformer_engine/pytorch/csrc/extensions.h @@ -290,6 +290,44 @@ 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..bce92fa3ab 100644 --- a/transformer_engine/pytorch/csrc/extensions/activation.cpp +++ b/transformer_engine/pytorch/csrc/extensions/activation.cpp @@ -342,5 +342,151 @@ 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); + + // 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(); + 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..944152f10b 100644 --- a/transformer_engine/pytorch/csrc/extensions/pybind.cpp +++ b/transformer_engine/pytorch/csrc/extensions/pybind.cpp @@ -284,6 +284,36 @@ 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/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 5ef0fa4339..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,7 +1241,7 @@ def _fuser_forward_split_quantize( saved.extend(ws) return out, tuple(saved) - def _fuser_forward_grouped_tensor( + def _fuser_forward_graph_safe( self, *, input_: torch.Tensor, @@ -1256,30 +1256,32 @@ def _fuser_forward_grouped_tensor( device: torch.device, out_buffer: Optional[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. - """ + """Build graph-safe grouped input storage and run grouped GEMM.""" 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. + 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) - 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) + 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, @@ -1287,8 +1289,59 @@ def _fuser_forward_grouped_tensor( quantizer=None, data=x.reshape(-1), first_dims=split_sizes, - tensor_offsets=base_split_offsets * self.in_features, + 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) @@ -1314,7 +1367,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 +1375,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 @@ -1388,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, ) @@ -1577,10 +1629,7 @@ 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( + def _fuser_backward_graph_safe( self, *, ctx: OperationContext, @@ -1590,63 +1639,44 @@ def _fuser_backward_grouped_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 - weights = self._get_weight_tensors() - 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, - # (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). - saved_tensors = ctx.saved_tensors - split_sizes = saved_tensors[0] - base_split_offsets = saved_tensors[1] - saved_tensors = saved_tensors[3:] - 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:] - - # Flatten grad_output to 2D (total_tokens, out_features) - # to figure out total tokens. + 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) - # 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] grad_output_quantizer.set_usage( - rowwise=ctx.input_requires_grad, columnwise=ctx.weight_requires_grad + 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 + dy_2d, + grad_output_quantizer, + num_groups, + split_sizes, ) else: grouped_dy = tex.group_quantize( - dy_2d, grad_output_quantizer, num_groups, split_sizes + 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, @@ -1657,6 +1687,71 @@ def _fuser_backward_grouped_tensor( 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: 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 = self._get_weight_tensors() + 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, + # (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). + saved_tensors = ctx.saved_tensors + split_sizes = saved_tensors[0] + base_split_offsets = saved_tensors[1] + saved_tensors = saved_tensors[3:] + 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 @@ -1681,7 +1776,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..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..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/forward_activation_grouped_linear.py b/transformer_engine/pytorch/ops/fused/forward_activation_grouped_linear.py new file mode 100644 index 0000000000..f004eb943f --- /dev/null +++ b/transformer_engine/pytorch/ops/fused/forward_activation_grouped_linear.py @@ -0,0 +1,228 @@ +# 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], + 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, + ) -> 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