From 5f623ad24eaf21a29d2e5b914abbc4a3e92f96e7 Mon Sep 17 00:00:00 2001 From: Cheng Wan <54331508+ch-wan@users.noreply.github.com> Date: Fri, 3 Jul 2026 19:05:20 -0700 Subject: [PATCH] Revert "Fix wrong RMSNorm fallback to old Flashinfer CUDA kernel when in PCG" (#30083) --- python/sglang/srt/layers/layernorm.py | 27 ++++------------- sgl-kernel/python/sgl_kernel/elementwise.py | 32 ++++++++++++++++++--- 2 files changed, 33 insertions(+), 26 deletions(-) diff --git a/python/sglang/srt/layers/layernorm.py b/python/sglang/srt/layers/layernorm.py index f86b7cacb..6504923a3 100644 --- a/python/sglang/srt/layers/layernorm.py +++ b/python/sglang/srt/layers/layernorm.py @@ -86,28 +86,11 @@ if _is_cuda or _is_xpu or _is_musa: else: _flashinfer_layernorm_available = False - from sgl_kernel import fused_add_rmsnorm as _sgl_fused_add_rmsnorm - from sgl_kernel import gemma_fused_add_rmsnorm as _sgl_gemma_fused_add_rmsnorm - from sgl_kernel import gemma_rmsnorm as _sgl_gemma_rmsnorm - from sgl_kernel import rmsnorm as _sgl_rmsnorm - - from sglang.srt.utils.custom_op import register_custom_op_from_extern - - rmsnorm = register_custom_op_from_extern( - _sgl_rmsnorm, op_name="sgl_rmsnorm", out_shape="input" - ) - fused_add_rmsnorm = register_custom_op_from_extern( - _sgl_fused_add_rmsnorm, - op_name="sgl_fused_add_rmsnorm", - mutates_args=["input", "residual"], - ) - gemma_rmsnorm = register_custom_op_from_extern( - _sgl_gemma_rmsnorm, op_name="sgl_gemma_rmsnorm", out_shape="input" - ) - gemma_fused_add_rmsnorm = register_custom_op_from_extern( - _sgl_gemma_fused_add_rmsnorm, - op_name="sgl_gemma_fused_add_rmsnorm", - mutates_args=["input", "residual"], + from sgl_kernel import ( + fused_add_rmsnorm, + gemma_fused_add_rmsnorm, + gemma_rmsnorm, + rmsnorm, ) _has_aiter_layer_norm = False _has_vllm_rms_norm = False diff --git a/sgl-kernel/python/sgl_kernel/elementwise.py b/sgl-kernel/python/sgl_kernel/elementwise.py index f365d5674..aa325b277 100644 --- a/sgl-kernel/python/sgl_kernel/elementwise.py +++ b/sgl-kernel/python/sgl_kernel/elementwise.py @@ -104,7 +104,19 @@ def rmsnorm( output: torch.Tensor Normalized tensor, shape (batch_size, hidden_size). """ - if _has_flashinfer and input.dtype in _FLASHINFER_NORM_SUPPORTED_DTYPES: + # torch.compiler.is_dynamo_compiling(): FlashInfer norm paths are not safe under + # torch.compile(..., fullgraph=True). Dynamo traces into FlashInfer's JIT module + # loading path, which calls Path.exists() / os.stat() — both untraceable — causing + # the entire compilation to fail. We fall back to the internal implementation while + # tracing as a temporary workaround. Once the upstream fix is merged and we upgrade + # FlashInfer, this check can be removed. + # See: https://github.com/flashinfer-ai/flashinfer/issues/2734 + # https://github.com/flashinfer-ai/flashinfer/pull/2733 + if ( + _has_flashinfer + and input.dtype in _FLASHINFER_NORM_SUPPORTED_DTYPES + and not torch.compiler.is_dynamo_compiling() + ): return _flashinfer_norm.rmsnorm(input, weight, eps, out, enable_pdl) else: return _rmsnorm_internal(input, weight, eps, out, enable_pdl) @@ -140,7 +152,11 @@ def fused_add_rmsnorm( `_ If None, will be automatically enabled on Hopper architecture. """ - if _has_flashinfer and input.dtype in _FLASHINFER_NORM_SUPPORTED_DTYPES: + if ( + _has_flashinfer + and input.dtype in _FLASHINFER_NORM_SUPPORTED_DTYPES + and not torch.compiler.is_dynamo_compiling() + ): _flashinfer_norm.fused_add_rmsnorm(input, residual, weight, eps, enable_pdl) else: _fused_add_rmsnorm_internal(input, residual, weight, eps, enable_pdl) @@ -177,7 +193,11 @@ def gemma_rmsnorm( output: torch.Tensor Gemma Normalized tensor, shape (batch_size, hidden_size). """ - if _has_flashinfer and input.dtype in _FLASHINFER_NORM_SUPPORTED_DTYPES: + if ( + _has_flashinfer + and input.dtype in _FLASHINFER_NORM_SUPPORTED_DTYPES + and not torch.compiler.is_dynamo_compiling() + ): return _flashinfer_norm.gemma_rmsnorm(input, weight, eps, out, enable_pdl) else: return _gemma_rmsnorm_internal(input, weight, eps, out, enable_pdl) @@ -213,7 +233,11 @@ def gemma_fused_add_rmsnorm( `_ If None, will be automatically enabled on Hopper architecture. """ - if _has_flashinfer and input.dtype in _FLASHINFER_NORM_SUPPORTED_DTYPES: + if ( + _has_flashinfer + and input.dtype in _FLASHINFER_NORM_SUPPORTED_DTYPES + and not torch.compiler.is_dynamo_compiling() + ): _flashinfer_norm.gemma_fused_add_rmsnorm( input, residual, weight, eps, enable_pdl )