Fix wrong RMSNorm fallback to old Flashinfer CUDA kernel when in PCG (#29702)

Co-authored-by: Brayden Zhong <brayden@radixark.ai>
This commit is contained in:
Brayden Zhong
2026-07-03 00:49:49 +08:00
committed by GitHub
co-authored by Brayden Zhong
parent 9588cacaa1
commit 1b6d1e9752
2 changed files with 26 additions and 33 deletions
+22 -5
View File
@@ -86,11 +86,28 @@ if _is_cuda or _is_xpu or _is_musa:
else:
_flashinfer_layernorm_available = False
from sgl_kernel import (
fused_add_rmsnorm,
gemma_fused_add_rmsnorm,
gemma_rmsnorm,
rmsnorm,
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"],
)
_has_aiter_layer_norm = False
_has_vllm_rms_norm = False