[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
@@ -186,7 +186,9 @@ LaunchKernel(num_blocks, kBlockSize, device.unwrap())(kernel, params);
|
||||
## Step 0 (optional): Generate a `.clangd` config for better IDE support
|
||||
|
||||
```bash
|
||||
python -m sglang.jit_kernel -h # for verbose help info about clangd configuration
|
||||
python -m sglang.jit_kernel
|
||||
python -m sglang.jit_kernel --dep cutlass flashinfer # with cutlass/flashinfer dependency
|
||||
```
|
||||
|
||||
---
|
||||
@@ -198,6 +200,8 @@ Create `python/sglang/jit_kernel/csrc/elementwise/scale.cuh`.
|
||||
The implementation fully uses the project abstractions described above:
|
||||
|
||||
```cpp
|
||||
// NOTE: Comments for headers are not common in practice.
|
||||
// It is only shown here for tutorial purposes to highlight the key abstractions.
|
||||
#include <sgl_kernel/tensor.h> // For TensorMatcher, SymbolicSize, SymbolicDevice
|
||||
#include <sgl_kernel/type.cuh> // For dtype_trait, fp16_t, bf16_t, fp32_t
|
||||
#include <sgl_kernel/utils.h> // For RuntimeCheck, div_ceil
|
||||
@@ -215,7 +219,7 @@ namespace {
|
||||
// kVecN = number of elements per vector load (e.g. 8 for fp16)
|
||||
// factor = runtime scale factor
|
||||
// ----------------------------------------------------------------
|
||||
template <typename T, int kVecN>
|
||||
template <typename T, int kVecN, bool kUsePDL>
|
||||
__global__ void scale_kernel(T* __restrict__ dst,
|
||||
const T* __restrict__ src,
|
||||
float factor,
|
||||
@@ -223,6 +227,10 @@ __global__ void scale_kernel(T* __restrict__ dst,
|
||||
using vec_t = device::AlignedVector<T, kVecN>;
|
||||
const uint32_t n_vecs = n_total / kVecN;
|
||||
|
||||
// If using PDL, wait for primary kernel before any global memory load.
|
||||
// This is NOT a synchronization point, which means some threads can early exit before this.
|
||||
device::PDLWaitPrimary<kUsePDL>();
|
||||
|
||||
// --- vectorised body ---
|
||||
const uint32_t vec_stride = blockDim.x * gridDim.x;
|
||||
for (uint32_t vi = blockIdx.x * blockDim.x + threadIdx.x;
|
||||
@@ -245,12 +253,16 @@ __global__ void scale_kernel(T* __restrict__ dst,
|
||||
i += scalar_stride) {
|
||||
dst[base + i] = static_cast<T>(static_cast<float>(src[base + i]) * factor);
|
||||
}
|
||||
|
||||
// If using PDL, signal for the secondary kernel to start after all threads have finished
|
||||
// This is NOT a synchronization point, which means some threads can early exit before this.
|
||||
device::PDLTriggerSecondary<kUsePDL>();
|
||||
}
|
||||
|
||||
// ----------------------------------------------------------------
|
||||
// Launcher: validates tensors, selects vector width, launches kernel
|
||||
// ----------------------------------------------------------------
|
||||
template <typename T>
|
||||
template <typename T, bool kUsePDL>
|
||||
void scale(tvm::ffi::TensorView dst, tvm::ffi::TensorView src, float factor) {
|
||||
using namespace host;
|
||||
|
||||
@@ -271,8 +283,12 @@ void scale(tvm::ffi::TensorView dst, tvm::ffi::TensorView src, float factor) {
|
||||
RuntimeCheck(n > 0, "scale: num_elements must be > 0, got ", n);
|
||||
|
||||
// 2. Choose vector width for 128-bit loads (16 bytes)
|
||||
// fp16/bf16: 8 elements × 2 bytes = 16 bytes
|
||||
// fp32: 4 elements × 4 bytes = 16 bytes
|
||||
// fp16/bf16: 8 elements x 2 bytes = 16 bytes
|
||||
// fp32: 4 elements x 4 bytes = 16 bytes
|
||||
// We encourage using `device::kMaxVecBytes`, which will change according to
|
||||
// the target architecture and can enable 256-bit vectorization on SM100+ if desired.
|
||||
// But 128-bit is more commonly adapted for better compatibility,
|
||||
// so it's still ok to hardcode 16 here just for simplicity.
|
||||
constexpr int kVecN = 16 / sizeof(T);
|
||||
const uint32_t n_work_items = div_ceil(n, static_cast<uint32_t>(kVecN));
|
||||
|
||||
@@ -280,8 +296,10 @@ void scale(tvm::ffi::TensorView dst, tvm::ffi::TensorView src, float factor) {
|
||||
constexpr uint32_t kBlockSize = 256;
|
||||
const uint32_t grid = div_ceil(n_work_items, kBlockSize);
|
||||
|
||||
LaunchKernel(grid, kBlockSize, device)(
|
||||
scale_kernel<T, kVecN>,
|
||||
// PDL feature is 100% optional. Without `enable_pdl`, the code should still be correct.
|
||||
// Try to enable it if profiling shows that it can benefit the performance of this kernel.
|
||||
LaunchKernel(grid, kBlockSize, device).enable_pdl(kUsePDL)(
|
||||
scale_kernel<T, kVecN, kUsePDL>,
|
||||
static_cast<T*>(dst.data_ptr()),
|
||||
static_cast<const T*>(src.data_ptr()),
|
||||
factor,
|
||||
@@ -302,7 +320,8 @@ void scale(tvm::ffi::TensorView dst, tvm::ffi::TensorView src, float factor) {
|
||||
- Prefer passing runtime scalars like `factor` directly unless compile-time specialisation is genuinely required
|
||||
- `fp16_t` / `bf16_t` / `fp32_t` are the project's type aliases (from `utils.cuh`)
|
||||
- `device::cast<To, From>` or `dtype_trait<T>::from(val)` for cross-type conversions
|
||||
- `device::math::` functions for device math instead of bare `__` intrinsics
|
||||
- `device::math::` functions for device math instead of bare `__` intrinsics if possible.
|
||||
- Try to use `PDL` feature. In some cases, this will benefit the performance.
|
||||
|
||||
---
|
||||
|
||||
@@ -317,7 +336,12 @@ from typing import TYPE_CHECKING
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args
|
||||
from sglang.jit_kernel.utils import (
|
||||
cache_once,
|
||||
is_arch_support_pdl,
|
||||
load_jit,
|
||||
make_cpp_args,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from tvm_ffi.module import Module
|
||||
@@ -326,7 +350,7 @@ if TYPE_CHECKING:
|
||||
@cache_once
|
||||
def _jit_scale_module(dtype: torch.dtype) -> Module:
|
||||
"""Compile and cache the JIT scale module for a given dtype."""
|
||||
args = make_cpp_args(dtype)
|
||||
args = make_cpp_args(dtype, is_arch_support_pdl())
|
||||
return load_jit(
|
||||
"scale",
|
||||
*args,
|
||||
@@ -351,24 +375,16 @@ def scale(src: torch.Tensor, factor: float, out: torch.Tensor | None = None) ->
|
||||
-------
|
||||
Scaled tensor (dst = src * factor).
|
||||
"""
|
||||
if not src.is_cuda:
|
||||
raise RuntimeError("src must be a CUDA tensor")
|
||||
# DO NOT add too much proactive validation here.
|
||||
# Keep the Python wrapper thin, only enforce the preconditions
|
||||
# that the current JIT/FFI path (C++ side) does not reject on its own.
|
||||
if src.dtype not in (torch.float16, torch.bfloat16, torch.float32):
|
||||
raise RuntimeError(
|
||||
f"Unsupported dtype {src.dtype}. Supported: float16, bfloat16, float32"
|
||||
)
|
||||
if out is None:
|
||||
out = torch.empty_like(src)
|
||||
else:
|
||||
if out.shape != src.shape:
|
||||
raise RuntimeError("out shape must match src")
|
||||
if out.dtype != src.dtype:
|
||||
raise RuntimeError("out dtype must match src")
|
||||
if out.device != src.device:
|
||||
raise RuntimeError("out device must match src")
|
||||
|
||||
# Keep the Python wrapper thin, but still enforce the basic preconditions
|
||||
# that the current JIT/FFI path does not reject safely on its own.
|
||||
module = _jit_scale_module(src.dtype)
|
||||
module.scale(out, src, factor)
|
||||
return out
|
||||
@@ -381,6 +397,7 @@ def scale(src: torch.Tensor, factor: float, out: torch.Tensor | None = None) ->
|
||||
- Only include compile-time specialisation knobs in the build marker; runtime values like `factor` should stay runtime unless the kernel truly needs templating
|
||||
- `cuda_wrappers`: `(export_name, kernel_symbol)` — `export_name` is called from Python
|
||||
- `make_cpp_args(dtype, ...)` converts `torch.dtype` to C++ type alias:
|
||||
- `is_arch_support_pdl()` checks if the current architecture supports PDL, which is typically passed as a template argument to the kernel.
|
||||
- Keep Python launchers thin, but still validate the basic invariants (`is_cuda`, supported dtype, `out` metadata). In the current JIT/FFI path, invalid tensors are not always rejected safely before launch
|
||||
|
||||
| `torch.dtype` | C++ type |
|
||||
@@ -393,6 +410,8 @@ def scale(src: torch.Tensor, factor: float, out: torch.Tensor | None = None) ->
|
||||
|
||||
## Step 3 (optional): Tune JIT build flags
|
||||
|
||||
If your kernel uses some math functions like `expf` or `sinf`, consider enabling `--use_fast_math` for better performance (with a potential precision tradeoff):
|
||||
|
||||
```python
|
||||
return load_jit(
|
||||
"scale",
|
||||
@@ -414,7 +433,7 @@ if torch.cuda.get_device_capability()[0] < 9:
|
||||
|
||||
## Step 4: Write tests (required)
|
||||
|
||||
JIT kernel tests live under `python/sglang/jit_kernel/tests/`. **CI does not run `pytest` in that directory directly.** The unified runner `test/run_suite.py` discovers every `test_*.py` there (and every `bench_*.py` under `benchmark/`), collects `register_*_ci(...)` calls by **statically parsing each file’s AST**, and executes the selected suite. Every test file must register at least one CUDA entry or the collector fails its sanity check.
|
||||
JIT kernel tests live under `python/sglang/jit_kernel/tests/`. **CI does not run `pytest` in that directory directly.** The unified runner `test/run_suite.py` discovers every `test_*.py` there (and every `bench_*.py` under `benchmark/`), collects `register_*_ci(...)` calls by **statically parsing each file's AST**, and executes the selected suite. Every test file must register at least one CUDA entry or the collector fails its sanity check.
|
||||
|
||||
- **PR / per-commit CUDA suites** (see `test/run_suite.py` → `PER_COMMIT_SUITES`): JIT unit tests use `stage-b-kernel-unit-1-gpu-large` (see `.github/workflows/pr-test-jit-kernel.yml`: `python3 run_suite.py --hw cuda --suite stage-b-kernel-unit-1-gpu-large`).
|
||||
- **Nightly kernel suite**: `nightly-kernel-1-gpu` with `--nightly` — typically used with `SGLANG_JIT_KERNEL_RUN_FULL_TESTS=1` in CI for expanded parameter grids (see `python/sglang/jit_kernel/utils.py` → `should_run_full_tests` / `get_ci_test_range`). Wired in `.github/workflows/nightly-test-nvidia.yml` (e.g. `python3 run_suite.py --hw cuda --suite nightly-kernel-1-gpu --nightly --continue-on-error`).
|
||||
|
||||
@@ -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:
|
||||
return gelu_and_mul(x)
|
||||
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
|
||||
|
||||
Reference in New Issue
Block a user