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
119 changes: 104 additions & 15 deletions python/sglang/jit_kernel/csrc/hisparse.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -8,14 +8,21 @@
#include <dlpack/dlpack.h>
#include <tvm/ffi/container/tensor.h>

#include <cuda_runtime.h>
#include <stdexcept>
#include <stdint.h>
#include <string>

namespace {

#ifdef USE_ROCM
constexpr int WARP_SIZE = 64;
using BallotMask = uint64_t;
constexpr BallotMask FULL_WARP_MASK = 0xFFFFFFFFFFFFFFFFull;
#else
constexpr int WARP_SIZE = 32;
using BallotMask = unsigned int;
constexpr BallotMask FULL_WARP_MASK = 0xFFFFFFFFu;
#endif
constexpr int32_t TOKEN_HIT = 0xFFFFFFFF;
constexpr int32_t HASH_EMPTY = -1;

Expand All @@ -24,6 +31,25 @@ __device__ __forceinline__ int hash_slot(int32_t key, int hash_size) {
return ((uint32_t)key * 2654435761u) % (uint32_t)hash_size;
}

#ifdef USE_ROCM
__device__ __forceinline__ void transfer_item_warp(
int32_t lane_id, const void* __restrict__ src_addr, void* __restrict__ dst_addr, int64_t item_size_bytes) {
const auto src = static_cast<const char*>(src_addr);
auto dst = static_cast<char*>(dst_addr);

const int64_t word_count = item_size_bytes / static_cast<int64_t>(sizeof(uint64_t));
const auto src_words = reinterpret_cast<const uint64_t*>(src);
auto dst_words = reinterpret_cast<uint64_t*>(dst);
for (int64_t i = lane_id; i < word_count; i += WARP_SIZE) {
dst_words[i] = src_words[i];
}

const int64_t tail_start = word_count * static_cast<int64_t>(sizeof(uint64_t));
for (int64_t i = tail_start + lane_id; i < item_size_bytes; i += WARP_SIZE) {
dst[i] = src[i];
}
}
#else
__device__ __forceinline__ void
transfer_item_warp(int32_t lane_id, const void* src_addr, void* dst_addr, int64_t item_size_bytes) {
// 128-bit bulk transfer via paired 64-bit loads (avoids alignment issues with uint4)
Expand Down Expand Up @@ -51,6 +77,15 @@ transfer_item_warp(int32_t lane_id, const void* src_addr, void* dst_addr, int64_
asm volatile("st.global.cg.b64 [%0],%1;" ::"l"(dst8 + lane_id), "l"(tmp) : "memory");
}
}
#endif

__device__ __forceinline__ int popc_mask(BallotMask mask) {
#ifdef USE_ROCM
return __popcll(mask);
#else
return __popc(mask);
#endif
}

template <int BLOCK_SIZE>
__global__ __launch_bounds__(BLOCK_SIZE, 1) void transfer_cache_dsv4_mla_kernel(
Expand Down Expand Up @@ -113,15 +148,15 @@ __device__ __forceinline__ int warp_inclusive_scan(int* s_data, int lane_id, int
int val = (idx < count) ? s_data[idx] : 0;

#pragma unroll
for (int i = 1; i < 32; i *= 2) {
int n = __shfl_up_sync(0xffffffff, val, i);
for (int i = 1; i < WARP_SIZE; i *= 2) {
int n = __shfl_up_sync(FULL_WARP_MASK, val, i);
if (lane_id >= i) val += n;
}
val += accumulator;
if (idx < count) {
s_data[idx] = val;
}
accumulator = __shfl_sync(0xffffffff, val, 31);
accumulator = __shfl_sync(FULL_WARP_MASK, val, WARP_SIZE - 1);
return accumulator;
}

Expand Down Expand Up @@ -188,7 +223,7 @@ __global__ void load_cache_to_device_buffer_kernel(
const int tid = threadIdx.x;
const int warp_id = tid / WARP_SIZE;
const int lane_id = tid % WARP_SIZE;
const unsigned int lanes_before = ((unsigned int)1 << lane_id) - 1;
const BallotMask lanes_before = (BallotMask(1) << lane_id) - BallotMask(1);

const int64_t rid = req_pool_indices[bid];
const int64_t seq_len = seq_lens[bid];
Expand Down Expand Up @@ -316,22 +351,35 @@ __global__ void load_cache_to_device_buffer_kernel(
int local_hit_offset = 0;
int local_evict_offset = 0;
if (has_valid_chunk) {
const unsigned int hit_mask = __ballot_sync(0xFFFFFFFF, is_hit);
const unsigned int evict_mask = __ballot_sync(0xFFFFFFFF, is_evictable);
local_hit_offset = __popc(hit_mask & lanes_before);
local_evict_offset = __popc(evict_mask & lanes_before);
const BallotMask hit_mask = __ballot_sync(FULL_WARP_MASK, is_hit);
const BallotMask evict_mask = __ballot_sync(FULL_WARP_MASK, is_evictable);
local_hit_offset = popc_mask(hit_mask & lanes_before);
local_evict_offset = popc_mask(evict_mask & lanes_before);
if (lane_id == 0) {
s_chunk_offset[chunk_idx + 1] = __popc(hit_mask);
s_evict_chunk_offset[chunk_idx + 1] = __popc(evict_mask);
s_chunk_offset[chunk_idx + 1] = popc_mask(hit_mask);
s_evict_chunk_offset[chunk_idx + 1] = popc_mask(evict_mask);
}
}
__syncthreads();

if (warp_id == 0) {
#ifdef USE_ROCM
// ROCm wavefront64: WARP_SIZE (64) > NUM_WARPS (16 at block_size=1024),
// so the wide-count form below would let lanes beyond this iteration's
// NUM_WARPS-wide window write the accumulator into s_chunk_offset
// positions belonging to future iterations, corrupting their reads.
// Bound the scan window to NUM_WARPS lanes.
const int scan_offset = iter * NUM_WARPS + 1;
const int scan_count = min(scan_offset + NUM_WARPS, NUM_BUFFER_CHUNKS + 1);
total_hit_count = warp_inclusive_scan(s_chunk_offset, lane_id, scan_offset, scan_count, total_hit_count);
total_evict_count =
warp_inclusive_scan(s_evict_chunk_offset, lane_id, scan_offset, scan_count, total_evict_count);
#else
total_hit_count =
warp_inclusive_scan(s_chunk_offset, lane_id, chunk_idx + 1, NUM_BUFFER_CHUNKS + 1, total_hit_count);
total_evict_count =
warp_inclusive_scan(s_evict_chunk_offset, lane_id, chunk_idx + 1, NUM_BUFFER_CHUNKS + 1, total_evict_count);
#endif
if (tid == 0) {
s_total_hits = total_hit_count;
}
Expand Down Expand Up @@ -380,17 +428,23 @@ __global__ void load_cache_to_device_buffer_kernel(
}

if (has_valid_chunk) {
const unsigned int miss_mask = __ballot_sync(0xFFFFFFFF, is_miss);
local_miss_offset = __popc(miss_mask & lanes_before);
const int warp_miss_count = __popc(miss_mask);
const BallotMask miss_mask = __ballot_sync(FULL_WARP_MASK, is_miss);
local_miss_offset = popc_mask(miss_mask & lanes_before);
const int warp_miss_count = popc_mask(miss_mask);
if (lane_id == 0) {
s_chunk_offset[chunk_idx + 1] = warp_miss_count;
}
}
__syncthreads();

if (warp_id == 0) {
#ifdef USE_ROCM
const int scan_offset = iter * NUM_WARPS + 1;
const int scan_count = min(scan_offset + NUM_WARPS, NUM_TOKEN_CHUNKS + 1);
total_misses = warp_inclusive_scan(s_chunk_offset, lane_id, scan_offset, scan_count, total_misses);
#else
total_misses = warp_inclusive_scan(s_chunk_offset, lane_id, chunk_idx + 1, NUM_TOKEN_CHUNKS + 1, total_misses);
#endif
}
__syncthreads();

Expand All @@ -410,6 +464,24 @@ __global__ void load_cache_to_device_buffer_kernel(
// Write back LRU order: evictables at front (LRU), hits at back (MRU).
{
const int total_evictable = HOT_BUFFER_SIZE - s_total_hits;
#ifdef USE_ROCM
// ROCm: cap writeback threads at 512 for large kernels.
constexpr int LRU_WRITEBACK_THREADS = (BLOCK_SIZE > 512) ? 512 : BLOCK_SIZE;
if (tid < LRU_WRITEBACK_THREADS) {
for (int i = tid; i < HOT_BUFFER_SIZE; i += LRU_WRITEBACK_THREADS) {
if (i < total_misses) {
// Misses: just loaded from host, place right before hits
req_lru_slots[total_evictable - total_misses + i] = s_lru_slots_out[HOT_BUFFER_SIZE - 1 - i];
} else if (i < total_evictable) {
// Remaining evictables: truly stale, dest at LRU front
req_lru_slots[i - total_misses] = s_lru_slots_out[HOT_BUFFER_SIZE - 1 - i];
} else {
// Hits: source at forward end, dest at MRU back
req_lru_slots[i] = s_lru_slots_out[i - total_evictable];
}
}
}
#else
for (int i = tid; i < HOT_BUFFER_SIZE; i += BLOCK_SIZE) {
if (i < total_misses) {
// Misses: just loaded from host, place right before hits
Expand All @@ -422,6 +494,7 @@ __global__ void load_cache_to_device_buffer_kernel(
req_lru_slots[i] = s_lru_slots_out[i - total_evictable];
}
}
#endif
}

// each warp copies one miss directly, can be separated into a new kernel if parallelism is a concern
Expand All @@ -433,14 +506,28 @@ __global__ void load_cache_to_device_buffer_kernel(
const int64_t dst_loc = static_cast<int64_t>(req_device_buffer_locs[evict_slot]);

if constexpr (IsDsv4Layout) {
// DSv4 path: page-padded device layout + page-padded host layout, K-only.
#ifdef USE_ROCM
// ROCm path: host cache and device buffer both use the page-padded C4
// layout (same as the write path and the CUDA branch). We can't reuse
// device::hisparse::transfer_item here because its warp logic is hardcoded
// to a 32-lane warp; on wavefront64 we use the warp-width-agnostic
// transfer_item_warp with paged source and destination addressing.
using namespace device::hisparse;
const auto [dst_value_ptr, dst_scale_ptr] = get_pointer_paged(device_buffer_k, static_cast<int32_t>(dst_loc));
const auto [src_value_ptr, src_scale_ptr] =
get_pointer_paged(const_cast<void*>(host_cache_k), static_cast<int32_t>(src_loc));

@amd-danli103 amd-danli103 Jun 12, 2026 •

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks for picking this up. Verified equivalent to our local fix on MI355X. This unblocks
the V4 HiSparse path on ROCm at the kernel level.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I have tested my PR with our daily 061226 image with DeepSeek-V4 Pro and HiSparse enabled and it runs fine.

transfer_item_warp(lane_id, src_value_ptr, dst_value_ptr, kValueBytes);
transfer_item_warp(lane_id, src_scale_ptr, dst_scale_ptr, kScaleBytes);
#else
// CUDA path: page-padded device layout + page-padded host layout, K-only.
// The host cache is pinned DRAM but uses the same row layout as the GPU C4
// cache, so use the page-padded address calculation for both ends.
device::hisparse::transfer_item(
/*dst_cache=*/device_buffer_k,
/*src_cache=*/const_cast<void*>(host_cache_k),
/*dst_index=*/static_cast<int32_t>(dst_loc),
/*src_index=*/static_cast<int32_t>(src_loc));
#endif
} else {
// Generic path: device + host both linear, stride = item_size_bytes.
const auto src_k = static_cast<const char*>(host_cache_k) + src_loc * item_size_bytes;
Expand Down Expand Up @@ -487,9 +574,11 @@ void load_cache_to_device_buffer(
// seq_lens and req_pool_indices; the correct combo is selected at runtime.
auto launch = [&](auto kernel_fn, const auto* seq_lens_ptr, const auto* req_pool_indices_ptr) {
constexpr size_t smem_bytes = SmemLayout<NUM_TOP_K, HOT_BUFFER_SIZE>::BYTES;
#ifndef USE_ROCM
if constexpr (smem_bytes > 48u * 1024u) {
cudaFuncSetAttribute(kernel_fn, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_bytes);
}
#endif
LaunchKernel(bs, BLOCK_SIZE, device, smem_bytes)(
kernel_fn,
static_cast<const int32_t*>(top_k_tokens.data_ptr()),
Expand Down
93 changes: 70 additions & 23 deletions python/sglang/srt/arg_groups/hisparse_hook.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,20 +8,34 @@

logger = logging.getLogger(__name__)


# Backend/dtype pairing: flashmla_sparse only takes BF16 KV;
# flashmla_kv only supports FP8 (it always reads KV as FP8 via
# is_fp8_kvcache=True, inline-quantizing BF16 would defeat HiSparse).
_HISPARSE_ALLOWED_BACKENDS_BY_DTYPE = {
HISPARSE_CUDA_DSA_BACKENDS_BY_DTYPE = {
"bfloat16": {"flashmla_sparse"},
"fp8_e4m3": {"flashmla_kv"},
}
HISPARSE_ROCM_DSA_BACKENDS = {"tilelang", "aiter"}
HISPARSE_KV_CACHE_DTYPES = ("bfloat16", "fp8_e4m3")


def _is_hip() -> bool:
from sglang.srt.server_args import is_hip

return is_hip()


def _hisparse_default_backend(kv_cache_dtype: str) -> str:
if _is_hip():
return "tilelang"
return "flashmla_kv" if kv_cache_dtype == "fp8_e4m3" else "flashmla_sparse"


def _hisparse_allowed_backends(kv_cache_dtype: str) -> set[str]:
if _is_hip():
return HISPARSE_ROCM_DSA_BACKENDS
return HISPARSE_CUDA_DSA_BACKENDS_BY_DTYPE.get(
kv_cache_dtype, {"flashmla_sparse", "flashmla_kv"}
)


def apply_hisparse_dsa_backend_defaults(
server_args: ServerArgs,
user_set_prefill: bool,
Expand All @@ -30,8 +44,8 @@ def apply_hisparse_dsa_backend_defaults(
) -> bool:
"""Pick DSA backends for --enable-hisparse based on KV dtype.

BF16 KV -> flashmla_sparse, FP8 KV -> flashmla_kv. Returns True if hisparse
handled backend selection (caller should skip its own default logic).
CUDA uses dtype-specific FlashMLA backends; ROCm uses TileLang. Returns
True if hisparse handled backend selection.
"""
if not server_args.enable_hisparse:
return False
Expand All @@ -48,6 +62,35 @@ def apply_hisparse_dsa_backend_defaults(
return True


def validate_hisparse_dsa_backend(
server_args: ServerArgs, attr: str, label: str
) -> None:
backend = getattr(server_args, attr)
allowed_backends = _hisparse_allowed_backends(server_args.kv_cache_dtype)
if backend is not None and backend not in allowed_backends:
raise ValueError(
f"HiSparse supports DSA {label} backend(s) {sorted(allowed_backends)} "
f"on this platform with --kv-cache-dtype={server_args.kv_cache_dtype}, "
f"but got --dsa-{label}-backend={backend}. "
f"Please use --dsa-{label}-backend="
f"{_hisparse_default_backend(server_args.kv_cache_dtype)} "
"or omit it."
)


def validate_hisparse_kv_cache_dtype(server_args: ServerArgs) -> None:
if server_args.kv_cache_dtype in HISPARSE_KV_CACHE_DTYPES:
return

choices = " or ".join(
f"--kv-cache-dtype={dtype}" for dtype in HISPARSE_KV_CACHE_DTYPES
)
raise ValueError(
f"HiSparse requires one of {HISPARSE_KV_CACHE_DTYPES} KV cache dtypes, "
f"but got --kv-cache-dtype={server_args.kv_cache_dtype}. Please use {choices}."
)


def validate_hisparse(server_args: ServerArgs) -> None:
"""Validate --enable-hisparse constraints (model class, radix cache, DSA backend)."""
if not server_args.enable_hisparse:
Expand All @@ -60,6 +103,7 @@ def validate_hisparse(server_args: ServerArgs) -> None:

hf_config = server_args.get_model_config().hf_config
is_v4_hisparse = is_deepseek_v4(hf_config)
is_hip = _is_hip()
assert is_deepseek_dsa(hf_config) or is_v4_hisparse, (
"--enable-hisparse is only supported for DSA (DeepSeek Sparse Attention) "
"models (e.g., DeepSeek V3.2, GLM-5) and DeepSeek V4 now. "
Expand All @@ -71,27 +115,30 @@ def validate_hisparse(server_args: ServerArgs) -> None:

# DSv4 hisparse handles its own dtype/backend pairing elsewhere; the dtype-
# aware checks below only apply to the DSA hisparse path.
if is_v4_hisparse:
if is_hip and is_v4_hisparse:
# TEMPORARY GUARD: DSv4 HiSparse is not supported on the unified-KV path.
# In unified-KV mode c4_kv_pool is None, so DeepSeekV4HiSparseTokenToKVPoolAllocator
# cannot attach and pool init dies with a cryptic AssertionError. Fail fast
# at startup with a clear message instead. Remove once unified-KV HiSparse lands.
from sglang.srt.layers.attention.dsv4.unified_kv_kernels.env_gate import (
is_unified_kv_triton,
)

if is_unified_kv_triton():
raise ValueError(
"--enable-hisparse is not supported with the unified-KV path on ROCm"
"(SGLANG_HACK_FLASHMLA_BACKEND=unified_kv_triton) for DeepSeek-V4: "
"HiSparse currently requires the separate packed KV layout. "
"Either set SGLANG_HACK_FLASHMLA_BACKEND=triton, or run without "
"--enable-hisparse."
)
return

if server_args.kv_cache_dtype not in ("bfloat16", "auto", "fp8_e4m3"):
raise ValueError(
f"HiSparse requires bfloat16 or fp8_e4m3 KV cache, "
f"but got --kv-cache-dtype={server_args.kv_cache_dtype}. "
f"Please use --kv-cache-dtype=bfloat16 or fp8_e4m3."
)
validate_hisparse_kv_cache_dtype(server_args)

allowed_backends = _HISPARSE_ALLOWED_BACKENDS_BY_DTYPE.get(
server_args.kv_cache_dtype, {"flashmla_sparse", "flashmla_kv"}
)
for attr, label in [
("dsa_prefill_backend", "prefill"),
("dsa_decode_backend", "decode"),
]:
backend = getattr(server_args, attr)
if backend is not None and backend not in allowed_backends:
raise ValueError(
f"HiSparse with --kv-cache-dtype={server_args.kv_cache_dtype} requires "
f"--dsa-{label}-backend in {sorted(allowed_backends)}, "
f"but got {backend}."
)
validate_hisparse_dsa_backend(server_args, attr, label)
Loading
Loading