add-jit-kernel
GitHub提供向 SGLang 添加轻量级 JIT CUDA 内核的分步教程,涵盖实现、编译及集成流程。
Trigger Scenarios
Install
npx skills add sgl-project/sglang --skill add-jit-kernel -g -y
SKILL.md
Frontmatter
{
"name": "add-jit-kernel",
"description": "Step-by-step tutorial for adding a new lightweight JIT CUDA kernel to sglang's jit_kernel module"
}
Tutorial: Adding a New JIT Kernel to SGLang
This tutorial walks through adding a simple element-wise scale operation as a JIT kernel. We'll implement scale(x, factor) = x * factor to demonstrate the complete workflow.
Goal
Add a new operation that scales each element of a tensor by a scalar factor:
- Input: tensor
x(CUDA) and scalarfactor(float, passed at runtime) - Output:
x * factor(element-wise), allocated internally - Supported dtypes: FP16 (
torch.float16), BF16 (torch.bfloat16), FP32 (torch.float32)
When to use JIT vs AOT (sgl-kernel)
- JIT (
jit_kernel): prefer this first for kernels that do not depend on CUTLASS or another large C++ project. It is the default choice for lightweight kernels that benefit from rapid iteration and first-use compilation. - AOT (
sgl-kernel): prefer this when the kernel does depend on CUTLASS or another large C++ project, or when it should live inpython/sglang/kernels/aot/and participate in the wheel build / torch op registration flow. - Exception: kernels that depend on
flashinfer, or on CUTLASS that is already provided throughflashinfer, can still be implemented asjit_kernel.
Conventions
These hold for every step below.
namespace sglangis where JIT code lives. Open it after the include block and close it at the end of the file, with the device kernels, traits and host wrapper inside. The sharedhost::/device::helpers are nested in it too, so they resolve unqualified.load_jitemits theTVM_FFI_DLL_EXPORT_TYPED_FUNCwrapper insidenamespace sglangas well, so thekernel_nameyou pass from Python needs nosglang::prefix.- Check where the check is cheapest:
static_assert> C++ host check > cached Python > per-call Python. Anything fixed at compile time is astatic_assert. Anything about the tensors is aTensorMatcher/CHECK_HOSTin the C++ launcher, free next to a kernel launch. A check Python cannot delegate goes inside the@cache_oncemodule factory, where it runs once per specialisation. What remains in the per-call entry point costs interpreter time on every forward, so it should be nothing but picking the module and allocatingout. - Fixed-width integer types. Prefer
int32_t/int64_t/uint32_t/size_toverint,long, orlong long, so an index has the same width on both sides of the FFI boundary. Bareintis fine only where the width plainly cannot matter — an unrolled loop counter over aconstexprbound, a templateintparameter. Shapes arrive asint64_t(SymbolicSize::unwrap()); narrowing touint32_tfor in-kernel indexing is a deliberate act, so write thestatic_castexplicitly and only where the range is known. - Doxygen comments in C++. Document exported entities with
///or/** ... */blocks using\brief,\param,\tparam,\return, the wayinclude/sgl_kernel/does.python -m sglang.kernels.jitwritesCommentFormat: Doxygeninto.clangdwhen clangd is 21 or newer, so these render on hover in the editor. Plain//remains fine for implementation notes inside a function body. - ASCII only in C++ and CUDA sources. Write
--,->,<=instead of—,→,≤, including in comments.grep -nP '[^\x00-\x7F]' <file>before committing. const T* __restrict__for read-only pointers. This is whatcsrc/does throughout, and it lets the compiler emit non-coherent (LDG) loads.- Watch the register budget. For memory-bound kernels, keep to roughly 64 registers per thread so occupancy does not become the limit. Build once with
extra_cuda_cflags=["-Xptxas", "-v"]to see the actual count, and prefer recomputing a value over letting it spill.
Common Abstractions in python/sglang/kernels/jit/include/sgl_kernel/
Always prefer these abstractions over raw CUDA primitives. They provide safety, readability, and consistency with the rest of the codebase. The only reason to drop to raw primitives is performance the abstraction cannot reach — a trade you make deliberately, and justify in a comment.
utils.h — Host-side utilities
#include <sgl_kernel/utils.h>
CHECK_HOST(cond) << "msg " << value— Preferred runtime check: stream-style, throwsPanicErrorwith file/line info on failure. Zero overhead on the true path — the message expressions are only evaluated when the check fails.host::RuntimeCheck(cond, args...)— Function-style alternative toCHECK_HOST. Note its message args are always evaluated (even when the check passes), so preferCHECK_HOST— especially on hot paths.host::Panic(args...)— Unconditionally throw aPanicErrorwith a descriptive message.host::div_ceil(a, b)— Integer ceiling division(a + b - 1) / b.host::irange(n)/host::irange(start, end)— Range views for cleaner loops.host::pointer::offset(ptr, offsets...)— Byte-safe pointer arithmetic onvoid*. Use this instead of raw casts.
utils.cuh — Device-side utilities + LaunchKernel
#include <sgl_kernel/utils.cuh>
-
Type aliases:
fp16_t,bf16_t,fp32_t,fp8_e4m3_t,fp8_e5m2_tand their packed variantsfp16x2_t,bf16x2_t,fp32x2_t, etc. -
SGL_DEVICE— Expands to__forceinline__ __device__. Use on all device functions. -
device::kWarpThreads— Constant32. -
device::load_as<T>(ptr, offset)/device::store_as<T>(ptr, val, offset)— Type-safe loads/stores fromvoid*. -
device::pointer::offset(ptr, offsets...)— Pointer arithmetic on device. -
host::LaunchKernel(grid, block, device_or_stream [, smem])— RAII kernel launcher that:- Resolves the CUDA stream from a
DLDevicevia TVM-FFI automatically. - Checks the CUDA error with file/line info after launch via
operator()(kernel, args...). - Supports
.enable_pdl(bool)for PDL (Programmatic Dependent Launch, SM90+).
- Resolves the CUDA stream from a
-
device::PDLWaitPrimary<kUsePDL>()/device::PDLTriggerSecondary<kUsePDL>()— The two halves of PDL, on sm_90+ (no-ops on older archs and ROCm). Their guarantees are not symmetric:PDLTriggerSecondary(griddepcontrol.launch_dependents) only lets the next kernel in the stream start early. It carries no memory ordering and publishes nothing — matching that, the header's asm has no"memory"clobber.PDLWaitPrimary(griddepcontrol.wait) is the ordering point: it waits until the preceding kernel has fully finished and its writes are visible.
So every read of data the preceding kernel produced must come after
PDLWaitPrimary(). What overlaps with the primary's tail is whatever you put before the wait — loading parameters, computing indices, touching buffers the primary never wrote — so a kernel that waits on its first line gains nothing. Neither call is a barrier: threads may reach or skip them independently. See "Programmatic Dependent Launch and Synchronization" in the CUDA C++ Programming Guide. -
CHECK_CUDA(expr) << "context"— Stream-style CUDA error check; evaluatesexpronce and throwsPanicErrorwithcudaGetErrorString+ file/line info if it is notcudaSuccess. Extra streamed context is optional. -
host::RuntimeDeviceCheck(cudaError_t)— Function-style alternative toCHECK_CUDA. It takes no context message, so preferCHECK_CUDA, which builds its error object only on the failure path.
tensor.h — Tensor validation (TensorMatcher, Symbolic types)
#include <sgl_kernel/tensor.h>
This is the primary validation API for all kernel launchers. Use it to validate every tvm::ffi::TensorView argument.
host::SymbolicSize{"name"}— A named symbolic dimension. Call.set_value(n)to pin it,.unwrap()to extract after verification.host::SymbolicDType— Symbolic dtype. Use.set_options<Ts...>()to restrict allowed types.host::SymbolicDevice— Symbolic device. Use.set_options<kDLCUDA>()to restrict to CUDA.host::TensorMatcher({dims...})— Fluent builder for tensor validation:.with_dtype<T>()— require a specific C++ type (e.g.fp16_t).with_dtype<T1, T2, ...>()— allow a set of types.with_device<kDLCUDA>(device_sym)— require CUDA and bind the checked device to aSymbolicDevice.with_strides({strides...})— validate strides (omit to require contiguous).verify(tensor_view)— execute the check; throwsPanicErrorwith full context on failure; chainable (verify(a).verify(b)to check multiple tensors with the same shape)
host::is_type<T>(dtype)— whether aDLDataTypedenotes the C++ typeT(e.g.fp16_t).
Typical pattern:
auto N = SymbolicSize{"num_elements"};
auto device = SymbolicDevice{};
device.set_options<kDLCUDA>();
TensorMatcher({N}) //
.with_dtype<fp16_t>()
.with_device<kDLCUDA>(device)
.verify(dst)
.verify(src); // same shape, dtype, device as dst
const int64_t n = N.unwrap();
const DLDevice dev = device.unwrap();
const int64_t last_dim = 128;
TensorMatcher({N, last_dim}) // a fixed dimension can be a plain integer
.with_dtype<fp16_t>()
.with_device<kDLCUDA>(device)
.verify(tensor_2d);
ffi.h — Tensor allocation and blob wrapping (host::ffi::)
#include <sgl_kernel/ffi.h>
The counterpart to tensor.h: that one validates what came in, this one produces new tvm::ffi::Tensor values. Allocation goes through the environment allocator (TVMFFIEnvTensorAlloc), so buffers come from PyTorch's caching allocator rather than a raw cudaMalloc.
host::alloc_workspace_tensor(nbytes, device)(declared inutils.cuh) — the way to get scratch memory: a 1-Duint8tensor ofnbytes, or an empty tensor whennbytes == 0. Hold the returnedTensorin a local across every launch that touches it — it frees on destruction.host::ffi::empty(shape, dtype, device)— Uninitialized tensor;shapeaccepts a braced list, soffi::empty({rows, sizeof(Plan)}, dtype, device)works for a typed scratch array.host::ffi::empty_like(tensor_view)— Same shape, dtype, and device as an existing tensor.host::ffi::from_blob(data, shape, dtype, device[, deleter, stride, byte_offset])/from_blob_like(data, tensor_view, ...)— View memory you already own as aTensor, no copy. The default deleter does nothing, so ownership stays with the caller; pass one only when theTensorshould own the block. Strides default to contiguous.
type.cuh — DTypeTrait<T>, packed_t<T>, and reduction traits
#include <sgl_kernel/type.cuh>
DTypeTrait<T>— Static trait struct, specialized for integral types,fp32_t,fp16_t,bf16_t,fp8_e4m3_t, and their packed x2/x4 variants. Provides:DTypeTrait<T>::from(value)— convert from another type via the right CUDA intrinsic (e.g.fp32_t→fp16_t)DTypeTrait<T>::abs/max/min— type-dispatched math (fp32, fp16/bf16 scalar and x2, integrals)DTypeTrait<T>::sqrt/rsqrt/exp/sin/cos(x)—fp32_tonly- Metadata:
packed_t/unpacked_t/kVecSize(packed layout),kFloatMax(dtype max as float, e.g. 448.0f for fp8-e4m3),kZeroBits
packed_t<T>— Two-element packed alias:packed_t<fp16_t>=fp16x2_t,packed_t<bf16_t>=bf16x2_t,packed_t<fp32_t>=fp32x2_t. Use for vectorized loads/stores.device::cast<To, From>(value)— Type-safe cast usingDTypeTrait, e.g.cast<fp32x2_t, fp16x2_t>(v).device::unpack(value)— View a packed value as anunpacked_t[kVecSize]array reference (e.g.fp32x2_t→fp32_t[2]); element writes propagate back to the packed value.device::ReductionOp(SUM/MAX/MIN) anddevice::ReductionTrait<Op, T>::reduce(x, y)— One binary reduction step, dispatched throughDTypeTrait(packed types reduce elementwise). This is the engine behindwarp::reduce; use it directly when writing custom reductions.
vec.cuh — Vectorized memory access (AlignedVector)
#include <sgl_kernel/vec.cuh>
device::AlignedVector<T, N>— Aligned storage for N elements of type T. N must be a power of two,sizeof(T)*N <= 32. Enables vectorized loads/stores for bandwidth efficiency. In terms of API/codegen constraints, the upper bound is 256-bit; in practice, 128-bit is the portable default, while 256-bit vectorization is typically only viable onSM100+and should be gated by an architecture check when needed..load(ptr, offset)— vectorized load fromptr[offset].store(ptr, offset)— vectorized store toptr[offset].fill(value)— fill all N elements withvalueoperator[](i)— element access
tile.cuh — tile::Memory (strided memory access pattern)
#include <sgl_kernel/tile.cuh>
tile::Memory<T>is fundamentally a 1D cooperative accessor over a contiguous region.device::tile::Memory<T>::cta(blockDim.x)— Creates a tile accessor where each thread handlestid = threadIdx.xwith stridetsize(forcta(blockDim.x), this isblockDim.x). Common for loops over a 1D array..load(ptr, offset)— loadsptr[tid + offset * tsize].store(ptr, val, offset)— stores toptr[tid + offset * tsize].in_bound(n, offset)— boundary check
For a 2D tile, either flatten (row, col) into a linear tile index first, or compute the address manually with ptr[row * stride + col] using your thread/block coordinates.
math.cuh — Device math (device::math::)
#include <sgl_kernel/math.cuh>
device::math::max/min<T>(a, b)— type-dispatched binary math viaDTypeTraitdevice::math::abs/sqrt/rsqrt/exp/sin/cos<T>(x)— type-dispatched unary math viaDTypeTrait
warp.cuh — Warp-level primitives
#include <sgl_kernel/warp.cuh>
device::warp::reduce<Op, kNumThreads, kInner>(value, active_mask)— generic warp reduction via__shfl_xor_sync.Opis adevice::ReductionOp(SUM/MAX/MIN);kNumThreadsis a power-of-two group size (default 32 = full warp);kInner=true(default) reduces within eachkNumThreads-sized group,kInner=falsereduces across groups (lanes at the same offset in different groups).device::warp::reduce_sum/reduce_max/reduce_min<kNumThreads, kInner>(value)— convenience wrappers overreduce. Work for any type with aReductionTrait: floats, integers, and packed x2 types.
cta.cuh — CTA-level primitives
#include <sgl_kernel/cta.cuh>
device::cta::reduce_max<T>(value, smem, min_value)— CTA-wide max using shared memory + warp reduction. Caller is responsible for a__syncthreads()after if the result insmem[0]is needed.
atomic.cuh — Atomic operations
#include <sgl_kernel/atomic.cuh>
device::atomic::max(float* addr, float value)— float atomic max (handles negative values correctly via bit tricks).
runtime.cuh — Occupancy and device info
#include <sgl_kernel/runtime.cuh>
host::runtime::get_blocks_per_sm(kernel, block_dim)— max active blocks per SM (occupancy)host::runtime::get_sm_count(device_id)— number of SMs on the devicehost::runtime::get_cc_major(device_id)— compute capability major version
Persistent kernel pattern (cap blocks to SM count × occupancy):
static const uint32_t max_occ = runtime::get_blocks_per_sm(kernel, kBlockSize);
static const uint32_t num_sm = runtime::get_sm_count(device.unwrap().device_id);
const auto num_blocks = std::min(num_sm * max_occ, div_ceil(n, kBlockSize));
LaunchKernel(num_blocks, kBlockSize, device.unwrap())(kernel, params);
Step 0 (optional): Generate a .clangd config for better IDE support
python -m sglang.kernels.jit -h # for verbose help info about clangd configuration
python -m sglang.kernels.jit
python -m sglang.kernels.jit --dep cutlass flashinfer # with cutlass/flashinfer dependency
Step 1: Implement the CUDA kernel in kernels/jit/csrc/
Create python/sglang/kernels/jit/csrc/elementwise/scale.cuh.
The implementation fully uses the project abstractions described above:
// 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 DTypeTrait, fp16_t, bf16_t, fp32_t
#include <sgl_kernel/utils.h> // For CHECK_HOST, div_ceil
#include <sgl_kernel/utils.cuh> // For LaunchKernel, SGL_DEVICE
#include <sgl_kernel/vec.cuh> // For AlignedVector
#include <dlpack/dlpack.h>
#include <tvm/ffi/container/tensor.h>
namespace sglang {
/**
* \brief Element-wise scale using vectorized 128-bit loads/stores.
*
* \tparam T Element type: fp16_t | bf16_t | fp32_t
* \tparam kVecN Elements per vector load (e.g. 8 for fp16)
* \tparam kUsePDL Whether to emit the PDL wait/trigger pair
* \param dst Output buffer, `n_total` elements
* \param src Input buffer, `n_total` elements
* \param factor Runtime scale factor
* \param n_total Number of elements to scale
*/
template <typename T, int kVecN, bool kUsePDL>
__global__ void scale_kernel(T* __restrict__ dst,
const T* __restrict__ src,
float factor,
uint32_t n_total) {
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;
vi < n_vecs;
vi += vec_stride) {
vec_t v;
v.load(src, vi);
#pragma unroll
for (int i = 0; i < kVecN; ++i) {
v[i] = static_cast<T>(static_cast<float>(v[i]) * factor);
}
v.store(dst, vi);
}
// --- scalar tail ---
const uint32_t base = n_vecs * kVecN;
const uint32_t scalar_stride = blockDim.x * gridDim.x;
for (uint32_t i = blockIdx.x * blockDim.x + threadIdx.x;
base + i < n_total;
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>();
}
/**
* \brief Validate the tensors, select the vector width, launch `scale_kernel`.
*
* \tparam T Element type: fp16_t | bf16_t | fp32_t
* \tparam kUsePDL Whether to launch with PDL enabled
* \param dst Output tensor; same shape / dtype / device as `src`
* \param src Input tensor on CUDA
* \param factor Runtime scale factor
*/
template <typename T, bool kUsePDL>
void scale(tvm::ffi::TensorView dst, tvm::ffi::TensorView src, float factor) {
using namespace host;
// 1. Validate input tensors with TensorMatcher
SymbolicSize N = {"num_elements"};
SymbolicDevice device_;
device_.set_options<kDLCUDA>();
TensorMatcher({N}) //
.with_dtype<T>()
.with_device<kDLCUDA>(device_)
.verify(dst)
.verify(src); // same shape / dtype / device as dst
const uint32_t n = static_cast<uint32_t>(N.unwrap());
const DLDevice device = device_.unwrap();
CHECK_HOST(n > 0) << "scale: num_elements must be > 0, got " << n;
// 2. Choose vector width for 128-bit loads (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));
// 3. Launch
constexpr uint32_t kBlockSize = 256;
const uint32_t grid = div_ceil(n_work_items, kBlockSize);
// 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,
n);
}
} // namespace sglang
Key points:
- Include headers from
sgl_kernel/— not raw CUDA headers for anything already covered - Use
TensorMatcherfor all tensor validation; never manually check shape/dtype/device - Use
AlignedVectorfor vectorised 128-bit loads/stores — significant bandwidth win - Use
LaunchKernel— it resolves the stream and checks errors automatically - Use
CHECK_HOST(cond) << ...for runtime assertions with useful error messages (zero overhead when the check passes) - Prefer passing runtime scalars like
factordirectly unless compile-time specialisation is genuinely required fp16_t/bf16_t/fp32_tare the project's type aliases (fromutils.cuh)device::cast<To, From>orDTypeTrait<T>::from(val)for cross-type conversionsdevice::math::functions for device math instead of bare__intrinsics if possible.- Consider PDL — it can help when the kernel has prologue work to overlap. Place
PDLWaitPrimary()right before the first read of upstream data, not at the top of the kernel
Step 2: Add the Python wrapper in kernels/ops/
The wrapper lives next to its functional group under python/sglang/kernels/ops/, not beside the CUDA source — kernels/jit/ holds only the JIT infrastructure (csrc/, include/, utils/, benchmark/). Create python/sglang/kernels/ops/elementwise/scale.py:
from __future__ import annotations
from typing import TYPE_CHECKING
import torch
from sglang.kernels.jit.utils import (
cache_once,
is_arch_support_pdl,
load_jit,
make_cpp_args,
)
if TYPE_CHECKING:
from tvm_ffi.module import Module
@cache_once
def _jit_scale_module(dtype: torch.dtype) -> Module:
"""Compile and cache the JIT scale module for a given dtype."""
# Checks on the compile key live here, not in `scale`: `cache_once` keys on
# `dtype`, so this runs once per specialisation instead of once per call.
if dtype not in (torch.float16, torch.bfloat16, torch.float32):
raise RuntimeError(
f"Unsupported dtype {dtype}. Supported: float16, bfloat16, float32"
)
args = make_cpp_args(dtype, is_arch_support_pdl())
return load_jit(
"scale",
*args,
cuda_files=["elementwise/scale.cuh"],
cuda_wrappers=[("scale", f"scale<{args}>")],
)
def scale(src: torch.Tensor, factor: float, out: torch.Tensor | None = None) -> torch.Tensor:
"""
Element-wise scale: dst = src * factor.
Supported dtypes: torch.float16, torch.bfloat16, torch.float32.
Parameters
----------
src : CUDA tensor (FP16 / BF16 / FP32)
factor : scale factor
out : optional pre-allocated output tensor (same shape/dtype as src)
Returns
-------
Scaled tensor (dst = src * factor).
"""
# DO NOT add proactive validation here: every check costs interpreter time
# on a per-forward path. Tensor invariants belong in the C++ launcher, and
# anything about the compile key belongs in `_jit_scale_module`.
if out is None:
out = torch.empty_like(src)
module = _jit_scale_module(src.dtype)
module.scale(out, src, factor)
return out
Key points:
- Use
cache_once— notfunctools.lru_cache(incompatible withtorch.compile) load_jitfirst arg(s) form the unique build marker; same marker = same cached binary- Only include compile-time specialisation knobs in the build marker; runtime values like
factorshould stay runtime unless the kernel truly needs templating cuda_wrappers:(export_name, kernel_symbol)—export_nameis called from Pythonmake_cpp_args(dtype, ...)convertstorch.dtypeto 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 the entry point thin (see Conventions). What Python must still check goes in the
@cache_oncemodule factory, not in the entry point:cache_oncekeys on its arguments, so a check there costs one evaluation per specialisation instead of one per call — that is where the supported-dtype guard lives. Tensor invariants belong in the C++ launcher; if something here never reaches a.verify(...), close that gap on the C++ side rather than in Python
torch.dtype |
C++ type |
|---|---|
torch.float16 |
fp16_t |
torch.bfloat16 |
bf16_t |
torch.float32 |
fp32_t |
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):
return load_jit(
"scale",
*args,
cuda_files=["elementwise/scale.cuh"],
cuda_wrappers=[("scale", f"scale<{args}>")],
extra_cuda_cflags=["-O3", "--use_fast_math"],
)
If your kernel requires SM90+, raise a clear Python error before calling load_jit. Arch gating is one of the checks that has to live in Python — it decides whether to compile at all, so the C++ launcher never gets to run:
if torch.cuda.get_device_capability()[0] < 9:
raise RuntimeError("This kernel requires SM90 (Hopper) or later")
Step 4: Write tests (required)
JIT kernel correctness tests and benchmarks live under test/registered/kernels/ops/<group>/ and test/registered/kernels/benchmark/<group>/, mirroring the wrapper's group under python/sglang/kernels/ops/ (NOT inside the sglang package -- a register_*_ci(...) call anywhere under python/sglang/ is rejected by the check-no-registered-tests-in-package pre-commit hook). Only their test-only helpers (e.g. benchmark/marker.py) stay alongside the kernel source under python/sglang/kernels/jit/ and are imported by absolute path. CI does not run pytest in those directories directly. The unified runner test/run_suite.py discovers every test_*.py and bench_*.py under test/registered/, 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 usebase-b-kernel-unit-test-1-gpu-largeon H100 andbase-b-kernel-unit-test-4-gpu-b200on B200/SM100 paths (see.github/workflows/pr-test-jit-kernel.yml). Multi-GPU JIT tests usebase-b-kernel-unit-test-8-gpu-h200. - Nightly kernel suite: register with
stage="nightly"plus therunner_configof the machine it needs (e.g.1-gpu-large), giving thenightly-test-1-gpu-largesuite..github/workflows/nightly-test-nvidia.ymlsetsSGLANG_JIT_KERNEL_RUN_FULL_TESTS=1for the whole nightly run, so the expanded parameter grids apply automatically (seepython/sglang/kernels/jit/utils/common.py→should_run_full_tests/get_ci_test_range). There is no separate kernel-only nightly job: every nightly test on one machine type shares that machine's suite.
Registration pattern (module level, literal est_time, stage, and runner_config values — required for AST parsing):
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="1-gpu-large")
# Optional B200/SM100 registration for tests that cover Blackwell-specific code paths
# register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="4-gpu-b200")
# Optional second registration: same file also runs nightly, same form,
# stage is just "nightly" there (and no `nightly=True`)
# register_cuda_ci(est_time=120, stage="nightly", runner_config="1-gpu-large")
CI generates the suite name as {stage}-test-{runner_config}, so stage="base-b-kernel-unit", runner_config="1-gpu-large" becomes the base-b-kernel-unit-test-1-gpu-large suite you pass to run_suite.py below — don't put the -test- infix in register_cuda_ci. Nightly uses the same shape with stage="nightly"; the single-string suite= form is left only for stress and non-CUDA pools.
Keep est_time, stage, runner_config, and suite as literal values. run_suite.py collects them from the file AST, so computed values and helper wrappers can break CI discovery.
Use register_cuda_ci(..., disabled="reason") if the file must stay in-tree but should be skipped in CI (e.g. multi-GPU only).
Run like CI (from repo root):
(cd test && python3 run_suite.py --hw cuda --suite base-b-kernel-unit-test-1-gpu-large)
# For B200/SM100-specific coverage:
(cd test && python3 run_suite.py --hw cuda --suite base-b-kernel-unit-test-4-gpu-b200)
For fast iteration you can still run pytest on a single file locally; CI coverage is via run_suite.py.
Create test/registered/kernels/ops/elementwise/test_scale.py:
import pytest
import torch
from sglang.kernels.ops.elementwise.scale import scale
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="1-gpu-large")
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16, torch.float32])
@pytest.mark.parametrize("size", [1, 127, 128, 1024, 4097]) # cover tail remainder
@pytest.mark.parametrize("factor", [0.5, 1.0, 2.0, 3.0])
def test_scale_correctness(dtype, size, factor):
src = torch.randn(size, dtype=dtype, device="cuda")
out = scale(src, factor)
expected = src * factor
rtol, atol = (1e-5, 1e-6) if dtype == torch.float32 else (1e-2, 1e-2)
torch.testing.assert_close(out, expected, rtol=rtol, atol=atol)
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16, torch.float32])
def test_scale_out_param(dtype):
src = torch.randn(1024, dtype=dtype, device="cuda")
out = torch.empty_like(src)
result = scale(src, 2.0, out=out)
assert result is out
torch.testing.assert_close(out, src * 2.0, rtol=1e-2, atol=1e-2)
def test_scale_cpu_error():
src = torch.randn(128, dtype=torch.float16) # CPU tensor
with pytest.raises(RuntimeError, match="CUDA"):
scale(src, 2.0)
def test_scale_unsupported_dtype():
src = torch.randint(0, 10, (128,), dtype=torch.int32, device="cuda")
with pytest.raises(RuntimeError, match="dtype"):
scale(src, 2.0)
if __name__ == "__main__":
import sys
sys.exit(pytest.main([__file__, "-v", "-s"]))
Step 5: Add a benchmark (required)
Benchmarks are bench_*.py files under test/registered/kernels/benchmark/<group>/. They are picked up by the same run_suite.py machinery as unit tests. Register them for base-b-kernel-benchmark-test-1-gpu-large (PR JIT benchmark job: python3 run_suite.py --hw cuda --suite base-b-kernel-benchmark-test-1-gpu-large).
Benchmarks use the project's own marker framework (in python/sglang/kernels/jit/benchmark/marker.py) — do not use triton.testing.perf_report / triton.testing.do_bench directly. The marker framework provides (public names: benchmark, parametrize, do_bench, skip, BenchResult, BenchSkip):
@marker.benchmark(line_arg, line_vals, *, unit="us")— the innermost decorator (bottom of the stack, directly abovedef benchmark). Declares the column axis: each value inline_valsbecomes a result column, andline_argis the parameter name passed into the benchmark function.unitis one of"us" | "ms" | "s".@marker.parametrize(names, vals, ci_vals=None)— stackable decorator that adds a row axis (pytest-style). Each@parametrizeadds one (or more, correlated) parameter the benchmark is swept over (Cartesian product across allparametrizedecorators).namesmay be a single name ("size") or a comma-separated correlated tuple axis ("h,d", withvalsthen a list of tuples like[(1, 64), (2, 128)]). Pass the optional thirdci_valsfor a smaller sweep that is auto-selected underis_in_ci()— this is the built-in CI-shrinking mechanism, so you usually don't needget_benchmark_rangefor swept axes.marker.do_bench(fn, *, input_args=(), input_kwargs={}, ...)— runsfnunder CUDA graph (default) or a naive loop, returns aBenchResult. Key knobs:memory_args: defaults to"all"(footprint derived from all input args/kwargs). Pass an explicit tuple of tensors (e.g.(k, v, indices)) to count only the inputs the kernel actually touches.memory_output: defaults to"out"— re-runsfnonce to capture its returned tensor and counts it. For in-place kernels (which returnNone), pass the written tensors explicitly (e.g.memory_output=(k, v)); the re-run is then skipped. Set toNoneto count no output.- Together
memory_args+memory_outputgive the GB/s column; with both defaults a functionout = f(src)already reportsbytes(src) + bytes(out). graph_clone_args/graph_clone_kwargs: which inputs to clone per CUDA-graph iteration to defeat L2 cache reuse. Defaults to"all"— pass an iterable of indices/keys to limit to the read args (writes don't need cloning).use_cuda_graph=Falsefor kernels that can't be captured.metrics=(0.5, "avg")controls reported quantiles (the first metric becomes the table latency column).disable_log_bandwidth(defaults fromSGLANG_KERNEL_DISABLE_LOG_BANDWIDTH=1) skips the bandwidth column entirely.
utils.create_random(*shape)/utils.create_empty(*shape)— shorthand fortorch.randn/torch.emptywithDEFAULT_DTYPE(bfloat16) andDEFAULT_DEVICE("cuda"). Override via thedtype=/device=kwargs.utils.get_benchmark_range(full_range, ci_range)— returns the smallerci_rangeunder CI (is_in_ci()), thefull_rangelocally. Still available for thebenchmark(...)column axis (which has noci_vals); forparametrizerow axes prefer the built-inci_valsargument.
Create test/registered/kernels/benchmark/elementwise/bench_scale.py:
import torch
from sglang.kernels.jit.benchmark import marker
from sglang.kernels.jit.benchmark.utils import create_random
from sglang.kernels.ops.elementwise.scale import scale as jit_scale
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=6, stage="base-b-kernel-benchmark", runner_config="1-gpu-large")
@torch.compile()
def torch_impl_scale(src: torch.Tensor, factor: float) -> torch.Tensor:
return src * factor
FN_MAP = {
"jit": jit_scale,
"torch": torch_impl_scale,
}
# `parametrize(name, full_vals, ci_vals)`: the 3rd arg is the smaller sweep
# auto-selected under CI; the full range runs locally.
@marker.parametrize("size", [2**n for n in range(10, 20)], [4096, 65536]) # 1K .. 512K
@marker.benchmark("impl", ["jit", "torch"])
def benchmark(size: int, impl: str):
src = create_random(size)
factor = 2.0
return marker.do_bench(
FN_MAP[impl],
input_args=(src, factor),
# `src` is read -> clone it per iter to avoid L2 reuse; factor is a scalar.
graph_clone_args=(0,),
# Defaults already report bandwidth: memory_args="all" counts src,
# memory_output="out" counts the returned tensor -> bytes(src)+bytes(out).
)
if __name__ == "__main__":
benchmark.run()
Key points:
- The
line_argname passed tobenchmark("impl"here) must match a parameter onbenchmark(...); same for everyparametrizename ("size"). - Stack
@parametrizeonce per swept axis. The required@marker.benchmarkis the innermost decorator (bottom of the stack, directly above the function) —@parametrizerows go above it. - Prefer
create_random/create_emptyfromutils.pyover open-codingtorch.randn(..., dtype=..., device=...). - The GB/s column appears by default (
memory_args="all"+memory_output="out"). For memory-bound kernels it's the most informative number; scopememory_args/memory_outputto the tensors actually touched if the defaults over- or under-count. For compute-bound kernels where bandwidth is misleading, setSGLANG_KERNEL_DISABLE_LOG_BANDWIDTH=1(ordisable_log_bandwidth=True). - For in-place kernels (which return
None), pass the written tensors viamemory_output=(...)since the"out"default would capture nothing. - Tune
graph_clone_args/graph_clone_kwargsto all the arguments that might be read by the kernel. We can only skip cloning for write-only args. For in-place modified args, we still need to clone them to get accurate timing (reusing the same buffer keeps it L2-hot and skews results). - Call
benchmark.run()(noprint_data=kwarg — the marker framework prints directly).
Run locally:
python test/registered/kernels/benchmark/elementwise/bench_scale.py
Run the benchmark suite the way CI does:
cd test && python3 run_suite.py --hw cuda --suite base-b-kernel-benchmark-test-1-gpu-large
Troubleshooting
No CI registry found in ...fromrun_suite.py: add a module-levelregister_cuda_ci(...)with literalest_time,stage, andrunner_config; starred args and non-literal values break AST collection- JIT compilation fails: ensure the
.cuhfile is underpython/sglang/kernels/jit/csrc/; reduce template argument combinations - CUDA crash / illegal memory access:
CUDA_LAUNCH_BLOCKING=1;compute-sanitizer --tool memcheck python ... - Unstable benchmark results:
marker.do_benchuses CUDA-graph-based timing by default; setuse_cuda_graph=Falseonly if the kernel can't be captured.graph_clone_argsdefaults to"all"; if you narrow it, it must still cover every read tensor — reusing a single buffer keeps it L2-hot and skews results. Keep write tensors in it too: they are what sets the rotation count, and a shared output buffer stays L2-hot the same way. - Missing GB/s column: the column is on by default; check that
SGLANG_KERNEL_DISABLE_LOG_BANDWIDTHis not1anddisable_log_bandwidthis notTrue. For in-place kernels (returnNone) thememory_output="out"default counts nothing — pass the written tensors viamemory_output=(...)
References
docs/docs/developer_guide/development_jit_kernel_guide.mdxtest/run_suite.py— suite names, discovery oftest/registered/, execution entrypoint for CIpython/sglang/test/ci/ci_register.py—register_cuda_ciand AST registration rulespython/sglang/kernels/jit/utils/compile.py—load_jit,make_cpp_argspython/sglang/kernels/jit/utils/common.py—cache_once,should_run_full_tests,get_ci_test_rangepython/sglang/kernels/jit/include/sgl_kernel/tensor.h—TensorMatcher,SymbolicSize/DType/Device,is_typepython/sglang/kernels/jit/include/sgl_kernel/ffi.h—ffi::empty,ffi::empty_like,ffi::from_blobpython/sglang/kernels/jit/include/sgl_kernel/utils.cuh— type aliases,LaunchKernel,SGL_DEVICEpython/sglang/kernels/jit/include/sgl_kernel/vec.cuh—AlignedVectorpython/sglang/kernels/jit/include/sgl_kernel/tile.cuh—tile::Memorypython/sglang/kernels/jit/include/sgl_kernel/type.cuh—DTypeTrait,packed_t,device::cast,device::unpack,ReductionTraitpython/sglang/kernels/jit/include/sgl_kernel/math.cuh—device::math::python/sglang/kernels/jit/include/sgl_kernel/warp.cuh—warp::reduce<Op>andreduce_sum/max/minwrapperspython/sglang/kernels/jit/include/sgl_kernel/cta.cuh—cta::reduce_maxpython/sglang/kernels/jit/include/sgl_kernel/atomic.cuh—atomic::maxpython/sglang/kernels/jit/include/sgl_kernel/runtime.cuh— occupancy / SM count helperspython/sglang/kernels/jit/csrc/add_constant.cuh— minimal runnable referencepython/sglang/kernels/jit/csrc/elementwise/rmsnorm.cuh— real example usingTensorMatcher+LaunchKernel+tile::Memorypython/sglang/kernels/jit/csrc/elementwise/qknorm.cuh— real example usingruntime::get_blocks_per_sm+ persistent kernel patternpython/sglang/kernels/jit/benchmark/marker.py—benchmark,parametrize,do_bench,BenchResultpython/sglang/kernels/jit/benchmark/utils.py—create_random/create_empty/get_benchmark_rangehelpers andDEFAULT_DTYPE/DEFAULT_DEVICEtest/registered/kernels/benchmark/layernorm/bench_qknorm.py— real example: multi-axisparametrize(withci_vals) + in-placememory_outputtest/registered/kernels/benchmark/kvcache/bench_store_cache.py— real example: scopedmemory_args/memory_output+ selectivegraph_clone_args
Summary of Files Created
python/sglang/kernels/jit/csrc/elementwise/scale.cuh # NEW: CUDA kernel
python/sglang/kernels/ops/elementwise/scale.py # NEW: Python wrapper
test/registered/kernels/ops/elementwise/test_scale.py # NEW: Tests
test/registered/kernels/benchmark/elementwise/bench_scale.py # NEW: Benchmark
Version History
- 1df78c2 Current 2026-08-20 08:17


