[Perf] Fast-path chain-style draft token organization in multi-layer EAGLE (#32887)
This commit is contained in:
@@ -3,6 +3,8 @@ from __future__ import annotations
|
|||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from typing import TYPE_CHECKING, Optional
|
from typing import TYPE_CHECKING, Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.managers.io_struct import (
|
from sglang.srt.managers.io_struct import (
|
||||||
UpdateWeightFromDiskReqInput,
|
UpdateWeightFromDiskReqInput,
|
||||||
@@ -13,6 +15,10 @@ if TYPE_CHECKING:
|
|||||||
|
|
||||||
|
|
||||||
class EagleDraftWorkerBase(ABC):
|
class EagleDraftWorkerBase(ABC):
|
||||||
|
# topk=1 chain constants for draft_forward's fast path; None when topk > 1.
|
||||||
|
_topk1_parents_prealloc: Optional[torch.Tensor] = None
|
||||||
|
_topk1_score_indices_prealloc: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def draft():
|
def draft():
|
||||||
pass
|
pass
|
||||||
@@ -41,6 +47,40 @@ class EagleDraftWorkerBase(ABC):
|
|||||||
self.draft_worker.init_cuda_graphs(capture_decode_cuda_graph=False)
|
self.draft_worker.init_cuda_graphs(capture_decode_cuda_graph=False)
|
||||||
self._capture_cuda_graphs()
|
self._capture_cuda_graphs()
|
||||||
|
|
||||||
|
def _rebuild_topk1_chain_buffers(self) -> None:
|
||||||
|
# For topk=1 the draft tree degenerates to a chain, so parent_list and
|
||||||
|
# top_scores_index are runtime-invariant. Must be rebuilt after any
|
||||||
|
# change to speculative_num_steps / speculative_num_draft_tokens.
|
||||||
|
if self.topk != 1:
|
||||||
|
return
|
||||||
|
# _override_worker_state can set both directly, bypassing the hook that
|
||||||
|
# pins this relation; the fast path is only valid when it holds.
|
||||||
|
assert self.speculative_num_draft_tokens == self.speculative_num_steps + 1, (
|
||||||
|
"topk=1 requires speculative_num_draft_tokens == speculative_num_steps + 1, "
|
||||||
|
f"got {self.speculative_num_draft_tokens} and {self.speculative_num_steps}"
|
||||||
|
)
|
||||||
|
num_steps = self.speculative_num_steps
|
||||||
|
sa = self.server_args
|
||||||
|
decode_max_bs = (
|
||||||
|
sa.cuda_graph_config.decode.max_bs
|
||||||
|
if sa.cuda_graph_config is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
max_bs = max(
|
||||||
|
decode_max_bs or 0,
|
||||||
|
sa.max_running_requests or 0,
|
||||||
|
1,
|
||||||
|
)
|
||||||
|
# A single-step chain has no parent entries (slow path drops the last
|
||||||
|
# step). repeat (not expand): the kernel reads these as contiguous.
|
||||||
|
parent_width = num_steps if num_steps > 1 else 0
|
||||||
|
self._topk1_parents_prealloc = torch.arange(
|
||||||
|
-1, parent_width - 1, dtype=torch.long, device=self.device
|
||||||
|
).repeat(max_bs, 1)
|
||||||
|
self._topk1_score_indices_prealloc = torch.arange(
|
||||||
|
num_steps, dtype=torch.long, device=self.device
|
||||||
|
).repeat(max_bs, 1)
|
||||||
|
|
||||||
|
|
||||||
class BaseSpecWorker(ABC):
|
class BaseSpecWorker(ABC):
|
||||||
@property
|
@property
|
||||||
|
|||||||
@@ -103,7 +103,9 @@ def organize_draft_results(
|
|||||||
parents_list: List[torch.Tensor],
|
parents_list: List[torch.Tensor],
|
||||||
num_draft_token: int,
|
num_draft_token: int,
|
||||||
):
|
):
|
||||||
|
# b, n, topk; n = 1 + (num_steps-1) * topk
|
||||||
score_list = torch.cat(score_list, dim=1).flatten(1)
|
score_list = torch.cat(score_list, dim=1).flatten(1)
|
||||||
|
# b, (topk + (num_steps-1) * topk)
|
||||||
ss_token_list = torch.cat(token_list, dim=1)
|
ss_token_list = torch.cat(token_list, dim=1)
|
||||||
top_scores = torch.topk(score_list, num_draft_token - 1, dim=-1)
|
top_scores = torch.topk(score_list, num_draft_token - 1, dim=-1)
|
||||||
top_scores_index = top_scores.indices
|
top_scores_index = top_scores.indices
|
||||||
|
|||||||
@@ -144,9 +144,6 @@ class EagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
server_args.speculative_algorithm
|
server_args.speculative_algorithm
|
||||||
)
|
)
|
||||||
|
|
||||||
# Pre-allocated constants for the topk=1 chain fast path in draft_forward.
|
|
||||||
self._topk1_parents_prealloc = None
|
|
||||||
self._topk1_score_indices_prealloc = None
|
|
||||||
self._rebuild_topk1_chain_buffers()
|
self._rebuild_topk1_chain_buffers()
|
||||||
|
|
||||||
# Load draft model weights only.
|
# Load draft model weights only.
|
||||||
@@ -251,40 +248,6 @@ class EagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
self.index_share_for_mtp_iteration and self.dsa_index_topk is not None
|
self.index_share_for_mtp_iteration and self.dsa_index_topk is not None
|
||||||
)
|
)
|
||||||
|
|
||||||
def _rebuild_topk1_chain_buffers(self) -> None:
|
|
||||||
# For topk=1 the draft tree degenerates to a chain, so parent_list and
|
|
||||||
# top_scores_index are runtime-invariant. Must be rebuilt after any
|
|
||||||
# change to speculative_num_steps / speculative_num_draft_tokens.
|
|
||||||
if self.topk != 1:
|
|
||||||
return
|
|
||||||
# _override_worker_state can set both directly, bypassing the hook that
|
|
||||||
# pins this relation; the fast path is only valid when it holds.
|
|
||||||
assert self.speculative_num_draft_tokens == self.speculative_num_steps + 1, (
|
|
||||||
"topk=1 requires speculative_num_draft_tokens == speculative_num_steps + 1, "
|
|
||||||
f"got {self.speculative_num_draft_tokens} and {self.speculative_num_steps}"
|
|
||||||
)
|
|
||||||
num_steps = self.speculative_num_steps
|
|
||||||
sa = self.server_args
|
|
||||||
decode_max_bs = (
|
|
||||||
sa.cuda_graph_config.decode.max_bs
|
|
||||||
if sa.cuda_graph_config is not None
|
|
||||||
else None
|
|
||||||
)
|
|
||||||
max_bs = max(
|
|
||||||
decode_max_bs or 0,
|
|
||||||
sa.max_running_requests or 0,
|
|
||||||
1,
|
|
||||||
)
|
|
||||||
# A single-step chain has no parent entries (slow path drops the last
|
|
||||||
# step). repeat (not expand): the kernel reads these as contiguous.
|
|
||||||
parent_width = num_steps if num_steps > 1 else 0
|
|
||||||
self._topk1_parents_prealloc = torch.arange(
|
|
||||||
-1, parent_width - 1, dtype=torch.long, device=self.device
|
|
||||||
).repeat(max_bs, 1)
|
|
||||||
self._topk1_score_indices_prealloc = torch.arange(
|
|
||||||
num_steps, dtype=torch.long, device=self.device
|
|
||||||
).repeat(max_bs, 1)
|
|
||||||
|
|
||||||
def init_token_map(self):
|
def init_token_map(self):
|
||||||
# Load hot token ids
|
# Load hot token ids
|
||||||
if self.speculative_algorithm.is_eagle3():
|
if self.speculative_algorithm.is_eagle3():
|
||||||
|
|||||||
@@ -50,6 +50,7 @@ from sglang.srt.speculative.eagle_info import (
|
|||||||
from sglang.srt.speculative.eagle_utils import (
|
from sglang.srt.speculative.eagle_utils import (
|
||||||
default_tree_mask_mode,
|
default_tree_mask_mode,
|
||||||
get_draft_recurrent_hidden_state_spec,
|
get_draft_recurrent_hidden_state_spec,
|
||||||
|
organize_draft_results,
|
||||||
)
|
)
|
||||||
from sglang.srt.speculative.eagle_worker_common import (
|
from sglang.srt.speculative.eagle_worker_common import (
|
||||||
build_eagle_verify_input,
|
build_eagle_verify_input,
|
||||||
@@ -133,6 +134,8 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
server_args.speculative_algorithm
|
server_args.speculative_algorithm
|
||||||
)
|
)
|
||||||
|
|
||||||
|
self._rebuild_topk1_chain_buffers()
|
||||||
|
|
||||||
# Set constant
|
# Set constant
|
||||||
EagleDraftInput.ALLOC_LEN_PER_DECODE = max(
|
EagleDraftInput.ALLOC_LEN_PER_DECODE = max(
|
||||||
self.speculative_num_steps * self.topk, self.speculative_num_draft_tokens
|
self.speculative_num_steps * self.topk, self.speculative_num_draft_tokens
|
||||||
@@ -442,6 +445,30 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
|
|
||||||
maybe_detect_nan(topk_p, "draft_forward: NaN in initial topk_p from spec_info")
|
maybe_detect_nan(topk_p, "draft_forward: NaN in initial topk_p from spec_info")
|
||||||
|
|
||||||
|
# Chain-style (topk=1, one token per draft step, all of them selected):
|
||||||
|
# _draft_forward_organize's slice/cat/topk/sort/gather is the identity on
|
||||||
|
# topk_index, and parent_list is the constant [-1, 0, .., S-2] per row.
|
||||||
|
parents_prealloc = self._topk1_parents_prealloc
|
||||||
|
if (
|
||||||
|
parents_prealloc is not None
|
||||||
|
and topk_index.shape[1] == self.speculative_num_steps
|
||||||
|
and topk_index.shape[0] <= parents_prealloc.shape[0]
|
||||||
|
):
|
||||||
|
bs = topk_index.shape[0]
|
||||||
|
return (
|
||||||
|
parents_prealloc[:bs],
|
||||||
|
self._topk1_score_indices_prealloc[:bs],
|
||||||
|
topk_index,
|
||||||
|
)
|
||||||
|
|
||||||
|
return self._draft_forward_organize(topk_p, topk_index, hidden_states)
|
||||||
|
|
||||||
|
def _draft_forward_organize(
|
||||||
|
self,
|
||||||
|
topk_p: torch.Tensor,
|
||||||
|
topk_index: torch.Tensor,
|
||||||
|
hidden_states: torch.Tensor,
|
||||||
|
):
|
||||||
# Return values
|
# Return values
|
||||||
score_list: List[torch.Tensor] = []
|
score_list: List[torch.Tensor] = []
|
||||||
token_list: List[torch.Tensor] = []
|
token_list: List[torch.Tensor] = []
|
||||||
@@ -473,33 +500,9 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
# Organize the results
|
return organize_draft_results(
|
||||||
score_list = torch.cat(score_list, dim=1).flatten(
|
score_list, token_list, parents_list, self.speculative_num_draft_tokens
|
||||||
1
|
|
||||||
) # b, n, topk; n= 1 + (num_steps-1) * self.topk
|
|
||||||
ss_token_list = torch.cat(
|
|
||||||
token_list, dim=1
|
|
||||||
) # b, (self.topk + (num_steps-1) * self.topk)
|
|
||||||
top_scores = torch.topk(
|
|
||||||
score_list, self.speculative_num_draft_tokens - 1, dim=-1
|
|
||||||
)
|
)
|
||||||
top_scores_index = top_scores.indices
|
|
||||||
top_scores_index = torch.sort(top_scores_index).values
|
|
||||||
maybe_detect_oob(
|
|
||||||
top_scores_index,
|
|
||||||
0,
|
|
||||||
ss_token_list.shape[1],
|
|
||||||
"draft_forward: top_scores_index OOB for gather on ss_token_list",
|
|
||||||
)
|
|
||||||
draft_tokens = torch.gather(ss_token_list, index=top_scores_index, dim=1)
|
|
||||||
|
|
||||||
if len(parents_list) > 1:
|
|
||||||
parent_list = torch.cat(parents_list[:-1], dim=1)
|
|
||||||
else:
|
|
||||||
batch_size = parents_list[0].shape[0]
|
|
||||||
parent_list = torch.empty(batch_size, 0, device=parents_list[0].device)
|
|
||||||
|
|
||||||
return parent_list, top_scores_index, draft_tokens
|
|
||||||
|
|
||||||
def draft_extend(self):
|
def draft_extend(self):
|
||||||
pass
|
pass
|
||||||
|
|||||||
@@ -51,9 +51,6 @@ class StandaloneDraftWorker(EagleDraftWorker):
|
|||||||
server_args.speculative_algorithm
|
server_args.speculative_algorithm
|
||||||
)
|
)
|
||||||
|
|
||||||
# Pre-allocated constants for the topk=1 chain fast path in draft_forward.
|
|
||||||
self._topk1_parents_prealloc = None
|
|
||||||
self._topk1_score_indices_prealloc = None
|
|
||||||
self._rebuild_topk1_chain_buffers()
|
self._rebuild_topk1_chain_buffers()
|
||||||
|
|
||||||
# Set constant
|
# Set constant
|
||||||
|
|||||||
Reference in New Issue
Block a user