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,
|
target_model_config: ModelConfig,
|
||||||
algo_label: str,
|
algo_label: str,
|
||||||
attention_backend_override: Optional[str] = None,
|
attention_backend_override: Optional[str] = None,
|
||||||
|
draft_worker_cls: type[TpModelWorker] = TpModelWorker,
|
||||||
) -> DraftWorkerBundle:
|
) -> DraftWorkerBundle:
|
||||||
# An override names a draft-specific backend the caller has already
|
# An override names a draft-specific backend the caller has already
|
||||||
# validated (e.g. a self-drafting architecture); it skips the generic
|
# 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
|
# workers run the draft outside speculative_moe_backend_context, so a
|
||||||
# construction-only swap would build and execute under different backends.
|
# construction-only swap would build and execute under different backends.
|
||||||
with draft_model_build_scope():
|
with draft_model_build_scope():
|
||||||
draft_worker = TpModelWorker(
|
draft_worker = draft_worker_cls(
|
||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
gpu_id=gpu_id,
|
gpu_id=gpu_id,
|
||||||
ps=ps,
|
ps=ps,
|
||||||
|
|||||||
@@ -82,6 +82,7 @@ class DSparkWorkerV2(BaseSpecWorker):
|
|||||||
ps: ParallelState,
|
ps: ParallelState,
|
||||||
nccl_port: int,
|
nccl_port: int,
|
||||||
target_worker: TpModelWorker,
|
target_worker: TpModelWorker,
|
||||||
|
draft_worker_cls: type[TpModelWorker] = TpModelWorker,
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
@@ -124,6 +125,7 @@ class DSparkWorkerV2(BaseSpecWorker):
|
|||||||
attention_backend_override=(
|
attention_backend_override=(
|
||||||
DSV4_DRAFT_ATTENTION_BACKEND if self._draft_is_moe else None
|
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_worker = bundle.draft_worker
|
||||||
self.draft_model_runner = bundle.draft_model_runner
|
self.draft_model_runner = bundle.draft_model_runner
|
||||||
|
|||||||
Reference in New Issue
Block a user