Revert "[Feature] JIT activation and update skills (by codex)" (#22078)

This commit is contained in:
Baizhou Zhang
2026-04-03 15:04:15 -07:00
committed by GitHub
parent 8cb337c8ea
commit ac1e437f6a
16 changed files with 64 additions and 489 deletions
-66
View File
@@ -1,66 +0,0 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Optional
import torch
from sglang.jit_kernel.utils import (
cache_once,
is_arch_support_pdl,
load_jit,
make_cpp_args,
)
from sglang.srt.utils.custom_op import register_custom_op
if TYPE_CHECKING:
from tvm_ffi.module import Module
@cache_once
def _jit_activation_module(dtype: torch.dtype) -> Module:
args = make_cpp_args(dtype, is_arch_support_pdl())
return load_jit(
"activation",
*args,
cuda_files=["elementwise/activation.cuh"],
extra_cuda_cflags=["--use_fast_math"],
cuda_wrappers=[
("run_activation", f"ActivationKernel<{args}>::run_activation"),
],
)
SUPPORTED_ACTIVATIONS = {"silu", "gelu", "gelu_tanh"}
@register_custom_op(mutates_args=["out"], out_shape="input")
def run_activation(
op_name: str, input: torch.Tensor, out: Optional[torch.Tensor]
) -> torch.Tensor:
assert op_name in SUPPORTED_ACTIVATIONS, f"Unsupported activation: {op_name}"
if out is None:
out = input.new_empty(*input.shape[:-1], input.shape[-1] // 2)
module = _jit_activation_module(input.dtype)
input_2d = input.view(-1, input.shape[-1])
out_2d = out.view(-1, out.shape[-1])
module.run_activation(input_2d, out_2d, op_name)
return out
def silu_and_mul(
input: torch.Tensor, out: Optional[torch.Tensor] = None
) -> torch.Tensor:
return run_activation("silu", input, out)
def gelu_and_mul(
input: torch.Tensor, out: Optional[torch.Tensor] = None
) -> torch.Tensor:
return run_activation("gelu", input, out)
def gelu_tanh_and_mul(
input: torch.Tensor, out: Optional[torch.Tensor] = None
) -> torch.Tensor:
return run_activation("gelu_tanh", input, out)
@@ -1,86 +0,0 @@
import itertools
import torch
import torch.nn.functional as F
import triton
import triton.testing
from sgl_kernel import gelu_and_mul as gelu_and_mul_aot
from sgl_kernel import gelu_tanh_and_mul as gelu_tanh_and_mul_aot
from sgl_kernel import silu_and_mul as silu_and_mul_aot
from sglang.jit_kernel.activation import gelu_and_mul as gelu_and_mul_jit
from sglang.jit_kernel.activation import gelu_tanh_and_mul as gelu_tanh_and_mul_jit
from sglang.jit_kernel.activation import silu_and_mul as silu_and_mul_jit
from sglang.jit_kernel.benchmark.utils import (
DEFAULT_DEVICE,
DEFAULT_DTYPE,
get_benchmark_range,
run_benchmark,
)
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=30, suite="stage-b-kernel-benchmark-1-gpu-large")
@torch.compile
def silu_and_mul(input: torch.Tensor) -> torch.Tensor:
lhs, rhs = input.split(input.shape[-1] // 2, dim=-1)
return F.silu(lhs) * rhs
@torch.compile
def gelu_and_mul(input: torch.Tensor) -> torch.Tensor:
lhs, rhs = input.split(input.shape[-1] // 2, dim=-1)
return F.gelu(lhs, approximate="none") * rhs
@torch.compile
def gelu_tanh_and_mul(input: torch.Tensor) -> torch.Tensor:
lhs, rhs = input.split(input.shape[-1] // 2, dim=-1)
return F.gelu(lhs, approximate="tanh") * rhs
OPS = {
"silu": (silu_and_mul_aot, silu_and_mul_jit, silu_and_mul),
"gelu": (gelu_and_mul_aot, gelu_and_mul_jit, gelu_and_mul),
"gelu_tanh": (gelu_tanh_and_mul_aot, gelu_tanh_and_mul_jit, gelu_tanh_and_mul),
}
BS_LIST = get_benchmark_range(full_range=[2**x for x in range(0, 15)], ci_range=[8])
DIM_LIST = get_benchmark_range(full_range=[1024, 4096, 6144, 8192], ci_range=[4096])
CONFIGS = list(itertools.product(OPS, DIM_LIST, BS_LIST))
NUM_LAYERS = 4 # to eliminate L2 effect
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["op_name", "dim", "batch_size"],
x_vals=CONFIGS,
line_arg="provider",
line_vals=["aot", "jit", "torch"],
line_names=["AOT (sgl-kernel)", "JIT (jit_kernel)", "torch.compile"],
styles=[("blue", "--"), ("orange", "-"), ("green", "-")],
ylabel="us",
plot_name="activation-aot-vs-jit",
args={},
)
)
def benchmark(op_name: str, dim: int, batch_size: int, provider: str):
x = torch.randn(
NUM_LAYERS,
batch_size,
2 * dim,
dtype=DEFAULT_DTYPE,
device=DEFAULT_DEVICE,
)
aot_op, jit_op, torch_op = OPS[op_name]
fn = {"aot": aot_op, "jit": jit_op, "torch": torch_op}[provider]
def f():
for i in range(NUM_LAYERS):
fn(x[i])
return run_benchmark(f, scale=NUM_LAYERS)
if __name__ == "__main__":
benchmark.run(print_data=True)
@@ -1,137 +0,0 @@
#include <sgl_kernel/tensor.h>
#include <sgl_kernel/utils.h>
#include <sgl_kernel/runtime.cuh>
#include <sgl_kernel/type.cuh>
#include <sgl_kernel/utils.cuh>
#include <sgl_kernel/vec.cuh>
#include <tvm/ffi/container/tensor.h>
#include <cmath>
#include <cstdint>
#include <limits>
#include <string>
namespace {
enum class ActivationKind : uint32_t {
kSiLU,
kGELU,
kGELUTanh,
};
template <typename T, ActivationKind kAct>
SGL_DEVICE T apply_activation(T x) {
const float x_f32 = device::cast<fp32_t>(x);
float y_f32 = 0.0f;
if constexpr (kAct == ActivationKind::kSiLU) {
y_f32 = x_f32 / (1.0f + expf(-x_f32));
} else if constexpr (kAct == ActivationKind::kGELU) {
constexpr auto kSqrt1Over2 = 0.7071067811865475f;
y_f32 = x_f32 * (0.5f * (1.0f + erff(x_f32 * kSqrt1Over2)));
} else if constexpr (kAct == ActivationKind::kGELUTanh) {
constexpr auto kGeluTanhAlpha = 0.044715f;
constexpr auto kGeluTanhBeta = 0.7978845608028654f;
const float x_cube = x_f32 * x_f32 * x_f32;
const float cdf = 0.5f * (1.0f + tanhf(kGeluTanhBeta * (x_f32 + kGeluTanhAlpha * x_cube)));
y_f32 = x_f32 * cdf;
} else {
static_assert(host::dependent_false_v<T>, "unsupported activation kind");
}
return device::cast<T>(y_f32);
}
struct ActivationParams {
const void* __restrict__ input;
void* __restrict__ out;
uint32_t hidden_dim;
uint32_t num_tokens;
};
template <typename T, ActivationKind kAct, bool kUsePDL>
__global__ void act_and_mul_kernel(const __grid_constant__ ActivationParams params) {
using namespace device;
constexpr auto kVecSize = kMaxVecBytes / sizeof(T);
using vec_t = AlignedVector<T, kMaxVecBytes / sizeof(T)>;
const auto num_vecs = params.hidden_dim / kVecSize; // per token
const auto tid = blockIdx.x * blockDim.x + threadIdx.x;
const auto token_id = tid / num_vecs;
if (token_id >= params.num_tokens) return;
const auto offset = tid % num_vecs;
const auto input_offset = token_id * (num_vecs * 2) + offset;
const auto output_offset = tid;
PDLWaitPrimary<kUsePDL>();
const auto gate = device::load_as<vec_t>(params.input, input_offset);
const auto up = device::load_as<vec_t>(params.input, input_offset + num_vecs);
vec_t out;
#pragma unroll
for (int i = 0; i < kVecSize; ++i) {
out[i] = apply_activation<T, kAct>(gate[i]) * up[i];
}
device::store_as<vec_t>(params.out, out, output_offset);
PDLTriggerSecondary<kUsePDL>();
}
template <typename T, bool kUsePDL>
struct ActivationKernel {
static constexpr auto kVecSize = device::kMaxVecBytes / sizeof(T);
static constexpr auto kBlockSize = 256u;
template <ActivationKind kAct>
static constexpr auto activation_kernel = act_and_mul_kernel<T, kAct, kUsePDL>;
static_assert(device::kMaxVecBytes % sizeof(T) == 0, "unsupported data type");
static void run_activation(const tvm::ffi::TensorView input, const tvm::ffi::TensorView out, std::string type) {
using namespace host;
auto N = SymbolicSize{"num_tokens"};
auto D_in = SymbolicSize{"input_width"};
auto D_out = SymbolicSize{"output_width"};
auto device_ = SymbolicDevice{};
device_.set_options<kDLCUDA>();
TensorMatcher({N, D_out}) //
.with_dtype<T>()
.with_device(device_)
.verify(out);
TensorMatcher({N, D_in}) //
.with_dtype<T>()
.with_device(device_)
.verify(input);
const auto hidden_size = D_out.unwrap();
const auto num_tokens = static_cast<uint32_t>(N.unwrap());
const auto device = device_.unwrap();
RuntimeCheck(hidden_size * 2 == D_in.unwrap(), "invalid activation dimension");
RuntimeCheck(hidden_size % kVecSize == 0, "hidden size must be divisible by vector size");
const auto kernel = [&]() -> decltype(activation_kernel<ActivationKind::kSiLU>) {
if (type == "silu") {
return activation_kernel<ActivationKind::kSiLU>;
} else if (type == "gelu") {
return activation_kernel<ActivationKind::kGELU>;
} else if (type == "gelu_tanh") {
return activation_kernel<ActivationKind::kGELUTanh>;
} else {
Panic("unsupported activation type: ", type);
}
return nullptr;
}();
// only get once to avoid overhead
const auto num_total_items = num_tokens * (hidden_size / kVecSize);
RuntimeCheck(num_total_items <= std::numeric_limits<uint32_t>::max(), "too many items for 32-bit indexing");
const auto num_blocks = div_ceil(static_cast<uint32_t>(num_total_items), kBlockSize);
const auto params = ActivationParams{
.input = input.data_ptr(),
.out = out.data_ptr(),
.hidden_dim = hidden_size,
.num_tokens = num_tokens,
};
LaunchKernel(num_blocks, kBlockSize, device).enable_pdl(kUsePDL)(kernel, params);
}
};
} // namespace
@@ -1,79 +0,0 @@
import sys
import pytest
import torch
import torch.nn.functional as F
from sglang.jit_kernel.activation import gelu_and_mul, gelu_tanh_and_mul, silu_and_mul
from sglang.jit_kernel.utils import get_ci_test_range
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=20, suite="stage-b-kernel-unit-1-gpu-large")
register_cuda_ci(est_time=120, suite="nightly-kernel-1-gpu", nightly=True)
OPS = {"silu": silu_and_mul, "gelu": gelu_and_mul, "gelu_tanh": gelu_tanh_and_mul}
DTYPES = [torch.float16, torch.bfloat16, torch.float32]
SHAPES = get_ci_test_range(
full_range=[
(7, 16),
(83, 1024),
(3, 5, 16),
(2, 3, 512),
(1, 17, 4096),
],
ci_range=[(7, 16), (2, 3, 512)],
)
def _reference(op_name: str, x: torch.Tensor) -> torch.Tensor:
d = x.shape[-1] // 2
lhs = x[..., :d].float()
rhs = x[..., d:]
if op_name == "silu":
act = F.silu(lhs)
elif op_name == "gelu":
act = F.gelu(lhs, approximate="none")
else:
act = F.gelu(lhs, approximate="tanh")
return act.to(dtype=x.dtype) * rhs
def _tolerances(dtype: torch.dtype) -> tuple[float, float]:
if dtype == torch.float32:
return 1e-4, 1e-4
return 1e-2, 1e-2
@pytest.mark.parametrize("op_name", OPS)
@pytest.mark.parametrize("dtype", DTYPES)
@pytest.mark.parametrize("shape", SHAPES)
def test_activation_correctness(
op_name: str, dtype: torch.dtype, shape: tuple[int, ...]
) -> None:
torch.manual_seed(42)
x = torch.randn(shape, dtype=dtype, device="cuda")
out = OPS[op_name](x)
expected = _reference(op_name, x)
atol, rtol = _tolerances(dtype)
torch.testing.assert_close(out, expected, atol=atol, rtol=rtol)
@pytest.mark.parametrize("op_name", OPS)
@pytest.mark.parametrize("dtype", DTYPES)
def test_activation_out_param(op_name: str, dtype: torch.dtype) -> None:
torch.manual_seed(0)
x = torch.randn((4, 7, 128), dtype=dtype, device="cuda")
out = torch.empty((4, 7, 64), dtype=dtype, device="cuda")
result = OPS[op_name](x, out)
assert result is out
expected = _reference(op_name, x)
atol, rtol = _tolerances(dtype)
torch.testing.assert_close(out, expected, atol=atol, rtol=rtol)
if __name__ == "__main__":
sys.exit(pytest.main([__file__, "-v", "-s"]))
@@ -16,14 +16,8 @@ from sglang.multimodal_gen.runtime.platforms import current_platform
_is_cuda = current_platform.is_cuda()
_is_hip = current_platform.is_hip()
_is_npu = current_platform.is_npu()
if _is_cuda:
from sglang.jit_kernel.activation import (
gelu_and_mul,
gelu_tanh_and_mul,
silu_and_mul,
)
elif _is_hip:
from sgl_kernel import gelu_and_mul, gelu_tanh_and_mul, silu_and_mul
if _is_cuda or _is_hip:
from sgl_kernel import silu_and_mul
if _is_npu:
import torch_npu
@@ -82,16 +76,8 @@ class GeluAndMul(CustomOp):
if approximate not in ("none", "tanh"):
raise ValueError(f"Unknown approximate mode: {approximate}")
def forward_cuda(self, x: torch.Tensor) -> Any:
d = x.shape[-1] // 2
output_shape = x.shape[:-1] + (d,)
out = torch.empty(output_shape, dtype=x.dtype, device=x.device)
if self.approximate == "tanh":
gelu_and_mul_fn = gelu_tanh_and_mul
else:
gelu_and_mul_fn = gelu_and_mul
gelu_and_mul_fn(x, out)
return out
def forward_cuda(self, *args, **kwargs) -> Any:
return self.forward_native(*args, **kwargs)
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
"""PyTorch-native implementation equivalent to forward()."""
+1 -7
View File
@@ -49,13 +49,7 @@ _is_cpu = is_cpu()
_is_hip = is_hip()
_is_xpu = is_xpu()
if _is_cuda:
from sglang.jit_kernel.activation import (
gelu_and_mul,
gelu_tanh_and_mul,
silu_and_mul,
)
elif _is_xpu:
if _is_cuda or _is_xpu:
from sgl_kernel import gelu_and_mul, gelu_tanh_and_mul, silu_and_mul
elif _is_hip:
from sgl_kernel import gelu_and_mul, gelu_quick, gelu_tanh_and_mul, silu_and_mul
+1 -1
View File
@@ -17,9 +17,9 @@ if _is_cuda:
fp8_blockwise_scaled_grouped_mm,
prepare_moe_input,
shuffle_rows,
silu_and_mul,
)
from sglang.jit_kernel.activation import silu_and_mul
from sglang.jit_kernel.nvfp4 import (
cutlass_fp4_group_mm,
scaled_fp4_experts_quant,
@@ -5,9 +5,8 @@ from typing import Optional
import torch
from sglang.srt.utils import is_cuda, is_cuda_alike
from sglang.srt.utils import is_cuda_alike
_is_cuda = is_cuda()
_is_cuda_alike = is_cuda_alike()
if _is_cuda_alike:
@@ -16,10 +15,7 @@ if _is_cuda_alike:
get_cutlass_w4a8_moe_mm_data,
)
if _is_cuda:
from sglang.jit_kernel.activation import silu_and_mul
else:
from sgl_kernel import silu_and_mul
from sgl_kernel import silu_and_mul
from sglang.jit_kernel.per_tensor_quant_fp8 import per_tensor_quant_fp8
from sglang.srt.distributed import get_moe_expert_parallel_world_size
@@ -8,9 +8,8 @@ from sglang.srt.utils.custom_op import register_custom_op
_is_cuda = is_cuda()
if _is_cuda:
from sgl_kernel import moe_sum_reduce
from sgl_kernel import moe_sum_reduce, silu_and_mul
from sglang.jit_kernel.activation import silu_and_mul
from sglang.jit_kernel.moe_wna16_marlin import moe_wna16_marlin_gemm
@@ -48,9 +48,7 @@ _use_sgl_xpu = use_intel_xpu_backend()
from sglang.srt.server_args import get_global_server_args
if _is_cuda:
from sgl_kernel import moe_sum_reduce
from sglang.jit_kernel.activation import gelu_and_mul, silu_and_mul
from sgl_kernel import gelu_and_mul, moe_sum_reduce, silu_and_mul
elif _is_cpu and _is_cpu_amx_available:
pass
elif _is_hip:
@@ -5,6 +5,7 @@ from __future__ import annotations
from typing import TYPE_CHECKING, Optional
import torch
from sgl_kernel import gelu_and_mul, silu_and_mul
from triton_kernels.matmul_ogs import (
FlexCtx,
FnSpecs,
@@ -16,13 +17,6 @@ from triton_kernels.numerics import InFlexData
from triton_kernels.routing import GatherIndx, RoutingData, ScatterIndx
from triton_kernels.swiglu import swiglu_fn
from sglang.srt.utils import is_cuda
if is_cuda():
from sglang.jit_kernel.activation import gelu_and_mul, silu_and_mul
else:
from sgl_kernel import gelu_and_mul, silu_and_mul
if TYPE_CHECKING:
from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig
from sglang.srt.layers.moe.topk import TopKOutput
@@ -44,7 +44,7 @@ _is_cuda = is_cuda()
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
if not (_is_npu or _is_hip) and _is_cuda:
from sglang.jit_kernel.activation import silu_and_mul
from sgl_kernel import silu_and_mul
_MASKED_GEMM_FAST_ACT = get_bool_env_var("SGLANG_MASKED_GEMM_FAST_ACT")
@@ -37,25 +37,26 @@ _is_xpu = is_xpu()
_MOE_PADDING_SIZE = 128 if bool(int(os.getenv("SGLANG_MOE_PADDING", "0"))) else 0
if _is_cuda:
from sglang.jit_kernel.activation import gelu_and_mul, silu_and_mul
elif _is_hip:
if _is_cuda or _is_hip:
from sgl_kernel import gelu_and_mul, silu_and_mul
_has_vllm = False
if _use_aiter:
try:
from aiter import moe_sum
except ImportError:
raise ImportError("aiter is required when SGLANG_USE_AITER is set to True")
else:
try:
from vllm import _custom_ops as vllm_ops # moe_sum
if _is_hip:
_has_vllm = False
if _use_aiter:
try:
from aiter import moe_sum
except ImportError:
raise ImportError(
"aiter is required when SGLANG_USE_AITER is set to True"
)
else:
try:
from vllm import _custom_ops as vllm_ops # moe_sum
_has_vllm = True
except ImportError:
# Fallback: vllm not available, will use triton moe_sum
_has_vllm = False
_has_vllm = True
except ImportError:
# Fallback: vllm not available, will use triton moe_sum
_has_vllm = False
elif _is_cpu and _is_cpu_amx_available:
pass
elif _is_xpu:
+9 -16
View File
@@ -33,19 +33,7 @@ _is_hip = is_hip()
_is_xpu = is_xpu()
_is_musa = is_musa()
if _is_cuda:
from sgl_kernel import moe_align_block_size, moe_sum
from sgl_kernel.quantization import (
ggml_dequantize,
ggml_moe_a8,
ggml_moe_a8_vec,
ggml_moe_get_block_size,
ggml_mul_mat_a8,
ggml_mul_mat_vec_a8,
)
from sglang.jit_kernel.activation import gelu_and_mul, silu_and_mul
elif _is_musa:
if _is_cuda or _is_musa:
from sgl_kernel import gelu_and_mul, moe_align_block_size, moe_sum, silu_and_mul
from sgl_kernel.quantization import (
ggml_dequantize,
@@ -200,11 +188,16 @@ def fused_moe_gguf(
activation: str,
) -> torch.Tensor:
def act(x: torch.Tensor):
d = x.shape[-1] // 2
output_shape = x.shape[:-1] + (d,)
out = torch.empty(output_shape, dtype=x.dtype, device=x.device)
if activation == "silu":
return silu_and_mul(x)
silu_and_mul(out, x)
elif activation == "gelu":
return gelu_and_mul(x)
raise ValueError(f"Unsupported activation: {activation}")
gelu_and_mul(out, x)
else:
raise ValueError(f"Unsupported activation: {activation}")
return out
out_hidden_states = torch.empty_like(x)
# unless we decent expert reuse we are better off running moe_vec kernel
+5 -4
View File
@@ -47,11 +47,12 @@ _use_aiter = bool(int(os.getenv("SGLANG_USE_AITER", "0")))
_is_xpu = is_xpu()
_MOE_PADDING_SIZE = 128 if bool(int(os.getenv("SGLANG_MOE_PADDING", "0"))) else 0
if _is_cuda:
from sglang.jit_kernel.activation import gelu_and_mul, silu_and_mul
elif _is_hip:
if _is_cuda or _is_hip:
from sgl_kernel import gelu_and_mul, silu_and_mul
from vllm import _custom_ops as vllm_ops # moe_sum
if _is_hip:
from vllm import _custom_ops as vllm_ops # moe_sum
elif _is_cpu and _is_cpu_amx_available:
pass
elif _is_xpu: