[Code sync] Fix registration of some ops in grok & Fix oss sync scripts (#13990)

Co-authored-by: Stefan He <hebiaobuaa@gmail.com>
This commit is contained in:
Lianmin Zheng
2025-11-26 13:11:52 -08:00
committed by GitHub
parent b6312e62ea
commit 0a186924ba
7 changed files with 79 additions and 30 deletions

View File

@@ -166,6 +166,7 @@ class Envs:
SGLANG_MIN_NEW_TOKEN_RATIO_FACTOR = EnvFloat(0.14)
SGLANG_NEW_TOKEN_RATIO_DECAY_STEPS = EnvInt(600)
SGLANG_RETRACT_DECODE_STEPS = EnvInt(20)
SGLANG_CLIP_MAX_NEW_TOKENS_ESTIMATION = EnvInt(4096)
# Scheduler: others:
SGLANG_EMPTY_CACHE_INTERVAL = EnvFloat(-1) # in seconds. Set if you observe high memory accumulation over a long serving period.

View File

@@ -1,10 +1,10 @@
from typing import Tuple
from typing import Optional, Tuple
import torch
import triton
import triton.language as tl
from sglang.srt.utils import is_hip
from sglang.srt.utils import direct_register_custom_op, is_hip
_is_hip = is_hip()
@@ -358,7 +358,11 @@ def experts_combine_kernel(
tl.store(out_hidden_states + start_index_mlp + offsets, combined_x, mask=mask)
def experts_combine_triton(moe_hidden_states, mlp_hidden_states, output_buffer=None):
def experts_combine_triton(
moe_hidden_states: torch.Tensor,
mlp_hidden_states: torch.Tensor,
output_buffer: Optional[torch.Tensor] = None,
) -> torch.Tensor:
assert moe_hidden_states.is_contiguous()
assert mlp_hidden_states.is_contiguous()
@@ -393,9 +397,26 @@ def experts_combine_triton(moe_hidden_states, mlp_hidden_states, output_buffer=N
hidden_dim,
**config,
)
return out_hidden_states
def experts_combine_triton_fake(
moe_hidden_states: torch.Tensor,
mlp_hidden_states: torch.Tensor,
output_buffer: Optional[torch.Tensor] = None,
) -> torch.Tensor:
return torch.empty_like(mlp_hidden_states)
direct_register_custom_op(
op_name="experts_combine_triton",
op_func=experts_combine_triton,
mutates_args=[],
fake_impl=experts_combine_triton_fake,
)
# gelu on first half of vector
@triton.jit
def gelu_and_mul_kernel(

View File

@@ -53,6 +53,7 @@ from sglang.srt.utils import (
is_hip,
is_npu,
)
from sglang.srt.utils.patch_torch import register_fake_if_exists
if TYPE_CHECKING:
from sglang.srt.layers.quantization import QuantizationConfig
@@ -72,30 +73,12 @@ _is_npu = is_npu()
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
if _is_cuda:
from sgl_kernel import kimi_k2_moe_fused_gate, moe_fused_gate
@torch.library.register_fake("sgl_kernel::kimi_k2_moe_fused_gate")
def _kimi_k2_moe_fused_gate(
input_tensor,
bias,
topk,
renormalize,
routed_scaling_factor,
apply_routed_scaling_factor_on_output,
):
num_rows = input_tensor.shape[0]
topk_weights = input_tensor.new_empty(
num_rows,
topk,
dtype=torch.float32,
)
topk_ids = input_tensor.new_empty(
num_rows,
topk,
dtype=torch.int32,
)
return topk_weights, topk_ids
from sgl_kernel import moe_fused_gate
try:
from sgl_kernel import kimi_k2_moe_fused_gate
except ImportError as e:
pass
if _is_cuda or _is_hip:
from sgl_kernel import topk_softmax
@@ -1044,7 +1027,7 @@ def select_experts(
if _is_cuda:
@torch.library.register_fake("sgl_kernel::moe_fused_gate")
def _(
def _moe_fused_gate(
input_tensor,
bias,
num_expert_group,
@@ -1062,3 +1045,25 @@ if _is_cuda:
(num_rows, topk), dtype=torch.int32, device=input_tensor.device
)
return topk_weights, topk_ids
@register_fake_if_exists("sgl_kernel::kimi_k2_moe_fused_gate")
def _kimi_k2_moe_fused_gate(
input_tensor,
bias,
topk,
renormalize,
routed_scaling_factor,
apply_routed_scaling_factor_on_output,
):
num_rows = input_tensor.shape[0]
topk_weights = input_tensor.new_empty(
num_rows,
topk,
dtype=torch.float32,
)
topk_ids = input_tensor.new_empty(
num_rows,
topk,
dtype=torch.int32,
)
return topk_weights, topk_ids

View File

@@ -49,6 +49,7 @@ from sglang.srt.utils.common import (
is_sm120_supported,
next_power_of_2,
)
from sglang.srt.utils.patch_torch import register_fake_if_exists
if TYPE_CHECKING:
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
@@ -127,7 +128,7 @@ def _sglang_fp4_gemm_fake(
if is_cuda() and (not is_sm120_supported()) and (fp4_quantize is not None):
@torch.library.register_fake("sgl_kernel::scaled_fp4_quant")
@register_fake_if_exists("sgl_kernel::scaled_fp4_quant")
def _sgl_kernel_scaled_fp4_quant_fake(
output, input, output_scale, input_global_scale
):

View File

@@ -2753,6 +2753,21 @@ def load_json_config(data: str):
def dispose_tensor(x: torch.Tensor):
"""
Dispose a tensor by freeing its memory.
During piecewise CUDA graph capture/replay, we skip disposal to avoid
interfering with torch.compile's memory tracking and graph recording.
"""
# Skip disposal during piecewise CUDA graph to avoid torch.compile issues
# we do local import to avoid circular import
from sglang.srt.compilation.piecewise_context_manager import (
is_in_piecewise_cuda_graph,
)
if is_in_piecewise_cuda_graph():
return
x.set_(torch.empty((0,), device=x.device, dtype=x.dtype))