[Feature] Spec-Overlap supporting DP-ATTN; PD-Disaggregation; npugraph mode (#12443)

This commit is contained in:
Even Zhou
2025-11-15 21:51:07 +08:00
committed by GitHub
parent 0d41ddfbd0
commit 2aec8b6e1b
24 changed files with 657 additions and 366 deletions
@@ -73,6 +73,6 @@ jobs:
push: ${{ github.repository == 'sgl-project/sglang' && github.event_name != 'pull_request' }} push: ${{ github.repository == 'sgl-project/sglang' && github.event_name != 'pull_request' }}
provenance: false provenance: false
build-args: | build-args: |
SGLANG_KERNEL_NPU_TAG=20251030 SGLANG_KERNEL_NPU_TAG=20251110
CANN_VERSION=${{ matrix.cann_version }} CANN_VERSION=${{ matrix.cann_version }}
DEVICE_TYPE=${{ matrix.device_type }} DEVICE_TYPE=${{ matrix.device_type }}
+1 -1
View File
@@ -69,6 +69,6 @@ jobs:
push: ${{ github.repository == 'sgl-project/sglang' && github.event_name != 'pull_request' }} push: ${{ github.repository == 'sgl-project/sglang' && github.event_name != 'pull_request' }}
provenance: false provenance: false
build-args: | build-args: |
SGLANG_KERNEL_NPU_TAG=20251030 SGLANG_KERNEL_NPU_TAG=20251110
CANN_VERSION=${{ matrix.cann_version }} CANN_VERSION=${{ matrix.cann_version }}
DEVICE_TYPE=${{ matrix.device_type }} DEVICE_TYPE=${{ matrix.device_type }}
+1 -1
View File
@@ -930,7 +930,7 @@ class SchedulerDisaggregationDecodeMixin:
# construct fake completed prefill # construct fake completed prefill
new_batch.prepare_for_prebuilt() new_batch.prepare_for_prebuilt()
new_batch.process_prebuilt(self.server_args, self.model_config) new_batch.process_prebuilt(self.server_args, self.future_map)
return new_batch return new_batch
@@ -14,7 +14,7 @@ from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.configs.model_config import ModelConfig from sglang.srt.managers.overlap_utils import FutureMap
from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
@@ -102,7 +102,9 @@ class ScheduleBatchDisaggregationDecodeMixin:
) )
def process_prebuilt( def process_prebuilt(
self: ScheduleBatch, server_args: ServerArgs, model_config: ModelConfig self: ScheduleBatch,
server_args: ServerArgs,
future_map: FutureMap,
): ):
"""Assign the buffered last input id to schedule batch""" """Assign the buffered last input id to schedule batch"""
self.output_ids = [] self.output_ids = []
@@ -166,7 +168,16 @@ class ScheduleBatchDisaggregationDecodeMixin:
topk_index=topk_index, topk_index=topk_index,
hidden_states=hidden_states, hidden_states=hidden_states,
verified_id=self.output_ids, verified_id=self.output_ids,
new_seq_lens=self.seq_lens,
allocate_lens=self.seq_lens,
) )
spec_info.prepare_for_extend(self) spec_info.prepare_for_extend(self)
spec_info.capture_hidden_mode = CaptureHiddenMode.LAST spec_info.capture_hidden_mode = CaptureHiddenMode.LAST
if self.enable_overlap:
spec_info.future_indices = future_map.alloc_future_indices(
len(self.seq_lens)
)
future_map.store_to_map_for_new_batch(
spec_info.future_indices, spec_info
)
self.spec_info = spec_info self.spec_info = spec_info
@@ -161,8 +161,24 @@ class AscendAttnBackend(AttentionBackend):
metadata.block_tables = self.graph_metadata["block_tables"][:bs, :] metadata.block_tables = self.graph_metadata["block_tables"][:bs, :]
metadata.seq_lens_cpu_list = seq_lens.cpu().int().tolist() metadata.seq_lens_cpu_list = seq_lens.cpu().int().tolist()
metadata.seq_lens = seq_lens metadata.seq_lens = seq_lens
if (
forward_mode.is_target_verify()
or forward_mode.is_draft_extend_v2()
or forward_mode.is_draft_extend()
):
metadata.actual_seq_lengths_q = torch.arange(
self.speculative_num_draft_tokens,
self.speculative_num_draft_tokens
+ bs * self.speculative_num_draft_tokens,
self.speculative_num_draft_tokens,
dtype=torch.int32,
device=seq_lens.device,
)
else:
metadata.actual_seq_lengths_q = torch.tensor( metadata.actual_seq_lengths_q = torch.tensor(
[1 + i * 1 for i in range(bs)], dtype=torch.int32, device=seq_lens.device [1 + i * 1 for i in range(bs)],
dtype=torch.int32,
device=seq_lens.device,
) )
self.graph_metadata[bs] = metadata self.graph_metadata[bs] = metadata
@@ -193,7 +209,8 @@ class AscendAttnBackend(AttentionBackend):
) )
metadata.block_tables[:bs, max_seq_pages:].fill_(0) metadata.block_tables[:bs, max_seq_pages:].fill_(0)
metadata.block_tables[bs:, :].fill_(0) metadata.block_tables[bs:, :].fill_(0)
if forward_mode.is_target_verify():
seq_lens = seq_lens + self.speculative_num_draft_tokens
metadata.seq_lens[:bs].copy_(seq_lens[:bs]) metadata.seq_lens[:bs].copy_(seq_lens[:bs])
self.forward_metadata = metadata self.forward_metadata = metadata
@@ -217,7 +234,12 @@ class AscendAttnBackend(AttentionBackend):
topk_indices: torch.Tensor = None, topk_indices: torch.Tensor = None,
): ):
is_prefill = forward_batch.forward_mode.is_extend() is_prefill = (
forward_batch.forward_mode.is_extend()
and not forward_batch.forward_mode.is_draft_extend_v2()
and not forward_batch.forward_mode.is_draft_extend()
and not forward_batch.forward_mode.is_target_verify()
)
if save_kv_cache: if save_kv_cache:
k = k.view(-1, layer.tp_k_head_num, self.kv_lora_rank) k = k.view(-1, layer.tp_k_head_num, self.kv_lora_rank)
@@ -232,6 +254,27 @@ class AscendAttnBackend(AttentionBackend):
actual_seq_qlen = torch.cumsum(forward_batch.seq_lens, dim=0) actual_seq_qlen = torch.cumsum(forward_batch.seq_lens, dim=0)
else: else:
if self.forward_metadata.actual_seq_lengths_q is None: if self.forward_metadata.actual_seq_lengths_q is None:
if (
forward_batch.forward_mode.is_draft_extend_v2()
or forward_batch.forward_mode.is_target_verify()
):
actual_seq_qlen = (
torch.arange(
self.speculative_num_draft_tokens,
self.speculative_num_draft_tokens + q.shape[0],
self.speculative_num_draft_tokens,
dtype=torch.int32,
)
.to(q.device)
.to(torch.int32)
)
elif forward_batch.forward_mode.is_draft_extend():
actual_seq_qlen = (
forward_batch.extend_seq_lens.cumsum()
.to(q.device)
.to(torch.int32)
)
else:
actual_seq_qlen = ( actual_seq_qlen = (
torch.arange(1, q.shape[0] + 1).to(q.device).to(torch.int32) torch.arange(1, q.shape[0] + 1).to(q.device).to(torch.int32)
) )
@@ -477,7 +520,7 @@ class AscendAttnBackend(AttentionBackend):
-1, layer.tp_v_head_num, self.page_size, self.kv_lora_rank -1, layer.tp_v_head_num, self.page_size, self.kv_lora_rank
) )
q_nope = q.view(-1, layer.tp_q_head_num, self.kv_lora_rank) q_nope = q.view(-1, layer.tp_q_head_num, self.kv_lora_rank).contiguous()
q_rope = q_rope.view(-1, layer.tp_q_head_num, self.qk_rope_head_dim) q_rope = q_rope.view(-1, layer.tp_q_head_num, self.qk_rope_head_dim)
if not self.graph_mode: if not self.graph_mode:
num_token_padding = q.shape[0] num_token_padding = q.shape[0]
@@ -919,7 +962,7 @@ class AscendAttnMultiStepDraftBackend:
encoder_lens=None, encoder_lens=None,
forward_mode=ForwardMode.DECODE, forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info, spec_info=forward_batch.spec_info,
seq_lens_cpu=None, seq_lens_cpu=forward_batch.seq_lens_cpu,
) )
self.common_template(forward_batch, call_fn) self.common_template(forward_batch, call_fn)
@@ -699,7 +699,12 @@ class Indexer(CustomOp):
enable_index_cp = ( enable_index_cp = (
get_bool_env_var("SGLANG_USE_AG_AFTER_QLORA") and layer_id >= 4 get_bool_env_var("SGLANG_USE_AG_AFTER_QLORA") and layer_id >= 4
) )
is_prefill = forward_batch.forward_mode.is_extend() is_prefill = (
forward_batch.forward_mode.is_extend()
and not forward_batch.forward_mode.is_draft_extend_v2()
and not forward_batch.forward_mode.is_target_verify()
and not forward_batch.forward_mode.is_draft_extend()
)
attention_tp_rank = get_attention_tp_rank() attention_tp_rank = get_attention_tp_rank()
attention_tp_size = get_attention_tp_size() attention_tp_size = get_attention_tp_size()
@@ -790,8 +795,26 @@ class Indexer(CustomOp):
else: else:
if forward_batch.attn_backend.forward_metadata.actual_seq_lengths_q is None: if forward_batch.attn_backend.forward_metadata.actual_seq_lengths_q is None:
if (
forward_batch.forward_mode.is_draft_extend_v2()
or forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend()
):
num_draft_tokens = (
forward_batch.attn_backend.speculative_num_draft_tokens
)
actual_seq_lengths_q = torch.arange(
num_draft_tokens,
num_draft_tokens + bs,
num_draft_tokens,
dtype=torch.int32,
device=k.device,
)
else:
actual_seq_lengths_q = torch.tensor( actual_seq_lengths_q = torch.tensor(
[1 + i * 1 for i in range(bs)], dtype=torch.int32, device=k.device [1 + i * 1 for i in range(bs)],
dtype=torch.int32,
device=k.device,
) )
else: else:
actual_seq_lengths_q = ( actual_seq_lengths_q = (
+19 -3
View File
@@ -114,17 +114,33 @@ class FutureMap:
else: else:
_resolve_future_token_ids(model_worker_batch.input_ids, self.token_ids_buf) _resolve_future_token_ids(model_worker_batch.input_ids, self.token_ids_buf)
def is_empty_slice(self, s: slice) -> bool:
start, stop, step = s.indices(self.future_buffer_len)
if step > 0:
return start >= stop
else:
return start <= stop
def store_to_map( def store_to_map(
self, future_indices: FutureIndices, batch_result: GenerationBatchResult self, future_indices: FutureIndices, batch_result: GenerationBatchResult
): ):
intv = future_indices.interval
if self.spec_algo.is_eagle(): if self.spec_algo.is_eagle():
draft_input: EagleDraftInput = batch_result.next_draft_input draft_input: EagleDraftInput = batch_result.next_draft_input
self.store_to_map_for_new_batch(future_indices, draft_input)
else:
intv = future_indices.interval
self.token_ids_buf[intv] = batch_result.next_token_ids
def store_to_map_for_new_batch(
self, future_indices: FutureIndices, draft_input: EagleDraftInput
):
intv = future_indices.interval
# idle indices do not need store info
if self.is_empty_slice(intv):
return
self._lazy_init_buf(draft_input) self._lazy_init_buf(draft_input)
self.topk_p_buf[intv] = draft_input.topk_p self.topk_p_buf[intv] = draft_input.topk_p
self.topk_index_buf[intv] = draft_input.topk_index self.topk_index_buf[intv] = draft_input.topk_index
self.hidden_states_buf[intv] = draft_input.hidden_states self.hidden_states_buf[intv] = draft_input.hidden_states
self.verified_id_buf[intv] = draft_input.verified_id self.verified_id_buf[intv] = draft_input.verified_id
self.new_seq_lens_buf[intv] = draft_input.new_seq_lens self.new_seq_lens_buf[intv] = draft_input.new_seq_lens
else:
self.token_ids_buf[intv] = batch_result.next_token_ids
+13 -10
View File
@@ -846,6 +846,15 @@ class Scheduler(
self.server_args.disaggregation_transfer_backend self.server_args.disaggregation_transfer_backend
) )
if self.draft_worker is None or self.spec_algorithm.is_ngram():
draft_token_to_kv_pool = None
elif self.spec_algorithm.is_eagle() and self.enable_overlap:
draft_token_to_kv_pool = (
self.draft_worker.draft_worker.draft_runner.token_to_kv_pool
)
else:
draft_token_to_kv_pool = self.draft_worker.model_runner.token_to_kv_pool
if ( if (
self.disaggregation_mode == DisaggregationMode.DECODE self.disaggregation_mode == DisaggregationMode.DECODE
): # *2 for the headroom. ): # *2 for the headroom.
@@ -874,11 +883,7 @@ class Scheduler(
self.disagg_decode_prealloc_queue = DecodePreallocQueue( self.disagg_decode_prealloc_queue = DecodePreallocQueue(
req_to_token_pool=self.req_to_token_pool, req_to_token_pool=self.req_to_token_pool,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator, token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
draft_token_to_kv_pool=( draft_token_to_kv_pool=draft_token_to_kv_pool,
None
if self.draft_worker is None or self.spec_algorithm.is_ngram()
else self.draft_worker.model_runner.token_to_kv_pool
),
req_to_metadata_buffer_idx_allocator=self.req_to_metadata_buffer_idx_allocator, req_to_metadata_buffer_idx_allocator=self.req_to_metadata_buffer_idx_allocator,
metadata_buffers=self.disagg_metadata_buffers, metadata_buffers=self.disagg_metadata_buffers,
scheduler=self, scheduler=self,
@@ -911,11 +916,7 @@ class Scheduler(
self.disagg_prefill_bootstrap_queue = PrefillBootstrapQueue( self.disagg_prefill_bootstrap_queue = PrefillBootstrapQueue(
token_to_kv_pool=self.token_to_kv_pool_allocator.get_kvcache(), token_to_kv_pool=self.token_to_kv_pool_allocator.get_kvcache(),
draft_token_to_kv_pool=( draft_token_to_kv_pool=draft_token_to_kv_pool,
None
if self.draft_worker is None or self.spec_algorithm.is_ngram()
else self.draft_worker.model_runner.token_to_kv_pool
),
req_to_metadata_buffer_idx_allocator=self.req_to_metadata_buffer_idx_allocator, req_to_metadata_buffer_idx_allocator=self.req_to_metadata_buffer_idx_allocator,
metadata_buffers=self.disagg_metadata_buffers, metadata_buffers=self.disagg_metadata_buffers,
tp_rank=self.tp_rank, tp_rank=self.tp_rank,
@@ -935,6 +936,8 @@ class Scheduler(
self.disagg_prefill_inflight_queue: List[Req] = [] self.disagg_prefill_inflight_queue: List[Req] = []
def init_overlap(self): def init_overlap(self):
self.future_map = None
if not self.enable_overlap: if not self.enable_overlap:
return return
@@ -866,6 +866,10 @@ class ForwardBatch:
self.spec_info.accept_length = self.spec_info.accept_length[:bs] self.spec_info.accept_length = self.spec_info.accept_length[:bs]
logits_output.next_token_logits = logits_output.next_token_logits[:bs] logits_output.next_token_logits = logits_output.next_token_logits[:bs]
logits_output.hidden_states = logits_output.hidden_states[:bs] logits_output.hidden_states = logits_output.hidden_states[:bs]
elif self.forward_mode.is_draft_extend_v2(): # draft extend_v2
bs = bs * self.spec_info.num_tokens_per_batch
logits_output.next_token_logits = logits_output.next_token_logits[:bs]
logits_output.hidden_states = logits_output.hidden_states[:bs]
elif self.forward_mode.is_extend() or self.forward_mode.is_idle(): elif self.forward_mode.is_extend() or self.forward_mode.is_idle():
logits_output.next_token_logits = logits_output.next_token_logits[:bs] logits_output.next_token_logits = logits_output.next_token_logits[:bs]
logits_output.hidden_states = logits_output.hidden_states[:bs] logits_output.hidden_states = logits_output.hidden_states[:bs]
@@ -2236,7 +2236,7 @@ class ModelRunner:
reinit_attn_backend=reinit_attn_backend, reinit_attn_backend=reinit_attn_backend,
forward_count=split_forward_count, forward_count=split_forward_count,
) )
elif forward_batch.forward_mode.is_extend(): elif forward_batch.forward_mode.is_extend(include_draft_extend_v2=True):
ret = self.forward_extend( ret = self.forward_extend(
forward_batch, forward_batch,
skip_attn_backend_init=skip_attn_backend_init, skip_attn_backend_init=skip_attn_backend_init,
@@ -89,7 +89,7 @@ class EAGLEDraftCudaGraphRunner:
set_torch_compile_config() set_torch_compile_config()
# Graph inputs # Graph inputs
with torch.device("cuda"): with torch.device(model_runner.device):
self.input_ids = torch.zeros((self.max_num_token,), dtype=torch.int64) self.input_ids = torch.zeros((self.max_num_token,), dtype=torch.int64)
self.req_pool_indices = torch.zeros((self.max_bs,), dtype=torch.int32) self.req_pool_indices = torch.zeros((self.max_bs,), dtype=torch.int32)
self.seq_lens = torch.full( self.seq_lens = torch.full(
@@ -158,13 +158,30 @@ class EAGLEDraftCudaGraphRunner:
return is_bs_supported return is_bs_supported
def _create_graph(self):
return torch.cuda.CUDAGraph()
def _capture_init(self, run_once_fn):
for _ in range(2):
torch.cuda.synchronize()
self.model_runner.tp_group.barrier()
run_once_fn()
def _capture_graph(self, graph, pool, stream, run_once_fn):
with torch.cuda.graph(graph, pool=pool, stream=stream):
out = run_once_fn()
return out
def _replay(self, forward_batch: ForwardBatch):
self.graphs[self.bs].replay()
def capture(self): def capture(self):
CudaGraphRunner.capture(self) CudaGraphRunner.capture(self)
def capture_one_batch_size( def capture_one_batch_size(
self, num_seqs: int, forward: Callable, stream_idx: int = 0 self, num_seqs: int, forward: Callable, stream_idx: int = 0
): ):
graph = torch.cuda.CUDAGraph() graph = self._create_graph()
stream = self.stream stream = self.stream
num_tokens = num_seqs * self.num_tokens_per_bs num_tokens = num_seqs * self.num_tokens_per_bs
@@ -285,16 +302,10 @@ class EAGLEDraftCudaGraphRunner:
self.deepep_adapter.capture(is_extend_in_batch=False) self.deepep_adapter.capture(is_extend_in_batch=False)
for _ in range(2): self._capture_init(run_once)
torch.cuda.synchronize() out = self._capture_graph(
self.model_runner.tp_group.barrier() graph, get_global_graph_memory_pool(), stream, run_once
)
run_once()
with torch.cuda.graph(
graph, pool=get_global_graph_memory_pool(), stream=stream
):
out = run_once()
set_global_graph_memory_pool(graph.pool()) set_global_graph_memory_pool(graph.pool())
return graph, out return graph, out
@@ -362,10 +373,12 @@ class EAGLEDraftCudaGraphRunner:
self.model_runner.draft_attn_backend.init_forward_metadata_replay_cuda_graph( self.model_runner.draft_attn_backend.init_forward_metadata_replay_cuda_graph(
forward_batch, bs forward_batch, bs
) )
self.raw_bs = raw_bs
self.bs = bs
# TODO: The forward_batch.seq_len_sum might need to be updated to reflect the padding in the cuda graph # TODO: The forward_batch.seq_len_sum might need to be updated to reflect the padding in the cuda graph
# Replay # Replay
self.graphs[bs].replay() self._replay(forward_batch)
out = self.output_buffers[bs] out = self.output_buffers[bs]
if bs != raw_bs: if bs != raw_bs:
@@ -43,8 +43,10 @@ class EAGLEDraftExtendCudaGraphRunner:
if not hasattr(eagle_worker, "model_runner"): if not hasattr(eagle_worker, "model_runner"):
# V2: EagleDraftWorker # V2: EagleDraftWorker
self.model_runner = model_runner = eagle_worker.draft_runner self.model_runner = model_runner = eagle_worker.draft_runner
self.forward_mode = ForwardMode.DRAFT_EXTEND_V2
else: else:
self.model_runner = model_runner = eagle_worker.model_runner self.model_runner = model_runner = eagle_worker.model_runner
self.forward_mode = ForwardMode.DRAFT_EXTEND
self.graphs = {} self.graphs = {}
self.output_buffers = {} self.output_buffers = {}
@@ -86,7 +88,7 @@ class EAGLEDraftExtendCudaGraphRunner:
set_torch_compile_config() set_torch_compile_config()
# Graph inputs # Graph inputs
with torch.device("cuda"): with torch.device(model_runner.device):
self.input_ids = torch.zeros((self.max_num_token,), dtype=torch.int64) self.input_ids = torch.zeros((self.max_num_token,), dtype=torch.int64)
self.req_pool_indices = torch.zeros((self.max_bs,), dtype=torch.int32) self.req_pool_indices = torch.zeros((self.max_bs,), dtype=torch.int32)
self.out_cache_loc = torch.ones((self.max_num_token,), dtype=torch.int64) self.out_cache_loc = torch.ones((self.max_num_token,), dtype=torch.int64)
@@ -116,8 +118,12 @@ class EAGLEDraftExtendCudaGraphRunner:
(self.max_num_token, self.model_runner.model_config.hidden_size), (self.max_num_token, self.model_runner.model_config.hidden_size),
dtype=self.model_runner.dtype, dtype=self.model_runner.dtype,
) )
self.seq_len_fill_value = (
self.seq_lens = torch.ones((self.max_bs,), dtype=torch.int32) self.model_runner.attn_backend.get_cuda_graph_seq_len_fill_value()
)
self.seq_lens = torch.full(
(self.max_bs,), self.seq_len_fill_value, dtype=torch.int32
)
self.extend_seq_lens = torch.ones((self.max_bs,), dtype=torch.int32) self.extend_seq_lens = torch.ones((self.max_bs,), dtype=torch.int32)
self.accept_length = torch.full( self.accept_length = torch.full(
(self.max_bs,), self.num_tokens_per_bs, dtype=torch.int32 (self.max_bs,), self.num_tokens_per_bs, dtype=torch.int32
@@ -153,7 +159,14 @@ class EAGLEDraftExtendCudaGraphRunner:
vocab_size = self.model_runner.model_config.vocab_size vocab_size = self.model_runner.model_config.vocab_size
self.next_token_logits_buffer = torch.zeros( self.next_token_logits_buffer = torch.zeros(
(self.max_bs, vocab_size), (
(
self.max_bs * self.num_tokens_per_bs
if self.forward_mode == ForwardMode.DRAFT_EXTEND_V2
else self.max_bs
),
vocab_size,
),
dtype=torch.float, dtype=torch.float,
) )
@@ -187,11 +200,28 @@ class EAGLEDraftExtendCudaGraphRunner:
return is_bs_supported return is_bs_supported
def _create_graph(self):
return torch.cuda.CUDAGraph()
def _capture_init(self, run_once_fn):
for _ in range(2):
torch.cuda.synchronize()
self.model_runner.tp_group.barrier()
run_once_fn()
def _capture_graph(self, graph, pool, stream, run_once_fn):
with torch.cuda.graph(graph, pool=pool, stream=stream):
out = run_once_fn()
return out
def _replay(self, forward_batch: ForwardBatch):
self.graphs[self.bs].replay()
def capture(self): def capture(self):
CudaGraphRunner.capture(self) CudaGraphRunner.capture(self)
def capture_one_batch_size(self, bs: int, forward: Callable, stream_idx: int = 0): def capture_one_batch_size(self, bs: int, forward: Callable, stream_idx: int = 0):
graph = torch.cuda.CUDAGraph() graph = self._create_graph()
stream = self.stream stream = self.stream
num_tokens = bs * self.num_tokens_per_bs num_tokens = bs * self.num_tokens_per_bs
@@ -207,7 +237,9 @@ class EAGLEDraftExtendCudaGraphRunner:
positions = self.positions[:num_tokens] positions = self.positions[:num_tokens]
mrope_positions = self.mrope_positions[:, :num_tokens] mrope_positions = self.mrope_positions[:, :num_tokens]
hidden_states = self.hidden_states[:num_tokens] hidden_states = self.hidden_states[:num_tokens]
next_token_logits_buffer = self.next_token_logits_buffer[:bs] next_token_logits_buffer = self.next_token_logits_buffer[
: bs if self.forward_mode == ForwardMode.DRAFT_EXTEND else num_tokens
]
if self.require_mlp_tp_gather: if self.require_mlp_tp_gather:
self.global_num_tokens_gpu.copy_( self.global_num_tokens_gpu.copy_(
@@ -254,7 +286,7 @@ class EAGLEDraftExtendCudaGraphRunner:
# Forward batch # Forward batch
forward_batch = ForwardBatch( forward_batch = ForwardBatch(
forward_mode=ForwardMode.DRAFT_EXTEND, forward_mode=self.forward_mode,
batch_size=bs, batch_size=bs,
input_ids=input_ids, input_ids=input_ids,
req_pool_indices=req_pool_indices, req_pool_indices=req_pool_indices,
@@ -287,7 +319,7 @@ class EAGLEDraftExtendCudaGraphRunner:
req_pool_indices=req_pool_indices, req_pool_indices=req_pool_indices,
seq_lens=seq_lens, seq_lens=seq_lens,
encoder_lens=None, encoder_lens=None,
forward_mode=ForwardMode.DRAFT_EXTEND, forward_mode=self.forward_mode,
spec_info=spec_info, spec_info=spec_info,
) )
@@ -318,16 +350,11 @@ class EAGLEDraftExtendCudaGraphRunner:
forward_batch.spec_info.hidden_states = hidden_states_backup forward_batch.spec_info.hidden_states = hidden_states_backup
return ret return ret
for _ in range(2): self._capture_init(run_once)
torch.cuda.synchronize()
self.model_runner.tp_group.barrier()
run_once() out = self._capture_graph(
graph, get_global_graph_memory_pool(), stream, run_once
with torch.cuda.graph( )
graph, pool=get_global_graph_memory_pool(), stream=stream
):
out = run_once()
set_global_graph_memory_pool(graph.pool()) set_global_graph_memory_pool(graph.pool())
return graph, out return graph, out
@@ -399,21 +426,32 @@ class EAGLEDraftExtendCudaGraphRunner:
seq_lens_sum=forward_batch.seq_lens_sum seq_lens_sum=forward_batch.seq_lens_sum
+ (bs - raw_bs) * self.seq_len_fill_value, + (bs - raw_bs) * self.seq_len_fill_value,
encoder_lens=None, encoder_lens=None,
forward_mode=ForwardMode.DRAFT_EXTEND, forward_mode=self.forward_mode,
spec_info=forward_batch.spec_info, spec_info=forward_batch.spec_info,
seq_lens_cpu=self.seq_lens_cpu, seq_lens_cpu=self.seq_lens_cpu,
) )
# Replay # Replay
self.graphs[bs].replay() self.raw_bs = raw_bs
self.bs = bs
self._replay(forward_batch)
out = self.output_buffers[bs] out = self.output_buffers[bs]
if bs != raw_bs:
if self.forward_mode == ForwardMode.DRAFT_EXTEND_V2:
# DRAFT_EXTEND_V2: all tokens calculations whether accepted or not.
unpadding_bs = num_tokens
elif bs != raw_bs:
forward_batch.spec_info.accept_length = self.accept_length[:raw_bs] forward_batch.spec_info.accept_length = self.accept_length[:raw_bs]
unpadding_bs = raw_bs
else:
unpadding_bs = None
if unpadding_bs is not None:
out_copy = out out_copy = out
out = LogitsProcessorOutput( out = LogitsProcessorOutput(
next_token_logits=out.next_token_logits[:raw_bs], next_token_logits=out.next_token_logits[:unpadding_bs],
hidden_states=out.hidden_states[:raw_bs], hidden_states=out.hidden_states[:unpadding_bs],
) )
out.topk_p = out_copy.topk_p[:raw_bs] out.topk_p = out_copy.topk_p[:unpadding_bs]
out.topk_index = out_copy.topk_index[:raw_bs] out.topk_index = out_copy.topk_index[:unpadding_bs]
return out return out
@@ -0,0 +1,68 @@
# Copyright 2024-2025 SGLang Team
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
"""Run the model with npu graph and torch.compile."""
from __future__ import annotations
import threading
from typing import TYPE_CHECKING
import torch
from sglang.srt.configs.model_config import is_deepseek_nsa
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.speculative.eagle_draft_extend_cuda_graph_runner import (
EAGLEDraftExtendCudaGraphRunner,
)
if TYPE_CHECKING:
from sglang.srt.speculative.eagle_worker import EAGLEWorker
class EAGLEDraftExtendNpuGraphRunner(EAGLEDraftExtendCudaGraphRunner):
def __init__(self, eagle_worker: EAGLEWorker):
super().__init__(eagle_worker)
def _create_graph(self):
return torch.npu.NPUGraph()
def _capture_init(self, run_once_fn):
for _ in range(2):
torch.npu.synchronize()
self.model_runner.tp_group.barrier()
run_once_fn()
def _capture_graph(self, graph, pool, stream, run_once_fn):
with torch.npu.graph(
graph, pool=pool, stream=stream, auto_dispatch_capture=True
):
out = run_once_fn()
return out
def _replay_update(self, seq_lens):
self.graphs[self.bs].update(
cpu_update_input=[{"actual_seq_lengths_kv": seq_lens}]
)
def _replay(self, forward_batch: ForwardBatch):
if not is_deepseek_nsa(self.model_runner.model_config.hf_config):
seq_lens = forward_batch.seq_lens_cpu.tolist() + [0] * (
self.bs - self.raw_bs
)
thread = threading.Thread(target=self._replay_update, args=(seq_lens,))
thread.start()
self.graphs[self.bs].replay()
thread.join()
else:
self.graphs[self.bs].replay()
@@ -0,0 +1,81 @@
# Copyright 2025 SGLang Team
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
""" Run the model with npu graph and torch.compile """
from __future__ import annotations
import logging
import threading
from typing import TYPE_CHECKING
import torch
from sglang.srt.configs.model_config import is_deepseek_nsa
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.speculative.eagle_draft_cuda_graph_runner import (
EAGLEDraftCudaGraphRunner,
)
if TYPE_CHECKING:
from sglang.srt.speculative.eagle_worker import EAGLEWorker
from sglang.srt.utils import is_npu
logger = logging.getLogger(__name__)
if is_npu():
torch.cuda.CUDAGraph = torch.npu.NPUGraph
torch.cuda.synchronize = torch.npu.synchronize
torch.cuda.graph = torch.npu.graph
torch.cuda.stream = torch.npu.stream
torch.cuda.Stream = torch.npu.Stream
torch.cuda.current_stream = torch.npu.current_stream
class EAGLEDraftNpuGraphRunner(EAGLEDraftCudaGraphRunner):
def __init__(self, eagle_worker: EAGLEWorker):
super().__init__(eagle_worker)
def _create_graph(self):
return torch.npu.NPUGraph()
def _capture_init(self, run_once_fn):
for _ in range(2):
torch.npu.synchronize()
self.model_runner.tp_group.barrier()
run_once_fn()
def _capture_graph(self, graph, pool, stream, run_once_fn):
with torch.npu.graph(
graph, pool=pool, stream=stream, auto_dispatch_capture=True
):
out = run_once_fn()
return out
def _replay_update(self, seq_lens):
self.graphs[self.bs].update(
cpu_update_input=[{"actual_seq_lengths_kv": seq_lens}]
)
def _replay(self, forward_batch: ForwardBatch):
if not is_deepseek_nsa(self.model_runner.model_config.hf_config):
seq_lens = forward_batch.seq_lens_cpu.tolist() + [0] * (
self.bs - self.raw_bs
)
thread = threading.Thread(target=self._replay_update, args=(seq_lens,))
thread.start()
self.graphs[self.bs].replay()
thread.join()
else:
self.graphs[self.bs].replay()
@@ -665,6 +665,8 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
topk_p=torch.empty((0, topk), device=device, dtype=torch.float32), topk_p=torch.empty((0, topk), device=device, dtype=torch.float32),
topk_index=torch.empty((0, topk), device=device, dtype=torch.int64), topk_index=torch.empty((0, topk), device=device, dtype=torch.int64),
capture_hidden_mode=capture_hidden_mode, capture_hidden_mode=capture_hidden_mode,
allocate_lens=torch.empty((0,), device=device, dtype=torch.int32),
new_seq_lens=torch.empty((0,), device=device, dtype=torch.int32),
accept_length=torch.empty((0,), device=device, dtype=torch.int32), accept_length=torch.empty((0,), device=device, dtype=torch.int32),
accept_length_cpu=[], accept_length_cpu=[],
) )
+28 -2
View File
@@ -142,6 +142,7 @@ class EagleDraftInputV2Mixin:
topk: int, topk: int,
num_steps: int, num_steps: int,
): ):
if not batch.forward_mode.is_idle():
bs = len(batch.seq_lens) bs = len(batch.seq_lens)
# Assign cache locations # Assign cache locations
@@ -162,6 +163,8 @@ class EagleDraftInputV2Mixin:
) )
# Get a forward batch # Get a forward batch
self.num_tokens_per_batch = topk
self.num_tokens_for_logprob_per_batch = topk
batch.capture_hidden_mode = CaptureHiddenMode.LAST batch.capture_hidden_mode = CaptureHiddenMode.LAST
self.positions = batch.seq_lens.repeat_interleave(topk, dim=0) self.positions = batch.seq_lens.repeat_interleave(topk, dim=0)
forward_batch = ForwardBatch.init_new(batch, draft_model_runner) forward_batch = ForwardBatch.init_new(batch, draft_model_runner)
@@ -174,6 +177,7 @@ class EagleDraftInputV2Mixin:
predict: torch.Tensor, predict: torch.Tensor,
num_draft_tokens: int, num_draft_tokens: int,
draft_model_runner: Any, draft_model_runner: Any,
cuda_graph_runner: Any,
): ):
seq_lens_cpu_ = batch.seq_lens_cpu seq_lens_cpu_ = batch.seq_lens_cpu
extend_num_tokens = len(batch.seq_lens) * num_draft_tokens extend_num_tokens = len(batch.seq_lens) * num_draft_tokens
@@ -187,8 +191,14 @@ class EagleDraftInputV2Mixin:
batch.extend_prefix_lens = seq_lens_cpu_.tolist() batch.extend_prefix_lens = seq_lens_cpu_.tolist()
batch.extend_num_tokens = extend_num_tokens batch.extend_num_tokens = extend_num_tokens
batch.capture_hidden_mode = CaptureHiddenMode.FULL batch.capture_hidden_mode = CaptureHiddenMode.FULL
batch.forward_mode = ForwardMode.DRAFT_EXTEND_V2 batch.forward_mode = (
ForwardMode.IDLE
if batch.forward_mode.is_idle()
else ForwardMode.DRAFT_EXTEND_V2
)
forward_batch = ForwardBatch.init_new(batch, draft_model_runner) forward_batch = ForwardBatch.init_new(batch, draft_model_runner)
can_cuda_graph = cuda_graph_runner and cuda_graph_runner.can_run(forward_batch)
if not batch.forward_mode.is_idle() and not can_cuda_graph:
draft_model_runner.attn_backend.init_forward_metadata(forward_batch) draft_model_runner.attn_backend.init_forward_metadata(forward_batch)
return forward_batch return forward_batch
@@ -201,6 +211,7 @@ class EagleVerifyInputV2Mixin:
batch: ModelWorkerBatch, batch: ModelWorkerBatch,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
if not batch.forward_mode.is_idle():
# Assign cache locations # Assign cache locations
bs = len(batch.req_pool_indices) bs = len(batch.req_pool_indices)
batch.input_ids = self.draft_token batch.input_ids = self.draft_token
@@ -216,7 +227,11 @@ class EagleVerifyInputV2Mixin:
) )
# Get a forward batch # Get a forward batch
batch.forward_mode = ForwardMode.TARGET_VERIFY batch.forward_mode = (
ForwardMode.IDLE
if batch.forward_mode.is_idle()
else ForwardMode.TARGET_VERIFY
)
batch.capture_hidden_mode = CaptureHiddenMode.FULL batch.capture_hidden_mode = CaptureHiddenMode.FULL
verify_forward_batch = ForwardBatch.init_new(batch, target_worker.model_runner) verify_forward_batch = ForwardBatch.init_new(batch, target_worker.model_runner)
@@ -228,6 +243,7 @@ class EagleVerifyInputV2Mixin:
if can_run_cuda_graph: if can_run_cuda_graph:
target_worker.model_runner.graph_runner.replay_prepare(verify_forward_batch) target_worker.model_runner.graph_runner.replay_prepare(verify_forward_batch)
else: else:
if not batch.forward_mode.is_idle():
target_worker.model_runner.attn_backend.init_forward_metadata( target_worker.model_runner.attn_backend.init_forward_metadata(
verify_forward_batch verify_forward_batch
) )
@@ -243,6 +259,16 @@ class EagleVerifyInputV2Mixin:
Verify and find accepted tokens based on logits output and batch Verify and find accepted tokens based on logits output and batch
(which contains spec decoding information). (which contains spec decoding information).
""" """
if batch.forward_mode.is_idle():
predict = torch.empty(0, dtype=torch.long, device=batch.input_ids.device)
accept_length = torch.empty(
0, dtype=torch.int32, device=batch.input_ids.device
)
accept_index = torch.empty(
0, dtype=torch.int32, device=batch.input_ids.device
)
return predict, accept_length, accept_index
bs = len(batch.seq_lens) bs = len(batch.seq_lens)
sampling_info = batch.sampling_info sampling_info = batch.sampling_info
next_token_logits = logits_output.next_token_logits next_token_logits = logits_output.next_token_logits
+12 -196
View File
@@ -16,116 +16,6 @@ if _is_cuda or _is_hip:
) )
def build_tree_efficient_native(
parent_list: torch.Tensor,
selected_index: torch.Tensor,
verified_seq_len: torch.Tensor,
tree_mask: torch.Tensor,
retrive_index: torch.Tensor,
retrive_next_token: torch.Tensor,
retrive_next_sibling: torch.Tensor,
topk: int,
draft_token_num: int,
tree_mask_mode: int,
bs: int,
):
# Generate batch and token index ranges
bs_range = torch.arange(bs, device=tree_mask.device).view(-1, 1)
draft_token_num_range = torch.arange(draft_token_num, device=tree_mask.device)
# Optimized common case for performance.
if draft_token_num == 2 and topk == 1 and tree_mask_mode == TreeMaskMode.FULL_MASK:
positions = verified_seq_len.repeat_interleave(draft_token_num)
positions = (positions.view(bs, -1) + draft_token_num_range).view(-1)
retrive_index[:] = bs_range * draft_token_num + draft_token_num_range
retrive_next_token[:, 0] = 1
retrive_next_token[:, 1] = -1
return (
positions,
retrive_index,
retrive_next_token,
retrive_next_sibling,
tree_mask,
)
# Precompute sequence tree indices
draft_token_num_range1 = torch.arange(draft_token_num - 1, device=tree_mask.device)
cum_seq_len = torch.cumsum(verified_seq_len * draft_token_num, dim=0)
cum_seq_len = torch.cat((torch.tensor([0], device=tree_mask.device), cum_seq_len))
cum_seq_len = cum_seq_len[:-1]
seq_tree_idx = (
draft_token_num * draft_token_num * torch.arange(bs, device=tree_mask.device)
+ cum_seq_len
)
# Batch processing for tree mask
if tree_mask_mode == TreeMaskMode.FULL_MASK:
token_tree_base = (
seq_tree_idx.view(-1, 1)
+ (verified_seq_len.view(-1, 1) + draft_token_num) * draft_token_num_range
)
token_tree_indices = token_tree_base + verified_seq_len.view(-1, 1) + 1
else:
token_tree_indices = (
bs_range * draft_token_num**2 + draft_token_num_range * draft_token_num + 1
)
tree_mask[token_tree_indices.flatten() - 1] = True
indices = token_tree_indices.unsqueeze(-1) + draft_token_num_range1.view(1, 1, -1)
tree_mask[indices.view(-1)] = False
positions = verified_seq_len.repeat_interleave(draft_token_num)
parent_tb_indices = selected_index // topk
retrive_index[:] = bs_range * draft_token_num + draft_token_num_range
tree_mask[token_tree_indices.view(-1, 1) + draft_token_num_range1] = True
for bid in range(bs):
for tid in range(draft_token_num):
position = 0
if tid == 0:
# Process root node
for i in range(draft_token_num - 1, 0, -1):
parent_position = 0
parent_tb_idx = parent_tb_indices[bid][i - 1]
if parent_tb_idx > 0:
parent_token_idx = parent_list[bid][parent_tb_idx]
loop_num = draft_token_num - parent_position
for _ in range(loop_num):
if selected_index[bid][parent_position] == parent_token_idx:
parent_position += 1
break
parent_position += 1
if parent_position == draft_token_num:
continue
if retrive_next_token[bid][parent_position] != -1:
retrive_next_sibling[bid][i] = retrive_next_token[bid][
parent_position
]
retrive_next_token[bid][parent_position] = i
else:
# Process no-root nodes
cur_position = tid - 1
while True:
position += 1
if cur_position >= draft_token_num:
tree_mask[token_tree_indices + cur_position] = True
parent_tb_idx = selected_index[bid][cur_position] // topk
else:
parent_tb_idx = parent_tb_indices[bid][cur_position]
if parent_tb_idx == 0:
break
token_idx = parent_list[bid][parent_tb_idx]
cur_position = 0
for _ in range(draft_token_num):
if selected_index[bid][cur_position] == token_idx:
break
cur_position += 1
positions[bid * draft_token_num + tid] += position
return positions, retrive_index, retrive_next_token, retrive_next_sibling, tree_mask
def organize_draft_results( def organize_draft_results(
score_list: List[torch.Tensor], score_list: List[torch.Tensor],
token_list: List[torch.Tensor], token_list: List[torch.Tensor],
@@ -229,24 +119,19 @@ def build_tree_kernel_efficient(
) )
if _is_npu: if _is_npu:
( torch.ops.npu.build_tree_kernel_efficient(
parent_list.to(dtype=torch.int64),
top_scores_index,
seq_lens,
tree_mask,
positions, positions,
retrive_index, retrive_index,
retrive_next_token, retrive_next_token,
retrive_next_sibling, retrive_next_sibling,
tree_mask,
) = build_tree_efficient_native(
parent_list,
top_scores_index,
seq_lens,
tree_mask,
retrive_index,
retrive_next_token,
retrive_next_sibling,
topk, topk,
spec_steps,
num_verify_tokens, num_verify_tokens,
tree_mask_mode, tree_mask_mode,
bs,
) )
else: else:
sgl_build_tree_kernel_efficient( sgl_build_tree_kernel_efficient(
@@ -273,75 +158,6 @@ def build_tree_kernel_efficient(
) )
def verify_tree_greedy_native(
predicts: torch.Tensor,
accept_index: torch.Tensor,
accept_token_num: torch.Tensor,
candidates: torch.Tensor,
retrive_index: torch.Tensor,
retrive_next_token: torch.Tensor,
retrive_next_sibling: torch.Tensor,
target_predict: torch.Tensor,
topk: int = -1,
):
batch_size, num_draft_tokens = candidates.shape
# Optimized common case for performance.
if num_draft_tokens == 2 and accept_index.shape[1] == 2 and topk == 1:
comparison_result = candidates[:, 1] == target_predict[:, 0]
predicts = target_predict.flatten()
accept_index = torch.arange(
0, num_draft_tokens * batch_size, device=candidates.device, dtype=torch.long
).reshape(batch_size, num_draft_tokens)
comparison_result = comparison_result.to(torch.int64)
accept_index_mask = accept_index[:, 1] * comparison_result
accept_index[:, 1] = accept_index_mask - (1 - comparison_result)
accept_token_num = comparison_result.int()
return predicts, accept_index, accept_token_num
# BFS
for bx in range(batch_size):
cur_candidates = candidates[bx]
cur_retrive_index = retrive_index[bx]
cur_next_token = retrive_next_token[bx]
cur_next_sibling = retrive_next_sibling[bx]
cur_target = target_predict[bx]
last_accepted_idx = cur_retrive_index[0]
accept_index[bx, 0] = last_accepted_idx
num_accepted = 0
cur_node = 0
for _ in range(1, num_draft_tokens):
cur_node = cur_next_token[cur_node]
found = False
while cur_node != -1:
draft_idx = cur_retrive_index[cur_node]
draft_token = cur_candidates[cur_node]
target_token = cur_target[last_accepted_idx - num_draft_tokens * bx]
if draft_token == target_token:
predicts[last_accepted_idx] = target_token
num_accepted += 1
accept_index[bx, num_accepted] = draft_idx
last_accepted_idx = draft_idx
found = True
break
else:
cur_node = cur_next_sibling[cur_node]
if not found:
break
accept_token_num[bx] = num_accepted
predicts[last_accepted_idx] = cur_target[
last_accepted_idx - num_draft_tokens * bx
]
return predicts, accept_index, accept_token_num
def verify_tree_greedy_func( def verify_tree_greedy_func(
predicts: torch.Tensor, predicts: torch.Tensor,
accept_index: torch.Tensor, accept_index: torch.Tensor,
@@ -368,16 +184,16 @@ def verify_tree_greedy_func(
) )
elif _is_npu: elif _is_npu:
predicts, accept_index, accept_token_num = verify_tree_greedy_native( from sgl_kernel_npu.sample.verify_tree_greedy import verify_tree_greedy
predicts=predicts, # mutable
accept_index=accept_index, # mutable verify_tree_greedy(
accept_token_num=accept_token_num, # mutable predicts=predicts,
accept_index=accept_index,
accept_token_num=accept_token_num,
candidates=candidates, candidates=candidates,
retrive_index=retrive_index, retrive_index=retrive_index,
retrive_next_token=retrive_next_token, retrive_next_token=retrive_next_token,
retrive_next_sibling=retrive_next_sibling, retrive_next_sibling=retrive_next_sibling,
target_predict=target_predict, target_predict=target_predict,
topk=topk,
) )
return predicts, accept_index, accept_token_num return predicts, accept_index, accept_token_num
+10 -3
View File
@@ -31,6 +31,7 @@ from sglang.srt.speculative.eagle_draft_cuda_graph_runner import (
from sglang.srt.speculative.eagle_draft_extend_cuda_graph_runner import ( from sglang.srt.speculative.eagle_draft_extend_cuda_graph_runner import (
EAGLEDraftExtendCudaGraphRunner, EAGLEDraftExtendCudaGraphRunner,
) )
from sglang.srt.speculative.eagle_draft_npu_graph_runner import EAGLEDraftNpuGraphRunner
from sglang.srt.speculative.eagle_info import ( from sglang.srt.speculative.eagle_info import (
EagleDraftInput, EagleDraftInput,
EagleVerifyInput, EagleVerifyInput,
@@ -214,9 +215,13 @@ class EAGLEWorker(TpModelWorker):
self.cuda_graph_runner = None self.cuda_graph_runner = None
self.cuda_graph_runner_for_draft_extend = None self.cuda_graph_runner_for_draft_extend = None
if self.server_args.disable_cuda_graph or _is_npu: if self.server_args.disable_cuda_graph:
return return
Device2DraftCudaGraphRunner = {
"npu": EAGLEDraftNpuGraphRunner,
"cuda": EAGLEDraftCudaGraphRunner,
}
# Capture draft # Capture draft
if self.speculative_num_steps > 1: if self.speculative_num_steps > 1:
tic = time.perf_counter() tic = time.perf_counter()
@@ -224,14 +229,16 @@ class EAGLEWorker(TpModelWorker):
logger.info( logger.info(
f"Capture draft cuda graph begin. This can take up to several minutes. avail mem={before_mem:.2f} GB" f"Capture draft cuda graph begin. This can take up to several minutes. avail mem={before_mem:.2f} GB"
) )
self.cuda_graph_runner = EAGLEDraftCudaGraphRunner(self) self.cuda_graph_runner = Device2DraftCudaGraphRunner[
self.target_worker.device
](self)
after_mem = get_available_gpu_memory(self.device, self.gpu_id) after_mem = get_available_gpu_memory(self.device, self.gpu_id)
logger.info( logger.info(
f"Capture draft cuda graph end. Time elapsed: {time.perf_counter() - tic:.2f} s. mem usage={(before_mem - after_mem):.2f} GB. avail mem={after_mem:.2f} GB." f"Capture draft cuda graph end. Time elapsed: {time.perf_counter() - tic:.2f} s. mem usage={(before_mem - after_mem):.2f} GB. avail mem={after_mem:.2f} GB."
) )
# Capture extend # Capture extend
if self.draft_extend_attn_backend: if self.draft_extend_attn_backend and not _is_npu:
tic = time.perf_counter() tic = time.perf_counter()
before_mem = get_available_gpu_memory(self.device, self.gpu_id) before_mem = get_available_gpu_memory(self.device, self.gpu_id)
logger.info( logger.info(
@@ -20,6 +20,10 @@ from sglang.srt.speculative.eagle_draft_cuda_graph_runner import (
from sglang.srt.speculative.eagle_draft_extend_cuda_graph_runner import ( from sglang.srt.speculative.eagle_draft_extend_cuda_graph_runner import (
EAGLEDraftExtendCudaGraphRunner, EAGLEDraftExtendCudaGraphRunner,
) )
from sglang.srt.speculative.eagle_draft_extend_npu_graph_runner import (
EAGLEDraftExtendNpuGraphRunner,
)
from sglang.srt.speculative.eagle_draft_npu_graph_runner import EAGLEDraftNpuGraphRunner
from sglang.srt.speculative.eagle_info import EagleDraftInput, EagleVerifyInput from sglang.srt.speculative.eagle_info import EagleDraftInput, EagleVerifyInput
from sglang.srt.speculative.eagle_info_v2 import ( from sglang.srt.speculative.eagle_info_v2 import (
assign_extend_cache_locs, assign_extend_cache_locs,
@@ -211,9 +215,13 @@ class EagleDraftWorker(BaseDraftWorker):
self.cuda_graph_runner = None self.cuda_graph_runner = None
self.cuda_graph_runner_for_draft_extend = None self.cuda_graph_runner_for_draft_extend = None
if self.server_args.disable_cuda_graph or _is_npu: if self.server_args.disable_cuda_graph:
return return
Device2DraftCudaGraphRunner = {
"npu": EAGLEDraftNpuGraphRunner,
"cuda": EAGLEDraftCudaGraphRunner,
}
# Capture draft # Capture draft
if self.speculative_num_steps > 1: if self.speculative_num_steps > 1:
tic = time.perf_counter() tic = time.perf_counter()
@@ -221,22 +229,29 @@ class EagleDraftWorker(BaseDraftWorker):
logger.info( logger.info(
f"Capture draft cuda graph begin. This can take up to several minutes. avail mem={before_mem:.2f} GB" f"Capture draft cuda graph begin. This can take up to several minutes. avail mem={before_mem:.2f} GB"
) )
self.cuda_graph_runner = EAGLEDraftCudaGraphRunner(self) self.cuda_graph_runner = Device2DraftCudaGraphRunner[
self.target_worker.device
](self)
after_mem = get_available_gpu_memory(self.device, self.gpu_id) after_mem = get_available_gpu_memory(self.device, self.gpu_id)
logger.info( logger.info(
f"Capture draft cuda graph end. Time elapsed: {time.perf_counter() - tic:.2f} s. mem usage={(before_mem - after_mem):.2f} GB. avail mem={after_mem:.2f} GB." f"Capture draft cuda graph end. Time elapsed: {time.perf_counter() - tic:.2f} s. mem usage={(before_mem - after_mem):.2f} GB. avail mem={after_mem:.2f} GB."
) )
Device2ExtendCudaGraphRunner = {
"npu": EAGLEDraftExtendNpuGraphRunner,
"cuda": EAGLEDraftExtendCudaGraphRunner,
}
# Capture extend # Capture extend
if self.draft_extend_attn_backend: # FIXME cuda not support draft_extend capture
if self.draft_extend_attn_backend and _is_npu:
tic = time.perf_counter() tic = time.perf_counter()
before_mem = get_available_gpu_memory(self.device, self.gpu_id) before_mem = get_available_gpu_memory(self.device, self.gpu_id)
logger.info( logger.info(
f"Capture draft extend cuda graph begin. This can take up to several minutes. avail mem={before_mem:.2f} GB" f"Capture draft extend cuda graph begin. This can take up to several minutes. avail mem={before_mem:.2f} GB"
) )
self.cuda_graph_runner_for_draft_extend = EAGLEDraftExtendCudaGraphRunner( self.cuda_graph_runner_for_draft_extend = Device2ExtendCudaGraphRunner[
self self.target_worker.device
) ](self)
after_mem = get_available_gpu_memory(self.device, self.gpu_id) after_mem = get_available_gpu_memory(self.device, self.gpu_id)
logger.info( logger.info(
f"Capture draft extend cuda graph end. Time elapsed: {time.perf_counter() - tic:.2f} s. mem usage={(before_mem - after_mem):.2f} GB. avail mem={after_mem:.2f} GB." f"Capture draft extend cuda graph end. Time elapsed: {time.perf_counter() - tic:.2f} s. mem usage={(before_mem - after_mem):.2f} GB. avail mem={after_mem:.2f} GB."
@@ -259,7 +274,10 @@ class EagleDraftWorker(BaseDraftWorker):
forward_batch, forward_batch,
) )
else: else:
if self.speculative_num_steps > 1: if (
not forward_batch.forward_mode.is_idle()
and self.speculative_num_steps > 1
):
# Skip attention backend init for 1-step draft, # Skip attention backend init for 1-step draft,
# `draft_forward` only does sample in this case. # `draft_forward` only does sample in this case.
self.draft_attn_backend.init_forward_metadata(forward_batch) self.draft_attn_backend.init_forward_metadata(forward_batch)
@@ -267,6 +285,13 @@ class EagleDraftWorker(BaseDraftWorker):
forward_batch forward_batch
) )
if model_worker_batch.forward_mode.is_idle():
return EagleVerifyInput.create_idle_input(
self.topk,
self.speculative_num_steps,
self.speculative_num_draft_tokens,
)
# Build tree mask # Build tree mask
# Directly write to cuda graph buffers for verify attn # Directly write to cuda graph buffers for verify attn
tree_mask_buf, position_buf = ( tree_mask_buf, position_buf = (
@@ -408,6 +433,7 @@ class EagleDraftWorker(BaseDraftWorker):
next_token_ids: Next token ids generated from the target forward. next_token_ids: Next token ids generated from the target forward.
""" """
# Construct input_ids # Construct input_ids
if not batch.forward_mode.is_idle():
pt = 0 pt = 0
for i, extend_len in enumerate(batch.extend_seq_lens): for i, extend_len in enumerate(batch.extend_seq_lens):
input_ids = batch.input_ids[pt : pt + extend_len] input_ids = batch.input_ids[pt : pt + extend_len]
@@ -422,7 +448,11 @@ class EagleDraftWorker(BaseDraftWorker):
verified_id=next_token_ids, verified_id=next_token_ids,
new_seq_lens=batch.seq_lens, new_seq_lens=batch.seq_lens,
allocate_lens=batch.seq_lens, allocate_lens=batch.seq_lens,
# draft mode is same with decode mode, only 1 num token per batch
num_tokens_per_batch=1,
num_tokens_for_logprob_per_batch=1,
) )
batch.spec_info = next_draft_input batch.spec_info = next_draft_input
# Run forward # Run forward
@@ -443,6 +473,8 @@ class EagleDraftWorker(BaseDraftWorker):
# Batch 2: Draft extend # Batch 2: Draft extend
draft_input = EagleDraftInput( draft_input = EagleDraftInput(
hidden_states=batch_result.logits_output.hidden_states, hidden_states=batch_result.logits_output.hidden_states,
num_tokens_per_batch=self.speculative_num_steps + 1,
num_tokens_for_logprob_per_batch=1,
) )
select_index = ( select_index = (
torch.arange(len(batch.seq_lens), device=self.device) torch.arange(len(batch.seq_lens), device=self.device)
@@ -458,6 +490,7 @@ class EagleDraftWorker(BaseDraftWorker):
batch_result.next_token_ids, batch_result.next_token_ids,
self.speculative_num_draft_tokens, self.speculative_num_draft_tokens,
self.draft_runner, self.draft_runner,
self.cuda_graph_runner_for_draft_extend,
) )
if self.plan_stream: if self.plan_stream:
@@ -466,8 +499,17 @@ class EagleDraftWorker(BaseDraftWorker):
) )
# Run draft extend batch in the main compute stream # Run draft extend batch in the main compute stream
draft_logits_output = self.draft_runner.model.forward( can_cuda_graph = (
forward_batch.input_ids, forward_batch.positions, forward_batch self.cuda_graph_runner_for_draft_extend
and self.cuda_graph_runner_for_draft_extend.can_run(forward_batch)
)
if can_cuda_graph:
draft_logits_output = self.cuda_graph_runner_for_draft_extend.replay(
forward_batch
)
else:
draft_logits_output, _ = self.draft_runner.forward(
forward_batch, skip_attn_backend_init=True
) )
# Reorganize the spec info for the next batch # Reorganize the spec info for the next batch
@@ -551,16 +593,10 @@ class EAGLEWorkerV2(BaseSpecWorker):
pass pass
def forward_batch_generation(self, model_worker_batch: ModelWorkerBatch): def forward_batch_generation(self, model_worker_batch: ModelWorkerBatch):
if model_worker_batch.forward_mode.is_decode(): if (
draft_input: EagleDraftInput = model_worker_batch.spec_info model_worker_batch.forward_mode.is_extend()
assert draft_input.is_draft_input() or model_worker_batch.is_extend_in_batch
verify_input: EagleVerifyInput = self.draft_worker.draft(model_worker_batch) ):
assert verify_input.is_verify_input()
model_worker_batch.spec_info = verify_input
batch_output = self.verify(model_worker_batch, draft_input.allocate_lens)
self.draft_worker._draft_extend_for_decode(model_worker_batch, batch_output)
return batch_output
else:
# Target prefill # Target prefill
model_worker_batch.capture_hidden_mode = CaptureHiddenMode.FULL model_worker_batch.capture_hidden_mode = CaptureHiddenMode.FULL
batch_output = self.target_worker.forward_batch_generation( batch_output = self.target_worker.forward_batch_generation(
@@ -575,6 +611,22 @@ class EAGLEWorkerV2(BaseSpecWorker):
batch_output.next_token_ids, batch_output.next_token_ids,
) )
return batch_output return batch_output
else:
if model_worker_batch.spec_info is None:
model_worker_batch.spec_info = EagleDraftInput.create_idle_input(
device=self.device,
hidden_size=self.target_worker.model_config.hidden_size,
dtype=self.target_worker.model_config.dtype,
topk=self.topk,
capture_hidden_mode=CaptureHiddenMode.LAST,
)
draft_input: EagleDraftInput = model_worker_batch.spec_info
verify_input: EagleVerifyInput = self.draft_worker.draft(model_worker_batch)
assert verify_input.is_verify_input()
model_worker_batch.spec_info = verify_input
batch_output = self.verify(model_worker_batch, draft_input.allocate_lens)
self.draft_worker._draft_extend_for_decode(model_worker_batch, batch_output)
return batch_output
def verify( def verify(
self, self,
@@ -605,7 +657,9 @@ class EAGLEWorkerV2(BaseSpecWorker):
# Correct some buffers due to the overlap plan # Correct some buffers due to the overlap plan
if self.plan_stream: if self.plan_stream:
torch.get_device_module().current_stream().wait_stream(self.plan_stream) torch.get_device_module(self.device).current_stream().wait_stream(
self.plan_stream
)
# Some values such as custom_mask and position depend on the output of draft, # Some values such as custom_mask and position depend on the output of draft,
# so the previous plan step used the wrong values. Here, we need to run the related # so the previous plan step used the wrong values. Here, we need to run the related
@@ -640,6 +694,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
verify_done = torch.get_device_module(self.device).Event() verify_done = torch.get_device_module(self.device).Event()
verify_done.record() verify_done.record()
if not batch.forward_mode.is_idle():
all_verified_id = predict[accept_index] all_verified_id = predict[accept_index]
verified_id = torch.empty_like(accept_length, dtype=torch.int32) verified_id = torch.empty_like(accept_length, dtype=torch.int32)
fill_new_verified_id[(bs,)]( fill_new_verified_id[(bs,)](
@@ -648,6 +703,8 @@ class EAGLEWorkerV2(BaseSpecWorker):
verified_id, verified_id,
self.speculative_num_draft_tokens, self.speculative_num_draft_tokens,
) )
else:
verified_id = torch.empty((0,), device=self.device, dtype=torch.int32)
# Construct the next draft input # Construct the next draft input
next_draft_input = EagleDraftInput( next_draft_input = EagleDraftInput(
+1 -1
View File
@@ -468,7 +468,7 @@ def select_top_k_tokens(
if hidden_states.shape[0] > 0: if hidden_states.shape[0] > 0:
selected_input_index = topk_cs_index.flatten() // topk + torch.arange( selected_input_index = topk_cs_index.flatten() // topk + torch.arange(
0, hidden_states.shape[0], step=topk, device="cuda" 0, hidden_states.shape[0], step=topk, device=topk_index.device
).repeat_interleave(topk) ).repeat_interleave(topk)
hidden_states = hidden_states[selected_input_index, :] hidden_states = hidden_states[selected_input_index, :]
+1 -1
View File
@@ -59,7 +59,7 @@ wget -O "${BISHENG_NAME}" "${BISHENG_URL}" && chmod a+x "${BISHENG_NAME}" && "./
### Install sgl-kernel-npu ### Install sgl-kernel-npu
SGL_KERNEL_NPU_TAG="20251106" SGL_KERNEL_NPU_TAG="20251110"
git clone --depth 1 https://github.com/sgl-project/sgl-kernel-npu.git --branch ${SGL_KERNEL_NPU_TAG} git clone --depth 1 https://github.com/sgl-project/sgl-kernel-npu.git --branch ${SGL_KERNEL_NPU_TAG}
# pin wheel to 0.45.1 ref: https://github.com/pypa/wheel/issues/662 # pin wheel to 0.45.1 ref: https://github.com/pypa/wheel/issues/662
pip install wheel==0.45.1 pip install wheel==0.45.1
+5 -20
View File
@@ -43,6 +43,9 @@ class TestAscendDeepSeekMTP(CustomTestCase):
32768, 32768,
"--tp-size", "--tp-size",
16, 16,
"--dp-size",
2,
"--enable-dp-attention",
"--speculative-algorithm", "--speculative-algorithm",
"NEXTN", "NEXTN",
"--speculative-num-steps", "--speculative-num-steps",
@@ -55,6 +58,8 @@ class TestAscendDeepSeekMTP(CustomTestCase):
cls.extra_envs = { cls.extra_envs = {
"SGLANG_NPU_USE_MLAPO": "1", "SGLANG_NPU_USE_MLAPO": "1",
"SGLANG_ENABLE_SPEC_V2": "1",
"SGLANG_ENABLE_OVERLAP_PLAN_STREAM": "1",
} }
os.environ.update(cls.extra_envs) os.environ.update(cls.extra_envs)
@@ -91,26 +96,6 @@ class TestAscendDeepSeekMTP(CustomTestCase):
finally: finally:
kill_process_tree(process.pid) kill_process_tree(process.pid)
def test_b_throughput(self):
for model in self.models:
with self.subTest(model=model):
print(f"##=== Testing throughput: {model} ===##")
output_throughput = run_bench_offline_throughput(
model,
[
*self.common_args,
],
)
print(f"##=== {model} throughput: {output_throughput} ===##")
if is_in_ci():
self.assertGreater(
output_throughput,
TEST_MODEL_MATRIX[model]["output_throughput"],
)
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()
+1
View File
@@ -181,6 +181,7 @@ suites = {
TestFile("test_flash_attention_4.py", 300), TestFile("test_flash_attention_4.py", 300),
TestFile("test_gpt_oss_4gpu.py", 600), TestFile("test_gpt_oss_4gpu.py", 600),
TestFile("test_llama31_fp4.py", 300), TestFile("test_llama31_fp4.py", 300),
TestFile("test_eagle_infer_beta_dp_attention.py", 200),
], ],
"per-commit-4-gpu-gb200": [ "per-commit-4-gpu-gb200": [
TestFile("test_cutedsl_moe.py", 300), TestFile("test_cutedsl_moe.py", 300),
@@ -0,0 +1,97 @@
import os
import unittest
from types import SimpleNamespace
import requests
from sglang.srt.utils import kill_process_tree
from sglang.test.few_shot_gsm8k import run_eval as run_eval_few_shot_gsm8k
from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
CustomTestCase,
is_in_ci,
popen_launch_server,
write_github_step_summary,
)
FULL_DEEPSEEK_V3_FP4_MODEL_PATH = "nvidia/DeepSeek-V3-0324-FP4"
class TestEagleDPAttnServerBase(CustomTestCase):
@classmethod
def setUpClass(cls):
os.environ["SGLANG_ENABLE_SPEC_V2"] = "1"
cls.model = FULL_DEEPSEEK_V3_FP4_MODEL_PATH
cls.base_url = DEFAULT_URL_FOR_TEST
other_args = [
"--tp-size",
"4",
"--dp-size",
"4",
"--enable-dp-attention",
"--attention-backend",
"trtllm_mla",
"--moe-runner-backend",
"flashinfer_trtllm",
"--quantization",
"modelopt_fp4",
"--speculative-algorithm",
"EAGLE",
"--speculative-num-steps",
"3",
"--speculative-eagle-topk",
"1",
"--speculative-num-draft-tokens",
"4",
"--kv-cache-dtype",
"fp8_e4m3",
]
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=other_args,
)
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
if "SGLANG_ENABLE_SPEC_V2" in os.environ:
del os.environ["SGLANG_ENABLE_SPEC_V2"]
def test_a_gsm8k(
self,
): # Append an "a" to make this test run first (alphabetically) to warm up the server
requests.get(self.base_url + "/flush_cache")
args = SimpleNamespace(
num_shots=5,
data_path=None,
num_questions=200,
max_new_tokens=512,
parallel=128,
host="http://127.0.0.1",
port=int(self.base_url.split(":")[-1]),
)
metrics = run_eval_few_shot_gsm8k(args)
print(f"{metrics=}")
server_info = requests.get(self.base_url + "/get_server_info")
avg_spec_accept_length = server_info.json()["internal_states"][0][
"avg_spec_accept_length"
]
print(f"{avg_spec_accept_length=}")
if is_in_ci():
write_github_step_summary(
f"### test_gsm8k (deepseek-v3-fp4 mtp)\n"
f'{metrics["accuracy"]=:.3f}\n'
f"{avg_spec_accept_length=:.2f}\n"
)
self.assertGreater(metrics["accuracy"], 0.94)
self.assertGreater(avg_spec_accept_length, 2.04)
if __name__ == "__main__":
unittest.main()