JIT 커널 개발 가이드
JIT 커널 개발 가이드
SGLang은 커널을 런타임에 컴파일하는 JIT(just-in-time) 방식을 지원해요. AOT로 미리 컴파일하는 sgl-kernel과 달리 정적 compile_commands.json을 만들 수 없어서, 코드 완성 같은 IDE 기능을 켜려면 별도 설정이 필요해요. 이 가이드가 그 환경 세팅부터 커널 추가까지 이어줄게요.
출처: JIT 커널 개발 가이드
환경 세팅
JIT 커널 개발에는 언어 서버로 clangd를 쓰는 걸 강력히 권장해요. Ubuntu/Debian에서는 apt.llvm.org에서 clangd를 받을 수 있어요. VS Code를 쓴다면 clangd 확장을 설치하면 IDE 통합이 더 좋아져요.
모든 JIT 관련 파일은 python/sglang/kernels/jit에 있어요. CUDA/C++ 바이너리를 미리(AOT) 컴파일하는 sgl-kernel과 달리, JIT 커널은 런타임에 컴파일돼요. 그래서 정적 compile_commands.json을 생성할 수 없어요. clangd에서 코드 완성이 되게 하려면 현재 디렉토리에 .clangd 구성 파일을 생성하도록 python -m sglang.kernels.jit를 실행하세요. 파일 생성 후 clangd 언어 서버를 재시작하면 모든 JIT 커널 파일을 인식할 거예요.
코드 구조
C++ 구현
C++ 소스 코드는 python/sglang/kernels/jit/csrc에 있어요. 재사용 가능한 함수는 python/sglang/kernels/jit/include에 두세요.
JIT C++는 namespace sglang 안에 있어요: include 블록 다음에 열고 파일 끝에서 닫으며, 디바이스 커널과 호스트 래퍼가 둘 다 그 안에 들어가요. 공유 host::·device:: 헬퍼도 그 안에 중첩돼 있어서 한정자 없이 해석되고 sglang:: 접두사가 필요 없어요.
외국어 바인딩에는 tvm-ffi를 써요. C++ 객체 내보내기 같은 고급 사용법은 문서를 참고하세요. 보통은 Python에서 PyTorch 텐서를 넘길 때 tvm::ffi::TensorView로 충분해요.
Python 인터페이스
Python 인터페이스는 python/sglang/kernels/jit에 정의돼 있어요. python/sglang/kernels/jit/utils/compile.py의 load_jit 유틸 함수가 컴파일된 모듈을 로드·반환해요. C++ 함수(예: cpp_func)를 내보내려면 load_jit에 cuda_wrappers=[("func", "cpp_func")]를 넘기세요. 그러면 Python에서 함수를 module.func로 호출할 수 있어요. load_jit는 namespace sglang 안에 export 래퍼를 만들어 주므로, cpp_func를 sglang:: 접두사 없이 작성하세요.
컴파일된 모듈 캐싱에는 functools.lru_cache보다 sglang.kernels.jit.utils.cache_once를 선호해요. functools.lru_cache는 torch.compile과 호환되지 않아요.
C++ 유틸리티
다음 C++ 유틸리티를 사용할 수 있어요.
정수 범위 (Integer Range)
PyTorch와 비슷하게 정수 범위를 나타내는 irange 함수를 제공해요.
#include <sgl_kernel/utils.h>
void test() {
for (auto i : host::irange(100)) { // [0, 100)
// do something
}
for (auto i : host::irange(0, 100)) { // [0, 100)
// do something
}
}
런타임 체크
CHECK_HOST는 선호되는 런타임 체크예요. 스트림 스타일이고, 체크가 통과하면 오버헤드가 0이에요 — 메시지 표현식은 실패할 때만 평가돼요. RuntimeCheck는 함수 스타일 대안인데, 체크가 통과해도 메시지 인자가 항상 평가된다는 점을 기억하세요. RuntimeDeviceCheck는 마지막 커널 실행의 상태를 검증하고, CHECK_CUDA는 추가 컨텍스트와 함께 cudaError_t를 검사하는 스트림 스타일 등가물이에요.
#include <sgl_kernel/utils.h>
#include <sgl_kernel/utils.cuh>
void test() {
CHECK_HOST(1 + 1 == 2) << 1 + 1 << " != " << 2; // preferred
host::RuntimeCheck(1 + 1 == 2, 1 + 1, " != ", 2);
host::RuntimeDeviceCheck();
// check the provided `cudaError_t`
host::RuntimeDeviceCheck(cudaGetLastError());
CHECK_CUDA(cudaGetLastError()) << "after my_kernel launch";
}
텐서 체크
TensorMatcher는 텐서 형태 정보를 검증하고 추출하는 읽기 쉬운 방법을 제공해요.
#include <sgl_kernel/tensor.h>
void test(const tvm::ffi::TensorView k_cache, const tvm::ffi::TensorView v_cache) {
using namespace host;
auto D = SymbolicSize{"D"}; // cache dimension
auto N = SymbolicSize{"N"}; // kvcache stride
auto dtype = SymbolicDType{};
auto device = SymbolicDevice{};
TensorMatcher({-1, D}) //
.with_strides({N, 1})
.with_dtype<int32_t, int64_t>(dtype)
.with_device<kDLCUDA, kDLCPU>(device)
.verify(k_cache)
.verify(v_cache);
}
검증 전에 TensorMatcher를 기대하는 stride·dtype·device 속성으로 구성하세요.
with_strides를 생략하면 텐서가 contiguous일 것으로 기대해요.with_dtype의 템플릿 인자는 허용 데이터 타입을 제한해요.with_device의 템플릿 인자는 허용 디바이스를 제한해요.with_xxx메서드에 전달된 값은 등가 검사를 강제해요.- size나 stride에
-1을 넘기면 어떤 값이든 매칭을 허용해요.
Symbolic 변수는 모든 검증에서 같은 값으로 해석돼야 해요. 검증 후 .unwrap()으로 매칭된 값을 가져올 수 있어요.
참고:
TensorMatcher는 임시 표현식이라 변수에 저장하면 안 돼요.
팁:
TensorMatcher체인 끝에//를 추가하면 들여쓰기가 제대로 유지돼요.
커널 실행
LaunchKernel::resolve_device는 PyTorch에서 현재 cudaStream을 가져와요. 커널은 LaunchKernel로 직접 실행할 수도 있어요.
#include <sgl_kernel/utils.cuh>
#include <dlpack/dlpack.h>
__global__ void kernel() {}
void test() {
const auto num_blocks = 1;
const auto num_threads = 32;
const auto dynamic_smem = 0;
DLDevice dev; // suppose this is initialized properly
host::LaunchKernel(num_blocks, num_threads, dev)(kernel);
cudaStream_t stream = host::LaunchKernel::resolve_device(dev);
host::LaunchKernel(num_blocks, num_threads, stream, dynamic_smem)(kernel);
}
새 커널 추가하기
이 절에서는 시스템에 새 JIT 커널을 추가하는 완전한 end-to-end 예시를 보여 줄게요. 입력 텐서의 모든 요소에 정수 상수를 더하는 간단한 add_constant 커널을 예로 들어요.
개념적으로 Python 인터페이스는 이렇게 생겼어요.
def add_constant(src: torch.Tensor, c: int):
return src + c
STEP 1: C++ 커널 작성하기
kernels/jit/csrc/elementwise/add_constant.cuh에 CUDA 커널을 작성하세요. 데모 목적으로 상수 값을 템플릿 파라미터로 넘겨요.
#include <sgl_kernel/tensor.h> // For TensorMatcher, SymbolicSize, SymbolicDevice
#include <sgl_kernel/utils.cuh> // For LaunchKernel
#include <sgl_kernel/utils.h> // For div_ceil, CHECK_HOST
#include <dlpack/dlpack.h>
#include <tvm/ffi/container/tensor.h>
#include <cstddef>
#include <cstdint>
namespace sglang {
template <int32_t kConstant>
__global__ void add_constant_kernel(int32_t* dst, const int32_t* src, size_t length) {
size_t idx = blockIdx.x * blockDim.x + threadIdx.x;
if (idx < length) {
dst[idx] = src[idx] + kConstant;
}
}
constexpr size_t kBlockSize = 256;
// You can also use struct with static method as an alternative
template <int32_t kConstant>
void add_constant(tvm::ffi::TensorView dst, tvm::ffi::TensorView src) {
using namespace host;
// 1. Validate input tensors
SymbolicSize N = {"num_elements"};
SymbolicDevice device_;
TensorMatcher({N}) // 1D tensor, must be contiguous
.with_dtype<int32_t>() // must be int32
.with_device<kDLCUDA>(device_) // must be on CUDA device
.verify(dst) // check tensor dst
.verify(src); // check tensor src
// 2. Extract required parameters, prepare for kernel launch
const size_t num_elements = N.unwrap();
const size_t grid_size = div_ceil(num_elements, kBlockSize);
const DLDevice device = device_.unwrap();
// some extra runtime checks using CHECK_HOST
CHECK_HOST(num_elements > 0) << "We only support non-empty tensors, got num_elements = " << num_elements;
// 3. Launch the kernel. Error code will be automatically checked.
LaunchKernel(grid_size, kBlockSize, device /*, dynamic_smem*/)(
// kernel function
add_constant_kernel<kConstant>,
// kernel arguments
static_cast<int32_t*>(dst.data_ptr()),
static_cast<int32_t*>(src.data_ptr()),
num_elements);
}
} // namespace sglang
STEP 2: Python 인터페이스 만들기
다음으로 Python 래퍼로 커널을 노출하세요. kernels/ops/elementwise/add_constant.py에 새 파일을 만들고 필요한 인터페이스를 노출하세요.
from __future__ import annotations
from typing import TYPE_CHECKING
import torch
from sglang.kernels.jit.utils import cache_once, load_jit, make_cpp_args
if TYPE_CHECKING:
from tvm_ffi.module import Module
@cache_once
def _jit_add_constant_module(constant: int) -> Module:
args = make_cpp_args(constant) # pass all the template argument
return load_jit(
"add_constant",
*args,
cuda_files=["elementwise/add_constant.cuh"],
cuda_wrappers=[("add_constant", f"add_constant<{args}>")],
)
def add_constant(src: torch.Tensor, constant: int) -> torch.Tensor:
if not src.is_cuda:
raise RuntimeError("src must be a CUDA tensor")
if src.dtype != torch.int32:
raise RuntimeError(f"Unsupported dtype {src.dtype}. Supported: int32")
dst = torch.empty_like(src)
module = _jit_add_constant_module(constant)
module.add_constant(dst, src)
return dst
Python 래퍼는 얇게 유지하되, 디스패치 전에 device·dtype 같은 기본 불변 조건은 검증하세요. 현재 JIT/FFI 경로에서 잘못된 텐서가 실행 전에 항상 안전하게 거부되지는 않아요.
STEP 3: 커널 사용하기
마지막으로, 커널을 일반 Python 함수처럼 import해서 사용하세요.
from sglang.kernels.ops.elementwise.add_constant import add_constant
완전하고 실행 가능한 예시는 test_add_constant.py를 참고하세요.
C++ Include 라이브러리 레퍼런스
JIT 커널 프레임워크는 python/sglang/kernels/jit/include/sgl_kernel/에 재사용 가능한 C++ 헤더 세트를 제공해요. 각 헤더는 가볍고 자립적이도록 설계됐어요. 아래는 각 헤더와 핵심 API 요약이에요.
코어 유틸리티
| 헤더 | 네임스페이스 | 용도 |
|---|---|---|
utils.h |
host |
호스트 측 필수 요소: RuntimeCheck, CHECK_HOST(cond) << ..., Panic, div_ceil, irange |
utils.cuh |
device / host |
타입 별칭(fp16_t, bf16_t, ...), SGL_DEVICE 매크로, PDL 헬퍼, LaunchKernel, RuntimeDeviceCheck, CHECK_CUDA(expr) << ... |
source_location.h |
(global) | 오류 보고용 이식 가능한 std::source_location 래퍼 |
runtime.cuh |
host::runtime |
CUDA 런타임 쿼리: get_blocks_per_sm, get_sm_count, get_cc_major, get_runtime_version, get_available_dynamic_smem_per_block |
텐서 검증
| 헤더 | 네임스페이스 | 용도 |
|---|---|---|
tensor.h |
host |
TensorMatcher, SymbolicSize, SymbolicDType, SymbolicDevice |
수학 & 타입 시스템
| 헤더 | 네임스페이스 | 용도 |
|---|---|---|
math.cuh |
device::math |
max, min, abs, sqrt, rsqrt, exp, sin, cos, 상수 |
type.cuh |
(global) / device |
DTypeTrait<T>, packed_t<T>, device::cast<To>(from) |
메모리 접근
| 헤더 | 네임스페이스 | 용도 |
|---|---|---|
vec.cuh |
device |
AlignedVector<T, N> - 벡터화된 load/store(최대 128-bit; 256-bit는 Blackwell GPU 필요) |
tile.cuh |
device::tile |
Memory<T> - 협조적 타일 메모리 I/O(thread/warp/CTA) |
병렬 프리미티브
| 헤더 | 네임스페이스 | 용도 |
|---|---|---|
warp.cuh |
device::warp |
reduce<Op, kNumThreads, kInner>(SUM/MAX/MIN, grouped 또는 inter-group) 및 __shfl_xor_sync를 통한 reduce_sum / reduce_max / reduce_min 래퍼 |
cta.cuh |
device::cta |
공유 메모리를 통한 warp 간 reduce_max |
atomic.cuh |
device::atomic |
max - 원자적 float max(CUDA + ROCm 폴백) |
재사용 가능 커널 템플릿
| 헤더 | 네임스페이스 | 용도 |
|---|---|---|
impl/norm.cuh |
host::norm / device::norm |
RMSNorm 빌딩 블록(warp & CTA 경로, StorageType) |