[Feature] Spec-Overlap supporting DP-ATTN; PD-Disaggregation; npugraph mode (#12443)
This commit is contained in:
@@ -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 }}
|
||||||
|
|||||||
@@ -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 }}
|
||||||
|
|||||||
@@ -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 = (
|
||||||
|
|||||||
@@ -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
|
|
||||||
|
|||||||
@@ -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=[],
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
@@ -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, :]
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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()
|
||||||
Reference in New Issue
Block a user