Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
64 commits
Select commit Hold shift + click to select a range
050dbac
init
At1a8 Jul 13, 2026
321ab96
Merge branch 'main' into fangyuan/dspark
At1a8 Jul 13, 2026
ee50a83
Merge branch 'main' into fangyuan/dspark
At1a8 Jul 13, 2026
d3de45d
support unified kv
At1a8 Jul 14, 2026
a3e723f
Merge branch 'main' into fangyuan/dspark
At1a8 Jul 14, 2026
058e7c0
hip guard
At1a8 Jul 14, 2026
909db6a
x
At1a8 Jul 14, 2026
888db41
ci
At1a8 Jul 14, 2026
c792ee9
Revert "ci"
At1a8 Jul 14, 2026
bbd7f31
ci test
At1a8 Jul 14, 2026
ac2470d
Revert "Delete CUTLASS FP8 blockwise for SM90 and SM100, move SM120 t…
At1a8 Jul 14, 2026
4661877
x
At1a8 Jul 14, 2026
0e6fcb1
Revert "Revert "Delete CUTLASS FP8 blockwise for SM90 and SM100, move…
At1a8 Jul 15, 2026
27b568c
fix conflict
Jul 15, 2026
9a3eb09
Merge branch 'main' into fangyuan/dspark
At1a8 Jul 15, 2026
dd23cc8
Merge branch 'main' into fangyuan/dspark
At1a8 Jul 15, 2026
b2d06a7
mirror the cuda backend about draft model
At1a8 Jul 16, 2026
e90a20a
Revert "mirror the cuda backend about draft model"
At1a8 Jul 16, 2026
0227778
fix kk comments
At1a8 Jul 17, 2026
7a9fd6a
ci test for kernel
At1a8 Jul 17, 2026
799b4c7
Merge branch 'main' into fangyuan/dspark
At1a8 Jul 17, 2026
63479f6
fix bug
At1a8 Jul 17, 2026
426d5b0
Merge branch 'main' into fangyuan/dspark
At1a8 Jul 18, 2026
5c810a1
update ci
At1a8 Jul 18, 2026
3b112ea
Merge branch 'main' into fangyuan/dspark
At1a8 Jul 18, 2026
592f914
Merge branch 'main' into fangyuan/dspark
At1a8 Jul 19, 2026
b64c2e2
Merge branch 'main' into fangyuan/dspark
At1a8 Jul 19, 2026
0976649
Merge branch 'main' into fangyuan/dspark
At1a8 Jul 21, 2026
043da94
Merge branch 'main' into fangyuan/dspark
At1a8 Jul 22, 2026
95c23ba
Merge branch 'main' into fangyuan/dspark
At1a8 Jul 22, 2026
a5a8e16
fix
At1a8 Jul 22, 2026
02c4b97
Merge branch 'main' into fangyuan/dspark
At1a8 Jul 22, 2026
bda5a32
Merge branch 'main' into fangyuan/dspark
At1a8 Jul 22, 2026
acfd443
Merge remote-tracking branch 'upstream/main' into fangyuan/dspark
At1a8 Jul 23, 2026
60a0d1a
Merge branch 'main' into fangyuan/dspark
At1a8 Jul 23, 2026
b76183d
Merge branch 'main' into fangyuan/dspark
At1a8 Jul 23, 2026
2ca14ad
Merge branch 'main' into fangyuan/dspark
At1a8 Jul 23, 2026
892bf4a
Merge branch 'main' into fangyuan/dspark
At1a8 Jul 25, 2026
0d5cedc
Merge branch 'main' into fangyuan/dspark
At1a8 Jul 31, 2026
95e8243
fix lint
Jul 31, 2026
ed35337
Merge branch 'main' into fangyuan/dspark
At1a8 Jul 31, 2026
058c424
fix bug
At1a8 Jul 31, 2026
172451e
fix
At1a8 Jul 31, 2026
6f406f1
make sure hip only changes
At1a8 Jul 31, 2026
f2e1f8d
Merge branch 'main' into fangyuan/dspark
At1a8 Aug 1, 2026
2bbcec7
fix
At1a8 Aug 1, 2026
18c666d
Merge branch 'main' into fangyuan/dspark
HaiShaw Aug 1, 2026
0ea8b5b
Merge branch 'main' into fangyuan/dspark
At1a8 Aug 1, 2026
64e8dac
Merge branch 'main' into fangyuan/dspark
At1a8 Aug 3, 2026
811ce46
Merge branch 'main' into fangyuan/dspark
At1a8 Aug 3, 2026
ea467c4
Merge branch 'main' into fangyuan/dspark
At1a8 Aug 3, 2026
ad9b7aa
Merge branch 'main' into fangyuan/dspark
At1a8 Aug 4, 2026
5808eb2
Merge branch 'main' into fangyuan/dspark
At1a8 Aug 4, 2026
6bd4ff4
Merge branch 'main' into fangyuan/dspark
At1a8 Aug 4, 2026
acdbd0a
Merge branch 'main' into fangyuan/dspark
At1a8 Aug 4, 2026
8eefbf4
Merge branch 'main' into fangyuan/dspark
At1a8 Aug 4, 2026
ed2752f
x
At1a8 Aug 4, 2026
4699fc1
Merge branch 'main' into fangyuan/dspark
At1a8 Aug 4, 2026
fe0d7d8
Merge branch 'main' into fangyuan/dspark
At1a8 Aug 4, 2026
4a2f3b2
Merge branch 'main' into fangyuan/dspark
At1a8 Aug 4, 2026
6f6d3d9
Merge branch 'main' into fangyuan/dspark
At1a8 Aug 4, 2026
0f1f6e5
Merge branch 'main' into fangyuan/dspark
At1a8 Aug 6, 2026
7c207d8
lint
At1a8 Aug 6, 2026
ecf9606
Merge branch 'main' into fangyuan/dspark
At1a8 Aug 6, 2026
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
Original file line number Diff line number Diff line change
Expand Up @@ -110,6 +110,57 @@ def store_swa_into_unified(
)


@triton.jit
def _scatter_loc_kernel(
kv_ptr, # [T, D] bf16
loc_ptr, # [T] int (unified row index; <0 => skip)
unified_ptr, # [pages, D] bf16
n_rows,
D: tl.constexpr,
BLOCK_D: tl.constexpr,
):
row = tl.program_id(0)
if row >= n_rows:
return
loc = tl.load(loc_ptr + row).to(tl.int64)
if loc < 0:
return
offs = tl.arange(0, BLOCK_D)
mask = offs < D
vals = tl.load(kv_ptr + row * D + offs, mask=mask, other=0.0)
tl.store(unified_ptr + loc * D + offs, vals, mask=mask)


def scatter_bf16_into_unified(
*,
kv: torch.Tensor, # [T, head_dim] bf16 (already norm+rope'd)
loc: torch.Tensor, # [T] int32/int64 unified ring row; <0 => skip
unified_kv: torch.Tensor, # [pages, head_dim] bf16
) -> None:
"""Scatter already-norm+rope'd bf16 K into ``unified_kv[loc]`` (skip loc < 0).

Companion to ``store_swa_into_unified`` for callers that already hold the
precomputed ring row index (the DSpark draft: ``get_unified_swa_loc`` for the
draft forward, or the commit-inject layout for target-hidden injection) and
need per-row commit masking expressed as ``loc == -1``.
"""
n_rows, D = kv.shape
if n_rows == 0:
return
assert kv.is_contiguous() and kv.dtype == unified_kv.dtype
assert loc.is_contiguous()
assert unified_kv.is_contiguous()
_scatter_loc_kernel[(n_rows,)](
kv,
loc,
unified_kv,
n_rows,
D=D,
BLOCK_D=triton.next_power_of_2(D),
num_warps=8,
)


# ---------------------------------------------------------------------------
# Ragged indptr helper (shared by the decode streams + prefill builders)
# ---------------------------------------------------------------------------
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -604,6 +604,37 @@ class CommitInjectLayoutResult(msgspec.Struct):
positions: torch.Tensor


def build_unified_commit_inject_layout(
*,
req_pool_indices: torch.Tensor,
prefix_lens: torch.Tensor,
block_pos_offsets: torch.Tensor,
commit_lens: torch.Tensor,
stride: int,
ring_stride: int,
) -> CommitInjectLayoutResult:
"""unified_kv counterpart of build_commit_inject_layout.

Non-unified injection translates the verify tokens' full cache locs through
``full_to_swa_mapping``; under unified_kv the SWA K lives in a ring addressed
directly by ``state_slot * ring_stride + pos % ring_stride``, so compute the
ring row here instead. Uncommitted tokens (col >= commit_len) get loc = -1 and
are skipped by the scatter. All ops are static-shape (CUDA-graph safe).
"""
bs = req_pool_indices.shape[0]
device = req_pool_indices.device
positions_2d = prefix_lens.unsqueeze(1) + block_pos_offsets[:stride]
positions = positions_2d.reshape(-1).to(torch.int64)
state_slot = (
req_pool_indices.to(torch.int64).view(-1, 1).expand(bs, stride).reshape(-1)
)
loc = state_slot * ring_stride + positions % ring_stride
col = torch.arange(stride, device=device).view(1, -1)
committed = (col < commit_lens.to(torch.long).view(-1, 1)).reshape(-1)
swa_loc = torch.where(committed, loc, torch.full_like(loc, -1)).to(torch.int32)
return CommitInjectLayoutResult(swa_loc=swa_loc, positions=positions)


class BuildCommitInjectLayout:
@classmethod
def execute(cls, *args, **kwargs) -> CommitInjectLayoutResult:
Expand Down
129 changes: 102 additions & 27 deletions python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py
Original file line number Diff line number Diff line change
Expand Up @@ -414,6 +414,7 @@ class DeepseekV4HipRadixBackend(
# both children and leaks ROCm HSA resources (HSA_STATUS_ERROR_OUT_OF_RESOURCES).
# TboAttnBackend reads this to skip children in the *_graph paths only.
tbo_supports_cuda_graph = False
supports_ragged_verify_graph: bool = True

def __init__(
self,
Expand Down Expand Up @@ -456,6 +457,18 @@ def __init__(
self.mtp_enabled = self.topk > 0
self.speculative_num_steps = speculative_num_steps
self.speculative_num_draft_tokens: int = get_spec().speculative_num_draft_tokens
self.is_dspark_draft = (
getattr(model_runner, "is_draft_worker", False)
and model_runner.spec_algorithm.is_dspark()
)
self.target_verify_num_draft_tokens = self.speculative_num_draft_tokens
if self.is_dspark_draft:
assert self.speculative_num_draft_tokens is not None
assert self.speculative_num_draft_tokens > 1
# DSpark draft workers verify gamma rows. The server arg keeps the
# CUDA-side convention gamma + 1, so use an explicit effective value
# instead of mutating speculative_num_draft_tokens in place.
self.target_verify_num_draft_tokens = self.speculative_num_draft_tokens - 1
self.speculative_step_id = speculative_step_id
self.forward_metadata: Union[
DSV4Metadata,
Expand Down Expand Up @@ -532,14 +545,34 @@ def init_forward_metadata_prefill(
extend_seq_lens_cpu: List[int],
need_compress: bool = True,
use_prefill_cuda_graph: bool = False,
compress_gpu_plan: bool = False,
extend_start_loc: Optional[torch.Tensor] = None,
) -> DSV4Metadata:
seq_lens_casual, req_pool_indices_repeated = self.expand_prefill_casually(
num_tokens=num_tokens,
seq_lens=seq_lens_cpu,
extend_seq_lens=extend_seq_lens_cpu,
req_pool_indices=req_pool_indices,
padded_num_tokens=out_cache_loc.shape[0],
)
if extend_start_loc is not None:
from sglang.kernels.ops.attention.dsv4_attn_metadata_kernels import (
ExpandPrefillCausally,
)

_expanded = ExpandPrefillCausally.execute(
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
extend_seq_lens=extend_seq_lens,
extend_start_loc=extend_start_loc,
seq_lens_cpu=None,
extend_seq_lens_cpu=None,
num_tokens=num_tokens,
padded_num_tokens=out_cache_loc.shape[0],
)
seq_lens_casual = _expanded.seq_lens_casual
req_pool_indices_repeated = _expanded.req_pool_indices_repeated
else:
seq_lens_casual, req_pool_indices_repeated = self.expand_prefill_casually(
num_tokens=num_tokens,
seq_lens=seq_lens_cpu,
extend_seq_lens=extend_seq_lens_cpu,
req_pool_indices=req_pool_indices,
padded_num_tokens=out_cache_loc.shape[0],
)
core_attn_metadata = self.make_core_attn_metadata(
req_to_token=self.req_to_token,
req_pool_indices_repeated=req_pool_indices_repeated,
Expand All @@ -559,6 +592,20 @@ def init_forward_metadata_prefill(
)
if not need_compress:
create = _create_dummy_paged_compress_data
elif compress_gpu_plan:
create = functools.partial(
create_paged_compressor_data,
is_prefill=True,
token_to_kv_pool=self.token_to_kv_pool,
req_to_token=self.req_to_token,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_cpu=None,
extend_lens=extend_seq_lens,
extend_lens_cpu=None,
num_q_tokens=num_tokens,
use_prefill_cuda_graph=use_prefill_cuda_graph,
)
else:
create = functools.partial(
create_paged_compressor_data,
Expand Down Expand Up @@ -588,6 +635,7 @@ def init_forward_metadata_target_verify(
extend_seq_lens: Optional[torch.Tensor] = None,
use_prefill_cuda_graph: bool = False,
seq_lens_cpu: Optional[List[int]] = None,
ragged_layout=None,
) -> Union[DSV4Metadata, DSV4RawVerifyMetadata]:
# HIP path: build target-verify metadata eagerly even when
# SGLANG_PREP_IN_CUDA_GRAPH is enabled. The raw/lazy-upgrade route can
Expand All @@ -601,6 +649,7 @@ def init_forward_metadata_target_verify(
seq_lens_cpu=seq_lens_cpu,
out_cache_loc=out_cache_loc,
use_prefill_cuda_graph=use_prefill_cuda_graph,
ragged_layout=ragged_layout,
)

def init_forward_metadata_target_verify_old(
Expand All @@ -611,13 +660,38 @@ def init_forward_metadata_target_verify_old(
seq_lens_cpu: Optional[List[int]] = None,
out_cache_loc: Optional[torch.Tensor] = None,
use_prefill_cuda_graph: bool = False,
ragged_layout=None,
) -> DSV4Metadata:
batch_size = len(seq_lens)
seq_lens = seq_lens + self.speculative_num_draft_tokens
seq_lens_cpu = [x + self.speculative_num_draft_tokens for x in seq_lens_cpu]
extend_seq_lens_cpu = [self.speculative_num_draft_tokens] * batch_size
extend_seq_lens = self._move_to_device(extend_seq_lens_cpu)
num_tokens = self.speculative_num_draft_tokens * batch_size
extend_start_loc = None
if ragged_layout is not None:
verify_lens_dev = ragged_layout.verify_lens.to(
device=seq_lens.device, dtype=torch.int32
)
extend_start_loc = ragged_layout.extend_start_loc.to(
device=seq_lens.device, dtype=torch.int32
)
extend_seq_lens = verify_lens_dev
seq_lens = seq_lens + verify_lens_dev.to(seq_lens.dtype)
# Total verify tokens to expand. For the graph path the padded layout
# sets total_verify_tokens == graph_num_tokens (tier); the eager path
# resolves a device-assembled layout whose total_verify_tokens is None,
# so fall back to sum(verify_lens) (== real total; padded == tier).
num_tokens = ragged_layout.total_verify_tokens
if num_tokens is None:
num_tokens = int(verify_lens_dev.sum().item())
else:
num_tokens = int(num_tokens)
extend_seq_lens_cpu = None
seq_lens_cpu = None
else:
seq_lens = seq_lens + self.target_verify_num_draft_tokens
seq_lens_cpu = [
x + self.target_verify_num_draft_tokens for x in seq_lens_cpu
]
extend_seq_lens_cpu = [self.target_verify_num_draft_tokens] * batch_size
num_tokens = self.target_verify_num_draft_tokens * batch_size
extend_seq_lens = self._move_to_device(extend_seq_lens_cpu)
if out_cache_loc is None:
out_cache_loc = seq_lens.new_zeros(num_tokens)
return self.init_forward_metadata_prefill(
Expand All @@ -631,6 +705,8 @@ def init_forward_metadata_target_verify_old(
extend_seq_lens_cpu=extend_seq_lens_cpu,
need_compress=True,
use_prefill_cuda_graph=use_prefill_cuda_graph,
compress_gpu_plan=ragged_layout is not None,
extend_start_loc=extend_start_loc,
)

def make_forward_metadata_from_raw_verify(
Expand All @@ -640,7 +716,7 @@ def make_forward_metadata_from_raw_verify(
seq_lens = raw_metadata.seq_lens
out_cache_loc = raw_metadata.out_cache_loc

bs, num_draft_tokens = len(seq_lens), self.speculative_num_draft_tokens
bs, num_draft_tokens = len(seq_lens), self.target_verify_num_draft_tokens
seq_lens = seq_lens + num_draft_tokens
extend_seq_lens = raw_metadata.extend_seq_lens
if extend_seq_lens is None or extend_seq_lens.numel() != bs:
Expand Down Expand Up @@ -846,6 +922,8 @@ def init_forward_metadata_out_graph(
chosen_max_seq_len = self.MAX_SEQ_LEN_FOR_CAPTURE
assert actual_max_seq_len <= chosen_max_seq_len

graph_key = bs

if bucket == _GraphBucket.DECODE_OR_IDLE:
assert out_cache_loc is not None
assert len(out_cache_loc.shape) == 1, f"{out_cache_loc.shape=}"
Expand All @@ -862,14 +940,14 @@ def init_forward_metadata_out_graph(
out_cache_loc=out_cache_loc_padded,
)
elif bucket == _GraphBucket.TARGET_VERIFY:
if resolve_ragged_verify_layout(forward_batch) is not None:
raise NotImplementedError(
"DSV4 ragged verify is not supported on the HIP backend "
"(DeepseekV4HipRadixBackend) cuda-graph path; disable "
"SGLANG_RAGGED_VERIFY_MODE or use a CUDA device."
)
assert out_cache_loc is not None
num_tokens_v = self.speculative_num_draft_tokens * bs
ragged_layout = resolve_ragged_verify_layout(forward_batch)
if ragged_layout is not None:
ragged_layout = ragged_layout.padded_to_bucket(padded_bs=bs)
num_tokens_v = ragged_layout.graph_num_tokens
graph_key = num_tokens_v
else:
num_tokens_v = self.target_verify_num_draft_tokens * bs
out_cache_loc_padded = torch.nn.functional.pad(
out_cache_loc,
pad=(0, num_tokens_v - len(out_cache_loc)),
Expand All @@ -885,6 +963,7 @@ def init_forward_metadata_out_graph(
# CPU mirror already available here (== seq_lens, no D2H);
# pass it so target_verify skips the per-iter seq_lens.tolist() sync.
seq_lens_cpu=seq_lens_cpu.tolist(),
ragged_layout=ragged_layout,
)
elif bucket == _GraphBucket.DRAFT_EXTEND:
num_tokens_per_req = self.draft_extend_num_tokens_per_req
Expand All @@ -910,7 +989,7 @@ def init_forward_metadata_out_graph(
raise NotImplementedError

self.replay_cuda_graph_metadata_from(
bs=bs, temp_metadata=temp_metadata, bucket=bucket
bs=graph_key, temp_metadata=temp_metadata, bucket=bucket
)

if in_capture:
Expand Down Expand Up @@ -955,12 +1034,7 @@ def init_forward_metadata(self, forward_batch: ForwardBatch) -> None:
out_cache_loc=out_cache_loc,
)
elif forward_batch.forward_mode.is_target_verify():
if resolve_ragged_verify_layout(forward_batch) is not None:
raise NotImplementedError(
"DSV4 ragged verify is not supported on the HIP backend "
"(DeepseekV4HipRadixBackend); disable SGLANG_RAGGED_VERIFY_MODE "
"or use a CUDA device."
)
ragged_layout = resolve_ragged_verify_layout(forward_batch)
metadata = self.init_forward_metadata_target_verify(
max_seq_len=max_seq_len,
req_pool_indices=req_pool_indices,
Expand All @@ -970,6 +1044,7 @@ def init_forward_metadata(self, forward_batch: ForwardBatch) -> None:
seq_lens_cpu=(
seq_lens_cpu.tolist() if seq_lens_cpu is not None else None
),
ragged_layout=ragged_layout,
)
elif forward_batch.forward_mode.is_prefill(include_draft_extend_v2=True):
extend_seq_lens_cpu = forward_batch.extend_seq_lens_cpu
Expand Down
28 changes: 28 additions & 0 deletions python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py
Original file line number Diff line number Diff line change
Expand Up @@ -1192,6 +1192,34 @@ def set_swa_key_buffer_radix_fused_norm_rope(
page_size=self.swa_kv_pool.page_size,
)

def set_unified_key_buffer_radix_fused_norm_rope(
self,
layer_id: int,
swa_loc: torch.Tensor,
kv: torch.Tensor,
kv_weight: torch.Tensor,
eps: float,
freqs_cis: torch.Tensor,
positions: torch.Tensor,
) -> None:
"""unified_kv counterpart of set_swa_key_buffer_radix_fused_norm_rope.

Under unified_kv the (fp8, paged) swa_kv_pool is None -- SWA K lives in
the shared bf16 unified_kv ring instead. Norm+RoPE the draft KV in place
(the same freqs_cis path the main model uses via _compute_kv_bf16) and
scatter it into ``unified_kv[swa_loc]``. Rows with swa_loc < 0
(uncommitted verify tokens) are skipped by the scatter.
"""
from sglang.kernels.ops.attention.dsv4 import fused_norm_rope_inplace
from sglang.kernels.ops.attention.dsv4.unified_kv_kernels import runtime

fused_norm_rope_inplace(kv, kv_weight, eps, freqs_cis, positions)
runtime.scatter_bf16_into_unified(
kv=kv,
loc=swa_loc,
unified_kv=self.get_unified_kv(layer_id),
)

def set_extra_key_buffer_fused(
self,
layer_id: int,
Expand Down
Loading
Loading