[Spec] Extract move_accept_tokens_to_target_kvcache into spec_utils (#27764)
This commit is contained in:
@@ -39,15 +39,9 @@ from sglang.srt.speculative.spec_utils import (
|
|||||||
from sglang.srt.speculative.triton_ops.cache_locs import (
|
from sglang.srt.speculative.triton_ops.cache_locs import (
|
||||||
assign_draft_cache_locs_contiguous as assign_draft_cache_locs_contiguous,
|
assign_draft_cache_locs_contiguous as assign_draft_cache_locs_contiguous,
|
||||||
)
|
)
|
||||||
from sglang.srt.speculative.triton_ops.cache_locs import (
|
|
||||||
assign_extend_cache_locs as assign_extend_cache_locs,
|
|
||||||
)
|
|
||||||
from sglang.srt.speculative.triton_ops.cache_locs import (
|
from sglang.srt.speculative.triton_ops.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.speculative.triton_ops.eagle import (
|
|
||||||
fill_accept_out_cache_loc as fill_accept_out_cache_loc,
|
|
||||||
)
|
|
||||||
from sglang.srt.speculative.triton_ops.eagle import (
|
from sglang.srt.speculative.triton_ops.eagle import (
|
||||||
fill_bonus_tokens as fill_bonus_tokens,
|
fill_bonus_tokens as fill_bonus_tokens,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -56,11 +56,7 @@ from sglang.srt.speculative.eagle_draft_extend_cuda_graph_runner import (
|
|||||||
EAGLEDraftExtendCudaGraphRunner,
|
EAGLEDraftExtendCudaGraphRunner,
|
||||||
)
|
)
|
||||||
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 fill_bonus_tokens
|
||||||
assign_extend_cache_locs,
|
|
||||||
fill_accept_out_cache_loc,
|
|
||||||
fill_bonus_tokens,
|
|
||||||
)
|
|
||||||
from sglang.srt.speculative.eagle_utils import (
|
from sglang.srt.speculative.eagle_utils import (
|
||||||
TreeMaskMode,
|
TreeMaskMode,
|
||||||
_eagle_prefill_tail_tokens,
|
_eagle_prefill_tail_tokens,
|
||||||
@@ -73,6 +69,7 @@ from sglang.srt.speculative.spec_utils import (
|
|||||||
draft_tp_context,
|
draft_tp_context,
|
||||||
generate_token_bitmask,
|
generate_token_bitmask,
|
||||||
load_token_map,
|
load_token_map,
|
||||||
|
move_accept_tokens_to_target_kvcache,
|
||||||
record_stream_each,
|
record_stream_each,
|
||||||
record_stream_for_v2_verify,
|
record_stream_for_v2_verify,
|
||||||
select_top_k_tokens,
|
select_top_k_tokens,
|
||||||
@@ -93,7 +90,6 @@ from sglang.srt.utils.common import (
|
|||||||
is_musa,
|
is_musa,
|
||||||
is_npu,
|
is_npu,
|
||||||
log_info_on_rank0,
|
log_info_on_rank0,
|
||||||
next_power_of_2,
|
|
||||||
)
|
)
|
||||||
from sglang.srt.utils.patch_torch import monkey_patch_torch_reductions
|
from sglang.srt.utils.patch_torch import monkey_patch_torch_reductions
|
||||||
|
|
||||||
@@ -1414,7 +1410,9 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
|||||||
downstream chain-layout code (draft-extend select_index, committed-KV reads)
|
downstream chain-layout code (draft-extend select_index, committed-KV reads)
|
||||||
assumes. Returns compacted predict; mutates logits_output.hidden_states
|
assumes. Returns compacted predict; mutates logits_output.hidden_states
|
||||||
(moved only when present)."""
|
(moved only when present)."""
|
||||||
self.move_accept_tokens_to_target_kvcache(batch, accept_index, accept_lens - 1)
|
move_accept_tokens_to_target_kvcache(
|
||||||
|
batch, accept_index, accept_lens - 1, self.token_to_kv_pool_allocator
|
||||||
|
)
|
||||||
predict = self._compact_accept_to_front(predict, accept_index, bs)
|
predict = self._compact_accept_to_front(predict, accept_index, bs)
|
||||||
if logits_output.hidden_states is not None:
|
if logits_output.hidden_states is not None:
|
||||||
logits_output.hidden_states = self._compact_accept_to_front(
|
logits_output.hidden_states = self._compact_accept_to_front(
|
||||||
@@ -1422,59 +1420,6 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
|||||||
)
|
)
|
||||||
return predict
|
return predict
|
||||||
|
|
||||||
def move_accept_tokens_to_target_kvcache(
|
|
||||||
self,
|
|
||||||
batch: ScheduleBatch,
|
|
||||||
accept_index: torch.Tensor,
|
|
||||||
num_correct_drafts: torch.Tensor,
|
|
||||||
):
|
|
||||||
"""
|
|
||||||
Move accepted tokens (drafts + bonus) to the target KV cache.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
batch: The batch to run.
|
|
||||||
accept_index: The index of the accepted tokens (incl. bonus).
|
|
||||||
num_correct_drafts: Per-req count of correct drafts (excludes bonus);
|
|
||||||
seq_lens is advanced by num_correct_drafts + 1 to cover the bonus slot.
|
|
||||||
"""
|
|
||||||
bs = len(batch.seq_lens)
|
|
||||||
# accept_index element count, NOT bs * num_draft_tokens: for topk > 1 the
|
|
||||||
# tree exceeds the accepted chain, over-reading accept_index (illegal memory).
|
|
||||||
size = bs * accept_index.shape[1]
|
|
||||||
|
|
||||||
# fill_accept_out_cache_loc reads out_cache_loc[accept_index]; -1 sentinel ok.
|
|
||||||
maybe_detect_oob(
|
|
||||||
accept_index,
|
|
||||||
-1,
|
|
||||||
batch.out_cache_loc.size(0),
|
|
||||||
"eagle v2 move_accepted_tokens accept_index",
|
|
||||||
)
|
|
||||||
|
|
||||||
tgt_cache_loc = torch.zeros(
|
|
||||||
size,
|
|
||||||
dtype=torch.int64,
|
|
||||||
device=self.device,
|
|
||||||
)
|
|
||||||
accept_out_cache_loc = torch.zeros(size, dtype=torch.int64, device=self.device)
|
|
||||||
assign_extend_cache_locs[(bs,)](
|
|
||||||
batch.req_pool_indices,
|
|
||||||
self.req_to_token_pool.req_to_token,
|
|
||||||
batch.seq_lens,
|
|
||||||
batch.seq_lens + num_correct_drafts + 1,
|
|
||||||
tgt_cache_loc,
|
|
||||||
self.req_to_token_pool.req_to_token.shape[1],
|
|
||||||
next_power_of_2(bs),
|
|
||||||
)
|
|
||||||
fill_accept_out_cache_loc[(size,)](
|
|
||||||
accept_index,
|
|
||||||
batch.out_cache_loc,
|
|
||||||
accept_out_cache_loc,
|
|
||||||
next_power_of_2(size),
|
|
||||||
)
|
|
||||||
self.token_to_kv_pool_allocator.get_kvcache().move_kv_cache(
|
|
||||||
tgt_cache_loc, accept_out_cache_loc
|
|
||||||
)
|
|
||||||
|
|
||||||
def _compact_accept_to_front(
|
def _compact_accept_to_front(
|
||||||
self, x: torch.Tensor, accept_index: torch.Tensor, bs: int
|
self, x: torch.Tensor, accept_index: torch.Tensor, bs: int
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
|
|||||||
@@ -18,6 +18,9 @@ from sglang.srt.server_args import get_global_server_args
|
|||||||
from sglang.srt.speculative.triton_ops.cache_locs import (
|
from sglang.srt.speculative.triton_ops.cache_locs import (
|
||||||
align_evict_mask_to_page_size as align_evict_mask_to_page_size,
|
align_evict_mask_to_page_size as align_evict_mask_to_page_size,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.speculative.triton_ops.cache_locs import (
|
||||||
|
assign_extend_cache_locs as assign_extend_cache_locs,
|
||||||
|
)
|
||||||
from sglang.srt.speculative.triton_ops.cache_locs import (
|
from sglang.srt.speculative.triton_ops.cache_locs import (
|
||||||
assign_req_to_token_pool as assign_req_to_token_pool,
|
assign_req_to_token_pool as assign_req_to_token_pool,
|
||||||
)
|
)
|
||||||
@@ -39,7 +42,11 @@ from sglang.srt.speculative.triton_ops.cache_locs import (
|
|||||||
from sglang.srt.speculative.triton_ops.cache_locs import (
|
from sglang.srt.speculative.triton_ops.cache_locs import (
|
||||||
get_target_cache_loc as get_target_cache_loc,
|
get_target_cache_loc as get_target_cache_loc,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils import is_cuda, is_hip, is_musa, is_npu
|
from sglang.srt.speculative.triton_ops.eagle import (
|
||||||
|
fill_accept_out_cache_loc as fill_accept_out_cache_loc,
|
||||||
|
)
|
||||||
|
from sglang.srt.utils import is_cuda, is_hip, is_musa, is_npu, next_power_of_2
|
||||||
|
from sglang.srt.utils.async_probe import maybe_detect_oob
|
||||||
|
|
||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
_is_hip = is_hip()
|
_is_hip = is_hip()
|
||||||
@@ -48,7 +55,8 @@ _is_musa = is_musa()
|
|||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.constrained.base_grammar_backend import BaseGrammarObject
|
from sglang.srt.constrained.base_grammar_backend import BaseGrammarObject
|
||||||
from sglang.srt.managers.schedule_batch import Req
|
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
|
||||||
|
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
from sglang.srt.speculative.eagle_info import EagleVerifyInput
|
from sglang.srt.speculative.eagle_info import EagleVerifyInput
|
||||||
|
|
||||||
@@ -473,3 +481,58 @@ def spec_stage_span(name: str):
|
|||||||
if torch.autograd._profiler_enabled():
|
if torch.autograd._profiler_enabled():
|
||||||
return torch.profiler.record_function(name)
|
return torch.profiler.record_function(name)
|
||||||
return nullcontext()
|
return nullcontext()
|
||||||
|
|
||||||
|
|
||||||
|
def move_accept_tokens_to_target_kvcache(
|
||||||
|
batch: ScheduleBatch,
|
||||||
|
accept_index: torch.Tensor,
|
||||||
|
num_correct_drafts: torch.Tensor,
|
||||||
|
token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Move accepted tokens (drafts + bonus) to the target KV cache.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
batch: The batch to run.
|
||||||
|
accept_index: The index of the accepted tokens (incl. bonus).
|
||||||
|
num_correct_drafts: Per-req count of correct drafts (excludes bonus);
|
||||||
|
seq_lens is advanced by ``num_correct_drafts + 1`` to cover the bonus slot.
|
||||||
|
"""
|
||||||
|
bs = len(batch.seq_lens)
|
||||||
|
device = batch.seq_lens.device
|
||||||
|
# accept_index element count, NOT bs * num_draft_tokens: for topk > 1 the
|
||||||
|
# tree exceeds the accepted chain, over-reading accept_index (illegal memory).
|
||||||
|
size = bs * accept_index.shape[1]
|
||||||
|
|
||||||
|
# fill_accept_out_cache_loc reads out_cache_loc[accept_index]; -1 sentinel ok.
|
||||||
|
maybe_detect_oob(
|
||||||
|
accept_index,
|
||||||
|
-1,
|
||||||
|
batch.out_cache_loc.size(0),
|
||||||
|
"spec v2 move_accept_tokens accept_index",
|
||||||
|
)
|
||||||
|
|
||||||
|
tgt_cache_loc = torch.zeros(
|
||||||
|
size,
|
||||||
|
dtype=torch.int64,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
accept_out_cache_loc = torch.zeros(size, dtype=torch.int64, device=device)
|
||||||
|
assign_extend_cache_locs[(bs,)](
|
||||||
|
batch.req_pool_indices,
|
||||||
|
batch.req_to_token_pool.req_to_token,
|
||||||
|
batch.seq_lens,
|
||||||
|
batch.seq_lens + num_correct_drafts + 1,
|
||||||
|
tgt_cache_loc,
|
||||||
|
batch.req_to_token_pool.req_to_token.shape[1],
|
||||||
|
next_power_of_2(bs),
|
||||||
|
)
|
||||||
|
fill_accept_out_cache_loc[(size,)](
|
||||||
|
accept_index,
|
||||||
|
batch.out_cache_loc,
|
||||||
|
accept_out_cache_loc,
|
||||||
|
next_power_of_2(size),
|
||||||
|
)
|
||||||
|
token_to_kv_pool_allocator.get_kvcache().move_kv_cache(
|
||||||
|
tgt_cache_loc, accept_out_cache_loc
|
||||||
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user