From 14203432b478cdc6d5b71ce5e339753cab75bed4 Mon Sep 17 00:00:00 2001 From: ishandhanani <82981111+ishandhanani@users.noreply.github.com> Date: Fri, 24 Oct 2025 11:00:12 -0700 Subject: [PATCH] fix(compile_utils, ep_moe): update environment variable and dtype check (#12034) --- docs/references/environment_variables.md | 2 +- .../sglang/srt/layers/deep_gemm_wrapper/compile_utils.py | 2 +- python/sglang/srt/layers/moe/ep_moe/layer.py | 7 ++++--- 3 files changed, 6 insertions(+), 5 deletions(-) diff --git a/docs/references/environment_variables.md b/docs/references/environment_variables.md index aa72aedd3..3fcee8212 100644 --- a/docs/references/environment_variables.md +++ b/docs/references/environment_variables.md @@ -36,7 +36,7 @@ SGLang supports various environment variables that can be used to configure its | `SGLANG_JIT_DEEPGEMM_PRECOMPILE` | Enable precompilation of DeepGEMM kernels | `"true"` | | `SGLANG_JIT_DEEPGEMM_COMPILE_WORKERS` | Number of workers for parallel DeepGEMM kernel compilation | `4` | | `SGL_IN_DEEPGEMM_PRECOMPILE_STAGE` | Indicator flag used during the DeepGEMM precompile script | `"false"` | -| `SGL_DG_CACHE_DIR` | Directory for caching compiled DeepGEMM kernels | `~/.cache/deep_gemm` | +| `SGLANG_DG_CACHE_DIR` | Directory for caching compiled DeepGEMM kernels | `~/.cache/deep_gemm` | | `SGL_DG_USE_NVRTC` | Use NVRTC (instead of Triton) for JIT compilation (Experimental) | `"0"` | | `SGL_USE_DEEPGEMM_BMM` | Use DeepGEMM for Batched Matrix Multiplication (BMM) operations | `"false"` | 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 202801b4e..121def4c6 100644 --- a/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py +++ b/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py @@ -26,7 +26,7 @@ _IN_PRECOMPILE_STAGE = get_bool_env_var("SGL_IN_DEEPGEMM_PRECOMPILE_STAGE", "fal # Force redirect deep_gemm cache_dir os.environ["DG_JIT_CACHE_DIR"] = os.getenv( - "SGL_DG_CACHE_DIR", os.path.join(os.path.expanduser("~"), ".cache", "deep_gemm") + "SGLANG_DG_CACHE_DIR", os.path.join(os.path.expanduser("~"), ".cache", "deep_gemm") ) # Refer to https://github.com/deepseek-ai/DeepGEMM/commit/d75b218b7b8f4a5dd5406ac87905039ead3ae42f diff --git a/python/sglang/srt/layers/moe/ep_moe/layer.py b/python/sglang/srt/layers/moe/ep_moe/layer.py index 12f04eb9e..46107ffa7 100644 --- a/python/sglang/srt/layers/moe/ep_moe/layer.py +++ b/python/sglang/srt/layers/moe/ep_moe/layer.py @@ -440,9 +440,10 @@ class DeepEPMoE(FusedMoE): hidden_states, hidden_states_scale, _, _, masked_m, expected_m = dispatch_output assert self.quant_method is not None assert self.moe_runner_config.activation == "silu" - assert ( - hidden_states_scale.dtype == torch.float32 - ), f"hidden_states_scale.dtype: {hidden_states_scale.dtype}" + assert hidden_states_scale.dtype == torch.float32 or ( + deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0 + and hidden_states_scale.dtype == torch.int32 + ), f"hidden_states_scale.dtype: {hidden_states_scale.dtype}, DEEPGEMM_SCALE_UE8M0: {deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0}" # GroupGemm-0 num_groups, m, k = hidden_states.size()