[Spec] Simplify compute_spec_v2_logprobs signature and skip identity gathers (#35058)
This commit is contained in:
@@ -13,6 +13,7 @@ from sglang.srt.runtime_context import get_exec
|
|||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.layers.logits_processor import LogitsMetadata, LogitsProcessorOutput
|
from sglang.srt.layers.logits_processor import LogitsMetadata, LogitsProcessorOutput
|
||||||
from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
||||||
|
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -348,25 +349,36 @@ def get_token_ids_logprobs_chunk(
|
|||||||
return next_split_pruned_len
|
return next_split_pruned_len
|
||||||
|
|
||||||
|
|
||||||
def compute_spec_v2_logprobs(
|
def compute_spec_logprobs(
|
||||||
batch,
|
batch: ScheduleBatch,
|
||||||
logits_output,
|
logits_output: LogitsProcessorOutput,
|
||||||
predict: torch.Tensor,
|
predict: torch.Tensor,
|
||||||
accept_index: torch.Tensor,
|
*,
|
||||||
speculative_num_steps: int,
|
accept_index: Optional[torch.Tensor] = None,
|
||||||
|
chain_stride: Optional[int] = None,
|
||||||
):
|
):
|
||||||
"""Compute logprobs for accepted tokens after spec v2 verify sampling.
|
assert (accept_index is None) != (
|
||||||
|
chain_stride is None
|
||||||
|
), "pass exactly one of accept_index / chain_stride"
|
||||||
|
|
||||||
Gathers logits at accepted positions, applies log_softmax (temperature-scaled
|
|
||||||
if not greedy), and populates logits_output.next_token_logprobs (plus optional
|
|
||||||
top-k / token-ids logprobs) so they flow through copy_to_cpu().
|
|
||||||
"""
|
|
||||||
bs = len(batch.seq_lens)
|
bs = len(batch.seq_lens)
|
||||||
max_accept = speculative_num_steps + 1
|
next_token_logits = logits_output.next_token_logits
|
||||||
device = predict.device
|
|
||||||
|
|
||||||
flat_accept_idx = accept_index.long().reshape(-1)
|
if accept_index is not None:
|
||||||
gathered_logits = logits_output.next_token_logits[flat_accept_idx]
|
max_accept = accept_index.shape[1]
|
||||||
|
flat_accept_idx = accept_index.long().reshape(-1)
|
||||||
|
gathered_logits = next_token_logits[flat_accept_idx]
|
||||||
|
accepted_token_ids = predict[flat_accept_idx]
|
||||||
|
else:
|
||||||
|
max_accept = chain_stride
|
||||||
|
# Guards the layout contract the identity gather rests on: out token
|
||||||
|
# (b, j) must come from logits row b * stride + j.
|
||||||
|
assert next_token_logits.shape[0] == bs * max_accept, (
|
||||||
|
f"chain layout expects {bs * max_accept} logits rows, got "
|
||||||
|
f"{next_token_logits.shape[0]}"
|
||||||
|
)
|
||||||
|
gathered_logits = next_token_logits
|
||||||
|
accepted_token_ids = predict
|
||||||
|
|
||||||
if batch.sampling_info.is_all_greedy or envs.SGLANG_RETURN_ORIGINAL_LOGPROB.get():
|
if batch.sampling_info.is_all_greedy or envs.SGLANG_RETURN_ORIGINAL_LOGPROB.get():
|
||||||
gathered_logprobs = torch.nn.functional.log_softmax(gathered_logits, dim=-1)
|
gathered_logprobs = torch.nn.functional.log_softmax(gathered_logits, dim=-1)
|
||||||
@@ -381,12 +393,9 @@ def compute_spec_v2_logprobs(
|
|||||||
)
|
)
|
||||||
gathered_logprobs.clamp_(min=torch.finfo(gathered_logprobs.dtype).min)
|
gathered_logprobs.clamp_(min=torch.finfo(gathered_logprobs.dtype).min)
|
||||||
|
|
||||||
accepted_token_ids = predict[flat_accept_idx]
|
logits_output.next_token_logprobs = gathered_logprobs.gather(
|
||||||
token_logprobs = gathered_logprobs[
|
1, accepted_token_ids.long().view(-1, 1)
|
||||||
torch.arange(bs * max_accept, device=device),
|
).view(bs, max_accept)
|
||||||
accepted_token_ids.long(),
|
|
||||||
]
|
|
||||||
logits_output.next_token_logprobs = token_logprobs.reshape(bs, max_accept)
|
|
||||||
|
|
||||||
if batch.top_logprobs_nums and any(x > 0 for x in batch.top_logprobs_nums):
|
if batch.top_logprobs_nums and any(x > 0 for x in batch.top_logprobs_nums):
|
||||||
top_logprobs_nums_expanded = [
|
top_logprobs_nums_expanded = [
|
||||||
|
|||||||
@@ -17,7 +17,7 @@ from sglang.srt.configs.hybrid_arch import mambaish_config
|
|||||||
from sglang.srt.distributed import get_tp_group
|
from sglang.srt.distributed import get_tp_group
|
||||||
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.logprob_processor import compute_spec_v2_logprobs
|
from sglang.srt.layers.logprob_processor import compute_spec_logprobs
|
||||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||||
from sglang.srt.managers.scheduler import GenerationBatchResult
|
from sglang.srt.managers.scheduler import GenerationBatchResult
|
||||||
from sglang.srt.managers.tp_worker import TpModelWorker
|
from sglang.srt.managers.tp_worker import TpModelWorker
|
||||||
@@ -1897,15 +1897,11 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
new_seq_lens = None
|
new_seq_lens = None
|
||||||
|
|
||||||
if batch.return_logprob:
|
if batch.return_logprob:
|
||||||
output_indices = torch.arange(
|
compute_spec_logprobs(
|
||||||
bs * block_size, dtype=torch.int64, device=device
|
|
||||||
).view(bs, block_size)
|
|
||||||
compute_spec_v2_logprobs(
|
|
||||||
batch,
|
batch,
|
||||||
logits_output,
|
logits_output,
|
||||||
out_tokens.reshape(-1),
|
out_tokens.reshape(-1),
|
||||||
output_indices,
|
chain_stride=block_size,
|
||||||
block_size - 1,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
if self._need_mamba_verify_commit:
|
if self._need_mamba_verify_commit:
|
||||||
|
|||||||
@@ -11,7 +11,7 @@ from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import (
|
|||||||
from sglang.srt.configs.hybrid_arch import mambaish_config
|
from sglang.srt.configs.hybrid_arch import mambaish_config
|
||||||
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.logprob_processor import compute_spec_v2_logprobs
|
from sglang.srt.layers.logprob_processor import compute_spec_logprobs
|
||||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||||
from sglang.srt.managers.scheduler import GenerationBatchResult
|
from sglang.srt.managers.scheduler import GenerationBatchResult
|
||||||
from sglang.srt.managers.tp_worker import TpModelWorker
|
from sglang.srt.managers.tp_worker import TpModelWorker
|
||||||
@@ -127,7 +127,6 @@ class DSparkWorkerV2(BaseSpecWorker):
|
|||||||
self.draft_model_runner = bundle.draft_model_runner
|
self.draft_model_runner = bundle.draft_model_runner
|
||||||
self.draft_model = bundle.draft_model
|
self.draft_model = bundle.draft_model
|
||||||
self._draft_sampler = None
|
self._draft_sampler = None
|
||||||
self._linear_accept_index_cache = None
|
|
||||||
|
|
||||||
# The mask token is input-only (it is embedded, never sampled), so its
|
# The mask token is input-only (it is embedded, never sampled), so its
|
||||||
# bound is the embedding-table row count: the PADDED vocab when the
|
# bound is the embedding-table row count: the PADDED vocab when the
|
||||||
@@ -406,19 +405,6 @@ class DSparkWorkerV2(BaseSpecWorker):
|
|||||||
def note_request_finished(self, *, rid: str, natural_stop: bool) -> None:
|
def note_request_finished(self, *, rid: str, natural_stop: bool) -> None:
|
||||||
self._observers.note_request_finished(rid=rid, natural_stop=natural_stop)
|
self._observers.note_request_finished(rid=rid, natural_stop=natural_stop)
|
||||||
|
|
||||||
def _linear_accept_indices(self, bs: int) -> torch.Tensor:
|
|
||||||
num_indices = bs * self.verify_num_draft_tokens
|
|
||||||
if (
|
|
||||||
self._linear_accept_index_cache is None
|
|
||||||
or self._linear_accept_index_cache.numel() < num_indices
|
|
||||||
):
|
|
||||||
self._linear_accept_index_cache = torch.arange(
|
|
||||||
num_indices, dtype=torch.int64, device=self.device
|
|
||||||
)
|
|
||||||
return self._linear_accept_index_cache[:num_indices].view(
|
|
||||||
bs, self.verify_num_draft_tokens
|
|
||||||
)
|
|
||||||
|
|
||||||
def forward_batch_generation(
|
def forward_batch_generation(
|
||||||
self,
|
self,
|
||||||
batch: ScheduleBatch,
|
batch: ScheduleBatch,
|
||||||
@@ -718,12 +704,11 @@ class DSparkWorkerV2(BaseSpecWorker):
|
|||||||
draft_tokens=draft_tokens,
|
draft_tokens=draft_tokens,
|
||||||
)
|
)
|
||||||
if batch.return_logprob:
|
if batch.return_logprob:
|
||||||
compute_spec_v2_logprobs(
|
compute_spec_logprobs(
|
||||||
batch,
|
batch,
|
||||||
logits_output,
|
logits_output,
|
||||||
accept.out_tokens.reshape(-1),
|
accept.out_tokens.reshape(-1),
|
||||||
self._linear_accept_indices(bs),
|
chain_stride=self.verify_num_draft_tokens,
|
||||||
self.verify_num_draft_tokens - 1,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
if on_publish is not None:
|
if on_publish is not None:
|
||||||
|
|||||||
@@ -8,7 +8,7 @@ from sglang.kernels.ops.speculative.cache_locs import (
|
|||||||
assign_draft_cache_locs_contiguous,
|
assign_draft_cache_locs_contiguous,
|
||||||
)
|
)
|
||||||
from sglang.kernels.ops.speculative.eagle import fill_bonus_tokens_func
|
from sglang.kernels.ops.speculative.eagle import fill_bonus_tokens_func
|
||||||
from sglang.srt.layers.logprob_processor import compute_spec_v2_logprobs
|
from sglang.srt.layers.logprob_processor import compute_spec_logprobs
|
||||||
from sglang.srt.managers.utils import GenerationBatchResult
|
from sglang.srt.managers.utils import GenerationBatchResult
|
||||||
from sglang.srt.model_executor.forward_batch_info import (
|
from sglang.srt.model_executor.forward_batch_info import (
|
||||||
CaptureHiddenMode,
|
CaptureHiddenMode,
|
||||||
@@ -467,7 +467,6 @@ def run_eagle_verify(
|
|||||||
plan_stream: Any,
|
plan_stream: Any,
|
||||||
plan_stream_ctx: Any,
|
plan_stream_ctx: Any,
|
||||||
topk: int,
|
topk: int,
|
||||||
num_steps: int,
|
|
||||||
num_draft_tokens: int,
|
num_draft_tokens: int,
|
||||||
device: str,
|
device: str,
|
||||||
metadata_ready_pre_pad: bool,
|
metadata_ready_pre_pad: bool,
|
||||||
@@ -625,7 +624,7 @@ def run_eagle_verify(
|
|||||||
bonus_tokens = torch.empty((0,), device=device, dtype=torch.int32)
|
bonus_tokens = torch.empty((0,), device=device, dtype=torch.int32)
|
||||||
|
|
||||||
if batch.return_logprob and not batch.forward_mode.is_idle():
|
if batch.return_logprob and not batch.forward_mode.is_idle():
|
||||||
compute_spec_v2_logprobs(batch, logits_output, predict, accept_index, num_steps)
|
compute_spec_logprobs(batch, logits_output, predict, accept_index=accept_index)
|
||||||
|
|
||||||
if finalize_tree_path and not batch.forward_mode.is_idle() and topk > 1:
|
if finalize_tree_path and not batch.forward_mode.is_idle() and topk > 1:
|
||||||
# topk == 1 needs nothing here: the accepted path is already the front
|
# topk == 1 needs nothing here: the accepted path is already the front
|
||||||
|
|||||||
@@ -1503,7 +1503,6 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
|||||||
plan_stream=self.plan_stream,
|
plan_stream=self.plan_stream,
|
||||||
plan_stream_ctx=self.plan_stream_ctx,
|
plan_stream_ctx=self.plan_stream_ctx,
|
||||||
topk=self.topk,
|
topk=self.topk,
|
||||||
num_steps=self.speculative_num_steps,
|
|
||||||
num_draft_tokens=self.speculative_num_draft_tokens,
|
num_draft_tokens=self.speculative_num_draft_tokens,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
metadata_ready_pre_pad=False,
|
metadata_ready_pre_pad=False,
|
||||||
|
|||||||
@@ -1038,7 +1038,6 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
|
|||||||
plan_stream=self.plan_stream,
|
plan_stream=self.plan_stream,
|
||||||
plan_stream_ctx=self.plan_stream_ctx,
|
plan_stream_ctx=self.plan_stream_ctx,
|
||||||
topk=self.topk,
|
topk=self.topk,
|
||||||
num_steps=self.speculative_num_steps,
|
|
||||||
num_draft_tokens=self.speculative_num_draft_tokens,
|
num_draft_tokens=self.speculative_num_draft_tokens,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
metadata_ready_pre_pad=False,
|
metadata_ready_pre_pad=False,
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ from sglang.kernels.ops.speculative.cache_locs import (
|
|||||||
assign_extend_cache_locs_func as assign_extend_cache_locs_func,
|
assign_extend_cache_locs_func as assign_extend_cache_locs_func,
|
||||||
)
|
)
|
||||||
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
||||||
from sglang.srt.layers.logprob_processor import compute_spec_v2_logprobs
|
from sglang.srt.layers.logprob_processor import compute_spec_logprobs
|
||||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||||
from sglang.srt.managers.scheduler import GenerationBatchResult
|
from sglang.srt.managers.scheduler import GenerationBatchResult
|
||||||
from sglang.srt.managers.tp_worker import TpModelWorker
|
from sglang.srt.managers.tp_worker import TpModelWorker
|
||||||
@@ -480,15 +480,11 @@ class NGRAMWorker(BaseSpecWorker):
|
|||||||
self.token_to_kv_pool_allocator,
|
self.token_to_kv_pool_allocator,
|
||||||
)
|
)
|
||||||
if batch.return_logprob:
|
if batch.return_logprob:
|
||||||
# The last arg is the accept_index row width minus 1. NGRAM's
|
compute_spec_logprobs(
|
||||||
# accept_index is (bs, draft_token_num) -- the tree depth is not
|
|
||||||
# bounded by spec_steps like EAGLE's (bs, spec_steps + 1).
|
|
||||||
compute_spec_v2_logprobs(
|
|
||||||
batch,
|
batch,
|
||||||
logits_output,
|
logits_output,
|
||||||
predict,
|
predict,
|
||||||
accept_index,
|
accept_index=accept_index,
|
||||||
self.draft_token_num - 1,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
if on_publish is not None:
|
if on_publish is not None:
|
||||||
|
|||||||
Reference in New Issue
Block a user