From 6b255351883e3b79eadf62c0a0294955677f46b3 Mon Sep 17 00:00:00 2001 From: mingyangw Date: Mon, 8 Jun 2026 17:09:46 -0700 Subject: [PATCH] make dgeglu config values compile time constants instead of runtime values --- .../grouped_gemm/grouped_gemm_dglu/api.py | 187 ++++++++---------- ...moe_blockscaled_grouped_gemm_dglu_dbias.py | 29 ++- 2 files changed, 91 insertions(+), 125 deletions(-) diff --git a/python/cudnn/grouped_gemm/grouped_gemm_dglu/api.py b/python/cudnn/grouped_gemm/grouped_gemm_dglu/api.py index 07ea46ce7..5a6a77dfa 100644 --- a/python/cudnn/grouped_gemm/grouped_gemm_dglu/api.py +++ b/python/cudnn/grouped_gemm/grouped_gemm_dglu/api.py @@ -59,7 +59,6 @@ from cudnn.datatypes import _convert_to_cutlass_data_type from cudnn.api_base import APIBase, TupleDict, ceil_div, is_power_of_2 - class GroupedGemmDgluSm100(APIBase): """Unified API for grouped GEMM dGLU backward operation on SM100+ GPUs. @@ -137,6 +136,10 @@ def __init__( b_major: str = "k", epilogue_op: Optional[str] = None, use_dynamic_sched: bool = False, + linear_offset: Optional[float] = None, + geglu_alpha: float = 1.702, + glu_clamp_max: float = 7.0, + glu_clamp_min: float = -7.0, ): """Initialize the GroupedGemmDgluSm100 API. @@ -171,6 +174,15 @@ def __init__( :param b_major: Major dimension for B tensor, one of "k" or "n" :param epilogue_op: Optional epilogue operation. Valid: None, "none", "identity", "relu", "srelu" :param use_dynamic_sched: Enable dynamic tile scheduling for load balancing + :param linear_offset: Compile-time linear offset for dGeGLU. When None, + defaults to 1.0 for dGeGLU and 0.0 for dSwiGLU. + Ignored when ``act_func == "dswiglu"``. + :param geglu_alpha: Compile-time dGeGLU GeGLU alpha. Ignored when + ``act_func == "dswiglu"``. + :param glu_clamp_max: Compile-time dGeGLU upper clamp. Ignored when + ``act_func == "dswiglu"``. + :param glu_clamp_min: Compile-time dGeGLU lower clamp. Ignored when + ``act_func == "dswiglu"``. """ super().__init__() @@ -254,6 +266,13 @@ def __init__( raise ValueError(f"Invalid epilogue operation: {epilogue_op}. " f"Valid values: None, 'none', 'identity', 'relu', 'srelu'") self.use_dynamic_sched = use_dynamic_sched + if linear_offset is None: + self.linear_offset = 1.0 if self.act_func == "dgeglu" else 0.0 + else: + self.linear_offset = float(linear_offset) + self.geglu_alpha = geglu_alpha + self.glu_clamp_max = glu_clamp_max + self.glu_clamp_min = glu_clamp_min self._interpret_uint8_as_fp4x2 = True self._has_dbias = self.dbias_desc is not None @@ -854,10 +873,6 @@ def _compile_dense(self, gemm_dglu, max_active_clusters, fake_stream) -> None: ) # Compile with keyword args (dense mode uses the unified __call__ positional order). - # linear_offset, geglu_alpha, glu_clamp_max, and glu_clamp_min are runtime - # cutlass.Float32 (not Constexpr), so the placeholder values below are - # irrelevant -- the values passed through tensor_api() at execute() time are - # what the kernel actually uses. dbias_fake = self._make_fake_cute_tensor_from_desc(self.dbias_desc, assumed_align=16) _compiled_kernel = cute.compile( @@ -884,13 +899,13 @@ def _compile_dense(self, gemm_dglu, max_active_clusters, fake_stream) -> None: prob=prob_cute_fake, dprob=dprob_cute_fake, dbias_tensor=dbias_fake, - linear_offset=cutlass.Float32(0.0), max_active_clusters=max_active_clusters, stream=fake_stream, epilogue_op=self.epilogue_op, - geglu_alpha=cutlass.Float32(1.702), - glu_clamp_max=cutlass.Float32(7.0), - glu_clamp_min=cutlass.Float32(-7.0), + linear_offset=self.linear_offset, + geglu_alpha=self.geglu_alpha, + glu_clamp_max=self.glu_clamp_max, + glu_clamp_min=self.glu_clamp_min, options="--enable-tvm-ffi", ) @@ -916,10 +931,6 @@ def tensor_api( dprob_tensor: torch.Tensor, dbias_tensor: Optional[torch.Tensor], stream: cuda.CUstream, - linear_offset: float = 0.0, - geglu_alpha: float = 1.702, - glu_clamp_max: float = 7.0, - glu_clamp_min: float = -7.0, ) -> None: norm_const_tensor = self._unpad_tensor_to_ndim(norm_const_tensor, 1, "norm_const") _compiled_kernel( @@ -943,12 +954,8 @@ def tensor_api( beta_tensor, prob_tensor, dprob_tensor, - cutlass.Float32(linear_offset), dbias_tensor, stream, - cutlass.Float32(geglu_alpha), - cutlass.Float32(glu_clamp_max), - cutlass.Float32(glu_clamp_min), ) self._compiled_kernel = tensor_api @@ -1053,42 +1060,38 @@ def _compile_discrete(self, gemm_dglu, max_active_clusters, fake_stream) -> None workspace_ptr_cute = from_dlpack(self._workspace, assumed_align=128).iterator - # linear_offset, geglu_alpha, glu_clamp_max, and glu_clamp_min are runtime - # cutlass.Float32 (not Constexpr), so the placeholders below are irrelevant - # -- the values passed through tensor_api() at execute() time are what the - # kernel actually uses. self._logger.debug("Compiling discrete grouped GEMM dGLU kernel") _compiled_kernel = cute.compile( gemm_dglu, - a_tensor, - b_ptrs_cute, - sfb_ptrs_cute, - cutlass.Int32(n), - cutlass.Int32(k), - cutlass.Int64(b_stride_size), - b_major_mode, - workspace_ptr_cute, - c_tensor, - d_row_tensor, - d_col_tensor, - sfa_tensor, - sfd_row_tensor, - sfd_col_tensor, - amax_tensor, - norm_const_tensor_cute, - padded_offsets_tensor, - alpha_tensor, - beta_tensor, - prob_tensor, - dprob_tensor, - cutlass.Float32(0.0), - dbias_tensor, - max_active_clusters, - fake_stream, - self.epilogue_op, - cutlass.Float32(1.702), - cutlass.Float32(7.0), - cutlass.Float32(-7.0), + a=a_tensor, + b=b_ptrs_cute, + sfb=sfb_ptrs_cute, + n=cutlass.Int32(n), + k=cutlass.Int32(k), + b_stride_size=cutlass.Int64(b_stride_size), + b_major_mode=b_major_mode, + workspace_ptr=workspace_ptr_cute, + c=c_tensor, + d=d_row_tensor, + d_col=d_col_tensor, + sfa=sfa_tensor, + sfd_row_tensor=sfd_row_tensor, + sfd_col_tensor=sfd_col_tensor, + amax_tensor=amax_tensor, + norm_const_tensor=norm_const_tensor_cute, + padded_offsets=padded_offsets_tensor, + alpha=alpha_tensor, + beta=beta_tensor, + prob=prob_tensor, + dprob=dprob_tensor, + dbias_tensor=dbias_tensor, + max_active_clusters=max_active_clusters, + stream=fake_stream, + epilogue_op=self.epilogue_op, + linear_offset=self.linear_offset, + geglu_alpha=self.geglu_alpha, + glu_clamp_max=self.glu_clamp_max, + glu_clamp_min=self.glu_clamp_min, options="--enable-tvm-ffi", ) @@ -1121,10 +1124,6 @@ def tensor_api( dprob_tensor: torch.Tensor, dbias_tensor: Optional[torch.Tensor], stream: cuda.CUstream, - linear_offset: float = 0.0, - geglu_alpha: float = 1.702, - glu_clamp_max: float = 7.0, - glu_clamp_min: float = -7.0, ) -> None: norm_const_tensor = self._unpad_tensor_to_ndim(norm_const_tensor, 1, "norm_const") b_ptrs_addr = int(b_ptrs_device.data_ptr()) @@ -1151,12 +1150,8 @@ def tensor_api( beta_tensor, prob_tensor, dprob_tensor, - cutlass.Float32(linear_offset), dbias_tensor, stream, - cutlass.Float32(geglu_alpha), - cutlass.Float32(glu_clamp_max), - cutlass.Float32(glu_clamp_min), ) self._compiled_kernel = tensor_api @@ -1189,10 +1184,6 @@ def execute( sfd_col_tensor: Optional[torch.Tensor] = None, amax_tensor: Optional[torch.Tensor] = None, norm_const_tensor: Optional[torch.Tensor] = None, - linear_offset: Optional[float] = None, - geglu_alpha: float = 1.702, - glu_clamp_max: float = 7.0, - glu_clamp_min: float = -7.0, current_stream: Optional[cuda.CUstream] = None, ) -> None: """Execute the compiled kernel. @@ -1219,20 +1210,6 @@ def execute( :param sfd_col_tensor: Optional column scale factor D :param amax_tensor: Optional amax tensor :param norm_const_tensor: Optional normalization constant - :param linear_offset: Linear offset matching the forward GeGLU activation. - Affects ``act_func == "dgeglu"``; ignored when - ``act_func == "dswiglu"``. When ``None`` (default), the offset is - chosen based on ``act_func`` for backwards compatibility: - ``1.0`` for ``"dgeglu"`` and ``0.0`` for ``"dswiglu"``. - :param geglu_alpha: Pre-sigmoid scaling factor for the GeGLU activation - being differentiated. Must match the value used in the forward; - defaults to ``1.702``. Ignored when ``act_func == "dswiglu"``. - :param glu_clamp_max: Upper clamp limit for ``up`` and ``gate`` in the - forward GeGLU. Default ``7.0``. The same limit also drives the - gradient mask. Ignored when ``act_func == "dswiglu"``. - :param glu_clamp_min: Lower clamp limit applied only to ``up`` in the - forward GeGLU. Default ``-7.0``. Ignored when - ``act_func == "dswiglu"``. :param current_stream: CUDA stream """ self._logger.debug("Entering execute") @@ -1246,12 +1223,6 @@ def execute( "Kernel not compiled; call compile() first", ) - # Resolve linear_offset default: None -> activation-derived legacy value - # (1.0 for dgeglu, 0.0 for dswiglu) for backwards compatibility with - # callers that pre-date the explicit linear_offset kwarg. - if linear_offset is None: - linear_offset = 1.0 if self.act_func == "dgeglu" else 0.0 - self._logger.debug("Executing grouped GEMM dGLU kernel") if self._has_dbias: self._value_error_if( @@ -1279,10 +1250,6 @@ def execute( dprob_tensor=dprob_tensor, dbias_tensor=dbias_tensor, stream=current_stream, - linear_offset=linear_offset, - geglu_alpha=geglu_alpha, - glu_clamp_max=glu_clamp_max, - glu_clamp_min=glu_clamp_min, ) else: self._compiled_kernel( @@ -1304,10 +1271,6 @@ def execute( dprob_tensor=dprob_tensor, dbias_tensor=dbias_tensor, stream=current_stream, - linear_offset=linear_offset, - geglu_alpha=geglu_alpha, - glu_clamp_max=glu_clamp_max, - glu_clamp_min=glu_clamp_min, ) self._logger.debug("Execute completed") @@ -1404,21 +1367,16 @@ def grouped_gemm_dglu_wrapper_sm100( ``act_func == "dgeglu"``; ignored when ``act_func == "dswiglu"``. When ``None`` (default), the offset is chosen based on ``act_func`` for backwards compatibility: ``1.0`` for ``"dgeglu"`` and ``0.0`` - for ``"dswiglu"``. Runtime parameter -- a single compiled kernel - serves any value, and ``linear_offset`` is intentionally not part - of the cache key. + for ``"dswiglu"``. geglu_alpha: Pre-sigmoid scaling factor for the GeGLU activation being differentiated. Must match the value used in the forward. - Default ``1.702``. Runtime parameter, intentionally not part of - the cache key. Ignored when ``act_func == "dswiglu"``. + Default ``1.702``. Ignored when ``act_func == "dswiglu"``. glu_clamp_max: Upper clamp limit applied to ``up`` and ``gate`` in the forward GeGLU; the same limit drives the gradient mask here. - Default ``7.0``. Runtime parameter, intentionally not part of the - cache key. Ignored when ``act_func == "dswiglu"``. + Default ``7.0``. Ignored when ``act_func == "dswiglu"``. glu_clamp_min: Lower clamp limit applied to ``up`` only in the forward GeGLU; the same limit drives the gradient mask here. - Default ``-7.0``. Runtime parameter, intentionally not part of the - cache key. Ignored when ``act_func == "dswiglu"``. + Default ``-7.0``. Ignored when ``act_func == "dswiglu"``. epilogue_op: Optional epilogue operation. Valid: None, "none", "identity", "relu", "srelu" use_dynamic_sched: Enable dynamic tile scheduling for load balancing current_stream: CUDA stream @@ -1429,11 +1387,18 @@ def grouped_gemm_dglu_wrapper_sm100( """ from cudnn.discrete_grouped_gemm.discrete_kernel_utils import _require_pointer_tensor - # Resolve linear_offset default: None means "use the activation-derived legacy - # default" (1.0 for dgeglu, 0.0 for dswiglu) for backwards compatibility with - # callers that have not been updated to pass linear_offset explicitly. + # Resolve linear_offset default: None means "use the activation-derived + # default" (1.0 for dgeglu, 0.0 for dswiglu). if linear_offset is None: linear_offset = 1.0 if act_func == "dgeglu" else 0.0 + dgeglu_cache_signature = None + if act_func == "dgeglu": + dgeglu_cache_signature = ( + float(linear_offset), + float(geglu_alpha), + float(glu_clamp_max), + float(glu_clamp_min), + ) # ---- Auto-detect weight mode ---- is_dense = b_tensor is not None @@ -1539,6 +1504,7 @@ def dynamic_m_tensor_signature( cache_key = ( weight_mode, act_func, + dgeglu_cache_signature, epilogue_op, use_full_dynamic, a_tensor.shape[1:] if not use_full_dynamic else None, @@ -1582,6 +1548,7 @@ def dynamic_m_tensor_signature( cache_key = ( weight_mode, act_func, + dgeglu_cache_signature, epilogue_op, *dynamic_m_tensor_signature(a_tensor, tuple(a_tensor.shape[1:]), dynamic_stride_dims=(2,)), b_shape, @@ -1652,6 +1619,10 @@ def dynamic_m_tensor_signature( act_func=act_func, epilogue_op=epilogue_op, use_dynamic_sched=use_dynamic_sched, + linear_offset=linear_offset, + geglu_alpha=geglu_alpha, + glu_clamp_max=glu_clamp_max, + glu_clamp_min=glu_clamp_min, ) else: api = GroupedGemmDgluSm100( @@ -1684,6 +1655,10 @@ def dynamic_m_tensor_signature( b_major=b_major, epilogue_op=epilogue_op, use_dynamic_sched=use_dynamic_sched, + linear_offset=linear_offset, + geglu_alpha=geglu_alpha, + glu_clamp_max=glu_clamp_max, + glu_clamp_min=glu_clamp_min, ) if not api.check_support(): @@ -1711,10 +1686,6 @@ def dynamic_m_tensor_signature( sfd_col_tensor=sfd_col_tensor, amax_tensor=amax_tensor, norm_const_tensor=norm_const_tensor, - linear_offset=linear_offset, - geglu_alpha=geglu_alpha, - glu_clamp_max=glu_clamp_max, - glu_clamp_min=glu_clamp_min, current_stream=current_stream, ) else: @@ -1736,10 +1707,6 @@ def dynamic_m_tensor_signature( sfd_col_tensor=sfd_col_tensor, amax_tensor=amax_tensor, norm_const_tensor=norm_const_tensor, - linear_offset=linear_offset, - geglu_alpha=geglu_alpha, - glu_clamp_max=glu_clamp_max, - glu_clamp_min=glu_clamp_min, current_stream=current_stream, ) diff --git a/python/cudnn/grouped_gemm/grouped_gemm_dglu/moe_blockscaled_grouped_gemm_dglu_dbias.py b/python/cudnn/grouped_gemm/grouped_gemm_dglu/moe_blockscaled_grouped_gemm_dglu_dbias.py index b31ced401..cd6f102b0 100644 --- a/python/cudnn/grouped_gemm/grouped_gemm_dglu/moe_blockscaled_grouped_gemm_dglu_dbias.py +++ b/python/cudnn/grouped_gemm/grouped_gemm_dglu/moe_blockscaled_grouped_gemm_dglu_dbias.py @@ -744,14 +744,14 @@ def __call__( beta: cute.Tensor, prob: cute.Tensor, dprob: cute.Tensor, - linear_offset: Float32, dbias_tensor: Optional[cute.Tensor], max_active_clusters: cutlass.Constexpr, stream: cuda.CUstream, epilogue_op: cutlass.Constexpr = lambda x: x, - geglu_alpha: Float32 = cutlass.Float32(1.702), - glu_clamp_max: Float32 = cutlass.Float32(7.0), - glu_clamp_min: Float32 = cutlass.Float32(-7.0), + linear_offset: cutlass.Constexpr = 1.0, + geglu_alpha: cutlass.Constexpr = 1.702, + glu_clamp_max: cutlass.Constexpr = 7.0, + glu_clamp_min: cutlass.Constexpr = -7.0, ): """Execute the GEMM. @@ -761,9 +761,8 @@ def __call__( ``b_major_mode`` describe the uniform per-expert layout. ``linear_offset``, ``geglu_alpha``, ``glu_clamp_max``, and - ``glu_clamp_min`` are runtime ``cutlass.Float32`` parameters that - configure the GeGLU activation that this kernel differentiates -- the - forward computed + ``glu_clamp_min`` are compile-time constants that configure the GeGLU + activation that this kernel differentiates -- the forward computed out = (clamp(up, min=glu_clamp_min, max=glu_clamp_max) + linear_offset) * silu(geglu_alpha * clamp(gate, max=glu_clamp_max)) and the backward consumes the same values plus the corresponding @@ -2069,10 +2068,10 @@ def kernel( beta: cute.Tensor, prob: cute.Tensor, dprob: cute.Tensor, - linear_offset: Float32, - geglu_alpha: Float32, - glu_clamp_max: Float32, - glu_clamp_min: Float32, + linear_offset: cutlass.Constexpr, + geglu_alpha: cutlass.Constexpr, + glu_clamp_max: cutlass.Constexpr, + glu_clamp_min: cutlass.Constexpr, mDbias_tensor: Optional[cute.Tensor], workspace_ptr, cluster_layout_vmnk: cute.Layout, @@ -3032,10 +3031,10 @@ def kernel( ab2_vec_load, mProb, square_alpha, - linear_offset, - geglu_alpha, - glu_clamp_max, - glu_clamp_min, + cutlass.Float32(linear_offset), + cutlass.Float32(geglu_alpha), + cutlass.Float32(glu_clamp_max), + cutlass.Float32(glu_clamp_min), dprob_swiglu, )