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
68 changes: 68 additions & 0 deletions python/sglang/kernels/ops/speculative/cache_locs.py
Original file line number Diff line number Diff line change
Expand Up @@ -258,6 +258,74 @@ def filter_finished_cache_loc_kernel(
)


@triton.jit
def rebuild_compact_draft_req_to_token(
draft_req_to_token,
target_req_to_token,
req_pool_indices,
suffix_start,
draft_prefix_lens,
verify_out_cache_loc,
verify_loc_stride,
draft_pool_len: tl.constexpr,
target_pool_len: tl.constexpr,
block_size: tl.constexpr,
):
"""Rebuild one request's draft-local compact req->token row in a single pass.

Row layout written: [0, prefix_len) = the committed target suffix window
(target_req_to_token[req, suffix_start : suffix_start + prefix_len]) and
[prefix_len, prefix_len + block_size) = the verify block slots. Fixed grid,
per-row data-dependent loop bound; no host reads, so the caller never syncs.
"""
BLOCK: tl.constexpr = 256
pid = tl.program_id(axis=0)
req = tl.load(req_pool_indices + pid).to(tl.int64)
start = tl.load(suffix_start + pid).to(tl.int64)
prefix_len = tl.load(draft_prefix_lens + pid).to(tl.int64)
total = prefix_len + block_size

src_row = target_req_to_token + req * target_pool_len
dst_row = draft_req_to_token + req * draft_pool_len
verify_row = verify_out_cache_loc + pid * verify_loc_stride

offs = tl.arange(0, BLOCK).to(tl.int64)
num_loop = tl.cdiv(total, BLOCK)
for i in range(num_loop):
col = offs + i * BLOCK
in_prefix = col < prefix_len
in_block = (col >= prefix_len) & (col < total)
src = tl.load(src_row + start + col, mask=in_prefix, other=0)
blk = tl.load(verify_row + (col - prefix_len), mask=in_block, other=0)
val = tl.where(in_prefix, src, blk)
tl.store(dst_row + col, val, mask=in_prefix | in_block)


def rebuild_compact_draft_req_to_token_func(
*,
draft_req_to_token: torch.Tensor,
target_req_to_token: torch.Tensor,
req_pool_indices: torch.Tensor,
suffix_start: torch.Tensor,
draft_prefix_lens: torch.Tensor,
verify_out_cache_loc_2d: torch.Tensor,
batch_size: int,
block_size: int,
) -> None:
rebuild_compact_draft_req_to_token[(batch_size,)](
draft_req_to_token,
target_req_to_token,
req_pool_indices,
suffix_start,
draft_prefix_lens,
verify_out_cache_loc_2d,
verify_out_cache_loc_2d.stride(0),
draft_req_to_token.shape[1],
target_req_to_token.shape[1],
block_size,
)


@triton.jit
def assign_extend_cache_locs(
req_pool_indices,
Expand Down
2 changes: 2 additions & 0 deletions python/sglang/srt/environ.py
Original file line number Diff line number Diff line change
Expand Up @@ -745,6 +745,8 @@ class Envs:

# Spec Config
SGLANG_SPEC_ENABLE_STRICT_FILTER_CHECK = EnvBool(True)
# A/B: keep the DFLASH draft greedy head eager (not folded in-graph).
SGLANG_DFLASH_EAGER_DRAFT_SAMPLER = EnvBool(False)
SGLANG_RAGGED_VERIFY_MODE = EnvStr("static")
SGLANG_DSPARK_CONFIDENCE_RELAY_LAG_STEPS = EnvInt(2)
SGLANG_TEST_RAGGED_VERIFY_FORCE_UNIFORM_CAPTURE = EnvBool(False)
Expand Down
6 changes: 6 additions & 0 deletions python/sglang/srt/layers/attention/hybrid_attn_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,12 @@ def __init__(
self.spec_attn_is_prefill = (
model_runner.server_args.speculative_attention_mode == "prefill"
)
# decide_needs_cpu_seq_lens ORs this flag across backends; without the
# delegation the base-class default (True) forces a per-step seq_lens
# D2H + host sync even when both sub-backends opted out.
self.needs_cpu_seq_lens = (
prefill_backend.needs_cpu_seq_lens or decode_backend.needs_cpu_seq_lens
)

def _select_backend(self, forward_mode: ForwardMode) -> AttentionBackend:
"""
Expand Down
11 changes: 7 additions & 4 deletions python/sglang/srt/layers/attention/trtllm_mla_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -384,8 +384,9 @@ def _apply_cuda_graph_metadata(
metadata = self.decode_cuda_graph_metadata[bs]

if forward_mode.is_target_verify():
seq_lens = seq_lens[:bs] + self.num_draft_tokens
metadata.seq_lens_k.copy_(seq_lens)
# Intentional int64 -> int32 same-kind out= downcast.
torch.add(seq_lens[:bs], self.num_draft_tokens, out=metadata.seq_lens_k)
seq_lens = metadata.seq_lens_k
elif forward_mode.is_draft_extend_v2():
num_tokens_per_req = self.num_draft_tokens
metadata.max_seq_len_q = num_tokens_per_req
Expand Down Expand Up @@ -514,11 +515,13 @@ def init_forward_metadata(self, forward_batch: ForwardBatch):
or forward_batch.forward_mode.is_draft_extend_v2()
):
self.forward_prefill_metadata = None
# Get maximum sequence length.
# Never read max_seq from the GPU tensor (.max().item() blocks the
# host on the stream backlog); max_seq only sizes the block table /
# scheduling hint, so the static context bound is a safe fallback.
if getattr(forward_batch, "seq_lens_cpu", None) is not None:
max_seq = forward_batch.seq_lens_cpu.max().item()
else:
max_seq = forward_batch.seq_lens.max().item()
max_seq = self.max_context_len

seq_lens = forward_batch.seq_lens

Expand Down
1 change: 1 addition & 0 deletions python/sglang/srt/managers/schedule_batch.py
Original file line number Diff line number Diff line change
Expand Up @@ -2929,6 +2929,7 @@ def filter_batch(
self.spec_info.filter_batch(
new_indices=keep_indices_device,
has_been_filtered=False,
new_indices_cpu=keep_indices,
)

def merge_batch(self, other: ScheduleBatch):
Expand Down
16 changes: 13 additions & 3 deletions python/sglang/srt/speculative/dflash_info_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@

import contextlib
from dataclasses import dataclass
from typing import Optional
from typing import List, Optional

import torch

Expand Down Expand Up @@ -212,9 +212,19 @@ def prepare_for_decode(self, batch: ScheduleBatch):
self.reserved_seq_lens_cpu = nxt_kv_lens_cpu_t
self.reserved_seq_lens_sum = reserved_seq_lens_sum

def filter_batch(self, new_indices: torch.Tensor, has_been_filtered: bool = True):
def filter_batch(
self,
new_indices: torch.Tensor,
has_been_filtered: bool = True,
new_indices_cpu: Optional[List[int]] = None,
):
if self.reserved_seq_lens_cpu is not None:
self.reserved_seq_lens_cpu = self.reserved_seq_lens_cpu[new_indices.cpu()]
if new_indices_cpu is not None:
self.reserved_seq_lens_cpu = self.reserved_seq_lens_cpu[new_indices_cpu]
else:
self.reserved_seq_lens_cpu = self.reserved_seq_lens_cpu[
new_indices.cpu()
]
self.reserved_seq_lens_sum = int(self.reserved_seq_lens_cpu.sum().item())

if self.future_indices is not None:
Expand Down
Loading
Loading