[JIT Kernel] Reland JIT activation (#22094)
Co-authored-by: Cheng Wan <54331508+ch-wan@users.noreply.github.com> Co-authored-by: Cheng Wan <chwan@rice.edu> Co-authored-by: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
Cheng Wan
Cheng Wan
Claude Opus 4.7
parent
0d224e5053
commit
82254bd9c5
@@ -0,0 +1,85 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.jit_kernel.utils import (
|
||||
cache_once,
|
||||
get_jit_cuda_arch,
|
||||
is_arch_support_pdl,
|
||||
is_hip_runtime,
|
||||
load_jit,
|
||||
make_cpp_args,
|
||||
)
|
||||
from sglang.srt.utils.custom_op import register_custom_op
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from tvm_ffi.module import Module
|
||||
|
||||
|
||||
def _fast_math_flags() -> list[str]:
|
||||
# Mirrors sgl-kernel's CMake policy: fast-math on SM90, precise on
|
||||
# SM100+ (Blackwell needs bit-exact expf), off on HIP (clang rejects).
|
||||
if is_hip_runtime():
|
||||
return []
|
||||
if get_jit_cuda_arch().major >= 10:
|
||||
return []
|
||||
return ["--use_fast_math"]
|
||||
|
||||
|
||||
@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=_fast_math_flags(),
|
||||
cuda_wrappers=[
|
||||
("run_activation", f"ActivationKernel<{args}>::run_activation"),
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
SUPPORTED_ACTIVATIONS = {"silu", "gelu", "gelu_tanh"}
|
||||
|
||||
|
||||
@register_custom_op(mutates_args=["out"])
|
||||
def _run_activation_inplace(
|
||||
op_name: str, input: torch.Tensor, out: torch.Tensor
|
||||
) -> None:
|
||||
hidden_size = input.shape[-1] // 2
|
||||
module = _jit_activation_module(input.dtype)
|
||||
input_2d = input.view(-1, hidden_size * 2)
|
||||
out_2d = out.view(-1, hidden_size)
|
||||
module.run_activation(input_2d, out_2d, op_name)
|
||||
|
||||
|
||||
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}"
|
||||
hidden_size = input.shape[-1] // 2
|
||||
if out is None:
|
||||
out = input.new_empty(*input.shape[:-1], hidden_size)
|
||||
_run_activation_inplace(op_name, input, out)
|
||||
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)
|
||||
@@ -0,0 +1,86 @@
|
||||
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)
|
||||
@@ -0,0 +1,135 @@
|
||||
#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 <ActivationKind kAct>
|
||||
SGL_DEVICE float apply_activation_f32(float x_f32) {
|
||||
if constexpr (kAct == ActivationKind::kSiLU) {
|
||||
return x_f32 / (1.0f + expf(-x_f32));
|
||||
} else if constexpr (kAct == ActivationKind::kGELU) {
|
||||
constexpr auto kSqrt1Over2 = 0.7071067811865475f;
|
||||
return 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 cdf = 0.5f * (1.0f + tanhf(kGeluTanhBeta * (x_f32 + kGeluTanhAlpha * x_f32 * x_f32 * x_f32)));
|
||||
return x_f32 * cdf;
|
||||
} else {
|
||||
static_assert(host::dependent_false_v<decltype(kAct)>, "unsupported activation kind");
|
||||
return 0.0f;
|
||||
}
|
||||
}
|
||||
|
||||
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) {
|
||||
const float gate_f32 = device::cast<fp32_t>(gate[i]);
|
||||
const float up_f32 = device::cast<fp32_t>(up[i]);
|
||||
out[i] = device::cast<T>(apply_activation_f32<kAct>(gate_f32) * up_f32);
|
||||
}
|
||||
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();
|
||||
if (num_tokens == 0) return;
|
||||
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
|
||||
@@ -0,0 +1,79 @@
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from sglang.jit_kernel.activation import SUPPORTED_ACTIVATIONS, run_activation
|
||||
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=30, suite="nightly-kernel-1-gpu", nightly=True)
|
||||
|
||||
|
||||
OPS = SUPPORTED_ACTIVATIONS
|
||||
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),
|
||||
*[(2**x, 2048) for x in range(0, 15, 2)],
|
||||
*[(2**x, 65536) for x in range(0, 5, 2)],
|
||||
],
|
||||
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:
|
||||
x = torch.randn(shape, dtype=dtype, device="cuda")
|
||||
out = run_activation(op_name, x, None)
|
||||
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)
|
||||
@pytest.mark.parametrize("shape", SHAPES)
|
||||
def test_activation_out_param(
|
||||
op_name: str, dtype: torch.dtype, shape: tuple[int, ...]
|
||||
) -> None:
|
||||
x = torch.randn(shape, dtype=dtype, device="cuda")
|
||||
out = torch.empty(shape[:-1] + (shape[-1] // 2,), dtype=dtype, device="cuda")
|
||||
result = run_activation(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,9 +16,12 @@ 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 or _is_hip:
|
||||
if _is_cuda:
|
||||
from sglang.jit_kernel.activation import silu_and_mul
|
||||
elif _is_hip:
|
||||
from sgl_kernel import silu_and_mul
|
||||
|
||||
|
||||
if _is_npu:
|
||||
import torch_npu
|
||||
# TODO (will): remove this dependency
|
||||
|
||||
@@ -51,7 +51,13 @@ _is_cpu = is_cpu()
|
||||
_is_hip = is_hip()
|
||||
_is_xpu = is_xpu()
|
||||
|
||||
if _is_cuda or _is_xpu:
|
||||
if _is_cuda:
|
||||
from sglang.jit_kernel.activation import (
|
||||
gelu_and_mul,
|
||||
gelu_tanh_and_mul,
|
||||
silu_and_mul,
|
||||
)
|
||||
elif _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,8 +5,9 @@ from typing import Optional
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.utils import is_cuda_alike
|
||||
from sglang.srt.utils import is_cuda, is_cuda_alike
|
||||
|
||||
_is_cuda = is_cuda()
|
||||
_is_cuda_alike = is_cuda_alike()
|
||||
|
||||
if _is_cuda_alike:
|
||||
@@ -15,7 +16,10 @@ if _is_cuda_alike:
|
||||
get_cutlass_w4a8_moe_mm_data,
|
||||
)
|
||||
|
||||
from sgl_kernel import silu_and_mul
|
||||
if _is_cuda:
|
||||
from sglang.jit_kernel.activation import silu_and_mul
|
||||
else:
|
||||
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,8 +8,9 @@ from sglang.srt.utils.custom_op import register_custom_op
|
||||
_is_cuda = is_cuda()
|
||||
|
||||
if _is_cuda:
|
||||
from sgl_kernel import moe_sum_reduce, silu_and_mul
|
||||
from sgl_kernel import moe_sum_reduce
|
||||
|
||||
from sglang.jit_kernel.activation import silu_and_mul
|
||||
from sglang.jit_kernel.moe_wna16_marlin import moe_wna16_marlin_gemm
|
||||
|
||||
|
||||
|
||||
@@ -5,7 +5,6 @@ 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,
|
||||
@@ -17,6 +16,13 @@ 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
|
||||
|
||||
@@ -46,7 +46,7 @@ _use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
|
||||
_is_musa = is_musa()
|
||||
|
||||
if not (_is_npu or _is_hip) and _is_cuda:
|
||||
from sgl_kernel import silu_and_mul
|
||||
from sglang.jit_kernel.activation import silu_and_mul
|
||||
|
||||
|
||||
_MASKED_GEMM_FAST_ACT = get_bool_env_var("SGLANG_MASKED_GEMM_FAST_ACT")
|
||||
|
||||
@@ -49,7 +49,9 @@ _is_musa = is_musa()
|
||||
|
||||
|
||||
if _is_cuda:
|
||||
from sgl_kernel import gelu_and_mul, moe_sum_reduce, silu_and_mul
|
||||
from sgl_kernel import moe_sum_reduce
|
||||
|
||||
from sglang.jit_kernel.activation import gelu_and_mul, silu_and_mul
|
||||
elif _is_cpu and _is_cpu_amx_available:
|
||||
pass
|
||||
elif _is_hip:
|
||||
|
||||
@@ -33,7 +33,19 @@ _is_hip = is_hip()
|
||||
_is_xpu = is_xpu()
|
||||
_is_musa = is_musa()
|
||||
|
||||
if _is_cuda or _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:
|
||||
from sgl_kernel import gelu_and_mul, moe_align_block_size, moe_sum, silu_and_mul
|
||||
from sgl_kernel.quantization import (
|
||||
ggml_dequantize,
|
||||
@@ -188,16 +200,11 @@ 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":
|
||||
silu_and_mul(out, x)
|
||||
return silu_and_mul(x)
|
||||
elif activation == "gelu":
|
||||
gelu_and_mul(out, x)
|
||||
else:
|
||||
raise ValueError(f"Unsupported activation: {activation}")
|
||||
return out
|
||||
return gelu_and_mul(x)
|
||||
raise ValueError(f"Unsupported activation: {activation}")
|
||||
|
||||
out_hidden_states = torch.empty_like(x)
|
||||
# unless we decent expert reuse we are better off running moe_vec kernel
|
||||
|
||||
Reference in New Issue
Block a user