[Spec] Relay ngram accept tokens through the FutureMap (#35198)
This commit is contained in:
@@ -16,6 +16,7 @@ if TYPE_CHECKING:
|
|||||||
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
from sglang.srt.speculative.eagle_info import EagleDraftInput
|
from sglang.srt.speculative.eagle_info import EagleDraftInput
|
||||||
|
from sglang.srt.speculative.ngram_info import NgramVerifyInput
|
||||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||||
|
|
||||||
|
|
||||||
@@ -131,12 +132,25 @@ class RelayPayload:
|
|||||||
`bonus_tokens`; which spec extras get relayed is decided by
|
`bonus_tokens`; which spec extras get relayed is decided by
|
||||||
`FutureMap.spec_algo`, not by this payload's shape."""
|
`FutureMap.spec_algo`, not by this payload's shape."""
|
||||||
|
|
||||||
bonus_tokens: torch.Tensor
|
bonus_tokens: Optional[torch.Tensor]
|
||||||
topk_p: Optional[torch.Tensor] = None
|
topk_p: Optional[torch.Tensor] = None
|
||||||
topk_index: Optional[torch.Tensor] = None
|
topk_index: Optional[torch.Tensor] = None
|
||||||
hidden_states: Optional[torch.Tensor] = None
|
hidden_states: Optional[torch.Tensor] = None
|
||||||
draft_probs: Optional[torch.Tensor] = None
|
draft_probs: Optional[torch.Tensor] = None
|
||||||
dsa_topk_indices: Optional[torch.Tensor] = None
|
dsa_topk_indices: Optional[torch.Tensor] = None
|
||||||
|
# ngram delays the draft extend (ngram update)
|
||||||
|
accept_tokens: Optional[torch.Tensor] = None
|
||||||
|
accept_lens: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_ngram(cls, draft_input: NgramVerifyInput) -> RelayPayload:
|
||||||
|
return cls(
|
||||||
|
bonus_tokens=None,
|
||||||
|
accept_tokens=draft_input.accept_tokens.reshape(
|
||||||
|
-1, draft_input.draft_token_num
|
||||||
|
),
|
||||||
|
accept_lens=draft_input.accept_lens,
|
||||||
|
)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_draft_input(cls, draft_input: EagleDraftInput) -> RelayPayload:
|
def from_draft_input(cls, draft_input: EagleDraftInput) -> RelayPayload:
|
||||||
@@ -283,6 +297,10 @@ class FutureMap:
|
|||||||
self._forward_buf_initialized = False
|
self._forward_buf_initialized = False
|
||||||
self.dsa_topk_indices_buf = None
|
self.dsa_topk_indices_buf = None
|
||||||
|
|
||||||
|
# ngram-only relay bufs
|
||||||
|
self.accept_tokens_buf: Optional[torch.Tensor] = None
|
||||||
|
self.accept_lens_buf: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
self.publish_ready = None # lazy device.Event(); only spec_v2 needs it
|
self.publish_ready = None # lazy device.Event(); only spec_v2 needs it
|
||||||
# Debug consume-once state: armed by a recording publish, consumed by
|
# Debug consume-once state: armed by a recording publish, consumed by
|
||||||
# resolve; arm/consume strictly alternate across all batch interleavings.
|
# resolve; arm/consume strictly alternate across all batch interleavings.
|
||||||
@@ -349,6 +367,22 @@ class FutureMap:
|
|||||||
device=self.device,
|
device=self.device,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _maybe_init_ngram_bufs(self, payload: RelayPayload) -> None:
|
||||||
|
if self.accept_tokens_buf is not None:
|
||||||
|
return
|
||||||
|
# zeros, not empty: an unstashed row resolves to accept_len 0 (empty
|
||||||
|
# splice at draft prep) instead of a garbage length.
|
||||||
|
self.accept_tokens_buf = torch.zeros(
|
||||||
|
(self.req_pool_size, payload.accept_tokens.shape[1]),
|
||||||
|
dtype=payload.accept_tokens.dtype,
|
||||||
|
device=self.device,
|
||||||
|
)
|
||||||
|
self.accept_lens_buf = torch.zeros(
|
||||||
|
(self.req_pool_size,),
|
||||||
|
dtype=payload.accept_lens.dtype,
|
||||||
|
device=self.device,
|
||||||
|
)
|
||||||
|
|
||||||
def resolve_confidence_cpu(
|
def resolve_confidence_cpu(
|
||||||
self, batch: ScheduleBatch
|
self, batch: ScheduleBatch
|
||||||
) -> Optional[ResolvedConfidence]:
|
) -> Optional[ResolvedConfidence]:
|
||||||
@@ -362,7 +396,14 @@ class FutureMap:
|
|||||||
|
|
||||||
def _resolve_spec_extras(self, batch: ScheduleBatch) -> None:
|
def _resolve_spec_extras(self, batch: ScheduleBatch) -> None:
|
||||||
if self.spec_algo.is_ngram():
|
if self.spec_algo.is_ngram():
|
||||||
# FIXME: remove once precomputed draft is supported.
|
draft_input = batch.spec_info
|
||||||
|
if draft_input is None or draft_input.future_indices is None:
|
||||||
|
return
|
||||||
|
indices = draft_input.future_indices
|
||||||
|
if indices.shape[0] == 0:
|
||||||
|
return
|
||||||
|
draft_input.accept_tokens = self.accept_tokens_buf[indices].flatten()
|
||||||
|
draft_input.accept_lens = self.accept_lens_buf[indices]
|
||||||
return
|
return
|
||||||
draft_input: EagleDraftInput = batch.spec_info
|
draft_input: EagleDraftInput = batch.spec_info
|
||||||
if draft_input is None:
|
if draft_input is None:
|
||||||
@@ -499,13 +540,15 @@ class FutureMap:
|
|||||||
)
|
)
|
||||||
|
|
||||||
def stash(self, future_indices: torch.Tensor, payload: RelayPayload) -> None:
|
def stash(self, future_indices: torch.Tensor, payload: RelayPayload) -> None:
|
||||||
if self.spec_algo.is_ngram():
|
|
||||||
# FIXME: remove once precomputed draft is supported.
|
|
||||||
return
|
|
||||||
indices = future_indices
|
indices = future_indices
|
||||||
if indices.shape[0] == 0:
|
if indices.shape[0] == 0:
|
||||||
# DP idle: payload is empty stub; lazy-init shape peek would IndexError.
|
# DP idle: payload is empty stub; lazy-init shape peek would IndexError.
|
||||||
return
|
return
|
||||||
|
if self.spec_algo.is_ngram():
|
||||||
|
self._maybe_init_ngram_bufs(payload)
|
||||||
|
self.accept_tokens_buf[indices] = payload.accept_tokens
|
||||||
|
self.accept_lens_buf[indices] = payload.accept_lens
|
||||||
|
return
|
||||||
if not self._forward_buf_initialized:
|
if not self._forward_buf_initialized:
|
||||||
self._lazy_init_forward_buf(payload)
|
self._lazy_init_forward_buf(payload)
|
||||||
self._maybe_init_dsa_topk_indices_buf(payload)
|
self._maybe_init_dsa_topk_indices_buf(payload)
|
||||||
|
|||||||
@@ -3874,9 +3874,11 @@ class Scheduler(
|
|||||||
def _relay_forward_payload(
|
def _relay_forward_payload(
|
||||||
self, future_indices: torch.Tensor, batch_result: GenerationBatchResult
|
self, future_indices: torch.Tensor, batch_result: GenerationBatchResult
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Stash this iter's relay payload for next iter's resolve_forward_inputs.
|
"""Stash this iter's relay payload for next iter's resolve_forward_inputs."""
|
||||||
ngram is skipped: it relays its draft via batch.spec_info, not the FutureMap."""
|
|
||||||
if self.spec_algorithm.is_ngram():
|
if self.spec_algorithm.is_ngram():
|
||||||
|
if batch_result.next_draft_input is not None:
|
||||||
|
payload = RelayPayload.from_ngram(batch_result.next_draft_input)
|
||||||
|
self.future_map.stash(future_indices, payload)
|
||||||
return
|
return
|
||||||
if batch_result.next_draft_input is not None:
|
if batch_result.next_draft_input is not None:
|
||||||
payload = RelayPayload.from_draft_input(batch_result.next_draft_input)
|
payload = RelayPayload.from_draft_input(batch_result.next_draft_input)
|
||||||
|
|||||||
@@ -120,6 +120,8 @@ class NgramVerifyInput(SpecInput):
|
|||||||
):
|
):
|
||||||
if self.future_indices is not None:
|
if self.future_indices is not None:
|
||||||
self.future_indices = self.future_indices[new_indices]
|
self.future_indices = self.future_indices[new_indices]
|
||||||
|
return
|
||||||
|
|
||||||
if self.new_seq_lens is not None:
|
if self.new_seq_lens is not None:
|
||||||
self.new_seq_lens = self.new_seq_lens[new_indices]
|
self.new_seq_lens = self.new_seq_lens[new_indices]
|
||||||
self.accept_tokens = self.accept_tokens.reshape(-1, self.draft_token_num)[
|
self.accept_tokens = self.accept_tokens.reshape(-1, self.draft_token_num)[
|
||||||
@@ -134,6 +136,8 @@ class NgramVerifyInput(SpecInput):
|
|||||||
self.future_indices = torch.cat(
|
self.future_indices = torch.cat(
|
||||||
(self.future_indices, spec_info.future_indices), dim=0
|
(self.future_indices, spec_info.future_indices), dim=0
|
||||||
)
|
)
|
||||||
|
return
|
||||||
|
|
||||||
if self.new_seq_lens is not None:
|
if self.new_seq_lens is not None:
|
||||||
assert spec_info.new_seq_lens is not None
|
assert spec_info.new_seq_lens is not None
|
||||||
self.new_seq_lens = torch.cat(
|
self.new_seq_lens = torch.cat(
|
||||||
|
|||||||
@@ -419,8 +419,6 @@ class NGRAMWorker(BaseSpecWorker):
|
|||||||
batch_result = self.target_worker.forward_batch_generation(
|
batch_result = self.target_worker.forward_batch_generation(
|
||||||
batch, is_verify=True
|
batch, is_verify=True
|
||||||
)
|
)
|
||||||
# Verify reads shared state past the in-graph marker; keep it coarse.
|
|
||||||
self.target_worker.model_runner.shared_read_done_event = None
|
|
||||||
|
|
||||||
logits_output, can_run_cuda_graph = (
|
logits_output, can_run_cuda_graph = (
|
||||||
batch_result.logits_output,
|
batch_result.logits_output,
|
||||||
|
|||||||
Reference in New Issue
Block a user