[Spec] Stage Inkling MTP draft metadata before verify (#38169)
Co-authored-by: Qiaolin-Yu <liin1211@outlook.com>
This commit is contained in:
co-authored by
Qiaolin-Yu
parent
a63efd9056
commit
bb15be6d79
@@ -129,6 +129,8 @@ class AttentionBackend(ABC):
|
|||||||
Default: no-op.
|
Default: no-op.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
supports_draft_extend_metadata_staging: bool = False
|
||||||
|
|
||||||
def draft_extend_metadata_captured_in_graph(self) -> bool:
|
def draft_extend_metadata_captured_in_graph(self) -> bool:
|
||||||
"""True when :py:meth:`init_forward_metadata_in_graph` fully rebuilds
|
"""True when :py:meth:`init_forward_metadata_in_graph` fully rebuilds
|
||||||
this backend's DRAFT_EXTEND_V2 replay metadata inside the captured
|
this backend's DRAFT_EXTEND_V2 replay metadata inside the captured
|
||||||
|
|||||||
@@ -451,6 +451,20 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def supports_draft_extend_metadata_staging(self) -> bool:
|
||||||
|
return (
|
||||||
|
self.topk == 1
|
||||||
|
and not self.kv_index_translator.is_translating
|
||||||
|
and self.draft_extend_metadata_captured_in_graph()
|
||||||
|
)
|
||||||
|
|
||||||
|
def stage_draft_extend_metadata(self, forward_batch: ForwardBatch):
|
||||||
|
self.forward_metadata = self.draft_extend_metadata[forward_batch.batch_size]
|
||||||
|
self.forward_metadata.max_seq_len_k = self.max_context_len
|
||||||
|
self.forward_metadata_spec_decode_expand = None
|
||||||
|
self.init_forward_metadata_in_graph(forward_batch)
|
||||||
|
|
||||||
def _in_graph_full_to_swa_index_mapping(self) -> Optional[torch.Tensor]:
|
def _in_graph_full_to_swa_index_mapping(self) -> Optional[torch.Tensor]:
|
||||||
# The in-graph SWA translation needs the raw mapping tensor; v2p-table
|
# The in-graph SWA translation needs the raw mapping tensor; v2p-table
|
||||||
# pools (UnifiedSWAKVPool) keep it None and must stay on the eager
|
# pools (UnifiedSWAKVPool) keep it None and must stay on the eager
|
||||||
|
|||||||
@@ -501,16 +501,22 @@ class InklingShortConvAttnBackend(ShortConvAttnBackend):
|
|||||||
mamba_track_indices: Optional[torch.Tensor],
|
mamba_track_indices: Optional[torch.Tensor],
|
||||||
mamba_steps_to_track: Optional[torch.Tensor],
|
mamba_steps_to_track: Optional[torch.Tensor],
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Commit the TARGET_VERIFY conv windows at each request's last accepted step.
|
"""Commit the TARGET_VERIFY conv windows at each request's last accepted step."""
|
||||||
|
|
||||||
Slot ids come from ``req_pool_indices``, not the per-step
|
|
||||||
``self._cache_indices``: this runs after the forward context exits, so that
|
|
||||||
buffer may already belong to a later forward.
|
|
||||||
"""
|
|
||||||
pool = self.req_to_token_pool
|
pool = self.req_to_token_pool
|
||||||
|
bs = req_pool_indices.shape[0]
|
||||||
|
if self._slot_gather_recordable:
|
||||||
|
assert (
|
||||||
|
self._cache_indices_buf is not None
|
||||||
|
and self._cache_indices_buf.shape[0] >= bs
|
||||||
|
)
|
||||||
|
slot_ids = self._cache_indices_buf[:bs]
|
||||||
|
else:
|
||||||
|
slot_ids = self._translate_mamba_indices(
|
||||||
|
pool.get_mamba_indices(req_pool_indices)
|
||||||
|
)
|
||||||
scatter_mamba_states_after_mtp_verify(
|
scatter_mamba_states_after_mtp_verify(
|
||||||
pool.get_speculative_mamba2_params_all_layers(),
|
pool.get_speculative_mamba2_params_all_layers(),
|
||||||
self._translate_mamba_indices(pool.get_mamba_indices(req_pool_indices)),
|
slot_ids,
|
||||||
last_correct_step_indices,
|
last_correct_step_indices,
|
||||||
mamba_track_indices,
|
mamba_track_indices,
|
||||||
mamba_steps_to_track,
|
mamba_steps_to_track,
|
||||||
@@ -548,6 +554,43 @@ class InklingShortConvHybridAttnBackend(ShortConvHybridAttnBackend):
|
|||||||
is Inkling's own, not the generic mamba scatter.
|
is Inkling's own, not the generic mamba scatter.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
@property
|
||||||
|
def supports_draft_extend_metadata_staging(self) -> bool:
|
||||||
|
return (
|
||||||
|
self.full_attn_backend.supports_draft_extend_metadata_staging
|
||||||
|
and self.short_conv_backend._slot_gather_recordable
|
||||||
|
)
|
||||||
|
|
||||||
|
def init_forward_metadata_out_graph(
|
||||||
|
self, forward_batch: ForwardBatch, in_capture: bool = False
|
||||||
|
):
|
||||||
|
if (
|
||||||
|
forward_batch.forward_mode.is_draft_extend_v2()
|
||||||
|
and self.supports_draft_extend_metadata_staging
|
||||||
|
):
|
||||||
|
if in_capture:
|
||||||
|
self.full_attn_backend.init_forward_metadata_out_graph(
|
||||||
|
forward_batch, in_capture=True
|
||||||
|
)
|
||||||
|
self.full_attn_backend.stage_draft_extend_metadata(forward_batch)
|
||||||
|
self.short_conv_backend._prepare_slot_indices(forward_batch)
|
||||||
|
else:
|
||||||
|
super().init_forward_metadata_out_graph(
|
||||||
|
forward_batch, in_capture=in_capture
|
||||||
|
)
|
||||||
|
|
||||||
|
def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch):
|
||||||
|
if (
|
||||||
|
forward_batch.forward_mode.is_draft_extend_v2()
|
||||||
|
and self.supports_draft_extend_metadata_staging
|
||||||
|
):
|
||||||
|
self.short_conv_backend._reset_step_state()
|
||||||
|
self.short_conv_backend._refresh_sconv_metadata(
|
||||||
|
forward_batch, on_graph_path=True
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
super().init_forward_metadata_in_graph(forward_batch)
|
||||||
|
|
||||||
def sconv_state(self, *, layer_id: int, stream: int) -> torch.Tensor:
|
def sconv_state(self, *, layer_id: int, stream: int) -> torch.Tensor:
|
||||||
return self.short_conv_backend.sconv_state(layer_id=layer_id, stream=stream)
|
return self.short_conv_backend.sconv_state(layer_id=layer_id, stream=stream)
|
||||||
|
|
||||||
@@ -610,4 +653,7 @@ class InklingShortConvHybridAttnBackend(ShortConvHybridAttnBackend):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def draft_extend_metadata_captured_in_graph(self) -> bool:
|
def draft_extend_metadata_captured_in_graph(self) -> bool:
|
||||||
return self.full_attn_backend.draft_extend_metadata_captured_in_graph()
|
return (
|
||||||
|
not self.supports_draft_extend_metadata_staging
|
||||||
|
and self.full_attn_backend.draft_extend_metadata_captured_in_graph()
|
||||||
|
)
|
||||||
|
|||||||
@@ -493,7 +493,10 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
out_cache_loc=buffers.out_cache_loc[:num_tokens],
|
out_cache_loc=buffers.out_cache_loc[:num_tokens],
|
||||||
spec_info=spec_info,
|
spec_info=spec_info,
|
||||||
)
|
)
|
||||||
if not self.metadata_captured_in_graph:
|
if (
|
||||||
|
not self.metadata_captured_in_graph
|
||||||
|
and not self.attn_backend.supports_draft_extend_metadata_staging
|
||||||
|
):
|
||||||
self.eagle_worker.draft_extend_attn_backend_list[
|
self.eagle_worker.draft_extend_attn_backend_list[
|
||||||
self.step
|
self.step
|
||||||
].init_forward_metadata_out_graph(fb_view)
|
].init_forward_metadata_out_graph(fb_view)
|
||||||
@@ -712,7 +715,47 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner:
|
|||||||
def _prepare_extra(self, forward_batch: ForwardBatch) -> None:
|
def _prepare_extra(self, forward_batch: ForwardBatch) -> None:
|
||||||
"""Hook for subclasses to populate extra per-call buffers (e.g. sconv)."""
|
"""Hook for subclasses to populate extra per-call buffers (e.g. sconv)."""
|
||||||
|
|
||||||
def prepare(self, forward_batch: ForwardBatch):
|
def stage_shared_reads(
|
||||||
|
self, *, seq_lens, req_pool_indices, out_cache_loc, positions=None
|
||||||
|
):
|
||||||
|
raw_bs = req_pool_indices.shape[0]
|
||||||
|
bs = self.get_runner(0)._pad_to_bucket(raw_bs, self.capture_bs)
|
||||||
|
buffers = self.buffers
|
||||||
|
buffers.seq_lens[:bs].fill_(self.seq_len_fill_value)
|
||||||
|
buffers.seq_lens[:raw_bs].copy_(seq_lens)
|
||||||
|
buffers.req_pool_indices[:bs].zero_()
|
||||||
|
buffers.req_pool_indices[:raw_bs].copy_(req_pool_indices)
|
||||||
|
num_tokens = raw_bs * self.captured_req_width
|
||||||
|
buffers.out_cache_loc[: bs * self.captured_req_width].zero_()
|
||||||
|
buffers.out_cache_loc[:num_tokens].copy_(out_cache_loc)
|
||||||
|
if positions is not None:
|
||||||
|
buffers.positions[:num_tokens].copy_(positions)
|
||||||
|
self._stage_metadata(bs, raw_bs)
|
||||||
|
self._staged_bs = bs
|
||||||
|
|
||||||
|
def _stage_metadata(self, bs: int, raw_bs: int):
|
||||||
|
backends = [
|
||||||
|
b
|
||||||
|
for b in self.draft_extend_attn_backend_list
|
||||||
|
if b.supports_draft_extend_metadata_staging
|
||||||
|
and not b.draft_extend_metadata_captured_in_graph()
|
||||||
|
]
|
||||||
|
if not backends:
|
||||||
|
return
|
||||||
|
buffers = self.buffers
|
||||||
|
buffers.req_pool_indices[raw_bs:bs].zero_()
|
||||||
|
batch = SimpleNamespace(
|
||||||
|
batch_size=bs,
|
||||||
|
forward_mode=ForwardMode.DRAFT_EXTEND_V2,
|
||||||
|
req_pool_indices=buffers.req_pool_indices[:bs],
|
||||||
|
seq_lens=buffers.seq_lens[:bs],
|
||||||
|
extend_seq_lens=buffers.extend_seq_lens[:bs],
|
||||||
|
out_cache_loc=buffers.out_cache_loc[: bs * self.captured_req_width],
|
||||||
|
)
|
||||||
|
for backend in backends:
|
||||||
|
backend.init_forward_metadata_out_graph(batch)
|
||||||
|
|
||||||
|
def prepare(self, forward_batch: ForwardBatch, *, staged: bool = False):
|
||||||
"""Populate the shared buffers once from ``forward_batch`` and bucketize
|
"""Populate the shared buffers once from ``forward_batch`` and bucketize
|
||||||
the batch size. Subsequent ``replay(step)`` calls reuse this state."""
|
the batch size. Subsequent ``replay(step)`` calls reuse this state."""
|
||||||
buffers = self.buffers
|
buffers = self.buffers
|
||||||
@@ -801,6 +844,10 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner:
|
|||||||
seq_lens_sum = seq_lens_sum + (bs - raw_bs) * self.seq_len_fill_value
|
seq_lens_sum = seq_lens_sum + (bs - raw_bs) * self.seq_len_fill_value
|
||||||
self.seq_lens_sum = seq_lens_sum
|
self.seq_lens_sum = seq_lens_sum
|
||||||
|
|
||||||
|
if staged:
|
||||||
|
assert bs == self._staged_bs
|
||||||
|
else:
|
||||||
|
self._stage_metadata(bs, raw_bs)
|
||||||
self._prepare_extra(forward_batch)
|
self._prepare_extra(forward_batch)
|
||||||
|
|
||||||
def replay(self, step: int):
|
def replay(self, step: int):
|
||||||
@@ -852,10 +899,9 @@ class OneGraphMultiLayerEagleMultiStepDraftExtendCudaGraphRunner(
|
|||||||
forwards + the inter-step input_ids rotation in ONE graph per bucket, instead
|
forwards + the inter-step input_ids rotation in ONE graph per bucket, instead
|
||||||
of one graph per step. The worker drops its per-step rotation (rotates_in_graph).
|
of one graph per step. The worker drops its per-step rotation (rotates_in_graph).
|
||||||
|
|
||||||
Each step's replay metadata is emitted in-graph via
|
Each step refreshes metadata in-graph or stages it before replay; no Python
|
||||||
init_forward_metadata_in_graph (no Python may run between
|
may run between captured steps. seq_lens / req_pool_indices / extend_seq_lens
|
||||||
captured steps). seq_lens / req_pool_indices / extend_seq_lens are chain-constant
|
are chain-constant (only input_ids rotates).
|
||||||
(only input_ids rotates), so per-step in-graph metadata is correct.
|
|
||||||
|
|
||||||
Rejection sampling is supported by sampling X ~ q inside the graph
|
Rejection sampling is supported by sampling X ~ q inside the graph
|
||||||
(_sample_draft_proposal, selected by the draft_probs buffer's presence):
|
(_sample_draft_proposal, selected by the draft_probs buffer's presence):
|
||||||
|
|||||||
@@ -114,6 +114,8 @@ logger = logging.getLogger(__name__)
|
|||||||
|
|
||||||
|
|
||||||
class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
|
class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
|
||||||
|
last_draft_extend_staged: bool = False
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
@@ -289,8 +291,8 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
|
|
||||||
def _compute_boundary_kv_locs_positions(self, batch):
|
def _compute_boundary_kv_locs_positions(self, batch):
|
||||||
if self.draft_extend_num_front_tokens == 0 or batch.forward_mode.is_idle():
|
if self.draft_extend_num_front_tokens == 0 or batch.forward_mode.is_idle():
|
||||||
return None, None, None
|
return None, None
|
||||||
locs, positions = compute_widened_draft_extend_locs_positions(
|
return compute_widened_draft_extend_locs_positions(
|
||||||
batch.seq_lens,
|
batch.seq_lens,
|
||||||
batch.req_pool_indices,
|
batch.req_pool_indices,
|
||||||
self.req_to_token_pool.req_to_token,
|
self.req_to_token_pool.req_to_token,
|
||||||
@@ -299,11 +301,6 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
self.draft_extend_num_front_tokens,
|
self.draft_extend_num_front_tokens,
|
||||||
self.draft_extend_num_warmup_tokens,
|
self.draft_extend_num_warmup_tokens,
|
||||||
)
|
)
|
||||||
ready_event = None
|
|
||||||
if self.plan_stream:
|
|
||||||
ready_event = torch.get_device_module(self.device).Event()
|
|
||||||
ready_event.record()
|
|
||||||
return locs, positions, ready_event
|
|
||||||
|
|
||||||
def _seed_boundary_kv_stash(self, forward_batch, target_hidden_states):
|
def _seed_boundary_kv_stash(self, forward_batch, target_hidden_states):
|
||||||
if (
|
if (
|
||||||
@@ -411,15 +408,11 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
before_mem = get_available_gpu_memory(self.device, self.gpu_id)
|
before_mem = get_available_gpu_memory(self.device, self.gpu_id)
|
||||||
|
|
||||||
if not _is_npu:
|
if not _is_npu:
|
||||||
# The single-CG runner replays with no Python between steps, so the
|
# Every step must refresh metadata in-graph or stage it before replay.
|
||||||
# attn backend must fully rebuild its per-step metadata as captured
|
|
||||||
# tensor ops; anything less gets capture-time-stale metadata (e.g.
|
|
||||||
# SWA translations, which only the eager replay path refreshes).
|
|
||||||
# Per-depth pools (banded MTP) mean per-depth backends — EVERY step
|
|
||||||
# must satisfy this, not just step 0.
|
|
||||||
draft_backend = self.draft_runner_list[0].attn_backend
|
draft_backend = self.draft_runner_list[0].attn_backend
|
||||||
backend_supports_single_cg = all(
|
backend_supports_single_cg = all(
|
||||||
runner.attn_backend.draft_extend_metadata_captured_in_graph()
|
runner.attn_backend.draft_extend_metadata_captured_in_graph()
|
||||||
|
or runner.attn_backend.supports_draft_extend_metadata_staging
|
||||||
for runner in self.draft_runner_list
|
for runner in self.draft_runner_list
|
||||||
)
|
)
|
||||||
if envs.SGLANG_ENABLE_SINGLE_CG_DRAFT.get() and backend_supports_single_cg:
|
if envs.SGLANG_ENABLE_SINGLE_CG_DRAFT.get() and backend_supports_single_cg:
|
||||||
@@ -429,8 +422,8 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
else:
|
else:
|
||||||
if envs.SGLANG_ENABLE_SINGLE_CG_DRAFT.get():
|
if envs.SGLANG_ENABLE_SINGLE_CG_DRAFT.get():
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"SGLANG_ENABLE_SINGLE_CG_DRAFT is on but %s does not fully "
|
"SGLANG_ENABLE_SINGLE_CG_DRAFT is on but %s cannot refresh "
|
||||||
"rebuild its draft-extend metadata in-graph; falling back "
|
"its draft-extend metadata for a combined graph; falling back "
|
||||||
"to per-step draft graphs.",
|
"to per-step draft graphs.",
|
||||||
type(draft_backend).__name__,
|
type(draft_backend).__name__,
|
||||||
)
|
)
|
||||||
@@ -724,8 +717,57 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
|
|
||||||
return next_draft_input
|
return next_draft_input
|
||||||
|
|
||||||
|
def _draft_extend_plan_for_decode(self, batch: ScheduleBatch) -> bool:
|
||||||
|
runner = self.cuda_graph_runner_for_draft_extend
|
||||||
|
if runner is None:
|
||||||
|
return False
|
||||||
|
assert self.draft_extend_attn_backend_list, (
|
||||||
|
"Draft graphs require initialized attention backends"
|
||||||
|
)
|
||||||
|
target_graph = self.target_worker.model_runner.decode_cuda_graph_runner
|
||||||
|
# In-graph mapping reads still require the final draft fence.
|
||||||
|
if (
|
||||||
|
batch.forward_mode.is_idle()
|
||||||
|
or getattr(target_graph, "in_graph_metadata_prep_done", None) is None
|
||||||
|
or not all(
|
||||||
|
b.supports_draft_extend_metadata_staging
|
||||||
|
and not b.draft_extend_metadata_captured_in_graph()
|
||||||
|
for b in self.draft_extend_attn_backend_list
|
||||||
|
)
|
||||||
|
or runner.require_mlp_tp_gather
|
||||||
|
or len(batch.seq_lens) > runner.max_bs
|
||||||
|
or (runner.disable_padding and len(batch.seq_lens) not in runner.capture_bs)
|
||||||
|
):
|
||||||
|
return False
|
||||||
|
assert self.topk == 1, "Draft-extend metadata staging requires topk=1"
|
||||||
|
locs, positions = self._compute_boundary_kv_locs_positions(batch)
|
||||||
|
if locs is None:
|
||||||
|
from sglang.kernels.ops.speculative.cache_locs import (
|
||||||
|
assign_extend_cache_locs_uniform_func,
|
||||||
|
)
|
||||||
|
|
||||||
|
locs = assign_extend_cache_locs_uniform_func(
|
||||||
|
req_pool_indices=batch.req_pool_indices,
|
||||||
|
req_to_token=self.req_to_token_pool.req_to_token,
|
||||||
|
start_offset=batch.seq_lens,
|
||||||
|
batch_size=len(batch.seq_lens),
|
||||||
|
draft_token_num=self.speculative_num_draft_tokens,
|
||||||
|
device=batch.device,
|
||||||
|
)
|
||||||
|
runner.stage_shared_reads(
|
||||||
|
seq_lens=batch.seq_lens + self.speculative_num_draft_tokens,
|
||||||
|
req_pool_indices=batch.req_pool_indices,
|
||||||
|
out_cache_loc=locs,
|
||||||
|
positions=positions,
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
|
||||||
def _draft_extend_for_decode(
|
def _draft_extend_for_decode(
|
||||||
self, batch: ScheduleBatch, batch_result: GenerationBatchResult
|
self,
|
||||||
|
batch: ScheduleBatch,
|
||||||
|
batch_result: GenerationBatchResult,
|
||||||
|
*,
|
||||||
|
staged: bool = False,
|
||||||
):
|
):
|
||||||
# Batch 2: Draft extend
|
# Batch 2: Draft extend
|
||||||
draft_extend_input = EagleDraftExtendInput(
|
draft_extend_input = EagleDraftExtendInput(
|
||||||
@@ -742,9 +784,19 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
|
|
||||||
# Prepare for draft extend in a separate stream
|
# Prepare for draft extend in a separate stream
|
||||||
# Notice that here we use batch_result.next_token_ids as the input ids
|
# Notice that here we use batch_result.next_token_ids as the input ids
|
||||||
boundary_kv_locs, boundary_kv_positions, boundary_kv_ready_event = (
|
if staged and self.draft_extend_num_front_tokens:
|
||||||
self._compute_boundary_kv_locs_positions(batch)
|
runner = self.cuda_graph_runner_for_draft_extend
|
||||||
)
|
num_tokens = len(batch.seq_lens) * runner.captured_req_width
|
||||||
|
boundary_kv_locs = runner.buffers.out_cache_loc[:num_tokens]
|
||||||
|
boundary_kv_positions = runner.buffers.positions[:num_tokens]
|
||||||
|
else:
|
||||||
|
boundary_kv_locs, boundary_kv_positions = (
|
||||||
|
self._compute_boundary_kv_locs_positions(batch)
|
||||||
|
)
|
||||||
|
boundary_kv_ready_event = None
|
||||||
|
if boundary_kv_locs is not None and self.plan_stream:
|
||||||
|
boundary_kv_ready_event = torch.get_device_module(self.device).Event()
|
||||||
|
boundary_kv_ready_event.record()
|
||||||
|
|
||||||
with self.plan_stream_ctx:
|
with self.plan_stream_ctx:
|
||||||
if boundary_kv_ready_event is not None:
|
if boundary_kv_ready_event is not None:
|
||||||
@@ -798,7 +850,7 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
cgr = self.cuda_graph_runner_for_draft_extend
|
cgr = self.cuda_graph_runner_for_draft_extend
|
||||||
# Populate the single shared buffer set once; each step replays
|
# Populate the single shared buffer set once; each step replays
|
||||||
# against it and the chain is advanced in place between steps.
|
# against it and the chain is advanced in place between steps.
|
||||||
cgr.prepare(forward_batch)
|
cgr.prepare(forward_batch, staged=staged)
|
||||||
rotates_in_graph = cgr.rotates_in_graph
|
rotates_in_graph = cgr.rotates_in_graph
|
||||||
for step in range(self.speculative_num_steps):
|
for step in range(self.speculative_num_steps):
|
||||||
_out, ret_topk_p, ret_topk_index = cgr.replay(step)
|
_out, ret_topk_p, ret_topk_index = cgr.replay(step)
|
||||||
@@ -936,6 +988,10 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
ret_draft_probs = torch.stack(ret_draft_probs_list, dim=1)
|
ret_draft_probs = torch.stack(ret_draft_probs_list, dim=1)
|
||||||
next_draft_input.draft_probs = ret_draft_probs
|
next_draft_input.draft_probs = ret_draft_probs
|
||||||
|
|
||||||
|
self.last_draft_extend_staged = bool(
|
||||||
|
staged and can_run_decode_cuda_graph and batch_result.can_run_cuda_graph
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class MultiLayerEagleWorkerV2(BaseSpecWorker):
|
class MultiLayerEagleWorkerV2(BaseSpecWorker):
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -979,6 +1035,8 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def last_shared_read_runner(self):
|
def last_shared_read_runner(self):
|
||||||
|
if self.draft_worker.last_draft_extend_staged:
|
||||||
|
return self.target_worker.model_runner
|
||||||
# Multi-layer eagle has no draft forward, only draft extend.
|
# Multi-layer eagle has no draft forward, only draft extend.
|
||||||
return self._draft_worker.draft_runner
|
return self._draft_worker.draft_runner
|
||||||
|
|
||||||
@@ -998,6 +1056,7 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
|
|||||||
def forward_batch_generation(
|
def forward_batch_generation(
|
||||||
self, batch: ScheduleBatch, on_publish=None, grammar_barrier=None
|
self, batch: ScheduleBatch, on_publish=None, grammar_barrier=None
|
||||||
):
|
):
|
||||||
|
self.draft_worker.last_draft_extend_staged = False
|
||||||
if batch.forward_mode.is_extend() or batch.is_extend_in_batch:
|
if batch.forward_mode.is_extend() or batch.is_extend_in_batch:
|
||||||
# Target prefill
|
# Target prefill
|
||||||
target_capture_mode = (
|
target_capture_mode = (
|
||||||
@@ -1044,11 +1103,14 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
|
|||||||
verify_input: EagleVerifyInput = self.draft_worker.draft(batch)
|
verify_input: EagleVerifyInput = self.draft_worker.draft(batch)
|
||||||
assert verify_input.is_verify_input()
|
assert verify_input.is_verify_input()
|
||||||
batch.spec_info = verify_input
|
batch.spec_info = verify_input
|
||||||
|
staged = self.draft_worker._draft_extend_plan_for_decode(batch)
|
||||||
batch_output = self.verify(batch, grammar_barrier=grammar_barrier)
|
batch_output = self.verify(batch, grammar_barrier=grammar_barrier)
|
||||||
# Publish before draft_extend so the fence is at verify-end.
|
# Publish before draft_extend so the fence is at verify-end.
|
||||||
if on_publish is not None:
|
if on_publish is not None:
|
||||||
on_publish(batch_output.new_seq_lens)
|
on_publish(batch_output.new_seq_lens)
|
||||||
self.draft_worker._draft_extend_for_decode(batch, batch_output)
|
self.draft_worker._draft_extend_for_decode(
|
||||||
|
batch, batch_output, staged=staged
|
||||||
|
)
|
||||||
return batch_output
|
return batch_output
|
||||||
|
|
||||||
def verify(self, batch: ScheduleBatch, grammar_barrier=None):
|
def verify(self, batch: ScheduleBatch, grammar_barrier=None):
|
||||||
|
|||||||
Reference in New Issue
Block a user