diff --git a/python/sglang/srt/speculative/draft_worker_common.py b/python/sglang/srt/speculative/draft_worker_common.py index 28a150f7375f..058f0f7824ed 100644 --- a/python/sglang/srt/speculative/draft_worker_common.py +++ b/python/sglang/srt/speculative/draft_worker_common.py @@ -71,6 +71,7 @@ def build_draft_tp_worker( target_model_config: ModelConfig, algo_label: str, attention_backend_override: Optional[str] = None, + draft_worker_cls: type[TpModelWorker] = TpModelWorker, ) -> DraftWorkerBundle: # An override names a draft-specific backend the caller has already # validated (e.g. a self-drafting architecture); it skips the generic @@ -88,7 +89,7 @@ def build_draft_tp_worker( # workers run the draft outside speculative_moe_backend_context, so a # construction-only swap would build and execute under different backends. with draft_model_build_scope(): - draft_worker = TpModelWorker( + draft_worker = draft_worker_cls( server_args=server_args, gpu_id=gpu_id, ps=ps, diff --git a/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py b/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py index 666cfd8c0b1d..d3024c4454b1 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py @@ -82,6 +82,7 @@ def __init__( ps: ParallelState, nccl_port: int, target_worker: TpModelWorker, + draft_worker_cls: type[TpModelWorker] = TpModelWorker, ): super().__init__() @@ -124,6 +125,7 @@ def __init__( attention_backend_override=( DSV4_DRAFT_ATTENTION_BACKEND if self._draft_is_moe else None ), + draft_worker_cls=draft_worker_cls, ) self._draft_worker = bundle.draft_worker self.draft_model_runner = bundle.draft_model_runner