From 99c12218c357c81a20b766e9e96d8c274240aaed Mon Sep 17 00:00:00 2001 From: Lianmin Zheng Date: Wed, 19 Aug 2026 17:48:17 -0700 Subject: [PATCH] Support custom draft worker classes in DSpark (#35397) Co-authored-by: Yichao Fu --- python/sglang/srt/speculative/draft_worker_common.py | 3 ++- .../srt/speculative/dspark_components/dspark_worker_v2.py | 2 ++ 2 files changed, 4 insertions(+), 1 deletion(-) diff --git a/python/sglang/srt/speculative/draft_worker_common.py b/python/sglang/srt/speculative/draft_worker_common.py index 28a150f73..058f0f782 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 666cfd8c0..d3024c445 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 @@ class DSparkWorkerV2(BaseSpecWorker): ps: ParallelState, nccl_port: int, target_worker: TpModelWorker, + draft_worker_cls: type[TpModelWorker] = TpModelWorker, ): super().__init__() @@ -124,6 +125,7 @@ class DSparkWorkerV2(BaseSpecWorker): 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