diff --git a/transformer_engine/pytorch/ops/fused/grouped_mlp.py b/transformer_engine/pytorch/ops/fused/grouped_mlp.py index 58f22c874c..bd8b82d3fa 100644 --- a/transformer_engine/pytorch/ops/fused/grouped_mlp.py +++ b/transformer_engine/pytorch/ops/fused/grouped_mlp.py @@ -239,10 +239,11 @@ def _group_quantize_for_grouped_mlp( split_sizes: Optional[torch.Tensor], *, tensor_offsets: Optional[torch.Tensor] = None, + use_dense_single_group: bool = False, ) -> GroupedTensor: """Quantize into grouped storage.""" - if num_groups != 1 or not isinstance(quantizer, (MXFP8Quantizer, NVFP4Quantizer)): + if not use_dense_single_group: return tex.group_quantize( tensor, quantizer, @@ -270,6 +271,7 @@ def _group_quantize_with_amax_for_grouped_mlp( columnwise_amax: torch.Tensor, *, tensor_offsets: Optional[torch.Tensor] = None, + use_dense_single_group: bool = False, ) -> GroupedTensor: """Quantize with precomputed NVFP4 amaxes into grouped storage.""" if not isinstance(quantizer, NVFP4Quantizer): @@ -279,9 +281,10 @@ def _group_quantize_with_amax_for_grouped_mlp( num_groups, split_sizes, tensor_offsets=tensor_offsets, + use_dense_single_group=use_dense_single_group, ) - if num_groups != 1: + if not use_dense_single_group: return tex.nvfp4_group_quantize_with_amax( tensor, quantizer, @@ -690,6 +693,7 @@ def _compute_grad_params( scale_view_dtype, sf_vec_size, offsets, + use_dense_single_group, ): """Compute weight gradients and build grad_params for a GroupedLinear layer. Returns the grad_params list in parameter registration order. @@ -757,7 +761,7 @@ def _compute_grad_params( "distributed-weight fused grouped-MLP requires delay_wgrad_compute=False." ) if ( - num_groups == 1 + use_dense_single_group and isinstance(grouped_x, (GroupedTensor, GroupedTensorStorage)) and isinstance(grouped_dy, (GroupedTensor, GroupedTensorStorage)) and isinstance(grouped_x.quantizer, (MXFP8Quantizer, NVFP4Quantizer)) @@ -1125,6 +1129,9 @@ def fuser_forward( raise ValueError(f"Unsupported input shape for fused grouped MLP ({in_shape=}).") num_groups = fc1_op.num_groups + use_dense_single_group = num_groups == 1 and ( + _cudnn_frontend_supports_single_group_runtime_offsets(type(activation_op)) + ) fc1_weight_param = fc1_op.weight if fc1_op.single_grouped_weight else fc1_op.weight0 fc2_weight_param = fc2_op.weight if fc2_op.single_grouped_weight else fc2_op.weight0 @@ -1195,7 +1202,7 @@ def fuser_forward( # Older cuDNN frontends do not expose this specialization, so use the # live generic offset calculation rather than caching CUDA metadata. use_offsetless_metadata = ( - num_groups == 1 + use_dense_single_group and unit_activation_scale and isinstance(fc1_input_quantizer, MXFP8Quantizer) and supports_single_group_runtime_offsets @@ -1276,6 +1283,7 @@ def fuser_forward( fc1_weight_quantizer, num_groups, None, + use_dense_single_group=use_dense_single_group, ) else: fc1_weights = [getattr(fc1_op, f"weight{idx}") for idx in range(num_groups)] @@ -1313,6 +1321,7 @@ def fuser_forward( fc2_weight_quantizer, num_groups, None, + use_dense_single_group=use_dense_single_group, ) else: fc2_weights = [getattr(fc2_op, f"weight{idx}") for idx in range(num_groups)] @@ -1360,6 +1369,7 @@ def fuser_forward( num_groups, split_sizes, tensor_offsets=fc1_x_tensor_offsets, + use_dense_single_group=use_dense_single_group, ) use_nvfp4 = isinstance(fc1_input_quantizer, NVFP4Quantizer) or isinstance( @@ -1495,7 +1505,7 @@ def fuser_forward( fc1_activation_kwargs["norm_const_tensor"] = fc1_norm_const_tensor fc1_activation_kwargs["discrete_col_sfd"] = not use_nvfp4 if supports_single_group_runtime_offsets: - fc1_activation_kwargs["use_single_group_runtime_offsets"] = num_groups == 1 + fc1_activation_kwargs["use_single_group_runtime_offsets"] = use_dense_single_group if self._pass_geglu_runtime_params: fc1_activation_kwargs.update( linear_offset=self._cudnn_linear_offset, @@ -1514,7 +1524,7 @@ def fuser_forward( if fc1_op.single_grouped_weight: # Clone and swizzle scales for GEMM. fc1_weight_for_gemm = grouped_fc1_weight.copy() - use_single_group_weight_swizzle = num_groups == 1 + use_single_group_weight_swizzle = use_dense_single_group if use_single_group_weight_swizzle: fc1_weight_single = _single_quantized_tensor_from_grouped(fc1_weight_for_gemm) fc1_weight_single._columnwise_data = None @@ -1551,7 +1561,7 @@ def fuser_forward( fc1_activation_kwargs["b_tensor"] = fc1_w_data fc1_activation_kwargs["sfb_tensor"] = fc1_w_scales else: - use_single_discrete_weight = num_groups == 1 + use_single_discrete_weight = use_dense_single_group if use_single_discrete_weight: fc1_weight_single = grouped_fc1_weight[0] original_rowwise_scale = fc1_weight_single._rowwise_scale_inv @@ -1653,6 +1663,7 @@ def fuser_forward( fc1_kernel_out["amax_tensor"].view(-1), fc1_kernel_out["post_rht_amax_tensor"].view(-1), tensor_offsets=fc2_x_tensor_offsets, + use_dense_single_group=use_dense_single_group, ) else: grouped_fc2_x = _group_quantize_for_grouped_mlp( @@ -1661,11 +1672,12 @@ def fuser_forward( num_groups, split_sizes, tensor_offsets=fc2_x_tensor_offsets, + use_dense_single_group=use_dense_single_group, ) fc2_out_buf = validate_or_alloc_output(output_buffer, fc2_out_shape, dtype, device) if ( - num_groups == 1 + use_dense_single_group and grouped_fc2_x.columnwise_data is not None and grouped_fc2_x.columnwise_scale_inv is not None ): @@ -1723,7 +1735,7 @@ def fuser_forward( with_gemm_swizzled_scales=True, ) - use_single_group_dense_fc2 = num_groups == 1 + use_single_group_dense_fc2 = use_dense_single_group fc2_out_buf = validate_or_alloc_output(output_buffer, fc2_out_shape, dtype, device) if use_single_group_dense_fc2: fc2_out = _single_group_fc2_gemm( @@ -1764,7 +1776,7 @@ def fuser_forward( } fc2_quant_kernel = self.grouped_gemm_quant_kernel() if supports_single_group_runtime_offsets: - fc2_quant_kwargs["use_single_group_runtime_offsets"] = num_groups == 1 + fc2_quant_kwargs["use_single_group_runtime_offsets"] = use_dense_single_group if fc2_op.single_grouped_weight: # Clone and swizzle scales for GEMM @@ -1935,6 +1947,9 @@ def fuser_backward( grad_output = grad_output.reshape(-1, fc2_weight_shape[0]) out_shape = list(grad_output.size()) num_groups = fc1_op.num_groups + use_dense_single_group = num_groups == 1 and ( + _cudnn_frontend_supports_single_group_runtime_offsets(type(activation_op)) + ) fc1_weight_param = fc1_op.weight if fc1_op.single_grouped_weight else fc1_op.weight0 fc2_weight_param = fc2_op.weight if fc2_op.single_grouped_weight else fc2_op.weight0 device = fc1_weight_param.device @@ -2041,6 +2056,7 @@ def fuser_backward( num_groups, split_sizes, tensor_offsets=fc2_out_tensor_offsets, + use_dense_single_group=use_dense_single_group, ) use_nvfp4 = ( @@ -2194,7 +2210,7 @@ def fuser_backward( # Never passed to a wrapper that would reject it -- the check above raises first. fc2_dactivation_kwargs["deterministic"] = True if _cudnn_frontend_supports_single_group_runtime_offsets(type(activation_op)): - fc2_dactivation_kwargs["use_single_group_runtime_offsets"] = num_groups == 1 + fc2_dactivation_kwargs["use_single_group_runtime_offsets"] = use_dense_single_group if self._cudnn_dact_func is not None: fc2_dactivation_kwargs["beta_tensor"] = fc2_beta_tensor fc2_dactivation_kwargs["act_func"] = self._cudnn_dact_func @@ -2248,7 +2264,7 @@ def fuser_backward( fc2_dactivation_kwargs["b_tensor"] = fc2_w_data fc2_dactivation_kwargs["sfb_tensor"] = fc2_w_scales else: - use_single_discrete_weight = num_groups == 1 + use_single_discrete_weight = use_dense_single_group if use_single_discrete_weight: fc2_weight_single = grouped_fc2_weight[0] original_rowwise_data = fc2_weight_single._rowwise_data @@ -2353,6 +2369,7 @@ def fuser_backward( num_groups, split_sizes, tensor_offsets=fc2_x_tensor_offsets, + use_dense_single_group=use_dense_single_group, ) else: sfd_col_d_srelu_tensor = fc2_dgrad_kernel_out.get("sfd_col_d_srelu_tensor") @@ -2430,6 +2447,7 @@ def fuser_backward( num_groups, split_sizes, tensor_offsets=fc1_dy_tensor_offsets, + use_dense_single_group=use_dense_single_group, ) else: grouped_fc1_dy = GroupedTensor( @@ -2466,6 +2484,7 @@ def fuser_backward( scale_view_dtype=scale_view_dtype, sf_vec_size=sf_vec_size, offsets=split_points, + use_dense_single_group=use_dense_single_group, ) # Clear FC2 input tensor if possible @@ -2491,7 +2510,7 @@ def fuser_backward( if is_distributed_weight(fc1_leader): grouped_fc1_weight = materialize_weight_for_backward(fc1_leader) - use_single_group_dense_dgrad = num_groups == 1 + use_single_group_dense_dgrad = use_dense_single_group if use_single_group_dense_dgrad: grad_input = validate_or_alloc_output(grad_input_buffer, in_shape, dtype, device) _single_group_dgrad_gemm( @@ -2545,7 +2564,7 @@ def fuser_backward( } fc1_dgrad_kernel = self.grouped_gemm_quant_kernel() if _cudnn_frontend_supports_single_group_runtime_offsets(type(activation_op)): - fc1_dgrad_kwargs["use_single_group_runtime_offsets"] = num_groups == 1 + fc1_dgrad_kwargs["use_single_group_runtime_offsets"] = use_dense_single_group if fc1_op.single_grouped_weight: # Clone and swizzle scales for GEMM @@ -2622,6 +2641,7 @@ def fuser_backward( scale_view_dtype=scale_view_dtype, sf_vec_size=sf_vec_size, offsets=split_points, + use_dense_single_group=use_dense_single_group, ) # Clear FC1 input tensor if possible