diff --git a/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py b/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py index 6aef866f5..c0c98a607 100644 --- a/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py +++ b/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py @@ -39,6 +39,10 @@ os.environ["DG_JIT_CACHE_DIR"] = os.getenv( # And NVCC JIT speed is also 9x faster in the ref commit os.environ["DG_JIT_USE_NVRTC"] = os.getenv("SGL_DG_USE_NVRTC", "0") +# Enable DeepGEMM kernel preloading if precompile is enabled +if _ENABLE_JIT_DEEPGEMM_PRECOMPILE: + os.environ["DG_PRELOAD_KERNELS"] = "1" + def update_deep_gemm_config(gpu_id: int, server_args: ServerArgs): global _BUILTIN_M_LIST diff --git a/python/sglang/srt/layers/dp_attention.py b/python/sglang/srt/layers/dp_attention.py index f5fcd1757..5aaf94106 100644 --- a/python/sglang/srt/layers/dp_attention.py +++ b/python/sglang/srt/layers/dp_attention.py @@ -501,6 +501,37 @@ def _dp_gather( forward_batch: ForwardBatch, is_partial: bool, ): + # Special case: handle empty local_tokens (e.g., IDLE batches) + # When local_tokens is empty, we can't perform gather operations + if local_tokens.shape[0] == 0: + global_tokens.fill_(0) + return + + # Initialize dp_padding_mode if it's None (e.g., for TBO sub-batches or idle batches) + if forward_batch.dp_padding_mode is None: + if forward_batch.global_num_tokens_cpu is not None: + # Compute dp_padding_mode based on actual global_num_tokens + forward_batch.dp_padding_mode = DpPaddingMode.get_dp_padding_mode( + forward_batch.is_extend_in_batch, + forward_batch.global_num_tokens_cpu + ) + else: + # For TBO sub-batches or cuda graph, infer mode from buffer sizes + # When global_num_tokens_cpu is None, we need to determine the mode based on + # the buffer allocation that was already done + dp_size = get_attention_dp_size() + expected_global_size = global_tokens.shape[0] + expected_local_size = local_tokens.shape[0] + + # Check if buffers are sized for MAX_LEN mode (all_gather_into_tensor) + # In MAX_LEN mode: global_size == local_size * dp_size + # In SUM_LEN mode: global_size != local_size * dp_size (typically) + if expected_global_size == expected_local_size * dp_size: + forward_batch.dp_padding_mode = DpPaddingMode.MAX_LEN + else: + # For other cases, use the default mode for cuda graph + forward_batch.dp_padding_mode = DpPaddingMode.get_default_mode_in_cuda_graph() + if forward_batch.dp_padding_mode.is_max_len(): _dp_gather_via_all_gather( global_tokens, local_tokens, forward_batch, is_partial diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index fc9afafac..3e5789b77 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -643,6 +643,7 @@ class ModelRunner(ModelRunnerKVCacheMixin): self.init_attention_backend() self.kernel_warmup() self.init_device_graphs() + self.init_deep_gemm() elif self.device in ["npu", "cpu"]: self.init_attention_backend() self.init_device_graphs() @@ -661,7 +662,6 @@ class ModelRunner(ModelRunnerKVCacheMixin): # Initialize piecewise CUDA graph self.init_piecewise_cuda_graphs() - self.prealloc_symmetric_memory_pool() def init_routed_experts_capturer(self): @@ -2456,6 +2456,81 @@ class ModelRunner(ModelRunnerKVCacheMixin): f"mem usage={mem_usage:.2f} GB. avail mem={after_mem:.2f} GB." ) + + def init_deep_gemm(self): + logger.info("[DeepGEMM Debug] Entering init_deep_gemm") + if not deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM: + return + try: + from sglang.srt.layers.deep_gemm_wrapper.compile_utils import ( + DeepGemmKernelType, + _maybe_compile_deep_gemm_one_type_all, + ) + import deep_gemm + except ImportError: + return + + # Find DeepSeek V3 MoE layers + model = self.model + if hasattr(model, "model"): + model = model.model + + if not hasattr(model, "layers"): + return + + # We only need to compile once for each unique shape. + # Since _maybe_compile_deep_gemm_one_type_all handles caching, we can just call it for every MoE layer. + for layer in model.layers: + if hasattr(layer, "mlp") and hasattr(layer.mlp, "experts"): + experts = layer.mlp.experts + w13 = None + w2 = None + + if hasattr(experts, "w13_weight_fp8") and hasattr( + experts, "w2_weight_fp8" + ): + w13 = experts.w13_weight_fp8[0] + w2 = experts.w2_weight_fp8[0] + elif hasattr(experts, "w13_weight") and hasattr(experts, "w2_weight"): + w13 = experts.w13_weight + w2 = experts.w2_weight + + if w13 is not None and w2 is not None: + # w13: [num_groups, n, k] + num_groups, n, k = w13.shape + _maybe_compile_deep_gemm_one_type_all( + DeepGemmKernelType.GROUPED_GEMM_NT_F8F8BF16_MASKED, + n, + k, + num_groups, + ) + + # w2: [num_groups, n, k] + num_groups, n, k = w2.shape + _maybe_compile_deep_gemm_one_type_all( + DeepGemmKernelType.GROUPED_GEMM_NT_F8F8BF16_MASKED, + n, + k, + num_groups, + ) + # Also compile Dense GEMM kernels for other layers (e.g. Attention) if they use FP8 and DeepGEMM + for module in self.model.modules(): + if hasattr(module, "weight") and module.weight is not None: + # Check if it is FP8 + is_fp8 = module.weight.dtype in [ + torch.float8_e4m3fn, + torch.float8_e4m3fnuz, + ] + if is_fp8 and module.weight.dim() == 2: + n, k = module.weight.shape + _maybe_compile_deep_gemm_one_type_all( + DeepGemmKernelType.GEMM_NT_F8F8BF16, + n, + k, + 1, # num_groups=1 for standard GEMM + ) + # deep_gemm.preload_kernels() + def init_threads_binding(self): omp_cpuids = os.environ.get("SGLANG_CPU_OMP_THREADS_BIND", "all") cpu_ids_by_node = get_cpu_ids_by_node()