[Feature] Introduce JIT Kernel in sglang (with hicache JIT kernel) (#13453)
This commit is contained in:
19
python/sglang/jit_kernel/.clang-format
Normal file
19
python/sglang/jit_kernel/.clang-format
Normal file
@@ -0,0 +1,19 @@
|
||||
BasedOnStyle: Google
|
||||
IndentWidth: 2
|
||||
ColumnLimit: 120
|
||||
AllowShortFunctionsOnASingleLine: Empty
|
||||
DerivePointerAlignment: false
|
||||
PointerAlignment: Left
|
||||
NamespaceIndentation: None
|
||||
SortIncludes: true
|
||||
AllowShortLoopsOnASingleLine: false
|
||||
BinPackParameters: false # Prevents packing parameters in declarations
|
||||
BinPackArguments: false # Prevents packing arguments in function calls
|
||||
AlignAfterOpenBracket: AlwaysBreak # Forces a break after the opening parenthesis
|
||||
AlignOperands: Align # Aligns arguments vertically
|
||||
PenaltyBreakBeforeFirstCallParameter: 1 # Encourages breaking before the first argument
|
||||
PenaltyReturnTypeOnItsOwnLine: 100 # Keeps return type with function name
|
||||
|
||||
IncludeCategories:
|
||||
- Regex: '^<sgl_kernel/.*>$'
|
||||
Priority: 0
|
||||
264
python/sglang/jit_kernel/csrc/hicache.cuh
Normal file
264
python/sglang/jit_kernel/csrc/hicache.cuh
Normal file
@@ -0,0 +1,264 @@
|
||||
#include <sgl_kernel/tensor.h>
|
||||
#include <sgl_kernel/utils.cuh>
|
||||
#include <sgl_kernel/utils.h>
|
||||
#include <sgl_kernel/warp.cuh>
|
||||
|
||||
#include <dlpack/dlpack.h>
|
||||
|
||||
#include <algorithm>
|
||||
#include <concepts>
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <type_traits>
|
||||
|
||||
namespace {
|
||||
|
||||
struct HicacheKernelParams {
|
||||
void* __restrict__ k_cache_dst;
|
||||
void* __restrict__ v_cache_dst;
|
||||
const void* __restrict__ indices_dst;
|
||||
void* __restrict__ k_cache_src;
|
||||
void* __restrict__ v_cache_src;
|
||||
const void* __restrict__ indices_src;
|
||||
std::size_t length;
|
||||
std::size_t kv_cache_src_stride;
|
||||
std::size_t kv_cache_dst_stride;
|
||||
std::size_t num_layers = 0; // only used in all_layer transfer
|
||||
};
|
||||
|
||||
template <
|
||||
std::integral T,
|
||||
std::size_t kElementSize,
|
||||
std::size_t kUnroll,
|
||||
std::size_t kBlockQuota,
|
||||
std::size_t kNumThreads,
|
||||
std::size_t kMaxOccupancy>
|
||||
__global__ __launch_bounds__(kNumThreads, kMaxOccupancy) void hicache_transfer_per_layer(
|
||||
const __grid_constant__ HicacheKernelParams params) {
|
||||
// each warp acts as a worker
|
||||
using namespace device;
|
||||
static_assert(kNumThreads % kWarpThreads == 0);
|
||||
static_assert(kWarpThreads % kUnroll == 0);
|
||||
|
||||
constexpr auto kWarpThreads = device::kWarpThreads / kUnroll;
|
||||
constexpr auto kWarpsPerBlock = kNumThreads / kWarpThreads;
|
||||
constexpr auto kWorkers = kWarpsPerBlock * kBlockQuota;
|
||||
|
||||
const auto& [
|
||||
k_cache_dst, v_cache_dst, indices_dst, // dst
|
||||
k_cache_src, v_cache_src, indices_src, // src
|
||||
length, kv_cache_src_stride, kv_cache_dst_stride, _ // metadata
|
||||
] = params;
|
||||
const auto warp_id = blockIdx.x * kWarpsPerBlock + threadIdx.x / kWarpThreads;
|
||||
|
||||
// force to transfer 128 bytes per iteration
|
||||
// since the PCIe transaction size is 128 bytes aligned
|
||||
constexpr auto kGranularity = 128 / kWarpThreads;
|
||||
|
||||
for (auto i = warp_id; i < length; i += kWorkers) {
|
||||
const auto pos_src = static_cast<const T*>(indices_src)[i];
|
||||
const auto pos_dst = static_cast<const T*>(indices_dst)[i];
|
||||
const auto src_k = pointer::offset(k_cache_src, pos_src * kv_cache_src_stride);
|
||||
const auto dst_k = pointer::offset(k_cache_dst, pos_dst * kv_cache_dst_stride);
|
||||
const auto src_v = pointer::offset(v_cache_src, pos_src * kv_cache_src_stride);
|
||||
const auto dst_v = pointer::offset(v_cache_dst, pos_dst * kv_cache_dst_stride);
|
||||
const auto vec_k = warp::load_vec<kElementSize, kGranularity, kWarpThreads>(src_k);
|
||||
const auto vec_v = warp::load_vec<kElementSize, kGranularity, kWarpThreads>(src_v);
|
||||
warp::store_vec<kElementSize, kGranularity, kWarpThreads>(dst_k, vec_k);
|
||||
warp::store_vec<kElementSize, kGranularity, kWarpThreads>(dst_v, vec_v);
|
||||
}
|
||||
}
|
||||
|
||||
template <
|
||||
std::integral T,
|
||||
std::size_t kElementSize,
|
||||
std::size_t kUnroll,
|
||||
std::size_t kBlockQuota,
|
||||
std::size_t kNumThreads,
|
||||
std::size_t kMaxOccupancy>
|
||||
__global__ __launch_bounds__(kNumThreads, kMaxOccupancy) void hicache_transfer_all_layer(
|
||||
const __grid_constant__ HicacheKernelParams params) {
|
||||
// each warp acts as a worker
|
||||
using namespace device;
|
||||
using src_ptr_t = std::add_pointer_t<const void* const>;
|
||||
using dst_ptr_t = std::add_pointer_t<void* const>;
|
||||
|
||||
static_assert(kNumThreads % kWarpThreads == 0);
|
||||
constexpr auto kWarpThreads = device::kWarpThreads / kUnroll;
|
||||
constexpr auto kWarpsPerBlock = static_cast<uint32_t>(kNumThreads) / kWarpThreads;
|
||||
constexpr auto kWorkers = kWarpsPerBlock * kBlockQuota;
|
||||
|
||||
const auto& [
|
||||
k_ptr_dst, v_ptr_dst, indices_dst, // dst
|
||||
k_ptr_src, v_ptr_src, indices_src, // src
|
||||
length, kv_cache_src_stride, kv_cache_dst_stride, num_layers // metadata
|
||||
] = params;
|
||||
const auto warp_id = blockIdx.x * kWarpsPerBlock + threadIdx.x / kWarpThreads;
|
||||
|
||||
// force to transfer 128 bytes per iteration
|
||||
// since the PCIe transaction size is 128 bytes aligned
|
||||
constexpr auto kGranularity = 128 / kWarpThreads;
|
||||
|
||||
for (auto i = warp_id; i < length; i += kWorkers) {
|
||||
const auto pos_src = static_cast<const T*>(indices_src)[i];
|
||||
const auto pos_dst = static_cast<const T*>(indices_dst)[i];
|
||||
for (std::size_t layer = 0; layer < num_layers; ++layer) {
|
||||
const auto k_cache_src = static_cast<src_ptr_t>(k_ptr_src)[layer];
|
||||
const auto v_cache_src = static_cast<src_ptr_t>(v_ptr_src)[layer];
|
||||
const auto k_cache_dst = static_cast<dst_ptr_t>(k_ptr_dst)[layer];
|
||||
const auto v_cache_dst = static_cast<dst_ptr_t>(v_ptr_dst)[layer];
|
||||
const auto src_k = pointer::offset(k_cache_src, pos_src * kv_cache_src_stride);
|
||||
const auto dst_k = pointer::offset(k_cache_dst, pos_dst * kv_cache_dst_stride);
|
||||
const auto src_v = pointer::offset(v_cache_src, pos_src * kv_cache_src_stride);
|
||||
const auto dst_v = pointer::offset(v_cache_dst, pos_dst * kv_cache_dst_stride);
|
||||
const auto vec_k = warp::load_vec<kElementSize, kGranularity, kWarpThreads>(src_k);
|
||||
const auto vec_v = warp::load_vec<kElementSize, kGranularity, kWarpThreads>(src_v);
|
||||
warp::store_vec<kElementSize, kGranularity, kWarpThreads>(dst_k, vec_k);
|
||||
warp::store_vec<kElementSize, kGranularity, kWarpThreads>(dst_v, vec_v);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <
|
||||
std::size_t kElementSize,
|
||||
std::size_t kUnroll,
|
||||
std::size_t kBlockQuota,
|
||||
std::size_t kNumThreads,
|
||||
std::size_t kMaxOccupancy>
|
||||
struct HiCacheKernel {
|
||||
template <typename T>
|
||||
static constexpr auto _kernel_one =
|
||||
hicache_transfer_per_layer<T, kElementSize, kUnroll, kBlockQuota, kNumThreads, kMaxOccupancy>;
|
||||
template <typename T>
|
||||
static constexpr auto _kernel_all =
|
||||
hicache_transfer_all_layer<T, kElementSize, kUnroll, kBlockQuota, kNumThreads, kMaxOccupancy>;
|
||||
|
||||
static void run_one(
|
||||
const tvm::ffi::TensorView k_cache_dst,
|
||||
const tvm::ffi::TensorView v_cache_dst,
|
||||
const tvm::ffi::TensorView indices_dst,
|
||||
const tvm::ffi::TensorView k_cache_src,
|
||||
const tvm::ffi::TensorView v_cache_src,
|
||||
const tvm::ffi::TensorView indices_src) {
|
||||
using namespace host;
|
||||
|
||||
auto D = SymbolicSize{"D"}; // cache dimension
|
||||
auto N = SymbolicSize{"N"}; // src kv stride
|
||||
auto M = SymbolicSize{"M"}; // dst kv stride
|
||||
auto L = SymbolicSize{"L"}; // indices length
|
||||
auto cache_dtype = SymbolicDType{};
|
||||
auto indices_dtype = SymbolicDType{};
|
||||
auto indices_device = SymbolicDevice{};
|
||||
|
||||
TensorMatcher({-1, D}) //
|
||||
.with_strides({N, 1})
|
||||
.with_dtype(cache_dtype)
|
||||
.with_device<kDLCUDA, kDLCUDAHost, kDLCPU>()
|
||||
.verify(k_cache_src)
|
||||
.verify(v_cache_src);
|
||||
TensorMatcher({-1, D}) //
|
||||
.with_strides({M, 1})
|
||||
.with_dtype(cache_dtype)
|
||||
.with_device<kDLCUDA, kDLCUDAHost, kDLCPU>()
|
||||
.verify(k_cache_dst)
|
||||
.verify(v_cache_dst);
|
||||
TensorMatcher({L}) //
|
||||
.with_dtype<int32_t, int64_t>(indices_dtype)
|
||||
.with_device<kDLCUDA>(indices_device)
|
||||
.verify(indices_src)
|
||||
.verify(indices_dst);
|
||||
|
||||
// verify dimension match
|
||||
const auto dtype_size = dtype_bytes(cache_dtype.unwrap());
|
||||
const auto element_bytes = D.unwrap() * dtype_size;
|
||||
RuntimeCheck(kElementSize == element_bytes, "HicacheKernel: cache dimension mismatch.");
|
||||
|
||||
const auto k_cache_dst_ptr = k_cache_dst.data_ptr();
|
||||
const auto v_cache_dst_ptr = v_cache_dst.data_ptr();
|
||||
const auto k_cache_src_ptr = k_cache_src.data_ptr();
|
||||
const auto v_cache_src_ptr = v_cache_src.data_ptr();
|
||||
const auto indices_dst_ptr = indices_dst.data_ptr();
|
||||
const auto indices_src_ptr = indices_src.data_ptr();
|
||||
const auto length = static_cast<std::size_t>(L.unwrap());
|
||||
const auto kv_cache_src_stride = static_cast<std::size_t>(N.unwrap()) * dtype_size;
|
||||
const auto kv_cache_dst_stride = static_cast<std::size_t>(M.unwrap()) * dtype_size;
|
||||
const auto use_int32 = indices_dtype.unwrap().bits == 32;
|
||||
const auto device = indices_device.unwrap();
|
||||
|
||||
constexpr auto kWorkersPerBlock = kNumThreads / (device::kWarpThreads / kUnroll);
|
||||
const auto num_blocks = std::min(div_ceil(length, kWorkersPerBlock), kBlockQuota);
|
||||
const auto params = HicacheKernelParams{
|
||||
.k_cache_dst = k_cache_dst_ptr,
|
||||
.v_cache_dst = v_cache_dst_ptr,
|
||||
.indices_dst = indices_dst_ptr,
|
||||
.k_cache_src = k_cache_src_ptr,
|
||||
.v_cache_src = v_cache_src_ptr,
|
||||
.indices_src = indices_src_ptr,
|
||||
.length = length,
|
||||
.kv_cache_src_stride = kv_cache_src_stride,
|
||||
.kv_cache_dst_stride = kv_cache_dst_stride,
|
||||
};
|
||||
const auto kernel = use_int32 ? _kernel_one<int32_t> : _kernel_one<int64_t>;
|
||||
LaunchKernel(num_blocks, kNumThreads, device)(kernel, params);
|
||||
}
|
||||
|
||||
static void run_all(
|
||||
const tvm::ffi::TensorView k_ptr_dst,
|
||||
const tvm::ffi::TensorView v_ptr_dst,
|
||||
const tvm::ffi::TensorView indices_dst,
|
||||
const tvm::ffi::TensorView k_ptr_src,
|
||||
const tvm::ffi::TensorView v_ptr_src,
|
||||
const tvm::ffi::TensorView indices_src,
|
||||
const std::size_t kv_src_stride,
|
||||
const std::size_t kv_dst_stride) {
|
||||
using namespace host;
|
||||
|
||||
auto N = SymbolicSize{"N"}; // num layers
|
||||
auto L = SymbolicSize{"L"}; // indices length
|
||||
auto dtype_ = SymbolicDType{};
|
||||
auto device_ = SymbolicDevice{};
|
||||
|
||||
TensorMatcher({N}) //
|
||||
.with_dtype<uint64_t>()
|
||||
.with_device<kDLCUDA>(device_)
|
||||
.verify(k_ptr_src)
|
||||
.verify(v_ptr_src)
|
||||
.verify(k_ptr_dst)
|
||||
.verify(v_ptr_dst);
|
||||
TensorMatcher({L}) //
|
||||
.with_dtype<int32_t, int64_t>(dtype_)
|
||||
.with_device<kDLCUDA>(device_)
|
||||
.verify(indices_src)
|
||||
.verify(indices_dst);
|
||||
|
||||
// verify dimension match
|
||||
const auto k_cache_dst_ptr = k_ptr_dst.data_ptr();
|
||||
const auto v_cache_dst_ptr = v_ptr_dst.data_ptr();
|
||||
const auto k_cache_src_ptr = k_ptr_src.data_ptr();
|
||||
const auto v_cache_src_ptr = v_ptr_src.data_ptr();
|
||||
const auto indices_dst_ptr = indices_dst.data_ptr();
|
||||
const auto indices_src_ptr = indices_src.data_ptr();
|
||||
const auto length = static_cast<std::size_t>(L.unwrap());
|
||||
const auto use_int32 = dtype_.unwrap().bits == 32;
|
||||
const auto device = device_.unwrap();
|
||||
|
||||
constexpr auto kWorkersPerBlock = kNumThreads / (device::kWarpThreads / kUnroll);
|
||||
const auto num_blocks = std::min(div_ceil(length, kWorkersPerBlock), kBlockQuota);
|
||||
const auto params = HicacheKernelParams{
|
||||
.k_cache_dst = k_cache_dst_ptr,
|
||||
.v_cache_dst = v_cache_dst_ptr,
|
||||
.indices_dst = indices_dst_ptr,
|
||||
.k_cache_src = k_cache_src_ptr,
|
||||
.v_cache_src = v_cache_src_ptr,
|
||||
.indices_src = indices_src_ptr,
|
||||
.length = length,
|
||||
.kv_cache_src_stride = kv_src_stride,
|
||||
.kv_cache_dst_stride = kv_dst_stride,
|
||||
.num_layers = static_cast<std::size_t>(N.unwrap()),
|
||||
};
|
||||
const auto kernel = use_int32 ? _kernel_all<int32_t> : _kernel_all<int64_t>;
|
||||
LaunchKernel(num_blocks, kNumThreads, device)(kernel, params);
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace
|
||||
113
python/sglang/jit_kernel/hicache.py
Normal file
113
python/sglang/jit_kernel/hicache.py
Normal file
@@ -0,0 +1,113 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from functools import lru_cache
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from sglang.jit_kernel.utils import load_jit, make_cpp_args
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import torch
|
||||
from tvm_ffi.module import Module
|
||||
|
||||
DEFAULT_BLOCK_QUOTA = 2
|
||||
|
||||
|
||||
@lru_cache(maxsize=None)
|
||||
def _jit_hicache_module(*, element_size: int, unroll: int, block_quota: int) -> Module:
|
||||
num_threads, occupancy = 1024, 1
|
||||
args = make_cpp_args(
|
||||
element_size,
|
||||
unroll,
|
||||
block_quota,
|
||||
num_threads,
|
||||
occupancy,
|
||||
)
|
||||
return load_jit(
|
||||
"hicache",
|
||||
*args,
|
||||
cuda_files=["hicache.cuh"],
|
||||
cuda_wrappers=[("launch_one", f"HiCacheKernel<{args}>::run_one")],
|
||||
)
|
||||
|
||||
|
||||
def _default_unroll(element_size: int) -> int:
|
||||
if element_size <= 512:
|
||||
return 4
|
||||
|
||||
if element_size <= 1024:
|
||||
return 2
|
||||
|
||||
# fallback: no unroll
|
||||
return 1
|
||||
|
||||
|
||||
def transfer_hicache_one_layer(
|
||||
k_cache_dst: torch.Tensor,
|
||||
v_cache_dst: torch.Tensor,
|
||||
indices_dst: torch.Tensor,
|
||||
k_cache_src: torch.Tensor,
|
||||
v_cache_src: torch.Tensor,
|
||||
indices_src: torch.Tensor,
|
||||
*,
|
||||
element_dim: int | None = None,
|
||||
unroll: int | None = None, # can be tuned for performance
|
||||
block_quota: int | None = None, # can be tuned for less interference
|
||||
) -> None:
|
||||
element_dim = element_dim or k_cache_dst.size(-1)
|
||||
k_cache_src = k_cache_src.view(-1, element_dim)
|
||||
v_cache_src = v_cache_src.view(-1, element_dim)
|
||||
k_cache_dst = k_cache_dst.view(-1, element_dim)
|
||||
v_cache_dst = v_cache_dst.view(-1, element_dim)
|
||||
element_size = element_dim * k_cache_dst.element_size()
|
||||
block_quota = block_quota or DEFAULT_BLOCK_QUOTA
|
||||
unroll = unroll or _default_unroll(element_size)
|
||||
module = _jit_hicache_module(
|
||||
element_size=element_size,
|
||||
unroll=unroll,
|
||||
block_quota=block_quota,
|
||||
)
|
||||
module.launch_one(
|
||||
k_cache_dst,
|
||||
v_cache_dst,
|
||||
indices_dst,
|
||||
k_cache_src,
|
||||
v_cache_src,
|
||||
indices_src,
|
||||
)
|
||||
|
||||
|
||||
def transfer_hicache_all_layer(
|
||||
k_ptr_dst: torch.Tensor,
|
||||
v_ptr_dst: torch.Tensor,
|
||||
indices_dst: torch.Tensor,
|
||||
k_ptr_src: torch.Tensor,
|
||||
v_ptr_src: torch.Tensor,
|
||||
indices_src: torch.Tensor,
|
||||
kv_cache_src_stride_bytes: int,
|
||||
kv_cache_dst_stride_bytes: int,
|
||||
*,
|
||||
element_size: int | None = None,
|
||||
unroll: int | None = None, # can be tuned for performance
|
||||
block_quota: int | None = None, # can be tuned for less interference
|
||||
) -> None:
|
||||
if element_size is None: # assume both contiguous
|
||||
assert kv_cache_dst_stride_bytes == kv_cache_src_stride_bytes
|
||||
element_size = kv_cache_dst_stride_bytes
|
||||
|
||||
block_quota = block_quota or DEFAULT_BLOCK_QUOTA
|
||||
unroll = unroll or _default_unroll(element_size)
|
||||
module = _jit_hicache_module(
|
||||
element_size=element_size,
|
||||
unroll=unroll,
|
||||
block_quota=block_quota,
|
||||
)
|
||||
module.launch_all(
|
||||
k_ptr_dst,
|
||||
v_ptr_dst,
|
||||
indices_dst,
|
||||
k_ptr_src,
|
||||
v_ptr_src,
|
||||
indices_src,
|
||||
kv_cache_src_stride_bytes,
|
||||
kv_cache_dst_stride_bytes,
|
||||
)
|
||||
488
python/sglang/jit_kernel/include/sgl_kernel/tensor.h
Normal file
488
python/sglang/jit_kernel/include/sgl_kernel/tensor.h
Normal file
@@ -0,0 +1,488 @@
|
||||
#pragma once
|
||||
#include <sgl_kernel/utils.h>
|
||||
|
||||
#include <dlpack/dlpack.h>
|
||||
#include <tvm/ffi/container/tensor.h>
|
||||
#include <tvm/ffi/dtype.h>
|
||||
|
||||
#include <algorithm>
|
||||
#include <array>
|
||||
#include <concepts>
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <initializer_list>
|
||||
#include <optional>
|
||||
#include <ranges>
|
||||
#include <source_location>
|
||||
#include <span>
|
||||
#include <sstream>
|
||||
#include <string>
|
||||
#include <string_view>
|
||||
#include <type_traits>
|
||||
#include <utility>
|
||||
|
||||
namespace host {
|
||||
|
||||
namespace stdr = std::ranges;
|
||||
namespace stdv = std::views;
|
||||
|
||||
namespace details {
|
||||
|
||||
struct SizeRef;
|
||||
struct DTypeRef;
|
||||
struct DeviceRef;
|
||||
|
||||
template <typename T>
|
||||
struct dtype_trait {};
|
||||
|
||||
template <std::integral T>
|
||||
struct dtype_trait<T> {
|
||||
inline static constexpr auto value = DLDataType{
|
||||
.code = std::is_signed_v<T> ? DLDataTypeCode::kDLInt : DLDataTypeCode::kDLUInt,
|
||||
.bits = static_cast<std::uint8_t>(sizeof(T) * 8),
|
||||
.lanes = 1};
|
||||
};
|
||||
|
||||
template <std::floating_point T>
|
||||
struct dtype_trait<T> {
|
||||
inline static constexpr auto value =
|
||||
DLDataType{.code = DLDataTypeCode::kDLFloat, .bits = static_cast<std::uint8_t>(sizeof(T) * 8), .lanes = 1};
|
||||
};
|
||||
|
||||
inline constexpr auto kAnyDeviceID = -1;
|
||||
inline constexpr auto kAnySize = static_cast<int64_t>(-1);
|
||||
inline constexpr auto kNullSize = static_cast<int64_t>(0);
|
||||
inline constexpr auto kNullDType = static_cast<DLDataTypeCode>(18u);
|
||||
inline constexpr auto kNullDevice = static_cast<DLDeviceType>(-1);
|
||||
|
||||
template <typename... Ts>
|
||||
inline constexpr auto kDTypeList = std::array{dtype_trait<Ts>::value...};
|
||||
|
||||
template <auto... Codes>
|
||||
inline constexpr auto kDeviceList = std::array<DLDevice, sizeof...(Codes)>{
|
||||
DLDevice{.device_type = static_cast<DLDeviceType>(Codes), .device_id = kAnyDeviceID}...};
|
||||
|
||||
template <typename T>
|
||||
struct PrintAbleSpan {
|
||||
explicit PrintAbleSpan(std::span<const T> data) : data(data) {}
|
||||
std::span<const T> data;
|
||||
};
|
||||
|
||||
// define DLDataType comparison and printing in root namespace
|
||||
template <void* = nullptr>
|
||||
inline constexpr auto kDeviceStringMap = [] {
|
||||
constexpr auto map = std::array{
|
||||
std::pair{DLDeviceType::kDLCPU, "cpu"},
|
||||
std::pair{DLDeviceType::kDLCUDA, "cuda"},
|
||||
std::pair{DLDeviceType::kDLCUDAHost, "cuda_host"},
|
||||
std::pair{DLDeviceType::kDLOpenCL, "opencl"},
|
||||
std::pair{DLDeviceType::kDLVulkan, "vulkan"},
|
||||
std::pair{DLDeviceType::kDLMetal, "metal"},
|
||||
std::pair{DLDeviceType::kDLVPI, "vpi"},
|
||||
std::pair{DLDeviceType::kDLROCM, "rocm"},
|
||||
std::pair{DLDeviceType::kDLROCMHost, "rocm_host"},
|
||||
std::pair{DLDeviceType::kDLExtDev, "ext_dev"},
|
||||
std::pair{DLDeviceType::kDLCUDAManaged, "cuda_managed"},
|
||||
std::pair{DLDeviceType::kDLOneAPI, "oneapi"},
|
||||
std::pair{DLDeviceType::kDLWebGPU, "webgpu"},
|
||||
std::pair{DLDeviceType::kDLHexagon, "hexagon"},
|
||||
std::pair{DLDeviceType::kDLMAIA, "maia"},
|
||||
std::pair{DLDeviceType::kDLTrn, "trn"},
|
||||
};
|
||||
constexpr auto max_type = stdr::max(map | stdv::keys);
|
||||
auto result = std::array<std::string_view, max_type + 1>{};
|
||||
for (const auto& [code, name] : map) {
|
||||
result[static_cast<std::size_t>(code)] = name;
|
||||
}
|
||||
return result;
|
||||
}();
|
||||
|
||||
struct PrintableDevice {
|
||||
DLDevice device;
|
||||
};
|
||||
|
||||
inline auto& operator<<(std::ostream& os, DLDevice device) {
|
||||
const auto& mapping = kDeviceStringMap<>;
|
||||
const auto entry = static_cast<std::size_t>(device.device_type);
|
||||
host::RuntimeCheck(entry < mapping.size());
|
||||
const auto name = mapping[entry];
|
||||
host::RuntimeCheck(!name.empty(), "Unknown device: ", int(device.device_type));
|
||||
os << name;
|
||||
if (device.device_id != kAnyDeviceID) os << "[" << device.device_id << "]";
|
||||
return os;
|
||||
}
|
||||
|
||||
inline auto& operator<<(std::ostream& os, PrintableDevice pd) {
|
||||
return os << pd.device;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
inline auto& operator<<(std::ostream& os, PrintAbleSpan<T> span) {
|
||||
os << "[";
|
||||
for (const auto i : stdv::iota(std::size_t{0}, span.data.size())) {
|
||||
if (i > 0) {
|
||||
os << ", ";
|
||||
}
|
||||
os << span.data[i];
|
||||
}
|
||||
os << "]";
|
||||
return os;
|
||||
}
|
||||
|
||||
} // namespace details
|
||||
|
||||
struct SymbolicSize {
|
||||
public:
|
||||
SymbolicSize(std::string_view annotation = {}) : m_value(details::kNullSize), m_annotation(annotation) {}
|
||||
|
||||
auto get_name() const -> std::string_view {
|
||||
return m_annotation;
|
||||
}
|
||||
auto set_value(int64_t value) -> void {
|
||||
host::RuntimeCheck(!this->has_value(), "Size value already set");
|
||||
m_value = value;
|
||||
}
|
||||
auto has_value() const -> bool {
|
||||
return m_value != 0;
|
||||
}
|
||||
auto get_value() const -> std::optional<int64_t> {
|
||||
return this->has_value() ? std::optional{m_value} : std::nullopt;
|
||||
}
|
||||
auto unwrap() const -> int64_t {
|
||||
host::RuntimeCheck(this->has_value(), "Size value is not set");
|
||||
return m_value;
|
||||
}
|
||||
|
||||
SymbolicSize(const SymbolicSize&) = delete;
|
||||
SymbolicSize& operator=(const SymbolicSize&) = delete;
|
||||
|
||||
auto verify(int64_t dim) -> void {
|
||||
if (this->has_value()) {
|
||||
host::RuntimeCheck(m_value == dim, "Size mismatch: expected ", m_value, " but got ", dim);
|
||||
} else {
|
||||
this->set_value(dim);
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
std::int64_t m_value;
|
||||
std::string_view m_annotation;
|
||||
};
|
||||
|
||||
inline auto operator==(DLDevice lhs, DLDevice rhs) -> bool {
|
||||
return lhs.device_type == rhs.device_type && lhs.device_id == rhs.device_id;
|
||||
}
|
||||
|
||||
struct SymbolicDType {
|
||||
public:
|
||||
SymbolicDType() : m_value({details::kNullDType, 0, 0}) {}
|
||||
|
||||
auto set_value(DLDataType value) -> void {
|
||||
host::RuntimeCheck(!this->has_value(), "Dtype value already set");
|
||||
host::RuntimeCheck(
|
||||
m_check(value), "Dtype value [", value, "] not in the allowed options: ", details::PrintAbleSpan{m_options});
|
||||
m_value = value;
|
||||
}
|
||||
auto has_value() const -> bool {
|
||||
return m_value.code != details::kNullDType;
|
||||
}
|
||||
auto get_value() const -> std::optional<DLDataType> {
|
||||
return this->has_value() ? std::optional{m_value} : std::nullopt;
|
||||
}
|
||||
auto unwrap() const -> DLDataType {
|
||||
host::RuntimeCheck(this->has_value(), "Dtype value is not set");
|
||||
return m_value;
|
||||
}
|
||||
|
||||
auto set_options(std::span<const DLDataType> options) -> void {
|
||||
m_options = options;
|
||||
}
|
||||
template <typename... Ts>
|
||||
auto set_options() -> void {
|
||||
m_options = details::kDTypeList<Ts...>;
|
||||
}
|
||||
|
||||
auto verify(DLDataType dtype) -> void {
|
||||
if (this->has_value()) {
|
||||
host::RuntimeCheck(m_value == dtype, "DType mismatch: expected ", m_value, " but got ", dtype);
|
||||
} else {
|
||||
this->set_value(dtype);
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
auto m_check(DLDataType value) const -> bool {
|
||||
return stdr::empty(m_options) || (stdr::find(m_options, value) != stdr::end(m_options));
|
||||
}
|
||||
|
||||
std::span<const DLDataType> m_options;
|
||||
DLDataType m_value;
|
||||
};
|
||||
|
||||
struct SymbolicDevice {
|
||||
public:
|
||||
SymbolicDevice() : m_value({details::kNullDevice, details::kAnyDeviceID}) {}
|
||||
|
||||
auto set_value(DLDevice value) -> void {
|
||||
host::RuntimeCheck(!this->has_value(), "Device value already set");
|
||||
host::RuntimeCheck(
|
||||
m_check(value),
|
||||
"Device value [",
|
||||
details::PrintableDevice{value},
|
||||
"] not in the allowed options: ",
|
||||
details::PrintAbleSpan{m_options});
|
||||
m_value = value;
|
||||
}
|
||||
auto has_value() const -> bool {
|
||||
return m_value.device_type != details::kNullDevice;
|
||||
}
|
||||
auto get_value() const -> std::optional<DLDevice> {
|
||||
return this->has_value() ? std::optional{m_value} : std::nullopt;
|
||||
}
|
||||
auto unwrap() const -> DLDevice {
|
||||
host::RuntimeCheck(this->has_value(), "Device value is not set");
|
||||
return m_value;
|
||||
}
|
||||
|
||||
auto set_options(std::span<const DLDevice> options) -> void {
|
||||
m_options = options;
|
||||
}
|
||||
template <DLDeviceType... Codes>
|
||||
auto set_options() -> void {
|
||||
m_options = details::kDeviceList<Codes...>;
|
||||
}
|
||||
|
||||
auto verify(DLDevice device) -> void {
|
||||
if (this->has_value()) {
|
||||
host::RuntimeCheck(
|
||||
m_value == device,
|
||||
"Device mismatch: expected ",
|
||||
details::PrintableDevice{m_value},
|
||||
" but got ",
|
||||
details::PrintableDevice{device});
|
||||
} else {
|
||||
this->set_value(device);
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
auto m_check(DLDevice value) const -> bool {
|
||||
return stdr::empty(m_options) || (stdr::any_of(m_options, [value](const DLDevice& opt) {
|
||||
// device type must exactly match
|
||||
if (opt.device_type != value.device_type) return false;
|
||||
// device id can be wildcarded
|
||||
return opt.device_id == details::kAnyDeviceID || opt.device_id == value.device_id;
|
||||
}));
|
||||
}
|
||||
|
||||
std::span<const DLDevice> m_options;
|
||||
DLDevice m_value;
|
||||
};
|
||||
|
||||
namespace details {
|
||||
|
||||
template <typename T>
|
||||
struct BaseRef {
|
||||
public:
|
||||
BaseRef(const BaseRef&) = delete;
|
||||
BaseRef& operator=(const BaseRef&) = delete;
|
||||
|
||||
auto operator->() const -> T* {
|
||||
return m_ref;
|
||||
}
|
||||
auto operator*() const -> T& {
|
||||
return *m_ref;
|
||||
}
|
||||
auto rebind(T& other) -> void {
|
||||
m_ref = &other;
|
||||
}
|
||||
|
||||
explicit BaseRef() : m_ref(&m_cache), m_cache() {}
|
||||
BaseRef(T& size) : m_ref(&size), m_cache() {}
|
||||
|
||||
private:
|
||||
T* m_ref;
|
||||
T m_cache;
|
||||
};
|
||||
|
||||
struct SizeRef : BaseRef<SymbolicSize> {
|
||||
using BaseRef::BaseRef;
|
||||
SizeRef(int64_t value) {
|
||||
if (value != kAnySize) {
|
||||
(**this).set_value(value);
|
||||
} else {
|
||||
// otherwise, we can match any size
|
||||
}
|
||||
}
|
||||
|
||||
auto value_or_name(std::size_t dim) const -> std::string {
|
||||
if (const auto value = (**this).get_value()) {
|
||||
return std::to_string(*value);
|
||||
} else {
|
||||
const auto annotation = (**this).get_name();
|
||||
if (annotation.empty()) {
|
||||
return "dim#" + std::to_string(dim);
|
||||
} else {
|
||||
return static_cast<std::string>(annotation);
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
struct DTypeRef : BaseRef<SymbolicDType> {
|
||||
using BaseRef::BaseRef;
|
||||
DTypeRef(DLDataType options) {
|
||||
(**this).set_value(options);
|
||||
}
|
||||
DTypeRef(std::initializer_list<DLDataType> options) {
|
||||
(**this).set_options(options);
|
||||
}
|
||||
DTypeRef(std::span<const DLDataType> options) {
|
||||
(**this).set_options(options);
|
||||
}
|
||||
};
|
||||
|
||||
struct DeviceRef : BaseRef<SymbolicDevice> {
|
||||
using BaseRef::BaseRef;
|
||||
DeviceRef(DLDevice options) {
|
||||
(**this).set_value(options);
|
||||
}
|
||||
DeviceRef(std::initializer_list<DLDevice> options) {
|
||||
(**this).set_options(options);
|
||||
}
|
||||
DeviceRef(std::span<const DLDevice> options) {
|
||||
(**this).set_options(options);
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace details
|
||||
|
||||
struct TensorMatcher {
|
||||
private:
|
||||
using SizeRef = details::SizeRef;
|
||||
using DTypeRef = details::DTypeRef;
|
||||
using DeviceRef = details::DeviceRef;
|
||||
using Loc_t = std::source_location;
|
||||
|
||||
public:
|
||||
TensorMatcher(const TensorMatcher&) = delete;
|
||||
TensorMatcher& operator=(const TensorMatcher&) = delete;
|
||||
|
||||
explicit TensorMatcher(std::initializer_list<SizeRef> shape) : m_shape(shape), m_strides(), m_dtype() {}
|
||||
|
||||
auto with_strides(std::initializer_list<SizeRef> strides) && -> TensorMatcher&& {
|
||||
// no partial update allowed
|
||||
host::RuntimeCheck(m_strides.size() == 0, "Strides already specified");
|
||||
host::RuntimeCheck(m_shape.size() == strides.size(), "Strides size must match shape size");
|
||||
m_strides = strides;
|
||||
return std::move(*this);
|
||||
}
|
||||
|
||||
template <typename... Ts>
|
||||
auto with_dtype(DTypeRef&& dtype) && -> TensorMatcher&& {
|
||||
m_init_dtype();
|
||||
m_dtype.rebind(*dtype);
|
||||
return std::move(*this);
|
||||
}
|
||||
|
||||
template <typename... Ts>
|
||||
auto with_dtype() && -> TensorMatcher&& {
|
||||
static_assert(sizeof...(Ts) > 0, "At least one dtype option must be specified");
|
||||
m_init_dtype();
|
||||
m_dtype->set_options<Ts...>();
|
||||
return std::move(*this);
|
||||
}
|
||||
|
||||
template <DLDeviceType... Codes>
|
||||
auto with_device(DeviceRef&& device) && -> TensorMatcher&& {
|
||||
m_init_device();
|
||||
m_device.rebind(*device);
|
||||
return std::move(*this);
|
||||
}
|
||||
|
||||
template <DLDeviceType... Codes>
|
||||
auto with_device() && -> TensorMatcher&& {
|
||||
static_assert(sizeof...(Codes) > 0, "At least one device option must be specified");
|
||||
m_init_device();
|
||||
m_device->set_options<Codes...>();
|
||||
return std::move(*this);
|
||||
}
|
||||
|
||||
// once we start verification, we cannot modify anymore
|
||||
auto verify(tvm::ffi::TensorView view, Loc_t loc = Loc_t::current()) const&& -> const TensorMatcher&& {
|
||||
try {
|
||||
this->m_verify_impl(view);
|
||||
} catch (PanicError& e) {
|
||||
auto oss = std::ostringstream{};
|
||||
oss << "Tensor match failed for " << this->debug_str() << " at " << loc.file_name() << ":" << loc.line()
|
||||
<< "\n- Root cause: " << e.detail();
|
||||
throw PanicError(std::move(oss).str());
|
||||
}
|
||||
return std::move(*this);
|
||||
}
|
||||
|
||||
auto debug_str() const -> std::string {
|
||||
auto oss = std::ostringstream{};
|
||||
oss << "Tensor<";
|
||||
std::size_t dim = 0;
|
||||
for (const auto& size_ref : m_shape) {
|
||||
if (dim > 0) {
|
||||
oss << ", ";
|
||||
}
|
||||
oss << size_ref.value_or_name(dim++);
|
||||
}
|
||||
oss << ">";
|
||||
if (m_strides.size() > 0) {
|
||||
oss << " [strides=<";
|
||||
dim = 0;
|
||||
for (const auto& stride_ref : m_strides) {
|
||||
if (dim > 0) {
|
||||
oss << ", ";
|
||||
}
|
||||
oss << stride_ref.value_or_name(dim++);
|
||||
}
|
||||
oss << ">]";
|
||||
}
|
||||
return std::move(oss).str();
|
||||
}
|
||||
|
||||
private:
|
||||
auto m_verify_impl(tvm::ffi::TensorView view) const -> void {
|
||||
const auto dim = static_cast<std::size_t>(view.dim());
|
||||
host::RuntimeCheck(dim == m_shape.size(), "Tensor dimension mismatch: expected ", m_shape.size(), " but got ", dim);
|
||||
for (const auto i : stdv::iota(std::size_t{0}, dim)) {
|
||||
m_shape[i]->verify(view.size(i));
|
||||
}
|
||||
if (this->m_has_strides()) {
|
||||
for (const auto i : stdv::iota(std::size_t{0}, dim)) {
|
||||
m_strides[i]->verify(view.stride(i));
|
||||
}
|
||||
} else {
|
||||
host::RuntimeCheck(view.is_contiguous(), "Tensor is not contiguous as expected");
|
||||
}
|
||||
// since we may double verify, we will force to check
|
||||
m_dtype->verify(view.dtype());
|
||||
m_device->verify(view.device());
|
||||
}
|
||||
|
||||
auto m_init_dtype() -> void {
|
||||
host::RuntimeCheck(!m_has_dtype, "DType already specified");
|
||||
m_has_dtype = true;
|
||||
}
|
||||
auto m_init_device() -> void {
|
||||
host::RuntimeCheck(!m_has_device, "Device already specified");
|
||||
m_has_device = true;
|
||||
}
|
||||
auto m_has_strides() const -> bool {
|
||||
return !m_strides.empty();
|
||||
}
|
||||
|
||||
std::span<const SizeRef> m_shape;
|
||||
std::span<const SizeRef> m_strides;
|
||||
DTypeRef m_dtype;
|
||||
DeviceRef m_device;
|
||||
bool m_has_dtype = false;
|
||||
bool m_has_device = false;
|
||||
};
|
||||
|
||||
} // namespace host
|
||||
101
python/sglang/jit_kernel/include/sgl_kernel/utils.cuh
Normal file
101
python/sglang/jit_kernel/include/sgl_kernel/utils.cuh
Normal file
@@ -0,0 +1,101 @@
|
||||
#pragma once
|
||||
|
||||
#include <sgl_kernel/utils.h>
|
||||
|
||||
#include <dlpack/dlpack.h>
|
||||
#include <tvm/ffi/extra/c_env_api.h>
|
||||
|
||||
#include <concepts>
|
||||
#include <cstddef>
|
||||
#include <source_location>
|
||||
#include <type_traits>
|
||||
|
||||
namespace device {
|
||||
|
||||
inline constexpr auto kWarpThreads = 32u;
|
||||
|
||||
namespace pointer {
|
||||
|
||||
// we only allow void * pointer arithmetic for safety
|
||||
|
||||
template <typename T, std::integral... U>
|
||||
__always_inline __device__ auto offset(T* ptr, U... offset) -> void* {
|
||||
static_assert(std::is_same_v<T, void>, "Pointer arithmetic is only allowed for void* pointers");
|
||||
return static_cast<char*>(ptr) + (... + offset);
|
||||
}
|
||||
|
||||
template <typename T, std::integral... U>
|
||||
__always_inline __device__ auto offset(const T* ptr, U... offset) -> const void* {
|
||||
static_assert(std::is_same_v<T, void>, "Pointer arithmetic is only allowed for void* pointers");
|
||||
return static_cast<const char*>(ptr) + (... + offset);
|
||||
}
|
||||
|
||||
} // namespace pointer
|
||||
|
||||
} // namespace device
|
||||
|
||||
namespace host {
|
||||
|
||||
inline auto
|
||||
RuntimeDeviceCheck(::cudaError_t error, std::source_location location = std::source_location::current()) -> void {
|
||||
if (error != ::cudaSuccess) {
|
||||
[[unlikely]];
|
||||
::host::panic(location, "CUDA error: ", ::cudaGetErrorString(error));
|
||||
}
|
||||
}
|
||||
|
||||
inline auto RuntimeCudaCheck(std::source_location location = std::source_location::current()) -> void {
|
||||
return RuntimeDeviceCheck(::cudaGetLastError(), location);
|
||||
}
|
||||
|
||||
template <auto F>
|
||||
inline void set_smem_once(std::size_t smem_size) {
|
||||
static const auto last_smem_size = [&] {
|
||||
RuntimeDeviceCheck(::cudaFuncSetAttribute(F, ::cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size));
|
||||
return smem_size;
|
||||
}();
|
||||
RuntimeCheck(
|
||||
smem_size <= last_smem_size,
|
||||
"Dynamic shared memory size exceeds the previously set maximum size: ",
|
||||
last_smem_size,
|
||||
" bytes");
|
||||
}
|
||||
|
||||
struct LaunchKernel {
|
||||
public:
|
||||
explicit LaunchKernel(
|
||||
dim3 grid_dim, dim3 block_dim, DLDevice device, std::size_t dynamic_shared_mem_bytes = 0) noexcept
|
||||
: m_config(s_make_config(grid_dim, block_dim, resolve_device(device), dynamic_shared_mem_bytes)) {}
|
||||
|
||||
explicit LaunchKernel(
|
||||
dim3 grid_dim, dim3 block_dim, cudaStream_t stream, std::size_t dynamic_shared_mem_bytes = 0) noexcept
|
||||
: m_config(s_make_config(grid_dim, block_dim, stream, dynamic_shared_mem_bytes)) {}
|
||||
|
||||
static auto resolve_device(DLDevice device) -> cudaStream_t {
|
||||
return static_cast<cudaStream_t>(::TVMFFIEnvGetStream(device.device_type, device.device_id));
|
||||
}
|
||||
|
||||
LaunchKernel(const LaunchKernel&) = delete;
|
||||
LaunchKernel& operator=(const LaunchKernel&) = delete;
|
||||
|
||||
template <typename T, typename... Args>
|
||||
auto operator()(T&& kernel, Args&&... args) const -> void {
|
||||
host::RuntimeDeviceCheck(::cudaLaunchKernelEx(&m_config, kernel, std::forward<Args>(args)...));
|
||||
}
|
||||
|
||||
private:
|
||||
static auto
|
||||
s_make_config(dim3 grid_dim, dim3 block_dim, cudaStream_t stream, std::size_t smem) -> cudaLaunchConfig_t {
|
||||
auto config = ::cudaLaunchConfig_t{};
|
||||
config.gridDim = grid_dim;
|
||||
config.blockDim = block_dim;
|
||||
config.dynamicSmemBytes = smem;
|
||||
config.stream = stream;
|
||||
config.numAttrs = 0;
|
||||
return config;
|
||||
}
|
||||
cudaLaunchConfig_t m_config;
|
||||
/// TODO: We can add a queue to store the attributes if needed in the future.
|
||||
};
|
||||
|
||||
} // namespace host
|
||||
88
python/sglang/jit_kernel/include/sgl_kernel/utils.h
Normal file
88
python/sglang/jit_kernel/include/sgl_kernel/utils.h
Normal file
@@ -0,0 +1,88 @@
|
||||
#pragma once
|
||||
|
||||
#include <dlpack/dlpack.h>
|
||||
|
||||
#include <concepts>
|
||||
#include <ostream>
|
||||
#include <source_location>
|
||||
#include <sstream>
|
||||
#include <utility>
|
||||
|
||||
namespace host {
|
||||
|
||||
struct PanicError : public std::runtime_error {
|
||||
public:
|
||||
// copy and move constructors
|
||||
explicit PanicError(std::string msg) : runtime_error(msg), m_message(std::move(msg)) {}
|
||||
auto detail() const -> std::string_view {
|
||||
const auto sv = std::string_view{m_message};
|
||||
const auto pos = sv.find(": ");
|
||||
return pos == std::string_view::npos ? sv : sv.substr(pos + 2);
|
||||
}
|
||||
|
||||
private:
|
||||
std::string m_message;
|
||||
};
|
||||
|
||||
template <typename... Args>
|
||||
[[noreturn]]
|
||||
inline auto panic(std::source_location location, Args&&... args) -> void {
|
||||
std::ostringstream os;
|
||||
os << "Runtime check failed at " << location.file_name() << ":" << location.line();
|
||||
if constexpr (sizeof...(args) > 0) {
|
||||
os << ": ";
|
||||
(os << ... << std::forward<Args>(args));
|
||||
} else {
|
||||
os << " in " << location.function_name();
|
||||
}
|
||||
throw PanicError(std::move(os).str());
|
||||
}
|
||||
|
||||
template <typename... Args>
|
||||
struct RuntimeCheck {
|
||||
using Loc_t = std::source_location;
|
||||
template <typename Cond>
|
||||
explicit RuntimeCheck(Cond&& condition, Args&&... args, Loc_t location = Loc_t::current()) {
|
||||
if (!condition) {
|
||||
[[unlikely]];
|
||||
::host::panic(location, std::forward<Args>(args)...);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
template <typename Cond, typename... Args>
|
||||
explicit RuntimeCheck(Cond&&, Args&&...) -> RuntimeCheck<Args...>;
|
||||
|
||||
template <std::signed_integral T, std::signed_integral U>
|
||||
inline constexpr auto div_ceil(T a, U b) {
|
||||
return (a + b - 1) / b;
|
||||
}
|
||||
|
||||
template <std::unsigned_integral T, std::unsigned_integral U>
|
||||
inline constexpr auto div_ceil(T a, U b) {
|
||||
return (a + b - 1) / b;
|
||||
}
|
||||
|
||||
inline auto dtype_bytes(DLDataType dtype) -> std::size_t {
|
||||
return static_cast<std::size_t>(dtype.bits / 8);
|
||||
}
|
||||
|
||||
namespace pointer {
|
||||
|
||||
// we only allow void * pointer arithmetic for safety
|
||||
|
||||
template <typename T, std::integral... U>
|
||||
inline auto offset(T* ptr, U... offset) -> void* {
|
||||
static_assert(std::is_same_v<T, void>, "Pointer arithmetic is only allowed for void* pointers");
|
||||
return static_cast<char*>(ptr) + (... + offset);
|
||||
}
|
||||
|
||||
template <typename T, std::integral... U>
|
||||
inline auto offset(const T* ptr, U... offset) -> const void* {
|
||||
static_assert(std::is_same_v<T, void>, "Pointer arithmetic is only allowed for void* pointers");
|
||||
return static_cast<const char*>(ptr) + (... + offset);
|
||||
}
|
||||
|
||||
} // namespace pointer
|
||||
|
||||
} // namespace host
|
||||
145
python/sglang/jit_kernel/include/sgl_kernel/warp.cuh
Normal file
145
python/sglang/jit_kernel/include/sgl_kernel/warp.cuh
Normal file
@@ -0,0 +1,145 @@
|
||||
#pragma once
|
||||
#include <sgl_kernel/utils.cuh>
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <type_traits>
|
||||
|
||||
namespace device::warp {
|
||||
|
||||
namespace details {
|
||||
|
||||
template <std::size_t kUnit>
|
||||
inline constexpr auto get_mem_package() {
|
||||
if constexpr (kUnit == 16) {
|
||||
return uint4{};
|
||||
} else if constexpr (kUnit == 8) {
|
||||
return uint2{};
|
||||
} else if constexpr (kUnit == 4) {
|
||||
return uint1{};
|
||||
} else {
|
||||
static_assert(kUnit == 16 || kUnit == 8 || kUnit == 4, "Unsupported memory package size");
|
||||
}
|
||||
}
|
||||
|
||||
inline constexpr auto default_unit_size(std::size_t x) -> std::size_t {
|
||||
if (x % (16 * kWarpThreads) == 0) return 16;
|
||||
if (x % (8 * kWarpThreads) == 0) return 8;
|
||||
if (x % (4 * kWarpThreads) == 0) return 4;
|
||||
return 0; // trigger static assert in _get_mem_package
|
||||
}
|
||||
|
||||
template <std::size_t kBytes, std::size_t kUnit>
|
||||
using mem_package_t = decltype(get_mem_package<kUnit>());
|
||||
|
||||
template <typename T, std::size_t N>
|
||||
struct storage_vec {
|
||||
T data[N];
|
||||
};
|
||||
|
||||
__always_inline __device__ auto load_nc(const uint1* __restrict__ src) -> uint1 {
|
||||
uint32_t tmp;
|
||||
asm volatile("ld.global.cs.b32 %0,[%1];" : "=r"(tmp) : "l"(src));
|
||||
return uint1{tmp};
|
||||
}
|
||||
|
||||
__always_inline __device__ auto load_nc(const uint2* __restrict__ src) -> uint2 {
|
||||
uint32_t tmp0, tmp1;
|
||||
asm volatile("ld.global.cs.v2.b32 {%0,%1},[%2];" : "=r"(tmp0), "=r"(tmp1) : "l"(src));
|
||||
return uint2{tmp0, tmp1};
|
||||
}
|
||||
|
||||
__always_inline __device__ auto load_nc(const uint4* __restrict__ src) -> uint4 {
|
||||
uint32_t tmp0, tmp1, tmp2, tmp3;
|
||||
asm volatile("ld.global.cs.v4.b32 {%0,%1,%2,%3},[%4];" : "=r"(tmp0), "=r"(tmp1), "=r"(tmp2), "=r"(tmp3) : "l"(src));
|
||||
return uint4{tmp0, tmp1, tmp2, tmp3};
|
||||
}
|
||||
|
||||
__always_inline __device__ void store_nc(uint1* __restrict__ dst, const uint1& value) {
|
||||
uint32_t tmp = value.x;
|
||||
asm volatile("st.global.cs.b32 [%0],%1;" ::"l"(dst), "r"(tmp));
|
||||
}
|
||||
|
||||
__always_inline __device__ void store_nc(uint2* __restrict__ dst, const uint2& value) {
|
||||
uint32_t tmp0 = value.x;
|
||||
uint32_t tmp1 = value.y;
|
||||
asm volatile("st.global.cs.v2.b32 [%0],{%1,%2};" ::"l"(dst), "r"(tmp0), "r"(tmp1));
|
||||
}
|
||||
|
||||
__always_inline __device__ void store_nc(uint4* __restrict__ dst, const uint4& value) {
|
||||
uint32_t tmp0 = value.x;
|
||||
uint32_t tmp1 = value.y;
|
||||
uint32_t tmp2 = value.z;
|
||||
uint32_t tmp3 = value.w;
|
||||
asm volatile("st.global.cs.v4.b32 [%0],{%1,%2,%3,%4};" ::"l"(dst), "r"(tmp0), "r"(tmp1), "r"(tmp2), "r"(tmp3));
|
||||
}
|
||||
|
||||
} // namespace details
|
||||
|
||||
template <
|
||||
std::size_t kBytes,
|
||||
std::size_t kUnit = details::default_unit_size(kBytes),
|
||||
std::size_t kThreads = ::device::kWarpThreads>
|
||||
__always_inline __device__ void copy(void* __restrict__ dst, const void* __restrict__ src) {
|
||||
using Package = details::mem_package_t<kBytes, kUnit>;
|
||||
constexpr auto kBytesPerLoop = sizeof(Package) * kThreads;
|
||||
constexpr auto kLoopCount = kBytes / kBytesPerLoop;
|
||||
static_assert(kBytes % kBytesPerLoop == 0, "kBytes must be multiple of 128 bytes");
|
||||
|
||||
const auto dst_packed = static_cast<Package*>(dst);
|
||||
const auto src_packed = static_cast<const Package*>(src);
|
||||
const auto lane_id = threadIdx.x % kThreads;
|
||||
|
||||
#pragma unroll kLoopCount
|
||||
for (std::size_t i = 0; i < kLoopCount; ++i) {
|
||||
const auto j = i * kThreads + lane_id;
|
||||
dst_packed[j] = src_packed[j];
|
||||
}
|
||||
}
|
||||
|
||||
template <
|
||||
std::size_t kBytes,
|
||||
std::size_t kUnit = details::default_unit_size(kBytes),
|
||||
std::size_t kThreads = ::device::kWarpThreads>
|
||||
__always_inline __device__ auto load_vec(const void* __restrict__ src) {
|
||||
using Package = details::mem_package_t<kBytes, kUnit>;
|
||||
constexpr auto kBytesPerLoop = sizeof(Package) * kThreads;
|
||||
constexpr auto kLoopCount = kBytes / kBytesPerLoop;
|
||||
static_assert(kBytes % kBytesPerLoop == 0, "kBytes must be multiple of 128 bytes");
|
||||
|
||||
const auto src_packed = static_cast<const Package*>(src);
|
||||
const auto lane_id = threadIdx.x % kThreads;
|
||||
details::storage_vec<Package, kLoopCount> vec;
|
||||
|
||||
#pragma unroll kLoopCount
|
||||
for (std::size_t i = 0; i < kLoopCount; ++i) {
|
||||
const auto j = i * kThreads + lane_id;
|
||||
vec.data[i] = details::load_nc(src_packed + j);
|
||||
}
|
||||
|
||||
return vec;
|
||||
}
|
||||
|
||||
template <
|
||||
std::size_t kBytes,
|
||||
std::size_t kUnit = details::default_unit_size(kBytes),
|
||||
std::size_t kThreads = ::device::kWarpThreads,
|
||||
typename Tp>
|
||||
__always_inline __device__ void store_vec(void* __restrict__ dst, const Tp& vec) {
|
||||
using Package = details::mem_package_t<kBytes, kUnit>;
|
||||
constexpr auto kBytesPerLoop = sizeof(Package) * kThreads;
|
||||
constexpr auto kLoopCount = kBytes / kBytesPerLoop;
|
||||
static_assert(kBytes % kBytesPerLoop == 0, "kBytes must be multiple of 128 bytes");
|
||||
static_assert(std::is_same_v<Tp, details::storage_vec<Package, kLoopCount>>);
|
||||
|
||||
const auto dst_packed = static_cast<Package*>(dst);
|
||||
const auto lane_id = threadIdx.x % kThreads;
|
||||
|
||||
#pragma unroll kLoopCount
|
||||
for (std::size_t i = 0; i < kLoopCount; ++i) {
|
||||
const auto j = i * kThreads + lane_id;
|
||||
details::store_nc(dst_packed + j, vec.data[i]);
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace device::warp
|
||||
103
python/sglang/jit_kernel/utils.py
Normal file
103
python/sglang/jit_kernel/utils.py
Normal file
@@ -0,0 +1,103 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import pathlib
|
||||
from functools import lru_cache
|
||||
from typing import TYPE_CHECKING, List, Tuple, TypeAlias, Union
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from tvm_ffi import Module
|
||||
|
||||
|
||||
def _make_wrapper(tup: Tuple[str, str]) -> str:
|
||||
export_name, kernel_name = tup
|
||||
return f"TVM_FFI_DLL_EXPORT_TYPED_FUNC({export_name}, ({kernel_name}));"
|
||||
|
||||
|
||||
@lru_cache()
|
||||
def _resolve_kernel_path() -> pathlib.Path:
|
||||
cur_dir = pathlib.Path(__file__).parent.resolve()
|
||||
|
||||
# first, try this directory structure
|
||||
def _environment_install():
|
||||
candidate = cur_dir.resolve()
|
||||
if (candidate / "include").exists() and (candidate / "csrc").exists():
|
||||
return candidate
|
||||
return None
|
||||
|
||||
def _package_install():
|
||||
# TODO: support find path by package
|
||||
return None
|
||||
|
||||
path = _environment_install() or _package_install()
|
||||
if path is None:
|
||||
raise RuntimeError("Cannot find sgl-kernel/jit path")
|
||||
return path
|
||||
|
||||
|
||||
KERNEL_PATH = _resolve_kernel_path()
|
||||
DEFAULT_INCLUDE = [str(KERNEL_PATH / "include")]
|
||||
DEFAULT_CFLAGS = ["-std=c++20", "-O3"]
|
||||
DEFAULT_CUDA_CFLAGS = ["-std=c++20", "-O3", "--expt-relaxed-constexpr"]
|
||||
DEFAULT_LDFLAGS = []
|
||||
CPP_TEMPLATE_TYPE: TypeAlias = Union[int, float, bool]
|
||||
|
||||
|
||||
class CPPArgList(list[str]):
|
||||
def __str__(self) -> str:
|
||||
return ", ".join(self)
|
||||
|
||||
|
||||
def make_cpp_args(*args: CPP_TEMPLATE_TYPE) -> CPPArgList:
|
||||
def _convert(arg: CPP_TEMPLATE_TYPE) -> str:
|
||||
if isinstance(arg, bool):
|
||||
return "true" if arg else "false"
|
||||
if isinstance(arg, (int, float)):
|
||||
return str(arg)
|
||||
raise TypeError(f"Unsupported argument type for cpp template: {type(arg)}")
|
||||
|
||||
return CPPArgList(_convert(arg) for arg in args)
|
||||
|
||||
|
||||
def load_jit(
|
||||
*args: str,
|
||||
cpp_files: List[str] | None = None,
|
||||
cuda_files: List[str] | None = None,
|
||||
cpp_wrappers: List[Tuple[str, str]] | None = None,
|
||||
cuda_wrappers: List[Tuple[str, str]] | None = None,
|
||||
extra_cflags: List[str] | None = None,
|
||||
extra_cuda_cflags: List[str] | None = None,
|
||||
extra_ldflags: List[str] | None = None,
|
||||
extra_include_paths: List[str] | None = None,
|
||||
build_directory: str | None = None,
|
||||
) -> Module:
|
||||
from tvm_ffi.cpp import load_inline
|
||||
|
||||
cpp_files = cpp_files or []
|
||||
cuda_files = cuda_files or []
|
||||
cpp_wrappers = cpp_wrappers or []
|
||||
cuda_wrappers = cuda_wrappers or []
|
||||
extra_cflags = extra_cflags or []
|
||||
extra_cuda_cflags = extra_cuda_cflags or []
|
||||
extra_ldflags = extra_ldflags or []
|
||||
extra_include_paths = extra_include_paths or []
|
||||
|
||||
# include cpp files
|
||||
cpp_paths = [(KERNEL_PATH / "csrc" / f).resolve() for f in cpp_files]
|
||||
cpp_sources = [f'#include "{path}"' for path in cpp_paths]
|
||||
cpp_sources += [_make_wrapper(tup) for tup in cpp_wrappers]
|
||||
|
||||
# include cuda files
|
||||
cuda_paths = [(KERNEL_PATH / "csrc" / f).resolve() for f in cuda_files]
|
||||
cuda_sources = [f'#include "{path}"' for path in cuda_paths]
|
||||
cuda_sources += [_make_wrapper(tup) for tup in cuda_wrappers]
|
||||
|
||||
return load_inline(
|
||||
"sgl_kernel_jit_" + "_".join(str(arg) for arg in args),
|
||||
cpp_sources=cpp_sources,
|
||||
cuda_sources=cuda_sources,
|
||||
extra_cflags=DEFAULT_CFLAGS + extra_cflags,
|
||||
extra_cuda_cflags=DEFAULT_CUDA_CFLAGS + extra_cuda_cflags,
|
||||
extra_ldflags=DEFAULT_LDFLAGS + extra_ldflags,
|
||||
extra_include_paths=DEFAULT_INCLUDE + extra_include_paths,
|
||||
build_directory=build_directory,
|
||||
)
|
||||
@@ -7,6 +7,7 @@ from typing import Optional
|
||||
import psutil
|
||||
import torch
|
||||
|
||||
from sglang.jit_kernel.hicache import transfer_hicache_one_layer
|
||||
from sglang.srt.mem_cache.memory_pool import KVCache, MHATokenToKVPool, MLATokenToKVPool
|
||||
from sglang.srt.utils import is_npu, is_xpu
|
||||
|
||||
@@ -14,14 +15,12 @@ _is_npu = is_npu()
|
||||
_is_xpu = is_xpu()
|
||||
if not (_is_npu or _is_xpu):
|
||||
from sgl_kernel.kvcacheio import (
|
||||
transfer_kv_all_layer,
|
||||
transfer_kv_all_layer_direct_lf_pf,
|
||||
transfer_kv_all_layer_lf_pf,
|
||||
transfer_kv_all_layer_lf_ph,
|
||||
transfer_kv_all_layer_mla,
|
||||
transfer_kv_all_layer_mla_lf_pf,
|
||||
transfer_kv_direct,
|
||||
transfer_kv_per_layer,
|
||||
transfer_kv_per_layer_direct_pf_lf,
|
||||
transfer_kv_per_layer_mla,
|
||||
transfer_kv_per_layer_mla_pf_lf,
|
||||
@@ -282,14 +281,14 @@ class MHATokenToKVPoolHost(HostKVCache):
|
||||
):
|
||||
if io_backend == "kernel":
|
||||
if self.layout == "layer_first":
|
||||
transfer_kv_per_layer(
|
||||
src_k=self.k_buffer[layer_id],
|
||||
dst_k=device_pool.k_buffer[layer_id],
|
||||
src_v=self.v_buffer[layer_id],
|
||||
dst_v=device_pool.v_buffer[layer_id],
|
||||
src_indices=host_indices,
|
||||
dst_indices=device_indices,
|
||||
item_size=self.token_stride_size,
|
||||
transfer_hicache_one_layer(
|
||||
k_cache_dst=device_pool.k_buffer[layer_id],
|
||||
v_cache_dst=device_pool.v_buffer[layer_id],
|
||||
k_cache_src=self.k_buffer[layer_id],
|
||||
v_cache_src=self.v_buffer[layer_id],
|
||||
indices_dst=device_indices,
|
||||
indices_src=host_indices,
|
||||
element_dim=self.head_num * self.head_dim,
|
||||
)
|
||||
elif self.layout == "page_first":
|
||||
transfer_kv_per_layer_pf_lf(
|
||||
@@ -369,16 +368,17 @@ class MHATokenToKVPoolHost(HostKVCache):
|
||||
):
|
||||
if io_backend == "kernel":
|
||||
if self.layout == "layer_first":
|
||||
transfer_kv_all_layer(
|
||||
src_k_layers=device_pool.k_data_ptrs,
|
||||
dst_k_layers=self.k_data_ptrs,
|
||||
src_v_layers=device_pool.v_data_ptrs,
|
||||
dst_v_layers=self.v_data_ptrs,
|
||||
src_indices=device_indices,
|
||||
dst_indices=host_indices,
|
||||
item_size=self.token_stride_size,
|
||||
num_layers=self.layer_num,
|
||||
)
|
||||
element_dim = self.head_num * self.head_dim
|
||||
for layer_id in range(self.layer_num):
|
||||
transfer_hicache_one_layer(
|
||||
k_cache_dst=self.k_buffer[layer_id],
|
||||
v_cache_dst=self.v_buffer[layer_id],
|
||||
k_cache_src=device_pool.k_buffer[layer_id],
|
||||
v_cache_src=device_pool.v_buffer[layer_id],
|
||||
indices_dst=host_indices,
|
||||
indices_src=device_indices,
|
||||
element_dim=element_dim,
|
||||
)
|
||||
elif self.layout == "page_first":
|
||||
transfer_kv_all_layer_lf_pf(
|
||||
src_k_layers=device_pool.k_data_ptrs,
|
||||
|
||||
Reference in New Issue
Block a user