Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
62 changes: 52 additions & 10 deletions python/sglang/jit_kernel/csrc/kvcacheio/hicache.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -37,44 +37,86 @@ inline constexpr auto get_mem_package() {
template <int kUnit>
using PackageType = decltype(get_mem_package<kUnit>());

// NVIDIA exposes an explicit "do not allocate in L1" cache hint via PTX. ROCm
// has no equivalent PTX, but non-temporal (streaming) loads/stores express the
// same intent for one-shot HiCache write-back traffic that should not pollute
// the cache. Guard the PTX behind USE_ROCM so the JIT module also compiles with
// hipcc; see python/sglang/jit_kernel/utils.py for the ROCm build flags.
#ifdef USE_ROCM
// Native Clang vector types so a single __builtin_nontemporal_{load,store} maps
// to one vectorized global_{load,store}_dwordx{2,4}. Issuing N independent
// 32-bit nontemporal ops instead leaves merging to the LoadStoreVectorizer,
// which is not guaranteed and may drop the nontemporal hint, throttling HiCache
// bandwidth. uint2/uint4 already carry 8B/16B alignment matching the vector
// types, so the pointer reinterpret_casts stay correctly aligned.
typedef uint32_t native_uint2 __attribute__((ext_vector_type(2)));
typedef uint32_t native_uint4 __attribute__((ext_vector_type(4)));
#endif

SGL_DEVICE uint1 load_nc(const uint1* __restrict__ src) {
#ifndef USE_ROCM
uint32_t tmp;
asm volatile("ld.global.L1::no_allocate.b32 %0,[%1];" : "=r"(tmp) : "l"(src));
return uint1{tmp};
#else
return uint1{__builtin_nontemporal_load(&src->x)};
#endif
}

SGL_DEVICE uint2 load_nc(const uint2* __restrict__ src) {
#ifndef USE_ROCM
uint32_t tmp0, tmp1;
asm volatile("ld.global.L1::no_allocate.v2.b32 {%0,%1},[%2];" : "=r"(tmp0), "=r"(tmp1) : "l"(src));
return uint2{tmp0, tmp1};
#else
native_uint2 tmp = __builtin_nontemporal_load(reinterpret_cast<const native_uint2*>(src));
return __builtin_bit_cast(uint2, tmp);
#endif
}

SGL_DEVICE uint4 load_nc(const uint4* __restrict__ src) {
#ifndef USE_ROCM
uint32_t tmp0, tmp1, tmp2, tmp3;
asm volatile("ld.global.L1::no_allocate.v4.b32 {%0,%1,%2,%3},[%4];"
: "=r"(tmp0), "=r"(tmp1), "=r"(tmp2), "=r"(tmp3)
: "l"(src));
return uint4{tmp0, tmp1, tmp2, tmp3};
#else
native_uint4 tmp = __builtin_nontemporal_load(reinterpret_cast<const native_uint4*>(src));
return __builtin_bit_cast(uint4, tmp);
#endif
}

SGL_DEVICE void store_nc(uint1* __restrict__ dst, const uint1& value) {
#ifndef USE_ROCM
uint32_t tmp = value.x;
asm volatile("st.global.L1::no_allocate.b32 [%0],%1;" ::"l"(dst), "r"(tmp));
#else
__builtin_nontemporal_store(value.x, &dst->x);
#endif
}

SGL_DEVICE void store_nc(uint2* __restrict__ dst, const uint2& value) {
#ifndef USE_ROCM
uint32_t tmp0 = value.x;
uint32_t tmp1 = value.y;
asm volatile("st.global.L1::no_allocate.v2.b32 [%0],{%1,%2};" ::"l"(dst), "r"(tmp0), "r"(tmp1));
#else
__builtin_nontemporal_store(__builtin_bit_cast(native_uint2, value), reinterpret_cast<native_uint2*>(dst));
#endif
}

SGL_DEVICE void store_nc(uint4* __restrict__ dst, const uint4& value) {
#ifndef USE_ROCM
uint32_t tmp0 = value.x;
uint32_t tmp1 = value.y;
uint32_t tmp2 = value.z;
uint32_t tmp3 = value.w;
asm volatile(
"st.global.L1::no_allocate.v4.b32 [%0],{%1,%2,%3,%4};" ::"l"(dst), "r"(tmp0), "r"(tmp1), "r"(tmp2), "r"(tmp3));
#else
__builtin_nontemporal_store(__builtin_bit_cast(native_uint4, value), reinterpret_cast<native_uint4*>(dst));
#endif
}

} // namespace details
Expand Down Expand Up @@ -256,18 +298,18 @@ struct HiCacheKernel {
TensorMatcher({-1, D}) //
.with_strides({N, 1})
.with_dtype(cache_dtype)
.with_device<kDLCUDA, kDLCUDAHost, kDLCPU>()
.with_device<kDLCUDA, kDLROCM, kDLCUDAHost, kDLROCMHost, 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>()
.with_device<kDLCUDA, kDLROCM, kDLCUDAHost, kDLROCMHost, kDLCPU>()
.verify(k_cache_dst)
.verify(v_cache_dst);
TensorMatcher({L}) //
.with_dtype<int32_t, int64_t>(indices_dtype)
.with_device<kDLCUDA>(indices_device)
.with_device<kDLCUDA, kDLROCM>(indices_device)
.verify(indices_src)
.verify(indices_dst);

Expand Down Expand Up @@ -323,14 +365,14 @@ struct HiCacheKernel {

TensorMatcher({N}) //
.with_dtype<uint64_t>()
.with_device<kDLCUDA>(device_)
.with_device<kDLCUDA, kDLROCM>(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_)
.with_device<kDLCUDA, kDLROCM>(device_)
.verify(indices_src)
.verify(indices_dst);

Expand Down Expand Up @@ -381,16 +423,16 @@ struct HiCacheKernel {
TensorMatcher({-1, D}) //
.with_strides({N, 1})
.with_dtype(cache_dtype)
.with_device<kDLCUDA, kDLCUDAHost, kDLCPU>()
.with_device<kDLCUDA, kDLROCM, kDLCUDAHost, kDLROCMHost, kDLCPU>()
.verify(cache_src);
TensorMatcher({-1, D}) //
.with_strides({M, 1})
.with_dtype(cache_dtype)
.with_device<kDLCUDA, kDLCUDAHost, kDLCPU>()
.with_device<kDLCUDA, kDLROCM, kDLCUDAHost, kDLROCMHost, kDLCPU>()
.verify(cache_dst);
TensorMatcher({L}) //
.with_dtype<int32_t, int64_t>(indices_dtype)
.with_device<kDLCUDA>(indices_device)
.with_device<kDLCUDA, kDLROCM>(indices_device)
.verify(indices_src)
.verify(indices_dst);

Expand Down Expand Up @@ -441,12 +483,12 @@ struct HiCacheKernel {

TensorMatcher({N}) //
.with_dtype<uint64_t>()
.with_device<kDLCUDA>(device_)
.with_device<kDLCUDA, kDLROCM>(device_)
.verify(ptr_src)
.verify(ptr_dst);
TensorMatcher({L}) //
.with_dtype<int32_t, int64_t>(dtype_)
.with_device<kDLCUDA>(device_)
.with_device<kDLCUDA, kDLROCM>(device_)
.verify(indices_src)
.verify(indices_dst);

Expand Down
16 changes: 8 additions & 8 deletions python/sglang/jit_kernel/csrc/kvcacheio/staged_write_back.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -210,41 +210,41 @@ struct HiCacheStagedWriteBackKernel {

TensorMatcher({T, N, D}) //
.with_dtype(cache_dtype)
.with_device<kDLCUDA>(device_)
.with_device<kDLCUDA, kDLROCM>(device_)
.verify(staging_k);
if constexpr (!kIsMLA) {
TensorMatcher({T, N, D}) //
.with_dtype(cache_dtype)
.with_device<kDLCUDA>(device_)
.with_device<kDLCUDA, kDLROCM>(device_)
.verify(staging_v);
}
TensorMatcher({-1, N, D}) //
.with_dtype(cache_dtype)
.with_device<kDLCPU, kDLCUDAHost>()
.with_device<kDLCPU, kDLCUDAHost, kDLROCMHost>()
.verify(k_cache_dst);
if constexpr (!kIsMLA) {
TensorMatcher({-1, N, D}) //
.with_dtype(cache_dtype)
.with_device<kDLCPU, kDLCUDAHost>()
.with_device<kDLCPU, kDLCUDAHost, kDLROCMHost>()
.verify(v_cache_dst);
}
TensorMatcher({N}) //
.with_dtype<uint64_t>()
.with_device<kDLCUDA>(device_)
.with_device<kDLCUDA, kDLROCM>(device_)
.verify(k_ptr_src);
if constexpr (!kIsMLA) {
TensorMatcher({N}) //
.with_dtype<uint64_t>()
.with_device<kDLCUDA>(device_)
.with_device<kDLCUDA, kDLROCM>(device_)
.verify(v_ptr_src);
}
TensorMatcher({P}) //
.with_dtype<int32_t, int64_t>(indices_dtype)
.with_device<kDLCUDA>(device_)
.with_device<kDLCUDA, kDLROCM>(device_)
.verify(page_indices_src);
TensorMatcher({T}) //
.with_dtype<int64_t>(dst_indices_dtype)
.with_device<kDLCPU, kDLCUDAHost>()
.with_device<kDLCPU, kDLCUDAHost, kDLROCMHost>()
.verify(dst_indices_cpu);

RuntimeCheck(page_size > 0, "HiCache staged relayout: page_size must be positive");
Expand Down
9 changes: 8 additions & 1 deletion python/sglang/srt/managers/cache_controller.py
Original file line number Diff line number Diff line change
Expand Up @@ -675,7 +675,14 @@ def start_writing(self) -> None:
return

op = CacheOperation.merge_ops(self.write_queue)
# Page-first write-back JIT kernels can keep destination host indices on CPU.
# Kernel write-back keeps host indices on CPU only for page_first AND only
# when the staged JIT write-back kernel is available (it stages through
# device memory and accepts CPU destination indices). Otherwise we fall back
# to the plain transfer kernel, whose CUDA/HIP implementation requires
# device-resident destination indices -- so the indices must be moved to the
# device first. Without the can_use_write_back_jit check this crashes on
# backends where the JIT kernel is unavailable, with
# "Destination indices must be a CUDA tensor".
if (
self.io_backend == "kernel"
and self.mem_pool_host.layout == "page_first"
Expand Down
32 changes: 28 additions & 4 deletions python/sglang/srt/mem_cache/memory_pool_host.py
Original file line number Diff line number Diff line change
Expand Up @@ -111,7 +111,13 @@ def __init__(
allocator_type,
)
self.element_dim = self.device_pool.head_num * self.device_pool.head_dim
self.can_use_jit = _is_cuda and can_use_hicache_jit_kernel(
# The JIT HiCache kernels also build with hipcc (ROCm): the PTX-only
# helpers in hicache.cuh are guarded by USE_ROCM and the staged
# write-back kernel has a ROCm path, so enable them on HIP too. This
# keeps the ROCm write-back path consistent with CUDA; without it, ROCm
# falls back to the C++ kernel that requires CUDA-resident destination
# indices and crashes when cache_controller keeps host indices on CPU.
self.can_use_jit = (_is_cuda or _is_hip) and can_use_hicache_jit_kernel(
element_size=self.element_dim * self.dtype.itemsize
)

Expand Down Expand Up @@ -193,7 +199,13 @@ def _init_write_back_staging_buffers(self):
if self.layout != "page_first" or (_is_npu or _is_xpu or _is_mps):
return

self.can_use_write_back_jit = _is_cuda and can_use_write_back_jit_kernel(
# The staged write-back JIT kernel builds with hipcc and has a ROCm path,
# so enable it on HIP too. This keeps the ROCm write-back path consistent
# with CUDA instead of falling back to the C++ kernel that requires
# device-resident destination indices.
self.can_use_write_back_jit = (
_is_cuda or _is_hip
) and can_use_write_back_jit_kernel(
element_size=self.element_dim * self.dtype.itemsize,
)
if not self.can_use_write_back_jit:
Expand Down Expand Up @@ -1282,7 +1294,13 @@ def __init__(
device,
allocator_type,
)
self.can_use_jit = _is_cuda and can_use_hicache_jit_kernel(
# The JIT HiCache kernels also build with hipcc (ROCm): the PTX-only
# helpers in hicache.cuh are guarded by USE_ROCM and the staged
# write-back kernel has a ROCm path, so enable them on HIP too. This
# keeps the ROCm write-back path consistent with CUDA; without it, ROCm
# falls back to the C++ kernel that requires CUDA-resident destination
# indices and crashes when cache_controller keeps host indices on CPU.
self.can_use_jit = (_is_cuda or _is_hip) and can_use_hicache_jit_kernel(
element_size=self.kv_cache_dim * self.dtype.itemsize
)

Expand Down Expand Up @@ -1402,7 +1420,13 @@ def _init_write_back_staging_buffers(self):
if self.layout != "page_first" or (_is_npu or _is_xpu or _is_mps):
return

self.can_use_write_back_jit = _is_cuda and can_use_write_back_jit_kernel(
# The staged write-back JIT kernel builds with hipcc and has a ROCm path,
# so enable it on HIP too. This keeps the ROCm write-back path consistent
# with CUDA instead of falling back to the C++ kernel that requires
# device-resident destination indices.
self.can_use_write_back_jit = (
_is_cuda or _is_hip
) and can_use_write_back_jit_kernel(
element_size=self.kv_cache_dim * self.dtype.itemsize,
)
if not self.can_use_write_back_jit:
Expand Down
14 changes: 0 additions & 14 deletions python/sglang/srt/server_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -5920,20 +5920,6 @@ def _resolve_layout_io_compatibility(self):
"Page first layout is not supported with direct IO backend, switching to page first direct layout"
)

# The page_first kernel write-back relies on the CUDA-only JIT staged
# kernel. On ROCm it falls back to a kernel that requires CUDA index
# tensors and crashes on host write-back, so use layer_first there.
if (
self.hicache_mem_layout == "page_first"
and self.hicache_io_backend == "kernel"
and is_hip()
):
self.hicache_mem_layout = "layer_first"
logger.warning(
"page_first kernel write-back requires the CUDA JIT kernel; "
"falling back to layer_first layout on ROCm."
)

def _resolve_storage_layout_compatibility(self):
if (
self.hicache_storage_backend != "mooncake"
Expand Down
Loading