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
6 changes: 6 additions & 0 deletions docs/docs/advanced_features/hisparse_guide.mdx
Original file line number Diff line number Diff line change
Expand Up @@ -118,6 +118,12 @@ Pass as a JSON string via `--hisparse-config`:

Example: `--hisparse-config='{"top_k": 2048, "device_buffer_size": 6144, "host_to_device_ratio": 10, "swap_in_block_size": 960}'`

### Shared-index prefetch (automatic)

When a model reuses one anchor layer's top-k selection across a run of subsequent "skip" layers (DSA `index_topk_freq` / `index_topk_pattern`; native in GLM-5.2 as IndexShare), the working set of every skip layer is known the moment the anchor's index is computed. HiSparse exploits this automatically: the anchor's swap-in kernel records its miss plan (which host slots go to which device-buffer slots), and each skip layer replays that plan with a copy-only kernel issued ahead on a side stream, so the skip layers' host→device IO overlaps the intervening layers' compute instead of sitting on the decode critical path. The replay kernel uses a small fixed grid to keep its SM footprint low while overlapped.

The prefetch is enabled automatically for eligible models (no pipeline parallelism, no speculative decoding) and can be turned off for A/B comparison with `SGLANG_DISABLE_HISPARSE_PREFETCH=1`.

## Deployment

HiSparse currently requires **PD disaggregation mode** and is enabled only on the **decode instance**.
Expand Down
251 changes: 209 additions & 42 deletions python/sglang/kernels/jit/csrc/hisparse.cuh

Large diffs are not rendered by default.

124 changes: 121 additions & 3 deletions python/sglang/kernels/ops/kvcache/hisparse.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,12 +19,29 @@ def _jit_sparse_module(
hot_buffer_size: int,
is_mla: bool = False,
is_dsv4_layout: bool = False,
record_miss_plan: bool = False,
skip_io: bool = False,
) -> Module:
# record_miss_plan / skip_io are compile-time kernel flags; the
# (False, False) production instantiation stays byte-identical.
template_args = make_cpp_args(
block_size, num_top_k, hot_buffer_size, is_mla, is_dsv4_layout
block_size,
num_top_k,
hot_buffer_size,
is_mla,
is_dsv4_layout,
record_miss_plan,
skip_io,
)
cache_args = make_cpp_args(
item_size_bytes, block_size, num_top_k, hot_buffer_size, is_mla, is_dsv4_layout
item_size_bytes,
block_size,
num_top_k,
hot_buffer_size,
is_mla,
is_dsv4_layout,
record_miss_plan,
skip_io,
)
return load_jit(
"sparse_cache",
Expand All @@ -39,6 +56,30 @@ def _jit_sparse_module(
)


@functools.cache
def _jit_copy_planned_module(
block_size: int,
is_mla: bool,
is_dsv4_layout: bool,
skip_io: bool,
) -> Module:
template_args = make_cpp_args(block_size, is_mla, is_dsv4_layout, skip_io)
return load_jit(
"sparse_copy_planned",
block_size,
is_mla,
is_dsv4_layout,
skip_io,
cuda_files=["hisparse.cuh"],
cuda_wrappers=[
(
"copy_cache_planned",
f"copy_cache_planned<{template_args}>",
)
],
)


@functools.cache
def _jit_dsv4_transfer_module(block_size: int) -> Module:
template_args = make_cpp_args(block_size)
Expand Down Expand Up @@ -91,18 +132,25 @@ def _load_cache_to_device_buffer_mla(
page_size: int,
block_size: int,
num_real_reqs: torch.Tensor | None,
miss_src: torch.Tensor | None,
miss_dst: torch.Tensor | None,
miss_count: torch.Tensor | None,
skip_io: bool,
) -> None:
assert (
hot_buffer_size >= num_top_k
), f"hot_buffer_size ({hot_buffer_size}) must be >= num_top_k ({num_top_k})"

record_miss_plan = miss_src is not None
module = _jit_sparse_module(
item_size_bytes,
block_size,
num_top_k,
hot_buffer_size,
is_mla=True,
is_dsv4_layout=is_dsv4_layout,
record_miss_plan=record_miss_plan,
skip_io=skip_io,
)

empty = torch.empty(0)
Expand All @@ -112,6 +160,16 @@ def _load_cache_to_device_buffer_mla(
[top_k_tokens.size(0)], dtype=torch.int32, device=top_k_tokens.device
)

if record_miss_plan:
assert miss_dst is not None and miss_count is not None
assert miss_src.dtype == torch.int64 and miss_dst.dtype == torch.int32
assert miss_count.dtype == torch.int32
# The kernel indexes both plan rows with one stride.
assert miss_src.stride(0) == miss_dst.stride(0)
else:
# Unused sentinels; the RecordMissPlan=false instantiation never reads them.
miss_src = miss_dst = miss_count = empty

module.load_cache_to_device_buffer(
top_k_tokens,
device_buffer_tokens,
Expand All @@ -128,6 +186,9 @@ def _load_cache_to_device_buffer_mla(
num_real_reqs,
page_size,
item_size_bytes,
miss_src,
miss_dst,
miss_count,
)


Expand All @@ -148,8 +209,16 @@ def load_cache_to_device_buffer_mla(
page_size: int = 1,
block_size: int = 256,
num_real_reqs: torch.Tensor | None = None,
miss_src: torch.Tensor | None = None,
miss_dst: torch.Tensor | None = None,
miss_count: torch.Tensor | None = None,
skip_io: bool = False,
) -> None:
"""Generic MLA hisparse swap-in: device + host both linear (stride=item_size_bytes)."""
"""Generic MLA hisparse swap-in: device + host both linear (stride=item_size_bytes).

Optional miss_src/miss_dst/miss_count record the miss plan for replay by
copy_cache_planned_mla; skip_io elides only the KV bytes (timing probe).
"""
_load_cache_to_device_buffer_mla(
is_dsv4_layout=False,
top_k_tokens=top_k_tokens,
Expand All @@ -168,6 +237,47 @@ def load_cache_to_device_buffer_mla(
page_size=page_size,
block_size=block_size,
num_real_reqs=num_real_reqs,
miss_src=miss_src,
miss_dst=miss_dst,
miss_count=miss_count,
skip_io=skip_io,
)


def copy_cache_planned_mla(
*,
miss_src: torch.Tensor,
miss_dst: torch.Tensor,
miss_count: torch.Tensor,
num_real_reqs: torch.Tensor,
host_cache: torch.Tensor,
device_buffer: torch.Tensor,
item_size_bytes: int,
num_blocks: int = 4,
block_size: int = 1024,
is_dsv4_layout: bool = False,
skip_io: bool = False,
) -> None:
"""Replay a recorded miss plan (host_cache -> device_buffer) for a skip layer.

IO-only, no planning; the small fixed grid keeps the SM footprint low while
overlapped on a side stream. The anchor's slot table stays valid (lockstep).
"""
assert miss_src.dtype == torch.int64 and miss_dst.dtype == torch.int32
assert miss_count.dtype == torch.int32
module = _jit_copy_planned_module(block_size, True, is_dsv4_layout, skip_io)
empty = torch.empty(0)
module.copy_cache_planned(
miss_src,
miss_dst,
miss_count,
num_real_reqs,
host_cache,
empty,
device_buffer,
empty,
num_blocks,
item_size_bytes,
)


Expand All @@ -188,6 +298,10 @@ def load_cache_to_device_buffer_dsv4_mla(
page_size: int = 1,
block_size: int = 256,
num_real_reqs: torch.Tensor | None = None,
miss_src: torch.Tensor | None = None,
miss_dst: torch.Tensor | None = None,
miss_count: torch.Tensor | None = None,
skip_io: bool = False,
) -> None:
"""DSv4 hisparse swap-in: page-padded device + page-padded host C4 layout."""
_load_cache_to_device_buffer_mla(
Expand All @@ -208,4 +322,8 @@ def load_cache_to_device_buffer_dsv4_mla(
page_size=page_size,
block_size=block_size,
num_real_reqs=num_real_reqs,
miss_src=miss_src,
miss_dst=miss_dst,
miss_count=miss_count,
skip_io=skip_io,
)
8 changes: 8 additions & 0 deletions python/sglang/srt/environ.py
Original file line number Diff line number Diff line change
Expand Up @@ -846,6 +846,14 @@ class Envs:
# Triton two_dot variant, 1.16-1.38x faster across GLM/DS shapes).
SGLANG_OPT_Q8KV8_QPREP_VARIANT = EnvStr("auto")

# HiSparse
# Kill-switch for the shared-index (IndexShare) swap-in prefetch
# (auto-enabled for GLM-5.2-style DSA); set True to A/B synchronous swap-in.
SGLANG_DISABLE_HISPARSE_PREFETCH = EnvBool(False)
# Timing probe: run the swap-in fully but skip the host->device KV bytes,
# measuring the "IO is free" floor. GARBAGE OUTPUT -- benchmarking only.
SGLANG_DEBUG_HISPARSE_SKIP_IO = EnvBool(False)

# TRT-LLM-gen fused MoE (SiTU) via sglang JIT: path to an unpacked SiTU
# cubin pool (cubins + flat ABI headers + overlay/; distributed as a
# single downloadable archive). Needs the public flashinfer package
Expand Down
Loading
Loading