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.pyload_jit 유틸 함수가 컴파일된 모듈을 로드·반환해요. C++ 함수(예: cpp_func)를 내보내려면 load_jitcuda_wrappers=[("func", "cpp_func")]를 넘기세요. 그러면 Python에서 함수를 module.func로 호출할 수 있어요. load_jitnamespace sglang 안에 export 래퍼를 만들어 주므로, cpp_funcsglang:: 접두사 없이 작성하세요.

컴파일된 모듈 캐싱에는 functools.lru_cache보다 sglang.kernels.jit.utils.cache_once를 선호해요. functools.lru_cachetorch.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)

더 알아보기 (Learn more)