Support custom draft worker classes in DSpark (#35397)

Co-authored-by: Yichao Fu <yichaofu@meta.com>
This commit is contained in:
Lianmin Zheng
2026-08-19 17:48:17 -07:00
committed by GitHub
co-authored by Yichao Fu
parent 9234e40aed
commit 99c12218c3
2 changed files with 4 additions and 1 deletions
@@ -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,
@@ -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