diff --git a/python/sglang/srt/layers/attention/aiter_backend.py b/python/sglang/srt/layers/attention/aiter_backend.py index ccf0c796a..273b71a83 100755 --- a/python/sglang/srt/layers/attention/aiter_backend.py +++ b/python/sglang/srt/layers/attention/aiter_backend.py @@ -27,6 +27,7 @@ from sglang.srt.layers.dp_attention import ( is_dp_attention_enabled, ) from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode +from sglang.srt.speculative.spec_utils import generate_draft_decode_kv_indices from sglang.srt.utils import is_gfx95_supported if TYPE_CHECKING: @@ -2776,8 +2777,6 @@ class AiterMultiStepDraftBackend: topk: int, speculative_num_steps: int, ): - from sglang.srt.speculative.spec_utils import generate_draft_decode_kv_indices - self.topk = topk self.speculative_num_steps = speculative_num_steps self.generate_draft_decode_kv_indices = generate_draft_decode_kv_indices diff --git a/python/sglang/srt/layers/attention/flashinfer_backend.py b/python/sglang/srt/layers/attention/flashinfer_backend.py index 70fd29d65..629cf4c25 100644 --- a/python/sglang/srt/layers/attention/flashinfer_backend.py +++ b/python/sglang/srt/layers/attention/flashinfer_backend.py @@ -27,6 +27,7 @@ from sglang.srt.layers.radix_attention import AttentionType from sglang.srt.mem_cache.allocator.swa import SWATokenToKVPoolAllocator from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.speculative.spec_info import SpecInput +from sglang.srt.speculative.spec_utils import generate_draft_decode_kv_indices from sglang.srt.utils import ( get_int_env_var, is_flashinfer_available, @@ -1538,8 +1539,6 @@ class FlashInferMultiStepDraftBackend: topk: int, speculative_num_steps: int, ): - from sglang.srt.speculative.spec_utils import generate_draft_decode_kv_indices - self.topk = topk self.speculative_num_steps = speculative_num_steps self.generate_draft_decode_kv_indices = generate_draft_decode_kv_indices diff --git a/python/sglang/srt/layers/attention/flashinfer_mla_backend.py b/python/sglang/srt/layers/attention/flashinfer_mla_backend.py index da4881e88..716c947c7 100644 --- a/python/sglang/srt/layers/attention/flashinfer_mla_backend.py +++ b/python/sglang/srt/layers/attention/flashinfer_mla_backend.py @@ -25,6 +25,7 @@ from sglang.srt.layers.dp_attention import get_attention_tp_size from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.server_args import get_global_server_args from sglang.srt.speculative.spec_info import SpecInput +from sglang.srt.speculative.spec_utils import generate_draft_decode_kv_indices from sglang.srt.utils import ( is_flashinfer_available, is_sm100_supported, @@ -892,8 +893,6 @@ class FlashInferMLAMultiStepDraftBackend: topk: int, speculative_num_steps: int, ): - from sglang.srt.speculative.spec_utils import generate_draft_decode_kv_indices - if topk > 1: raise ValueError( "Currently Flashinfer MLA only supports topk=1 for speculative decoding" diff --git a/python/sglang/srt/speculative/spec_utils.py b/python/sglang/srt/speculative/spec_utils.py index 8a7b370d8..6938a5918 100644 --- a/python/sglang/srt/speculative/spec_utils.py +++ b/python/sglang/srt/speculative/spec_utils.py @@ -9,15 +9,13 @@ from typing import TYPE_CHECKING, List, Optional import torch from huggingface_hub import snapshot_download -from sglang.srt.constrained.base_grammar_backend import BaseGrammarObject from sglang.srt.distributed.parallel_state import ( GroupCoordinator, patch_tensor_parallel_group, ) from sglang.srt.environ import envs -from sglang.srt.managers.schedule_batch import Req from sglang.srt.mem_cache.common import get_last_loc -from sglang.srt.server_args import ServerArgs, get_global_server_args +from sglang.srt.server_args import get_global_server_args from sglang.srt.speculative.triton_ops.cache_locs import ( align_evict_mask_to_page_size as align_evict_mask_to_page_size, ) @@ -53,6 +51,9 @@ _is_npu = is_npu() _is_musa = is_musa() if TYPE_CHECKING: + from sglang.srt.constrained.base_grammar_backend import BaseGrammarObject + from sglang.srt.managers.schedule_batch import Req + from sglang.srt.server_args import ServerArgs from sglang.srt.speculative.eagle_info import EagleVerifyInput