Revert "[Feature] JIT activation and update skills (by codex)" (#22078)
This commit is contained in:
@@ -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()."""
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user