[AMD] Enable JIT staged HiCache write-back and fix CPU-index crash (#28534)

Co-authored-by: Duyi-Wang <duyi.wang@amd.com>
This commit is contained in:
AMD-yanfeiwang
2026-07-09 01:22:37 -07:00
committed by GitHub
co-authored by Duyi-Wang
parent 61602b95fb
commit d74619b373
7 changed files with 349 additions and 23 deletions
@@ -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
@@ -256,18 +298,18 @@ struct HiCacheKernel {
TensorMatcher({-1, D}) //
.with_strides({N, 1})
.with_dtype(cache_dtype)
.with_device<kDLCUDA, kDLCUDAHost, kDLCPU>()
.with_device<kDLGPU, kDLGPUHost, 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<kDLGPU, kDLGPUHost, 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<kDLGPU>(indices_device)
.verify(indices_src)
.verify(indices_dst);
@@ -323,14 +365,14 @@ struct HiCacheKernel {
TensorMatcher({N}) //
.with_dtype<uint64_t>()
.with_device<kDLCUDA>(device_)
.with_device<kDLGPU>(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<kDLGPU>(device_)
.verify(indices_src)
.verify(indices_dst);
@@ -381,16 +423,16 @@ struct HiCacheKernel {
TensorMatcher({-1, D}) //
.with_strides({N, 1})
.with_dtype(cache_dtype)
.with_device<kDLCUDA, kDLCUDAHost, kDLCPU>()
.with_device<kDLGPU, kDLGPUHost, kDLCPU>()
.verify(cache_src);
TensorMatcher({-1, D}) //
.with_strides({M, 1})
.with_dtype(cache_dtype)
.with_device<kDLCUDA, kDLCUDAHost, kDLCPU>()
.with_device<kDLGPU, kDLGPUHost, kDLCPU>()
.verify(cache_dst);
TensorMatcher({L}) //
.with_dtype<int32_t, int64_t>(indices_dtype)
.with_device<kDLCUDA>(indices_device)
.with_device<kDLGPU>(indices_device)
.verify(indices_src)
.verify(indices_dst);
@@ -441,12 +483,12 @@ struct HiCacheKernel {
TensorMatcher({N}) //
.with_dtype<uint64_t>()
.with_device<kDLCUDA>(device_)
.with_device<kDLGPU>(device_)
.verify(ptr_src)
.verify(ptr_dst);
TensorMatcher({L}) //
.with_dtype<int32_t, int64_t>(dtype_)
.with_device<kDLCUDA>(device_)
.with_device<kDLGPU>(device_)
.verify(indices_src)
.verify(indices_dst);
@@ -210,41 +210,41 @@ struct HiCacheStagedWriteBackKernel {
TensorMatcher({T, N, D}) //
.with_dtype(cache_dtype)
.with_device<kDLCUDA>(device_)
.with_device<kDLGPU>(device_)
.verify(staging_k);
if constexpr (!kIsMLA) {
TensorMatcher({T, N, D}) //
.with_dtype(cache_dtype)
.with_device<kDLCUDA>(device_)
.with_device<kDLGPU>(device_)
.verify(staging_v);
}
TensorMatcher({-1, N, D}) //
.with_dtype(cache_dtype)
.with_device<kDLCPU, kDLCUDAHost>()
.with_device<kDLCPU, kDLGPUHost>()
.verify(k_cache_dst);
if constexpr (!kIsMLA) {
TensorMatcher({-1, N, D}) //
.with_dtype(cache_dtype)
.with_device<kDLCPU, kDLCUDAHost>()
.with_device<kDLCPU, kDLGPUHost>()
.verify(v_cache_dst);
}
TensorMatcher({N}) //
.with_dtype<uint64_t>()
.with_device<kDLCUDA>(device_)
.with_device<kDLGPU>(device_)
.verify(k_ptr_src);
if constexpr (!kIsMLA) {
TensorMatcher({N}) //
.with_dtype<uint64_t>()
.with_device<kDLCUDA>(device_)
.with_device<kDLGPU>(device_)
.verify(v_ptr_src);
}
TensorMatcher({P}) //
.with_dtype<int32_t, int64_t>(indices_dtype)
.with_device<kDLCUDA>(device_)
.with_device<kDLGPU>(device_)
.verify(page_indices_src);
TensorMatcher({T}) //
.with_dtype<int64_t>(dst_indices_dtype)
.with_device<kDLCPU, kDLCUDAHost>()
.with_device<kDLCPU, kDLGPUHost>()
.verify(dst_indices_cpu);
RuntimeCheck(page_size > 0, "HiCache staged relayout: page_size must be positive");
@@ -89,8 +89,10 @@ using fp32x4_t = float4;
// DLPack device type for the current platform
#ifndef USE_ROCM
inline constexpr auto kDLGPU = kDLCUDA;
inline constexpr auto kDLGPUHost = kDLCUDAHost;
#else
inline constexpr auto kDLGPU = kDLROCM;
inline constexpr auto kDLGPUHost = kDLROCMHost;
#endif
namespace device {
@@ -675,7 +675,14 @@ class HiCacheController:
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"
@@ -94,7 +94,11 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
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.
self.can_use_jit = (_is_cuda or _is_hip) and can_use_hicache_jit_kernel(
element_size=self.kv_cache_dim * self.dtype.itemsize
)
@@ -214,7 +218,11 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
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 (consistent with the CUDA path).
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:
+10 -2
View File
@@ -88,7 +88,11 @@ class MHATokenToKVPoolHost(HostKVCache):
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.
self.can_use_jit = (_is_cuda or _is_hip) and can_use_hicache_jit_kernel(
element_size=self.element_dim * self.dtype.itemsize
)
@@ -170,7 +174,11 @@ class MHATokenToKVPoolHost(HostKVCache):
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 (consistent with the CUDA path).
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: