--- 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 scalar `factor` (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 in `python/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 through `flashinfer`, can still be implemented as `jit_kernel`. --- ## Conventions These hold for every step below. - **`namespace sglang` is 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 shared `host::` / `device::` helpers are nested in it too, so they resolve unqualified. `load_jit` emits the `TVM_FFI_DLL_EXPORT_TYPED_FUNC` wrapper inside `namespace sglang` as well, so the `kernel_name` you pass from Python needs no `sglang::` prefix. - **Tuning policy belongs in Python; C++ gets the answer, not the decision.** For a hyperparameter — split factor, threads per item, block size, vector width, an algorithm variant — prefer (not required) making it a **template parameter** and letting the `@cache_once` module factory choose the value, over a runtime `if`/`switch` in the launcher that picks among pre-instantiated kernels. The kernel then compiles exactly one specialisation and asserts its own preconditions with `static_assert`, while the heuristic that produced the value sits in Python where it is readable, adjustable, and inspectable without recompiling CUDA. `store_cache`'s `num_threads` is the worked example: a `get_kernel(num_split)` ladder guarded by byte-alignment `if constexpr`s became one template argument plus a heuristic in `_jit_kvcache_module`. The exception is a knob that genuinely varies per call with the runtime shape — that has to stay a kernel argument. - **Check where the check is cheapest: `static_assert` > C++ host check > cached Python > per-call Python.** Anything fixed at compile time is a `static_assert`. Anything about the tensors is a `TensorMatcher` / `CHECK_HOST` in the C++ launcher, free next to a kernel launch. A check Python cannot delegate goes inside the `@cache_once` module 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 allocating `out`. - **Fixed-width integer types.** Prefer `int32_t` / `int64_t` / `uint32_t` / `size_t` over `int`, `long`, or `long long`, so an index has the same width on both sides of the FFI boundary. Bare `int` is fine only where the width plainly cannot matter — an unrolled loop counter over a `constexpr` bound, a template `int` parameter. Shapes arrive as `int64_t` (`SymbolicSize::unwrap()`); narrowing to `uint32_t` for in-kernel indexing is a deliberate act, so write the `static_cast` explicitly and only where the range is known. - **Doxygen comments in C++.** Document exported entities with `///` or `/** ... */` blocks using `\brief`, `\param`, `\tparam`, `\return`, the way `include/sgl_kernel/` does. `python -m sglang.kernels.jit` writes `CommentFormat: Doxygen` into `.clangd` when 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]' ` before committing. - **`namespace details` is private.** Anything inside one is an implementation detail of its own header, free to change without notice. Do not name `details::` from another module, and never `using namespace details`. If you find yourself wanting something in there, that is the signal to promote it to a documented name instead. - **`const T* __restrict__` for read-only pointers.** This is what `csrc/` 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; prefer recomputing a value over letting it spill. Run `SGLANG_JIT_LOG_RESOURCE_USAGE=1 SGLANG_JIT_FORCE_RECOMPILE=1` to log per-kernel registers, spills and shared memory at INFO. Both halves are load-bearing: a cache hit has no compiler output to report, and passing `-Xptxas -v` by hand does nothing on its own because the build captures compiler output and replays it only on failure. - **Alignment is a contract you enforce, not a property you hope for.** A vectorized copy derives its access width from the row *size* (`load_bytes` picks 16B for a 1024B row), but the address it lands on is `base + index * row_stride` — and a stride is not constrained by the size. `TensorMatcher` admits any stride unless you say otherwise, so a padded cache row silently produces `misaligned address`. Every kernel that vectorizes is highly recommended to end its validation with `.ensure_alignment(w)`, where `w` is the width the kernel actually uses. --- ## 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 ```cpp #include ``` - **`CHECK_HOST(cond) << "msg " << value`** — **Preferred** runtime check: stream-style, throws `PanicError` with 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 to `CHECK_HOST`. Note its message args are always evaluated (even when the check passes), so prefer `CHECK_HOST` — especially on hot paths. - **`host::Panic(args...)`** — Unconditionally throw a `PanicError` with 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 on `void*`. Use this instead of raw casts. ### `utils.cuh` — Device-side utilities + `LaunchKernel` ```cpp #include ``` - **Type aliases**: `fp16_t`, `bf16_t`, `fp32_t`, `fp8_e4m3_t`, `fp8_e5m2_t` and their packed variants `fp16x2_t`, `bf16x2_t`, `fp32x2_t`, etc. - **`SGL_DEVICE`** — Expands to `__forceinline__ __device__`. Use on all device functions. - **`SGL_DEVICE_HOST`** — `__forceinline__ __device__ __host__`. Use it when a `constexpr` helper must be callable from both the kernel and its launcher, so the two cannot drift. - **`device::kWarpThreads`** — Constant `32` (the codebase's *logical* warp width). `kWarpSize` is an alias. Note this is a group size, not the hardware wave: on gfx950 `warpSize` is 64, and `warp::kFullWidth` reflects that. - **`device::get_lane_id()`** — this thread's index within its logical group. **Prefer it over `threadIdx.x % kNumThreads` whenever the value feeds an address.** On CUDA the lane register is one read that folds into the address computation, whereas the modulo is a derived value the compiler re-materialises at *every* address scale — 8 instructions and 2 registers on a two-tile warp copy, measured on sm_100a. It picks the cheaper form per platform, so callers never need to care. Only equals the true in-warp lane when `blockDim.x` is a multiple of `kNumThreads`. - **`device::load_as(ptr, offset)`** / **`device::store_as(ptr, val, offset)`** — Type-safe loads/stores from `void*`. - **`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 `DLDevice` via 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+). - **`device::PDLWaitPrimary()`** / **`device::PDLTriggerSecondary()`** — 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; evaluates `expr` once and throws `PanicError` with `cudaGetErrorString` + file/line info if it is not `cudaSuccess`. Extra streamed context is optional. - **`host::RuntimeDeviceCheck(cudaError_t)`** — Function-style alternative to `CHECK_CUDA`. It takes no context message, so prefer `CHECK_CUDA`, which builds its error object only on the failure path. ### `tensor.h` — Tensor validation (`TensorMatcher`, Symbolic types) ```cpp #include ``` 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()` to restrict allowed types. - **`host::SymbolicDevice`** — Symbolic device. Use `.set_options()` to restrict to CUDA. - **`host::TensorMatcher({dims...})`** — Fluent builder for tensor validation: - `.with_dtype()` — require a specific C++ type (e.g. `fp16_t`) - `.with_dtype()` — allow a set of types - `.with_device(device_sym)` — require CUDA and bind the checked device to a `SymbolicDevice` - `.with_strides({strides...})` — validate strides (omit to require contiguous) - `.ensure_alignment(bytes)` — require the data pointer **and every stride except the innermost** to be a multiple of `bytes` (a power of two). Size-1 dimensions are skipped, since their strides are arbitrary. This is how a vectorized kernel states its precondition; see the alignment convention above. - `.verify(tensor_view)` — execute the check; throws `PanicError` with full context on failure; **chainable** (`verify(a).verify(b)` to check multiple tensors with the same shape) - **`host::is_type(dtype)`** — whether a `DLDataType` denotes the C++ type `T` (e.g. `fp16_t`). **Typical pattern:** ```cpp auto N = SymbolicSize{"num_elements"}; auto device = SymbolicDevice{}; device.set_options(); TensorMatcher({N}) // .with_dtype() .with_device(device) .ensure_alignment(16) // e.g. for 128-bit vectorized load/store .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() .with_device(device) .verify(tensor_2d); ``` ### `ffi.h` — Tensor allocation and blob wrapping (`host::ffi::`) ```cpp #include ``` 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 in `utils.cuh`) — **the way to get scratch memory**: a 1-D `uint8` tensor of `nbytes`, or an empty tensor when `nbytes == 0`. Hold the returned `Tensor` in a local across every launch that touches it — it frees on destruction. - **`host::ffi::empty(shape, dtype, device)`** — Uninitialized tensor; `shape` accepts a braced list, so `ffi::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 a `Tensor`, no copy. The default deleter does nothing, so ownership stays with the caller; pass one only when the `Tensor` should own the block. Strides default to contiguous. ### `type.cuh` — `DTypeTrait`, `packed_t`, and reduction traits ```cpp #include ``` - **`DTypeTrait`** — Static trait struct, specialized for integral types, `fp32_t`, `fp16_t`, `bf16_t`, `fp8_e4m3_t`, and their packed x2/x4 variants. Provides: - `DTypeTrait::from(value)` — convert from another type via the right CUDA intrinsic (e.g. `fp32_t` → `fp16_t`) - `DTypeTrait::abs/max/min` — type-dispatched math (fp32, fp16/bf16 scalar and x2, integrals) - `DTypeTrait::sqrt/rsqrt/exp/sin/cos(x)` — `fp32_t` only - Metadata: `packed_t` / `unpacked_t` / `kVecSize` (packed layout), `kFloatMax` (dtype max as float, e.g. 448.0f for fp8-e4m3), `kZeroBits` - **`packed_t`** — Two-element packed alias: `packed_t` = `fp16x2_t`, `packed_t` = `bf16x2_t`, `packed_t` = `fp32x2_t`. Use for vectorized loads/stores. - **`device::cast(value)`** — Type-safe cast using `DTypeTrait`, e.g. `cast(v)`. - **`device::unpack(value)`** — View a packed value as an `unpacked_t[kVecSize]` array reference (e.g. `fp32x2_t` → `fp32_t[2]`); element writes propagate back to the packed value. - **`device::ReductionOp` (`SUM`/`MAX`/`MIN`) and `device::ReductionTrait::reduce(x, y)`** — One binary reduction step, dispatched through `DTypeTrait` (packed types reduce elementwise). This is the engine behind `warp::reduce`; use it directly when writing custom reductions. ### `vec.cuh` — Vectorized memory access (`AlignedVector`) ```cpp #include ``` - **`device::AlignedVector`** — 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 on `SM100+` and should be gated by an architecture check when needed. - `.load(ptr, offset)` — vectorized load from `ptr[offset]` - `.store(ptr, offset)` — vectorized store to `ptr[offset]` - `.fill(value)` — fill all N elements with `value` - `operator[](i)` — element access - **`device::LoadStoreBytes`** — named widths to pass around instead of bare integers: `MAX_GMEM` (arch-dependent, 32 on Blackwell), `MAX_SMEM` (16), `MAX_PORTABLE` (16, safe on CUDA and HIP), `MIN_COALALESCED` (4), and `RAW_1B` .. `RAW_32B`. ### `tile.cuh` — `tile::Memory` (strided memory access pattern) ```cpp #include ``` - `tile::Memory` is fundamentally a **1D cooperative accessor** over a contiguous region. - **`device::tile::Memory::cta(blockDim.x)`** — Creates a tile accessor where each thread handles `tid = threadIdx.x` with stride `tsize` (for `cta(blockDim.x)`, this is `blockDim.x`). Common for loops over a 1D array. - **`device::tile::Memory::warp()`** — the 32-lane flavour, and what `warp::load_bytes` is built on. It takes its lane index from `get_lane_id()` rather than `threadIdx.x % 32`, for the addressing reason described under `utils.cuh`. Use `warp(int n)` for a narrower sub-group; that overload keeps the modulo, since a hardware lane index is the wrong grouping below 32. - **`device::tile::Memory::thread()`** — single-thread accessor, no cooperation. - **`.load(ptr, offset)`** — loads `ptr[tid + offset * tsize]` - **`.store(ptr, val, offset)`** — stores to `ptr[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. ### `bits.h` — Compile-time bit helpers (`host::`) ```cpp #include ``` `constexpr` wrappers over ``, usable in `static_assert` and in template arguments: `host::is_pow2(x)`, `log2_floor(x)`, `log2_ceil(x)` (both `-1` for `x == 0`), `round_up_pow2(x)`, `round_down_pow2(x)`. They live in `host::` but are constant-expression-only, so device-side `static_assert` can call them. ### `math.cuh` — Device math (`device::math::`) ```cpp #include ``` - `device::math::max/min(a, b)` — type-dispatched binary math via `DTypeTrait` - `device::math::abs/sqrt/rsqrt/exp/sin/cos(x)` — type-dispatched unary math via `DTypeTrait` ### `warp.cuh` — Warp-level primitives ```cpp #include ``` **Vectorized row copy — reach for this before hand-rolling.** A warp cooperatively moving one contiguous row is the single most common shape in this tree, and it used to be re-implemented per kernel. - **`warp::load_bytes(src)`** / **`warp::store_bytes(dst, val)`** — the whole warp moves `kBytes` from/to one contiguous row. `kBytes` is arbitrary, down to 1; the helper picks the widest vector that divides it and handles the ragged tail. The returned value is opaque and encodes the width, so a load and its matching store **must use the same ``**. - **`warp::LoadStorePattern`** — the `kPattern` values, and the sign carries the meaning. A **negative** one (`WARP_UNIFORM_16B` and friends: `_GMEM`, `_SMEM`, `_4B`/`_8B`/`_32B`) says *the whole warp splits this row*, so the width is derived from the per-lane share (`gcd(kBytes / 32, |kPattern|)`), falling back to `gcd(kBytes, 4)` when the row is too narrow to give every lane 4 bytes. A **positive** one (a bare `LoadStoreBytes` value) makes no warp-splitting assumption and just caps the per-thread vector at `gcd(kBytes, kPattern)`. `WARP_UNIFORM_16B` is the usual choice for a row copy. - **`LoadStorePattern::get_vec_bytes()`** — the width actually chosen. `SGL_DEVICE_HOST`, so **the launcher must call this to feed `.ensure_alignment(...)`**. Pass it the same value the kernel passes `load_bytes` — for a split row that is the *per-warp share*, not the whole row; the full row can resolve to a different width and would under-constrain the strides. **Reductions.** - `warp::reduce(value, active_mask)` — `Op` is a `device::ReductionOp` (`SUM`/`MAX`/`MIN`). `kStart` and `kFinish` bound a range of lane-index bits: `` reduces within contiguous groups of `N` lanes, `` reduces *across* groups at the same offset. The operation is **symmetric** — `` and `` reduce the same lane set. - `warp::reduce_sum/reduce_max/reduce_min(value)` — wrappers. Work for any type with a `ReductionTrait`: floats, integers, packed x2. This is more widely used in practice than generic `reduce`. - `warp::inclusive_reduce(val, lane_id, mask)` and `inclusive_sum` / `inclusive_max` / `inclusive_min` — segmented inclusive scan: every lane keeps its own running total, unlike `reduce`. Forward when `kStart < kFinish`, backward when `kStart > kFinish`. `lane_id` defaults to `get_lane_id()`, which is correct by construction — only pass it if you already have the **segment-relative** index (`threadIdx.x % kWidth`); a warp-relative one silently corrupts every segment past the first when `kWidth < 32`. - `warp::broadcast(value, src_lane)` — `src_lane` is **segment-relative**: each `kWidth` segment reads its own lane, not one warp-wide source. - `warp::get_lane_id`, `warp::elect_one_lane()` (one elected lane, for gating a single-thread TMA issue; CUDA-only), `warp::kFullMask` and `warp::kFullWidth` (32 on CUDA, 64 on HIP — note this differs from `kWarpThreads`). ### `cta.cuh` — CTA-level primitives ```cpp #include ``` - `device::cta::reduce_max(value, smem, min_value)` — CTA-wide max using shared memory + warp reduction. Caller is responsible for a `__syncthreads()` after if the result in `smem[0]` is needed. ### `atomic.cuh` — Atomic operations ```cpp #include ``` - `device::atomic::max(float* addr, float value)` — float atomic max (handles negative values correctly via bit tricks). - **`device::atomic::Event`** — a cross-CTA arrive/wait counter in one 32-bit word. Producers call `arrive()`; consumers call `wait(num_producers)` (exactly one consumer) or `wait_multi(num_producers, num_consumers, n)` (several). The last released consumer resets the word, so one `Event` is reusable across launches with no host re-zero. Preconditions, all of them load-bearing: - the word must be **zeroed by the host** before first use, and must live in **global** memory — a shared or local `Event` is an illegal address; - `arrive()` is a release and the waits are acquires, so a producer's writes before `arrive()` are visible after the wait; - **generations must be explicitly ordered** — the word has no phase bit, so overlapping generation N+1 with a consumer still in generation N is undefined behaviour (a lagging consumer never exits); - `wait()` and `wait_multi()` lay the word out incompatibly; mixing them on one `Event` is undefined behaviour. ### `runtime.cuh` — Occupancy and device info ```cpp #include ``` - `host::runtime::get_blocks_per_sm(kernel, block_dim)` — max active blocks per SM (occupancy) - `host::runtime::get_sm_count(device_id [, use_cache])` — number of SMs on the device. Memoized per device ordinal, so it is cheap enough to call on a launch path; pass `use_cache=false` to force a driver query. - `host::runtime::get_cc_major` / `get_cc_minor` / `get_sm_version` — compute capability, same caching. **Do not query the architecture at runtime.** The JIT compiles for the exact local GPU, so the arch is a *compile-time* fact: `SGL_CUDA_ARCH` is injected by `load_jit`, and `SGL_ARCH_HOPPER_OR_GREATER` / `SGL_ARCH_BLACKWELL_OR_GREATER` / `device::kMaxVecBytes` are derived from it. Reach for `get_cc_*` only for something the arch genuinely does not determine. SM *count* is the opposite case — it varies within an arch (B200 vs B300 vs a MIG slice), so it has to stay a runtime query. **Persistent kernel pattern** (cap blocks to SM count × occupancy): ```cpp 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 ```bash 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: ```cpp // NOTE: Comments for headers are not common in practice. // It is only shown here for tutorial purposes to highlight the key abstractions. #include // For TensorMatcher, SymbolicSize, SymbolicDevice #include // For DTypeTrait, fp16_t, bf16_t, fp32_t #include // For CHECK_HOST, div_ceil #include // For LaunchKernel, SGL_DEVICE #include // For AlignedVector #include #include 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 __global__ void scale_kernel(T* __restrict__ dst, const T* __restrict__ src, float factor, uint32_t n_total) { using vec_t = device::AlignedVector; 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(); // --- 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(static_cast(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(static_cast(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(); } /** * \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 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(); TensorMatcher({N}) // .with_dtype() .with_device(device_) .verify(dst) .verify(src); // same shape / dtype / device as dst const uint32_t n = static_cast(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(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, static_cast(dst.data_ptr()), static_cast(src.data_ptr()), factor, n); } } // namespace sglang ``` **Key points:** - Include headers from `sgl_kernel/` — **not** raw CUDA headers for anything already covered - Use `TensorMatcher` for all tensor validation; never manually check shape/dtype/device - Use `AlignedVector` for 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 `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` or `DTypeTrait::from(val)` for cross-type conversions - `device::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`: ```python 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` — **not** `functools.lru_cache` (incompatible with `torch.compile`) - `load_jit` first 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 `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 the entry point thin (see Conventions). What Python must still check goes in the `@cache_once` module factory, not in the entry point: `cache_once` keys 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): ```python 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: ```python 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//` and `test/registered/kernels/benchmark//`, 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 use `base-b-kernel-unit-test-1-gpu-large` on H100 and `base-b-kernel-unit-test-4-gpu-b200` on B200/SM100 paths (see `.github/workflows/pr-test-jit-kernel.yml`). Multi-GPU JIT tests use `base-b-kernel-unit-test-8-gpu-h200`. - **Nightly kernel suite**: register with `stage="nightly"` plus the `runner_config` of the machine it needs (e.g. `1-gpu-large`), giving the `nightly-test-1-gpu-large` suite. `.github/workflows/nightly-test-nvidia.yml` sets `SGLANG_JIT_KERNEL_RUN_FULL_TESTS=1` for the whole nightly run, so the expanded parameter grids apply automatically (see `python/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): ```python 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): ```bash (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`: ```python 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//`. 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 above `def benchmark`). Declares the column axis: each value in `line_vals` becomes a result column, and `line_arg` is the parameter name passed into the benchmark function. `unit` is one of `"us" | "ms" | "s"`. - **`@marker.parametrize(names, vals, ci_vals=None)`** — stackable decorator that adds a row axis (pytest-style). Each `@parametrize` adds one (or more, correlated) parameter the benchmark is swept over (Cartesian product across all `parametrize` decorators). `names` may be a single name (`"size"`) or a comma-separated correlated tuple axis (`"h,d"`, with `vals` then a list of tuples like `[(1, 64), (2, 128)]`). Pass the optional third `ci_vals` for a smaller sweep that is auto-selected under `is_in_ci()` — this is the built-in CI-shrinking mechanism, so you usually don't need `get_benchmark_range` for swept axes. - **`marker.do_bench(fn, *, input_args=(), input_kwargs={}, ...)`** — runs `fn` under CUDA graph (default) or a naive loop, returns a `BenchResult`. 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-runs `fn` once to capture its **returned** tensor and counts it. For in-place kernels (which return `None`), pass the written tensors explicitly (e.g. `memory_output=(k, v)`); the re-run is then skipped. Set to `None` to count no output. - Together `memory_args` + `memory_output` give the GB/s column; with both defaults a function `out = f(src)` already reports `bytes(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=False` for 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 from `SGLANG_JIT_BENCHMARK_DISABLE_LOG_BANDWIDTH=1`) skips the bandwidth column entirely. - **`utils.create_random(*shape)` / `utils.create_empty(*shape)`** — shorthand for `torch.randn` / `torch.empty` with `DEFAULT_DTYPE` (`bfloat16`) and `DEFAULT_DEVICE` (`"cuda"`). Override via the `dtype=` / `device=` kwargs. - **`utils.get_benchmark_range(full_range, ci_range)`** — returns the smaller `ci_range` under CI (`is_in_ci()`), the `full_range` locally. Still available for the `benchmark(...)` column axis (which has no `ci_vals`); for `parametrize` row axes prefer the built-in `ci_vals` argument. Create `test/registered/kernels/benchmark/elementwise/bench_scale.py`: ```python 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_arg` name passed to `benchmark` (`"impl"` here) must match a parameter on `benchmark(...)`; same for every `parametrize` name (`"size"`). - Stack `@parametrize` once per swept axis. The required `@marker.benchmark` is the **innermost** decorator (bottom of the stack, directly above the function) — `@parametrize` rows go above it. - Prefer `create_random` / `create_empty` from `utils.py` over open-coding `torch.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; scope `memory_args` / `memory_output` to the tensors actually touched if the defaults over- or under-count. For compute-bound kernels where bandwidth is misleading, set `SGLANG_JIT_BENCHMARK_DISABLE_LOG_BANDWIDTH=1` (or `disable_log_bandwidth=True`). - For in-place kernels (which return `None`), pass the written tensors via `memory_output=(...)` since the `"out"` default would capture nothing. - Tune `graph_clone_args` / `graph_clone_kwargs` to 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()` (no `print_data=` kwarg — the marker framework prints directly). Run locally: ```bash python test/registered/kernels/benchmark/elementwise/bench_scale.py ``` Run the benchmark suite the way CI does: ```bash cd test && python3 run_suite.py --hw cuda --suite base-b-kernel-benchmark-test-1-gpu-large ``` --- ## Troubleshooting - **`No CI registry found in ...` from `run_suite.py`**: add a module-level `register_cuda_ci(...)` with literal `est_time`, `stage`, and `runner_config`; starred args and non-literal values break AST collection - **JIT compilation fails**: ensure the `.cuh` file is under `python/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_bench` uses CUDA-graph-based timing by default; set `use_cuda_graph=False` only if the kernel can't be captured. `graph_clone_args` defaults 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_JIT_BENCHMARK_DISABLE_LOG_BANDWIDTH` is not `1` and `disable_log_bandwidth` is not `True`. For in-place kernels (return `None`) the `memory_output="out"` default counts nothing — pass the written tensors via `memory_output=(...)` --- ## References - `docs/docs/developer_guide/development_jit_kernel_guide.mdx` - `test/run_suite.py` — suite names, discovery of `test/registered/`, execution entrypoint for CI - `python/sglang/test/ci/ci_register.py` — `register_cuda_ci` and AST registration rules - `python/sglang/kernels/jit/utils/compile.py` — `load_jit`, `make_cpp_args` - `python/sglang/kernels/jit/utils/common.py` — `cache_once`, `should_run_full_tests`, `get_ci_test_range` - `python/sglang/kernels/jit/include/sgl_kernel/tensor.h` — `TensorMatcher`, `SymbolicSize/DType/Device`, `is_type` - `python/sglang/kernels/jit/include/sgl_kernel/ffi.h` — `ffi::empty`, `ffi::empty_like`, `ffi::from_blob` - `python/sglang/kernels/jit/include/sgl_kernel/utils.cuh` — type aliases, `LaunchKernel`, `SGL_DEVICE` - `python/sglang/kernels/jit/include/sgl_kernel/vec.cuh` — `AlignedVector` - `python/sglang/kernels/jit/include/sgl_kernel/tile.cuh` — `tile::Memory` - `python/sglang/kernels/jit/include/sgl_kernel/bits.h` — compile-time bit helpers - `python/sglang/kernels/jit/include/sgl_kernel/type.cuh` — `DTypeTrait`, `packed_t`, `device::cast`, `device::unpack`, `ReductionTrait` - `python/sglang/kernels/jit/include/sgl_kernel/math.cuh` — `device::math::` - `python/sglang/kernels/jit/include/sgl_kernel/warp.cuh` — `warp::reduce` and `reduce_sum/max/min` wrappers - `python/sglang/kernels/jit/include/sgl_kernel/cta.cuh` — `cta::reduce_max` - `python/sglang/kernels/jit/include/sgl_kernel/atomic.cuh` — `atomic::max` - `python/sglang/kernels/jit/include/sgl_kernel/runtime.cuh` — occupancy / SM count helpers - `python/sglang/kernels/jit/csrc/add_constant.cuh` — minimal runnable reference - `python/sglang/kernels/jit/csrc/elementwise/rmsnorm.cuh` — real example using `TensorMatcher` + `LaunchKernel` + `tile::Memory` - `python/sglang/kernels/jit/csrc/elementwise/qknorm.cuh` — real example using `runtime::get_blocks_per_sm` + persistent kernel pattern - `python/sglang/kernels/jit/benchmark/marker.py` — `benchmark`, `parametrize`, `do_bench`, `BenchResult` - `python/sglang/kernels/jit/benchmark/utils.py` — `create_random` / `create_empty` / `get_benchmark_range` helpers and `DEFAULT_DTYPE` / `DEFAULT_DEVICE` - `test/registered/kernels/benchmark/layernorm/bench_qknorm.py` — real example: multi-axis `parametrize` (with `ci_vals`) + in-place `memory_output` - `test/registered/kernels/benchmark/kvcache/bench_store_cache.py` — real example: scoped `memory_args` / `memory_output` + selective `graph_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 ```