用 Codex 或 Claude 帮你安装 复制这段 Prompt,粘贴到 Codex、Claude 或其他助手里,让它检查 Skill 页面并帮你完成安装。
直接命令不会经过审查 Prompt;运行前请先检查来源。
npx skills add https://github.com/rand/ananke-sglang --skill add-jit-kernel命令会保持在同一行。复制前请横向滚动并检查完整内容。
想先保存到本地?可下载 SkillsMP 当前能够提供的文件。
基于 SOC 职业分类
正在显示 SKILL.md
| name | add-jit-kernel |
| description | Step-by-step tutorial for adding a new lightweight JIT CUDA kernel to sglang's jit_kernel module |
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.
Add a new operation that scales each element of a tensor by a scalar factor:
x (CUDA) and scalar factor (float, passed as C++ template argument)x * factor (element-wise), allocated internallytorch.float16), BF16 (torch.bfloat16), FP32 (torch.float32)sgl-kernel)jit_kernel): lightweight, few dependencies, rapid iteration, compiled on first usesgl-kernel): depends on CUTLASS / FlashInfer / DeepGEMM, needs pre-built wheelpython/sglang/jit_kernel/include/sgl_kernel/Always prefer these abstractions over raw CUDA primitives. They provide safety, readability, and consistency with the rest of the codebase.
utils.h — Host-side utilities#include <sgl_kernel/utils.h>
host::RuntimeCheck(cond, args...) — Assert a condition at runtime; throws PanicError with file/line info on failure. Prefer this over bare assert.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#include <sgl_kernel/utils.cuh>
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.device::kWarpThreads — Constant 32.device::load_as<T>(ptr, offset) / device::store_as<T>(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:
DLDevice via TVM-FFI automatically.operator()(kernel, args...)..enable_pdl(bool) for PDL (Programmatic Dependent Launch, SM90+).host::RuntimeDeviceCheck(cudaError_t) — Check a CUDA error; throw on failure.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, bind device to symbol.with_strides({strides...}) — validate strides (omit to require contiguous).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)Typical pattern:
auto N = SymbolicSize{"num_elements"};
auto device = SymbolicDevice{};
device.set_options<kDLCUDA>();
TensorMatcher({N}) //
.with_dtype<fp16_t>()
.with_device(device)
.verify(dst)
.verify(src); // same shape, dtype, device as dst
const size_t n = N.unwrap();
const DLDevice dev = device.unwrap();
type.cuh — dtype_trait<T> and packed_t<T>#include <sgl_kernel/type.cuh>
dtype_trait<T> — Static trait struct for each scalar type. Provides:
dtype_trait<T>::from(value) — convert from another type (e.g. fp32_t → fp16_t)dtype_trait<T>::abs/sqrt/rsqrt/max/min(x) — type-dispatched math (for fp32_t)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 using dtype_trait, e.g. cast<fp32x2_t, fp16x2_t>(v).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 128-bit vector loads/stores for bandwidth efficiency.
.load(ptr, offset) — vectorized load from ptr[offset].store(ptr, offset) — vectorized store to ptr[offset].fill(value) — fill all lanesoperator[](i) — element accesstile.cuh — tile::Memory (strided memory access pattern)#include <sgl_kernel/tile.cuh>
device::tile::Memory<T>::cta(blockDim.x) — Creates a tile accessor where each thread handles tid = threadIdx.x with stride blockDim.x. Common for loops over a 1D array..load(ptr, offset) — loads ptr[tid + offset * blockDim.x].store(ptr, val, offset) — stores to ptr[tid + offset * blockDim.x].in_bound(n, offset) — boundary checkmath.cuh — Device math (device::math::)#include <sgl_kernel/math.cuh>
device::math::max/min/abs/sqrt/rsqrt<T>(a, b) — type-dispatched math via dtype_traitdevice::math::exp/sin/cos(float) — fast float math wrapperswarp.cuh — Warp-level primitives#include <sgl_kernel/warp.cuh>
device::warp::reduce_sum<T>(value) — warp-level sum reduction via __shfl_xor_syncdevice::warp::reduce_max<T>(value) — warp-level max reductioncta.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 in smem[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 versionPersistent 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);
.clangd config for better IDE supportpython -m sglang.jit_kernel
jit_kernel/csrc/Create python/sglang/jit_kernel/csrc/elementwise/scale.cuh.
The implementation fully uses the project abstractions described above:
#include <sgl_kernel/tensor.h> // TensorMatcher, SymbolicSize, SymbolicDevice
#include <sgl_kernel/type.cuh> // dtype_trait, fp16_t, bf16_t, fp32_t
#include <sgl_kernel/utils.h> // RuntimeCheck, div_ceil
#include <sgl_kernel/utils.cuh> // LaunchKernel, SGL_DEVICE
#include <sgl_kernel/vec.cuh> // AlignedVector
#include <dlpack/dlpack.h>
#include <tvm/ffi/container/tensor.h>
namespace {
// ----------------------------------------------------------------
// Kernel: element-wise scale using vectorized 128-bit loads/stores
// T = fp16_t | bf16_t | fp32_t
// kVecN = number of elements per vector load (e.g. 8 for fp16)
// kFactor = scale factor encoded as kFactorNumer / kFactorDenom
// ----------------------------------------------------------------
template <typename T, int kVecN, int32_t kFactorNumer, int32_t kFactorDenom>
__global__ void scale_kernel(T* __restrict__ dst,
const T* __restrict__ src,
uint32_t n_vecs,
uint32_t n_remainder,
uint32_t n_total) {
kFactor = <>(kFactorNumer)
/ <>(kFactorDenom);
= device::AlignedVector<T, kVecN>;
vec_stride = blockDim.x * gridDim.x;
( vi = blockIdx.x * blockDim.x + threadIdx.x;
vi < n_vecs;
vi += vec_stride) {
v;
v.(src, vi);
( i = ; i < kVecN; ++i) {
v[i] = <T>(<>(v[i]) * kFactor);
}
v.(dst, vi);
}
base = n_vecs * kVecN;
scalar_stride = blockDim.x * gridDim.x;
( i = blockIdx.x * blockDim.x + threadIdx.x;
i < n_remainder;
i += scalar_stride) {
dst[base + i] = <T>(<>(src[base + i]) * kFactor);
}
}
< T, kFactorNumer, kFactorDenom>
{
host;
SymbolicSize N = {};
SymbolicDevice device_;
device_.<kDLCUDA>();
({N})
.<T>()
.(device_)
.(dst)
.(src);
n = <>(N.());
DLDevice device = device_.();
(n > , , n);
kVecN = / (T);
n_vecs = n / kVecN;
n_remainder = n % kVecN;
kBlockSize = ;
grid = (std::(n_vecs, n_remainder), kBlockSize);
(grid, kBlockSize, device)(
scale_kernel<T, kVecN, kFactorNumer, kFactorDenom>,
<T*>(dst.()),
< T*>(src.()),
n_vecs,
n_remainder,
n);
}
}
Key points:
sgl_kernel/ — not raw CUDA headers for anything already coveredTensorMatcher for all tensor validation; never manually check shape/dtype/deviceAlignedVector for vectorised 128-bit loads/stores — significant bandwidth winLaunchKernel — it resolves the stream and checks errors automaticallyRuntimeCheck for runtime assertions with useful error messagesfp16_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 conversionsdevice::math:: functions for device math instead of bare __ intrinsicsjit_kernel/Create python/sglang/jit_kernel/scale.py:
from __future__ import annotations
from typing import TYPE_CHECKING
import torch
from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args
if TYPE_CHECKING:
from tvm_ffi.module import Module
@cache_once
def _jit_scale_module(dtype: torch.dtype, factor_numer: int, factor_denom: int) -> Module:
"""Compile and cache the JIT scale module for a given dtype and factor."""
args = make_cpp_args(dtype, factor_numer, factor_denom)
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).
"""
assert src.is_cuda, "src must be a CUDA tensor"
assert src.dtype in (torch.float16, torch.bfloat16, torch.float32), (
f"Unsupported dtype {src.dtype}. Supported: float16, bfloat16, float32"
)
if out is None:
out = torch.empty_like(src)
:
out.shape == src.shape,
out.dtype == src.dtype,
factor_denom =
factor_numer = (factor * factor_denom)
module = _jit_scale_module(src.dtype, factor_numer, factor_denom)
module.scale(out, src)
out
Key points:
cache_once — not functools.lru_cache (incompatible with torch.compile)load_jit first arg(s) form the unique build marker; same marker = same cached binarycuda_wrappers: (export_name, kernel_symbol) — export_name is called from Pythonmake_cpp_args(dtype, ...) converts torch.dtype to C++ type alias:torch.dtype | C++ type |
|---|---|
torch.float16 | fp16_t |
torch.bfloat16 | bf16_t |
torch.float32 | fp32_t |
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:
if torch.cuda.get_device_capability()[0] < 9:
raise RuntimeError("This kernel requires SM90 (Hopper) or later")
Create python/sglang/jit_kernel/tests/test_scale.py:
import pytest
import torch
from sglang.jit_kernel.scale import scale
@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 ():
src = torch.randn(, dtype=torch.float16)
pytest.raises(AssertionError, =):
scale(src, )
():
src = torch.randint(, , (,), dtype=torch.int32, device=)
pytest.raises(AssertionError, =):
scale(src, )
__name__ == :
pytest.main([__file__, ])
Run:
pytest python/sglang/jit_kernel/tests/test_scale.py -q
Create python/sglang/jit_kernel/benchmark/bench_scale.py:
import itertools
import torch
import triton
import triton.testing
from sglang.jit_kernel.benchmark.utils import (
DEFAULT_DEVICE,
DEFAULT_DTYPE,
get_benchmark_range,
run_benchmark,
)
from sglang.jit_kernel.scale import scale as jit_scale
SIZE_LIST = get_benchmark_range(
full_range=[2**n for n in range(10, 20)], # 1K … 512K elements
ci_range=[4096, 65536],
)
configs = list(itertools.product(SIZE_LIST))
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["size"],
x_vals=configs,
line_arg="provider",
line_vals=["jit", "torch"],
line_names=["SGL JIT Kernel", "PyTorch"],
styles=[("blue", "-"), ("red", "--")],
ylabel="us",
plot_name="scale-performance",
args={},
)
)
def benchmark(size: int, provider: str):
src = torch.randn(size, dtype=DEFAULT_DTYPE, device=DEFAULT_DEVICE)
factor = 2.0
if provider == "jit":
fn = lambda: jit_scale(src, factor)
else:
fn = : src * factor
run_benchmark(fn)
__name__ == :
benchmark.run(print_data=)
Run:
python python/sglang/jit_kernel/benchmark/bench_scale.py
.cuh file is under python/sglang/jit_kernel/csrc/; reduce template argument combinationsCUDA_LAUNCH_BLOCKING=1; compute-sanitizer --tool memcheck python ...run_benchmark uses CUDA-graph-based timing by defaultdocs/developer_guide/development_jit_kernel_guide.mdpython/sglang/jit_kernel/utils.py — cache_once, load_jit, make_cpp_argspython/sglang/jit_kernel/include/sgl_kernel/tensor.h — TensorMatcher, SymbolicSize/DType/Devicepython/sglang/jit_kernel/include/sgl_kernel/utils.cuh — type aliases, LaunchKernel, SGL_DEVICEpython/sglang/jit_kernel/include/sgl_kernel/vec.cuh — AlignedVectorpython/sglang/jit_kernel/include/sgl_kernel/tile.cuh — tile::Memorypython/sglang/jit_kernel/include/sgl_kernel/type.cuh — dtype_trait, packed_t, device::castpython/sglang/jit_kernel/include/sgl_kernel/math.cuh — device::math::python/sglang/jit_kernel/include/sgl_kernel/warp.cuh — warp::reduce_sum/maxpython/sglang/jit_kernel/include/sgl_kernel/cta.cuh — cta::reduce_maxpython/sglang/jit_kernel/include/sgl_kernel/atomic.cuh — atomic::maxpython/sglang/jit_kernel/include/sgl_kernel/runtime.cuh — occupancy / SM count helperspython/sglang/jit_kernel/csrc/add_constant.cuh — minimal runnable referencepython/sglang/jit_kernel/csrc/elementwise/rmsnorm.cuh — real example using TensorMatcher + LaunchKernel + tile::Memorypython/sglang/jit_kernel/csrc/elementwise/qknorm.cuh — real example using runtime::get_blocks_per_sm + persistent kernel patternpython/sglang/jit_kernel/benchmark/utils.py — benchmark helperspython/sglang/jit_kernel/csrc/elementwise/scale.cuh # NEW: CUDA kernel
python/sglang/jit_kernel/scale.py # NEW: Python wrapper
python/sglang/jit_kernel/tests/test_scale.py # NEW: Tests
python/sglang/jit_kernel/benchmark/bench_scale.py # NEW: Benchmark