From 05d1ab51e87bc1a9e3f38065de1c646073ac23a1 Mon Sep 17 00:00:00 2001 From: Brayden Zhong Date: Sat, 9 May 2026 06:42:56 -0400 Subject: [PATCH] Enable PDL for various kernels in DSV32/GLM5 (#23965) Co-authored-by: b8zhong --- .../srt/layers/attention/nsa/nsa_indexer.py | 8 +++- python/sglang/srt/layers/attention/utils.py | 3 ++ python/sglang/srt/mem_cache/utils.py | 21 ++++++++++ sgl-kernel/csrc/elementwise/pos_enc.cuh | 8 ++-- sgl-kernel/csrc/gemm/dsv3_fused_a_gemm.cu | 4 +- .../csrc/gemm/dsv3_router_gemm_bf16_out.cu | 4 +- .../csrc/gemm/dsv3_router_gemm_float_out.cu | 4 +- .../gemm/per_token_group_quant_8bit_v2.cu | 41 ++++++++++++++----- 8 files changed, 71 insertions(+), 22 deletions(-) diff --git a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py index 84a32b30c..28854f2f6 100644 --- a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py +++ b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py @@ -299,6 +299,12 @@ class Indexer(MultiPlatformOp): weights = weights.unsqueeze(-1) * q_scale * self.softmax_scale return weights + @torch.compile(dynamic=True) + def _apply_q_scale_and_softmax_scale( + self, weights: torch.Tensor, q_scale: torch.Tensor + ): + return weights.unsqueeze(-1) * q_scale * self.softmax_scale + def _get_q_k_bf16( self, q_lora: torch.Tensor, @@ -1161,7 +1167,7 @@ class Indexer(MultiPlatformOp): act_quant=act_quant, ) current_stream.wait_stream(self.alt_stream) - weights = weights.unsqueeze(-1) * q_scale * self.softmax_scale + weights = self._apply_q_scale_and_softmax_scale(weights, q_scale) else: query, key = self._get_q_k_bf16( q_lora, x, positions, enable_dual_stream, forward_batch=forward_batch diff --git a/python/sglang/srt/layers/attention/utils.py b/python/sglang/srt/layers/attention/utils.py index e0774c9a4..277d46054 100644 --- a/python/sglang/srt/layers/attention/utils.py +++ b/python/sglang/srt/layers/attention/utils.py @@ -12,6 +12,8 @@ _is_cuda = is_cuda() if _is_cuda: from sgl_kernel import concat_mla_absorb_q +from sglang.jit_kernel.utils import is_arch_support_pdl + @triton.jit def create_flashinfer_kv_indices_triton( @@ -462,6 +464,7 @@ def mla_quantize_and_rope_for_fp8( # Quantization scales (set to 1.0 for no additional scaling) quant_scale_q=1.0, quant_scale_kv=1.0, + enable_pdl=is_arch_support_pdl(), ) return q_out, k_nope_out, k_rope_out diff --git a/python/sglang/srt/mem_cache/utils.py b/python/sglang/srt/mem_cache/utils.py index 65b7b165c..bf443661d 100644 --- a/python/sglang/srt/mem_cache/utils.py +++ b/python/sglang/srt/mem_cache/utils.py @@ -20,6 +20,7 @@ import torch import triton import triton.language as tl +from sglang.jit_kernel.utils import is_arch_support_pdl from sglang.srt.environ import envs @@ -35,6 +36,7 @@ def set_mla_kv_buffer_kernel( nope_dim: tl.constexpr, rope_dim: tl.constexpr, BLOCK: tl.constexpr, + USE_GDC: tl.constexpr = False, ): pid_loc = tl.program_id(0) pid_blk = tl.program_id(1) @@ -44,6 +46,9 @@ def set_mla_kv_buffer_kernel( total_dim = nope_dim + rope_dim mask = offs < total_dim + if USE_GDC: + tl.extra.cuda.gdc_wait() + loc = tl.load(loc_ptr + pid_loc).to(tl.int64) dst_ptr = kv_buffer_ptr + loc * buffer_stride + offs @@ -82,6 +87,9 @@ def set_mla_kv_buffer_kernel( tl.store(dst_ptr, src, mask=mask) + if USE_GDC: + tl.extra.cuda.gdc_launch_dependents() + def set_mla_kv_buffer_triton( kv_buffer: torch.Tensor, @@ -96,6 +104,8 @@ def set_mla_kv_buffer_triton( n_loc = loc.numel() grid = (n_loc, triton.cdiv(total_dim, BLOCK)) + pdl_kwargs = {"USE_GDC": True, "launch_pdl": True} if is_arch_support_pdl() else {} + set_mla_kv_buffer_kernel[grid]( kv_buffer, cache_k_nope, @@ -107,6 +117,7 @@ def set_mla_kv_buffer_triton( nope_dim, rope_dim, BLOCK=BLOCK, + **pdl_kwargs, ) @@ -122,6 +133,7 @@ def set_mla_kv_buffer_fp8_quant_kernel( nope_dim: tl.constexpr, rope_dim: tl.constexpr, BLOCK: tl.constexpr, + USE_GDC: tl.constexpr = False, ): """Fuse BF16/FP16->FP8 cast with paged KV write.""" pid_loc = tl.program_id(0) @@ -132,6 +144,9 @@ def set_mla_kv_buffer_fp8_quant_kernel( total_dim = nope_dim + rope_dim mask = offs < total_dim + if USE_GDC: + tl.extra.cuda.gdc_wait() + loc = tl.load(loc_ptr + pid_loc).to(tl.int64) dst_ptr = kv_buffer_fp8_ptr + loc * buffer_stride + offs @@ -165,6 +180,9 @@ def set_mla_kv_buffer_fp8_quant_kernel( # Destination pointer is FP8-typed view; tl.store performs downcast. tl.store(dst_ptr, src, mask=mask) + if USE_GDC: + tl.extra.cuda.gdc_launch_dependents() + def set_mla_kv_buffer_triton_fp8_quant( kv_buffer: torch.Tensor, @@ -183,6 +201,8 @@ def set_mla_kv_buffer_triton_fp8_quant( n_loc = loc.numel() grid = (n_loc, triton.cdiv(total_dim, BLOCK)) + pdl_kwargs = {"USE_GDC": True, "launch_pdl": True} if is_arch_support_pdl() else {} + set_mla_kv_buffer_fp8_quant_kernel[grid]( kv_buffer_fp8, cache_k_nope, @@ -194,6 +214,7 @@ def set_mla_kv_buffer_triton_fp8_quant( nope_dim, rope_dim, BLOCK=BLOCK, + **pdl_kwargs, ) diff --git a/sgl-kernel/csrc/elementwise/pos_enc.cuh b/sgl-kernel/csrc/elementwise/pos_enc.cuh index a2e4e2ebb..34124c3e7 100644 --- a/sgl-kernel/csrc/elementwise/pos_enc.cuh +++ b/sgl-kernel/csrc/elementwise/pos_enc.cuh @@ -105,7 +105,7 @@ __global__ void BatchQKApplyRotaryPosIdsCosSinCacheEnhancedHeadParallelismKernel const uint32_t bdy = blockDim.y; #if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) - asm volatile("griddepcontrol.wait;"); + cudaGridDependencySynchronize(); #endif vec_t cos, sin; @@ -184,7 +184,7 @@ __global__ void BatchQKApplyRotaryPosIdsCosSinCacheEnhancedHeadParallelismKernel } #if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) - asm volatile("griddepcontrol.launch_dependents;"); + cudaTriggerProgrammaticLaunchCompletion(); #endif } @@ -229,7 +229,7 @@ __global__ void BatchQKApplyRotaryPosIdsCosSinCacheEnhancedKernel( const uint32_t bdy = blockDim.y; #if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) - asm volatile("griddepcontrol.wait;"); + cudaGridDependencySynchronize(); #endif vec_t cos, sin; @@ -310,7 +310,7 @@ __global__ void BatchQKApplyRotaryPosIdsCosSinCacheEnhancedKernel( } #if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) - asm volatile("griddepcontrol.launch_dependents;"); + cudaTriggerProgrammaticLaunchCompletion(); #endif } diff --git a/sgl-kernel/csrc/gemm/dsv3_fused_a_gemm.cu b/sgl-kernel/csrc/gemm/dsv3_fused_a_gemm.cu index 4dc0a796a..c393b5a58 100644 --- a/sgl-kernel/csrc/gemm/dsv3_fused_a_gemm.cu +++ b/sgl-kernel/csrc/gemm/dsv3_fused_a_gemm.cu @@ -285,7 +285,7 @@ struct GmemLoaderB { __device__ void issue_mainloop() { #if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900 - asm volatile("griddepcontrol.wait;"); + cudaGridDependencySynchronize(); #pragma unroll 1 for (int loop_idx = 0; loop_idx < k_iter_cnt; loop_idx++) { if (need_wait) { @@ -571,7 +571,7 @@ __global__ __launch_bounds__(256, 1) void fused_a_gemm_kernel( mma_computer.issue_mainloop(); mma_computer.epi(); } - asm volatile("griddepcontrol.launch_dependents;"); + cudaTriggerProgrammaticLaunchCompletion(); #endif } diff --git a/sgl-kernel/csrc/gemm/dsv3_router_gemm_bf16_out.cu b/sgl-kernel/csrc/gemm/dsv3_router_gemm_bf16_out.cu index e613bd75c..e60a83db6 100644 --- a/sgl-kernel/csrc/gemm/dsv3_router_gemm_bf16_out.cu +++ b/sgl-kernel/csrc/gemm/dsv3_router_gemm_bf16_out.cu @@ -75,7 +75,7 @@ __launch_bounds__(128, 1) void router_gemm_kernel_bf16_output(__nv_bfloat16* out } #if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) - asm volatile("griddepcontrol.wait;"); + cudaGridDependencySynchronize(); #endif // Process the GEMM in chunks @@ -159,7 +159,7 @@ __launch_bounds__(128, 1) void router_gemm_kernel_bf16_output(__nv_bfloat16* out } } #if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) - asm volatile("griddepcontrol.launch_dependents;"); + cudaTriggerProgrammaticLaunchCompletion(); #endif } diff --git a/sgl-kernel/csrc/gemm/dsv3_router_gemm_float_out.cu b/sgl-kernel/csrc/gemm/dsv3_router_gemm_float_out.cu index 88a364e2c..0abcaf1c1 100644 --- a/sgl-kernel/csrc/gemm/dsv3_router_gemm_float_out.cu +++ b/sgl-kernel/csrc/gemm/dsv3_router_gemm_float_out.cu @@ -74,7 +74,7 @@ __global__ __launch_bounds__(128, 1) void router_gemm_kernel_float_output(float* } #if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) - asm volatile("griddepcontrol.wait;"); + cudaGridDependencySynchronize(); #endif // Process the GEMM in chunks @@ -158,7 +158,7 @@ __global__ __launch_bounds__(128, 1) void router_gemm_kernel_float_output(float* } } #if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) - asm volatile("griddepcontrol.launch_dependents;"); + cudaTriggerProgrammaticLaunchCompletion(); #endif } diff --git a/sgl-kernel/csrc/gemm/per_token_group_quant_8bit_v2.cu b/sgl-kernel/csrc/gemm/per_token_group_quant_8bit_v2.cu index c886569f9..4fbaf08ef 100644 --- a/sgl-kernel/csrc/gemm/per_token_group_quant_8bit_v2.cu +++ b/sgl-kernel/csrc/gemm/per_token_group_quant_8bit_v2.cu @@ -271,6 +271,10 @@ __global__ void per_token_group_quant_8bit_kernel( using scale_element_t = std::conditional_t; static_assert(sizeof(scale_packed_t) % sizeof(scale_element_t) == 0); +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900 + cudaGridDependencySynchronize(); +#endif + SCHEDULER::execute( subwarps_per_block, hidden_dim_num_groups, @@ -398,6 +402,10 @@ __global__ void per_token_group_quant_8bit_kernel( reinterpret_cast(output_q + offset_num_groups * GROUP_SIZE + lane_id * INPUT_PRIMARY_VEC_SIZE), output_buf); }); + +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900 + cudaTriggerProgrammaticLaunchCompletion(); +#endif } void sgl_per_token_group_quant_8bit_v2( @@ -445,17 +453,28 @@ void sgl_per_token_group_quant_8bit_v2( SCHEDULER::compute_exec_config( \ THREADS_PER_SUBWARP, num_local_experts, hidden_dim_num_groups, num_groups, subwarps_per_block, grid, block); \ \ - per_token_group_quant_8bit_kernel \ - <<>>( \ - static_cast(input.data_ptr()), \ - static_cast(output_q.data_ptr()), \ - static_cast(output_s.data_ptr()), \ - static_cast(masked_m.has_value() ? masked_m->data_ptr() : 0), \ - subwarps_per_block, \ - hidden_dim_num_groups, \ - scale_expert_stride, \ - scale_hidden_stride, \ - num_tokens_per_expert); \ + cudaLaunchConfig_t config; \ + config.gridDim = grid; \ + config.blockDim = block; \ + config.dynamicSmemBytes = 0; \ + config.stream = stream; \ + cudaLaunchAttribute attrs[1]; \ + attrs[0].id = cudaLaunchAttributeProgrammaticStreamSerialization; \ + attrs[0].val.programmaticStreamSerializationAllowed = getEnvEnablePDL(); \ + config.numAttrs = 1; \ + config.attrs = attrs; \ + cudaLaunchKernelEx( \ + &config, \ + per_token_group_quant_8bit_kernel, \ + static_cast(input.data_ptr()), \ + static_cast(output_q.data_ptr()), \ + static_cast(output_s.data_ptr()), \ + static_cast(masked_m.has_value() ? masked_m->data_ptr() : 0), \ + subwarps_per_block, \ + hidden_dim_num_groups, \ + scale_expert_stride, \ + scale_hidden_stride, \ + num_tokens_per_expert); \ } while (0) #define LAUNCH_KERNEL(GROUP_SIZE, T, DST_DTYPE) \