diff --git a/python/sglang/srt/kv_canary/plan_input.py b/python/sglang/srt/kv_canary/plan_input.py index d4e6bbdf8..e8de87090 100644 --- a/python/sglang/srt/kv_canary/plan_input.py +++ b/python/sglang/srt/kv_canary/plan_input.py @@ -111,14 +111,14 @@ def _extract_prefix_lens_and_extend_seq_lens( out_prefix_lens[: positions.shape[0]].copy_(positions.to(torch.int64)) out_extend_seq_lens.fill_(1) elif forward_mode.is_target_verify(): - # Evidence: EagleVerifyInputV2Mixin.prepare_for_v2_verify assigns out_cache_loc in + # Evidence: EagleVerifyInputV2Mixin.prepare_for_verify assigns out_cache_loc in # [seq_lens, seq_lens + draft_token_num) without bumping seq_lens. The target-verify # branch in TRTLLMHAAttnBackend.init_forward_metadata uses seq_lens as the prefix and # tokens_per_req as the query length, so mirror that as seq_lens plus draft_token_num. out_prefix_lens.copy_(forward_batch.seq_lens[:bs].to(torch.int64)) out_extend_seq_lens.fill_(int(spec_info.draft_token_num)) elif forward_mode.is_draft_extend_v2(): - # Evidence: EagleDraftExtendInputV2Mixin.prepare_for_extend_to_fill_draft_kvcache bumps + # Evidence: EagleDraftWorkerBase.prepare_for_draft_extend bumps # seq_lens by num_draft_tokens. FlashAttentionBackend.init_forward_metadata reads the # draft-extend-v2 query length from spec_info.extend_seq_lens_tensor when available. # CUDA-graph replay passes extend_seq_lens but omits extend_prefix_lens, so derive the diff --git a/python/sglang/srt/speculative/base_spec_worker.py b/python/sglang/srt/speculative/base_spec_worker.py index b1ba3d3bc..953458ff7 100644 --- a/python/sglang/srt/speculative/base_spec_worker.py +++ b/python/sglang/srt/speculative/base_spec_worker.py @@ -1,13 +1,17 @@ from __future__ import annotations from abc import ABC, abstractmethod -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Any + +import torch if TYPE_CHECKING: + from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.managers.tp_worker import TpModelWorker + from sglang.srt.speculative.eagle_info import EagleDraftExtendInput -class BaseDraftWorker(ABC): +class EagleDraftWorkerBase(ABC): @abstractmethod def draft(): pass @@ -29,6 +33,84 @@ class BaseDraftWorker(ABC): self.init_attention_backend() self.init_cuda_graphs() + def prepare_for_draft_extend( + self, + draft_extend_input: EagleDraftExtendInput, + batch: ScheduleBatch, + predict: torch.Tensor, + num_draft_tokens: int, + draft_model_runner: Any, + cuda_graph_runner: Any, + ): + from sglang.srt.model_executor.forward_batch_info import ( + CaptureHiddenMode, + ForwardBatch, + ForwardMode, + ) + from sglang.srt.utils.async_probe import maybe_detect_oob + from sglang.srt.utils.common import is_npu + + bs = len(batch.seq_lens) + extend_num_tokens = bs * num_draft_tokens + # When seq_lens_cpu is absent, stay on GPU-only path -- no .tolist()/.cpu(). + gpu_only = batch.seq_lens_cpu is None + + batch.spec_info = draft_extend_input + batch.input_ids = predict + maybe_detect_oob( + batch.input_ids, + 0, + batch.model_config.vocab_size, + "v2 prepare_for_draft_extend input_ids", + ) + # init_new requires both list or both Tensor; + # gpu_only emits device tensors to skip H2D. + if gpu_only: + batch.prefix_lens = batch.seq_lens.to(torch.int32) + batch.extend_lens = torch.full( + (bs,), num_draft_tokens, dtype=torch.int32, device=batch.seq_lens.device + ) + else: + batch.prefix_lens = batch.seq_lens_cpu.tolist() + batch.extend_lens = [num_draft_tokens] * bs + batch.extend_num_tokens = extend_num_tokens + capture_mode = ( + CaptureHiddenMode.NULL + if draft_model_runner.spec_algorithm.is_standalone() + else CaptureHiddenMode.FULL + ) + batch.forward_mode = ( + ForwardMode.IDLE + if batch.forward_mode.is_idle() + else ForwardMode.DRAFT_EXTEND_V2 + ) + batch.capture_hidden_mode = capture_mode + forward_batch = ForwardBatch.init_new(batch, draft_model_runner) + # Forward sees post-write length (draft extend writes num_draft_tokens + # slots); mutation stays on forward_batch to preserve SB.seq_lens. + forward_batch.seq_lens = forward_batch.seq_lens + num_draft_tokens + if not gpu_only: + forward_batch.seq_lens_cpu = forward_batch.seq_lens_cpu + num_draft_tokens + forward_batch.seq_lens_sum = int(forward_batch.seq_lens_cpu.sum()) + else: + # Supply CPU mirror (extend_seq_lens are all num_draft_tokens) so + # backend max() reads from list without a per-iter D2H sync. + forward_batch.extend_seq_lens_cpu = [num_draft_tokens] * bs + can_cuda_graph = cuda_graph_runner and cuda_graph_runner.can_run(forward_batch) + if not batch.forward_mode.is_idle() and not can_cuda_graph: + draft_model_runner.attn_backend.init_forward_metadata(forward_batch) + # Planned pre-pad; do NOT opt into post-pad re-plan. DSA's indexer + # cannot rebuild its deep_gemm schedule_meta on a DP-padded batch + # (the `_batch_size == batch_size` assertion, see #27091); the + # marked pre-pad metadata is used as-is, matching the proven + # skip_attn_backend_init=True behavior. + # On NPU with --disable-cuda-graph, block_table shape won't match + # after prepare_mlp_sync_batch padding; defer re-init to + # forward_extend (post-pad) instead. + if not is_npu() or can_cuda_graph: + forward_batch.mark_forward_metadata_ready() + return forward_batch + class BaseSpecWorker(ABC): @property @@ -38,7 +120,7 @@ class BaseSpecWorker(ABC): @property @abstractmethod - def draft_worker(self) -> BaseDraftWorker: + def draft_worker(self) -> EagleDraftWorkerBase: pass @property diff --git a/python/sglang/srt/speculative/dflash_info.py b/python/sglang/srt/speculative/dflash_info.py index 51912240c..120b2b8c0 100644 --- a/python/sglang/srt/speculative/dflash_info.py +++ b/python/sglang/srt/speculative/dflash_info.py @@ -49,7 +49,7 @@ class DFlashVerifyInput(SpecInput): def get_spec_adjust_token_coefficient(self) -> Tuple[int, int]: return self.draft_token_num, self.draft_token_num - def prepare_for_v2_verify( + def prepare_for_verify( self, batch: ScheduleBatch, target_worker: TpModelWorker, diff --git a/python/sglang/srt/speculative/dflash_worker_v2.py b/python/sglang/srt/speculative/dflash_worker_v2.py index 605516575..1184fb2d1 100644 --- a/python/sglang/srt/speculative/dflash_worker_v2.py +++ b/python/sglang/srt/speculative/dflash_worker_v2.py @@ -262,7 +262,7 @@ class DFlashWorkerV2(BaseSpecWorker): def draft_worker(self): # DFLASH drives the draft model through a plain TpModelWorker: the # draft KV is materialized from target hidden states, so there is no - # BaseDraftWorker draft/draft_extend split to wrap it in. + # EagleDraftWorkerBase draft/draft_extend split to wrap it in. return self._draft_worker @property @@ -1524,7 +1524,7 @@ class DFlashWorkerV2(BaseSpecWorker): model_worker_batch.seq_lens_cpu = draft_input.reserved_seq_lens_cpu model_worker_batch.seq_lens_sum = int(draft_input.reserved_seq_lens_sum) - verify_forward_batch, _ = verify_input.prepare_for_v2_verify( + verify_forward_batch, _ = verify_input.prepare_for_verify( model_worker_batch, self.target_worker ) model_worker_batch.seq_lens_cpu = seq_lens_cpu_backup diff --git a/python/sglang/srt/speculative/eagle_info.py b/python/sglang/srt/speculative/eagle_info.py index a3391bb16..2e4056dfa 100644 --- a/python/sglang/srt/speculative/eagle_info.py +++ b/python/sglang/srt/speculative/eagle_info.py @@ -9,7 +9,6 @@ from sglang.srt.environ import envs from sglang.srt.layers.attention.utils import create_flashinfer_kv_indices_triton from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode from sglang.srt.speculative.eagle_info_v2 import ( - EagleDraftExtendInputV2Mixin, EagleDraftInputV2Mixin, EagleVerifyInputV2Mixin, ) @@ -273,7 +272,7 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin): @dataclass -class EagleDraftExtendInput(SpecInput, EagleDraftExtendInputV2Mixin): +class EagleDraftExtendInput(SpecInput): """Inputs to the draft-extend forward (the fill-draft-kvcache pass after target prefill / verify). diff --git a/python/sglang/srt/speculative/eagle_info_v2.py b/python/sglang/srt/speculative/eagle_info_v2.py index 7b8a402d6..10dfb998f 100644 --- a/python/sglang/srt/speculative/eagle_info_v2.py +++ b/python/sglang/srt/speculative/eagle_info_v2.py @@ -1,7 +1,7 @@ from __future__ import annotations from dataclasses import dataclass -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING import torch import torch.nn.functional as F @@ -61,7 +61,6 @@ if TYPE_CHECKING: EAGLEDraftCudaGraphRunner, ) from sglang.srt.speculative.eagle_info import ( - EagleDraftExtendInput, EagleDraftInput, EagleVerifyInput, ) @@ -220,7 +219,7 @@ class EagleDraftInputV2Mixin: bs, ) - def prepare_for_v2_draft( + def prepare_for_draft( self: EagleDraftInput, req_to_token_pool: ReqToTokenPool, batch: ScheduleBatch, @@ -307,77 +306,6 @@ class EagleDraftInputV2Mixin: return forward_batch, can_cuda_graph -class EagleDraftExtendInputV2Mixin: - def prepare_for_extend_to_fill_draft_kvcache( - self: EagleDraftExtendInput, - batch: ScheduleBatch, - predict: torch.Tensor, - num_draft_tokens: int, - draft_model_runner: Any, - cuda_graph_runner: Any, - ): - bs = len(batch.seq_lens) - extend_num_tokens = bs * num_draft_tokens - # When seq_lens_cpu is absent, stay on GPU-only path -- no .tolist()/.cpu(). - gpu_only = batch.seq_lens_cpu is None - - batch.spec_info = self - batch.input_ids = predict - maybe_detect_oob( - batch.input_ids, - 0, - batch.model_config.vocab_size, - "v2 prepare_for_extend_to_fill_draft_kvcache input_ids", - ) - # init_new requires both list or both Tensor; - # gpu_only emits device tensors to skip H2D. - if gpu_only: - batch.prefix_lens = batch.seq_lens.to(torch.int32) - batch.extend_lens = torch.full( - (bs,), num_draft_tokens, dtype=torch.int32, device=batch.seq_lens.device - ) - else: - batch.prefix_lens = batch.seq_lens_cpu.tolist() - batch.extend_lens = [num_draft_tokens] * bs - batch.extend_num_tokens = extend_num_tokens - capture_mode = ( - CaptureHiddenMode.NULL - if draft_model_runner.spec_algorithm.is_standalone() - else CaptureHiddenMode.FULL - ) - batch.forward_mode = ( - ForwardMode.IDLE - if batch.forward_mode.is_idle() - else ForwardMode.DRAFT_EXTEND_V2 - ) - batch.capture_hidden_mode = capture_mode - forward_batch = ForwardBatch.init_new(batch, draft_model_runner) - # Forward sees post-write length (draft extend writes num_draft_tokens - # slots); mutation stays on forward_batch to preserve SB.seq_lens. - forward_batch.seq_lens = forward_batch.seq_lens + num_draft_tokens - if not gpu_only: - forward_batch.seq_lens_cpu = forward_batch.seq_lens_cpu + num_draft_tokens - forward_batch.seq_lens_sum = int(forward_batch.seq_lens_cpu.sum()) - else: - # Supply CPU mirror (extend_seq_lens are all num_draft_tokens) so - # backend max() reads from list without a per-iter D2H sync. - forward_batch.extend_seq_lens_cpu = [num_draft_tokens] * bs - can_cuda_graph = cuda_graph_runner and cuda_graph_runner.can_run(forward_batch) - if not batch.forward_mode.is_idle() and not can_cuda_graph: - draft_model_runner.attn_backend.init_forward_metadata(forward_batch) - # Planned pre-pad; do NOT opt into post-pad re-plan. DSA's indexer - # cannot rebuild its deep_gemm schedule_meta on a DP-padded batch - # (the `_batch_size == batch_size` assertion, see #27091); the - # marked pre-pad metadata is used as-is, matching the proven - # skip_attn_backend_init=True behavior. - # On NPU with --disable-cuda-graph, block_table shape won't match - # after prepare_mlp_sync_batch padding; defer re-init to - # forward_extend (post-pad) instead. - if not _is_npu or can_cuda_graph: - forward_batch.mark_forward_metadata_ready() - return forward_batch - - @dataclass class EagleVerifyInputV2Mixin: @property @@ -393,7 +321,7 @@ class EagleVerifyInputV2Mixin: irregular tree (no fixed per-level branching).""" return self.topk - def prepare_for_v2_verify( + def prepare_for_verify( self: EagleVerifyInput, req_to_token_pool: ReqToTokenPool, batch: ScheduleBatch, diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index 03b378faa..571f1ecc0 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -47,7 +47,7 @@ from sglang.srt.speculative.adaptive_runtime_state import ( AdaptiveController, SpecRuntimeState, ) -from sglang.srt.speculative.base_spec_worker import BaseDraftWorker, BaseSpecWorker +from sglang.srt.speculative.base_spec_worker import BaseSpecWorker, EagleDraftWorkerBase from sglang.srt.speculative.draft_utils import DraftBackendFactory from sglang.srt.speculative.eagle_draft_cuda_graph_runner import ( EAGLEDraftCudaGraphRunner, @@ -117,7 +117,7 @@ def _get_plan_stream( return None, contextlib.nullcontext() -class EagleDraftWorker(BaseDraftWorker): +class EagleDraftWorker(EagleDraftWorkerBase): def __init__( self, server_args: ServerArgs, @@ -408,7 +408,7 @@ class EagleDraftWorker(BaseDraftWorker): def draft(self, batch: ScheduleBatch): draft_input: EagleDraftInput = batch.spec_info - forward_batch, can_cuda_graph = draft_input.prepare_for_v2_draft( + forward_batch, can_cuda_graph = draft_input.prepare_for_draft( self.req_to_token_pool, batch, self.cuda_graph_runner, @@ -730,7 +730,8 @@ class EagleDraftWorker(BaseDraftWorker): # Prepare for draft extend in a separate stream with self.plan_stream_ctx: - forward_batch = draft_extend_input.prepare_for_extend_to_fill_draft_kvcache( + forward_batch = self.prepare_for_draft_extend( + draft_extend_input, batch, batch_result.next_token_ids, self.speculative_num_draft_tokens, @@ -1213,12 +1214,10 @@ class EAGLEWorkerV2(BaseSpecWorker): # Batch 1: Target verify # Prepare for target verify in a separate stream with self.plan_stream_ctx: - verify_forward_batch, can_run_cuda_graph = ( - verify_input.prepare_for_v2_verify( - self.req_to_token_pool, - batch, - self.target_worker, - ) + verify_forward_batch, can_run_cuda_graph = verify_input.prepare_for_verify( + self.req_to_token_pool, + batch, + self.target_worker, ) # Cover post-prepare rebinds: draft_token, plan_stream-allocated out_cache_loc. @@ -1264,7 +1263,7 @@ class EAGLEWorkerV2(BaseSpecWorker): # Run target verify batch in the main compute stream (GPU compute). # Metadata init is skipped iff cuda-graph already ran replay_prepare — - # prepare_for_v2_verify marked the batch in exactly that case; the + # prepare_for_verify marked the batch in exactly that case; the # non-cuda-graph path stays unmarked and gets forward_extend's init # (post-pad). forward_batch_output = self.target_worker.forward_batch_generation( @@ -1343,7 +1342,7 @@ class EAGLEWorkerV2(BaseSpecWorker): # verify_forward_batch transitively holds verify-time GPU tensors # (draft_token / out_cache_loc / ...) that must outlive the imminent - # batch.input_ids rebind in prepare_for_extend_to_fill_draft_kvcache. + # batch.input_ids rebind in prepare_for_draft_extend. # Scheduler pins it in batch_record_buf for the 2-iter window. return GenerationBatchResult( logits_output=logits_output, diff --git a/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py b/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py index 07afe21c1..6bb8fdbad 100644 --- a/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py +++ b/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py @@ -41,7 +41,7 @@ from sglang.srt.model_executor.forward_batch_info import ( from sglang.srt.model_executor.forward_context import ForwardContext, forward_context from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig from sglang.srt.server_args import ServerArgs -from sglang.srt.speculative.base_spec_worker import BaseDraftWorker +from sglang.srt.speculative.base_spec_worker import EagleDraftWorkerBase from sglang.srt.speculative.eagle_utils import ( build_tree_kernel_efficient, organize_draft_results, @@ -77,7 +77,7 @@ from sglang.srt.utils.async_probe import ( logger = logging.getLogger(__name__) -class FrozenKVMTPDraftWorker(BaseDraftWorker, TpModelWorker): +class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker): """Frozen-KV MTP draft worker. The assistant reads target KV only. It reuses EAGLE's verify input/output @@ -130,7 +130,7 @@ class FrozenKVMTPDraftWorker(BaseDraftWorker, TpModelWorker): with ( empty_context() ), speculative_moe_backend_context(), speculative_moe_a2a_backend_context(): - # NOTE: call TpModelWorker.__init__ explicitly -- BaseDraftWorker is + # NOTE: call TpModelWorker.__init__ explicitly -- EagleDraftWorkerBase is # an ABC with no __init__, so cooperative super() would be ambiguous. TpModelWorker.__init__( self, @@ -180,7 +180,7 @@ class FrozenKVMTPDraftWorker(BaseDraftWorker, TpModelWorker): req_to_token_pool=None, token_to_kv_pool_allocator=None, ): - # NOTE: call TpModelWorker explicitly -- BaseDraftWorker precedes it in + # NOTE: call TpModelWorker explicitly -- EagleDraftWorkerBase precedes it in # the MRO and its alloc_memory_pool is a no-op stub. TpModelWorker.alloc_memory_pool( self, @@ -583,7 +583,7 @@ class FrozenKVMTPDraftWorker(BaseDraftWorker, TpModelWorker): ) def draft_extend(self): - # BaseDraftWorker contract. Frozen has no draft-KV extend forward; the + # EagleDraftWorkerBase contract. Frozen has no draft-KV extend forward; the # orchestrator calls `_draft_extend_for_{prefill,decode}` directly. pass diff --git a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py index 38e91f3a4..27d379156 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py @@ -41,7 +41,7 @@ from sglang.srt.model_executor.forward_batch_info import ( ForwardBatch, ) from sglang.srt.server_args import ServerArgs -from sglang.srt.speculative.base_spec_worker import BaseDraftWorker, BaseSpecWorker +from sglang.srt.speculative.base_spec_worker import BaseSpecWorker, EagleDraftWorkerBase from sglang.srt.speculative.draft_utils import DraftBackendFactory from sglang.srt.speculative.eagle_info import ( EagleDraftExtendInput, @@ -92,7 +92,7 @@ def _get_plan_stream( return None, contextlib.nullcontext() -class MultiLayerEagleDraftWorker(BaseDraftWorker): +class MultiLayerEagleDraftWorker(EagleDraftWorkerBase): def __init__( self, server_args: ServerArgs, @@ -253,7 +253,7 @@ class MultiLayerEagleDraftWorker(BaseDraftWorker): def draft(self, batch: ScheduleBatch): draft_input: EagleDraftInput = batch.spec_info - forward_batch, can_cuda_graph = draft_input.prepare_for_v2_draft( + forward_batch, can_cuda_graph = draft_input.prepare_for_draft( self.req_to_token_pool, batch, self.cuda_graph_runner, @@ -518,7 +518,8 @@ class MultiLayerEagleDraftWorker(BaseDraftWorker): # Prepare for draft extend in a separate stream # Notice that here we use batch_result.next_token_ids as the input ids with self.plan_stream_ctx: - forward_batch = draft_extend_input.prepare_for_extend_to_fill_draft_kvcache( + forward_batch = self.prepare_for_draft_extend( + draft_extend_input, batch, batch_result.next_token_ids, self.speculative_num_draft_tokens, @@ -805,12 +806,10 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker): # Batch 1: Target verify # Prepare for target verify in a separate stream with self.plan_stream_ctx: - verify_forward_batch, can_run_cuda_graph = ( - verify_input.prepare_for_v2_verify( - self.req_to_token_pool, - batch, - self.target_worker, - ) + verify_forward_batch, can_run_cuda_graph = verify_input.prepare_for_verify( + self.req_to_token_pool, + batch, + self.target_worker, ) # Cover post-prepare rebinds: draft_token, plan_stream-allocated out_cache_loc. @@ -834,7 +833,7 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker): ), ) # NOTE: metadata init is skipped here unconditionally, although - # prepare_for_v2_verify only plans when cuda-graph replay_prepare ran. + # prepare_for_verify only plans when cuda-graph replay_prepare ran. # eagle_worker_v2 re-inits the non-graph path instead (post-pad); this # worker has not adopted that fix, so preserve its behavior verbatim. # On NPU with --disable-cuda-graph, non-graph verify needs metadata init diff --git a/python/sglang/srt/speculative/ngram_worker.py b/python/sglang/srt/speculative/ngram_worker.py index f2d7c405a..b762641ad 100644 --- a/python/sglang/srt/speculative/ngram_worker.py +++ b/python/sglang/srt/speculative/ngram_worker.py @@ -12,7 +12,7 @@ from sglang.srt.managers.tp_worker import TpModelWorker from sglang.srt.model_executor.forward_batch_info import ForwardMode from sglang.srt.observability.req_time_stats import set_time_batch from sglang.srt.server_args import ServerArgs -from sglang.srt.speculative.base_spec_worker import BaseDraftWorker, BaseSpecWorker +from sglang.srt.speculative.base_spec_worker import BaseSpecWorker, EagleDraftWorkerBase from sglang.srt.speculative.cpp_ngram.ngram_corpus import NgramCorpus from sglang.srt.speculative.ngram_info import NgramVerifyInput from sglang.srt.speculative.spec_utils import ( @@ -110,7 +110,7 @@ class NGRAMWorker(BaseSpecWorker): return self._target_worker @property - def draft_worker(self) -> Optional[BaseDraftWorker]: + def draft_worker(self) -> Optional[EagleDraftWorkerBase]: # NGRAM has no draft model; drafts come from the CPU-side corpus. return None diff --git a/python/sglang/srt/speculative/spec_utils.py b/python/sglang/srt/speculative/spec_utils.py index fb3e7b3de..604db2faf 100644 --- a/python/sglang/srt/speculative/spec_utils.py +++ b/python/sglang/srt/speculative/spec_utils.py @@ -120,7 +120,7 @@ def record_stream_each(tensors, stream): def record_stream_for_v2_verify(batch, verify_input, fwd_stream): """Mark pre-prepare SB / verify_input GPU tensors as used on `fwd_stream`. - Spec V2 mutates SB mid-forward (`prepare_for_v2_verify` rebinds + Spec V2 mutates SB mid-forward (`prepare_for_verify` rebinds `batch.input_ids` / `out_cache_loc`; `_draft_extend_for_decode` later replaces `batch.input_ids` again). Each rebind drops the only SB Python ref to the old tensor while the verify forward kernel may still be diff --git a/test/registered/unit/spec/test_decode_bookkeeping_ownership.py b/test/registered/unit/spec/test_decode_bookkeeping_ownership.py index d2b45d8c9..09f967780 100644 --- a/test/registered/unit/spec/test_decode_bookkeeping_ownership.py +++ b/test/registered/unit/spec/test_decode_bookkeeping_ownership.py @@ -154,7 +154,7 @@ def _scan_srt(): def _draft_worker_classes(): - """All transitive BaseDraftWorker subclasses under speculative/.""" + """All transitive EagleDraftWorkerBase subclasses under speculative/.""" by_name = {} for path in sorted(_SPECULATIVE_DIR.glob("*.py")): rel = path.relative_to(_SRT_DIR).as_posix() @@ -166,7 +166,7 @@ def _draft_worker_classes(): } by_name[node.name] = (rel, node, bases) - workers = {"BaseDraftWorker"} + workers = {"EagleDraftWorkerBase"} changed = True while changed: changed = False @@ -177,7 +177,7 @@ def _draft_worker_classes(): return [ (rel, node) for name, (rel, node, _) in sorted(by_name.items()) - if name in workers and name != "BaseDraftWorker" + if name in workers and name != "EagleDraftWorkerBase" ]