Fix DFlash mamba verify init ordering (#30680)
This commit is contained in:
@@ -118,12 +118,7 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
self.nccl_port = nccl_port
|
self.nccl_port = nccl_port
|
||||||
self._target_worker = target_worker
|
self._target_worker = target_worker
|
||||||
self.model_runner = target_worker.model_runner
|
self.model_runner = target_worker.model_runner
|
||||||
self._need_mamba_verify_commit = (
|
self._need_mamba_verify_commit = False
|
||||||
self.model_runner.mambaish_config is not None
|
|
||||||
and hasattr(
|
|
||||||
self.model_runner.attn_backend, "update_mamba_state_after_mtp_verify"
|
|
||||||
)
|
|
||||||
)
|
|
||||||
self.page_size = server_args.page_size
|
self.page_size = server_args.page_size
|
||||||
# Normalized in arg_groups.speculative_hook.handle_speculative_decoding.
|
# Normalized in arg_groups.speculative_hook.handle_speculative_decoding.
|
||||||
self.draft_window_size: Optional[int] = (
|
self.draft_window_size: Optional[int] = (
|
||||||
@@ -286,6 +281,13 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
|
|
||||||
def init_attention_backends(self):
|
def init_attention_backends(self):
|
||||||
self._draft_worker.init_attention_backends()
|
self._draft_worker.init_attention_backends()
|
||||||
|
self._need_mamba_verify_commit = (
|
||||||
|
self.model_runner.mambaish_config is not None
|
||||||
|
and hasattr(
|
||||||
|
self.model_runner.attn_backend,
|
||||||
|
"update_mamba_state_after_mtp_verify",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
def init_cuda_graphs(self):
|
def init_cuda_graphs(self):
|
||||||
capture_decode_cuda_graph = (
|
capture_decode_cuda_graph = (
|
||||||
|
|||||||
Reference in New Issue
Block a user