Support custom draft worker classes in DSpark (#35397)
Co-authored-by: Yichao Fu <yichaofu@meta.com>
This commit is contained in:
co-authored by
Yichao Fu
parent
9234e40aed
commit
99c12218c3
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user