[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:
co-authored by
Stefan He
parent
b6312e62ea
commit
0a186924ba
@@ -166,6 +166,7 @@ class Envs:
|
|||||||
SGLANG_MIN_NEW_TOKEN_RATIO_FACTOR = EnvFloat(0.14)
|
SGLANG_MIN_NEW_TOKEN_RATIO_FACTOR = EnvFloat(0.14)
|
||||||
SGLANG_NEW_TOKEN_RATIO_DECAY_STEPS = EnvInt(600)
|
SGLANG_NEW_TOKEN_RATIO_DECAY_STEPS = EnvInt(600)
|
||||||
SGLANG_RETRACT_DECODE_STEPS = EnvInt(20)
|
SGLANG_RETRACT_DECODE_STEPS = EnvInt(20)
|
||||||
|
SGLANG_CLIP_MAX_NEW_TOKENS_ESTIMATION = EnvInt(4096)
|
||||||
|
|
||||||
# Scheduler: others:
|
# Scheduler: others:
|
||||||
SGLANG_EMPTY_CACHE_INTERVAL = EnvFloat(-1) # in seconds. Set if you observe high memory accumulation over a long serving period.
|
SGLANG_EMPTY_CACHE_INTERVAL = EnvFloat(-1) # in seconds. Set if you observe high memory accumulation over a long serving period.
|
||||||
|
|||||||
@@ -1,10 +1,10 @@
|
|||||||
from typing import Tuple
|
from typing import Optional, Tuple
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import triton
|
import triton
|
||||||
import triton.language as tl
|
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()
|
_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)
|
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 moe_hidden_states.is_contiguous()
|
||||||
assert mlp_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,
|
hidden_dim,
|
||||||
**config,
|
**config,
|
||||||
)
|
)
|
||||||
|
|
||||||
return out_hidden_states
|
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
|
# gelu on first half of vector
|
||||||
@triton.jit
|
@triton.jit
|
||||||
def gelu_and_mul_kernel(
|
def gelu_and_mul_kernel(
|
||||||
|
|||||||
@@ -53,6 +53,7 @@ from sglang.srt.utils import (
|
|||||||
is_hip,
|
is_hip,
|
||||||
is_npu,
|
is_npu,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.utils.patch_torch import register_fake_if_exists
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.layers.quantization import QuantizationConfig
|
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
|
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
|
||||||
|
|
||||||
if _is_cuda:
|
if _is_cuda:
|
||||||
from sgl_kernel import kimi_k2_moe_fused_gate, moe_fused_gate
|
from sgl_kernel import 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
|
|
||||||
|
|
||||||
|
try:
|
||||||
|
from sgl_kernel import kimi_k2_moe_fused_gate
|
||||||
|
except ImportError as e:
|
||||||
|
pass
|
||||||
|
|
||||||
if _is_cuda or _is_hip:
|
if _is_cuda or _is_hip:
|
||||||
from sgl_kernel import topk_softmax
|
from sgl_kernel import topk_softmax
|
||||||
@@ -1044,7 +1027,7 @@ def select_experts(
|
|||||||
if _is_cuda:
|
if _is_cuda:
|
||||||
|
|
||||||
@torch.library.register_fake("sgl_kernel::moe_fused_gate")
|
@torch.library.register_fake("sgl_kernel::moe_fused_gate")
|
||||||
def _(
|
def _moe_fused_gate(
|
||||||
input_tensor,
|
input_tensor,
|
||||||
bias,
|
bias,
|
||||||
num_expert_group,
|
num_expert_group,
|
||||||
@@ -1062,3 +1045,25 @@ if _is_cuda:
|
|||||||
(num_rows, topk), dtype=torch.int32, device=input_tensor.device
|
(num_rows, topk), dtype=torch.int32, device=input_tensor.device
|
||||||
)
|
)
|
||||||
return topk_weights, topk_ids
|
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
|
||||||
|
|||||||
@@ -49,6 +49,7 @@ from sglang.srt.utils.common import (
|
|||||||
is_sm120_supported,
|
is_sm120_supported,
|
||||||
next_power_of_2,
|
next_power_of_2,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.utils.patch_torch import register_fake_if_exists
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
|
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):
|
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(
|
def _sgl_kernel_scaled_fp4_quant_fake(
|
||||||
output, input, output_scale, input_global_scale
|
output, input, output_scale, input_global_scale
|
||||||
):
|
):
|
||||||
|
|||||||
@@ -2753,6 +2753,21 @@ def load_json_config(data: str):
|
|||||||
|
|
||||||
|
|
||||||
def dispose_tensor(x: torch.Tensor):
|
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))
|
x.set_(torch.empty((0,), device=x.device, dtype=x.dtype))
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -44,12 +44,15 @@ folder_names = [
|
|||||||
"docs",
|
"docs",
|
||||||
"examples",
|
"examples",
|
||||||
"python/sglang/lang",
|
"python/sglang/lang",
|
||||||
|
"python/sglang/jit_kernel",
|
||||||
"python/sglang/srt",
|
"python/sglang/srt",
|
||||||
"python/sglang/test",
|
"python/sglang/test",
|
||||||
"python/sglang/utils.py",
|
"python/sglang/utils.py",
|
||||||
"python/sglang/README.md",
|
"python/sglang/README.md",
|
||||||
"sgl-kernel",
|
"sgl-kernel",
|
||||||
"test/lang",
|
"test/manual",
|
||||||
|
"test/nightly",
|
||||||
|
"test/registered",
|
||||||
"test/srt",
|
"test/srt",
|
||||||
"test/README.md",
|
"test/README.md",
|
||||||
"README.md",
|
"README.md",
|
||||||
|
|||||||
@@ -44,12 +44,15 @@ folder_names = [
|
|||||||
"docs",
|
"docs",
|
||||||
"examples",
|
"examples",
|
||||||
"python/sglang/lang",
|
"python/sglang/lang",
|
||||||
|
"python/sglang/jit_kernel",
|
||||||
"python/sglang/srt",
|
"python/sglang/srt",
|
||||||
"python/sglang/test",
|
"python/sglang/test",
|
||||||
"python/sglang/utils.py",
|
"python/sglang/utils.py",
|
||||||
"python/sglang/README.md",
|
"python/sglang/README.md",
|
||||||
"sgl-kernel",
|
"sgl-kernel",
|
||||||
"test/lang",
|
"test/manual",
|
||||||
|
"test/nightly",
|
||||||
|
"test/registered",
|
||||||
"test/srt",
|
"test/srt",
|
||||||
"test/README.md",
|
"test/README.md",
|
||||||
"README.md",
|
"README.md",
|
||||||
|
|||||||
Reference in New Issue
Block a user