Revert "[Spec] Split init_backends; account draft weights in --mem-fraction-static" (#28841)
This commit is contained in:
@@ -324,8 +324,7 @@ def load_model(server_args, port_args, gpu_id, tp_rank):
|
|||||||
else:
|
else:
|
||||||
model_runner = ModelRunner(**runner_kwargs)
|
model_runner = ModelRunner(**runner_kwargs)
|
||||||
model_runner.alloc_memory_pool()
|
model_runner.alloc_memory_pool()
|
||||||
model_runner.init_attention_backends()
|
model_runner.init_backends()
|
||||||
model_runner.init_cuda_graphs()
|
|
||||||
rank_print(f"max_total_num_tokens={model_runner.max_total_num_tokens}")
|
rank_print(f"max_total_num_tokens={model_runner.max_total_num_tokens}")
|
||||||
tokenizer = get_tokenizer(
|
tokenizer = get_tokenizer(
|
||||||
server_args.tokenizer_path,
|
server_args.tokenizer_path,
|
||||||
|
|||||||
@@ -840,17 +840,11 @@ class Scheduler(
|
|||||||
token_to_kv_pool_allocator=allocator,
|
token_to_kv_pool_allocator=allocator,
|
||||||
)
|
)
|
||||||
|
|
||||||
def init_all_attention_backends(self):
|
def init_all_backends(self):
|
||||||
"""Initialize attention backends for all workers."""
|
"""Initialize attention backends and capture cuda graphs for all workers."""
|
||||||
self.tp_worker.init_attention_backends()
|
self.tp_worker.init_backends()
|
||||||
if self.draft_worker is not None:
|
if self.draft_worker is not None:
|
||||||
self.draft_worker.init_attention_backends()
|
self.draft_worker.init_backends()
|
||||||
|
|
||||||
def init_all_cuda_graphs(self):
|
|
||||||
"""Capture cuda graphs for all workers."""
|
|
||||||
self.tp_worker.init_cuda_graphs()
|
|
||||||
if self.draft_worker is not None:
|
|
||||||
self.draft_worker.init_cuda_graphs()
|
|
||||||
|
|
||||||
def init_model_worker(self):
|
def init_model_worker(self):
|
||||||
# Load model weights.
|
# Load model weights.
|
||||||
@@ -861,11 +855,12 @@ class Scheduler(
|
|||||||
self.maybe_init_draft_worker()
|
self.maybe_init_draft_worker()
|
||||||
|
|
||||||
# Allocate KV cache pools for all workers.
|
# Allocate KV cache pools for all workers.
|
||||||
|
# Memory profiling now sees all loaded weights.
|
||||||
self.init_memory_pools()
|
self.init_memory_pools()
|
||||||
|
|
||||||
|
# Initialize attention backends and capture cuda graphs.
|
||||||
# TODO: make memory profile consider cuda graph memory as well
|
# TODO: make memory profile consider cuda graph memory as well
|
||||||
self.init_all_attention_backends()
|
self.init_all_backends()
|
||||||
self.init_all_cuda_graphs()
|
|
||||||
|
|
||||||
# Dispatch the model worker
|
# Dispatch the model worker
|
||||||
if self.spec_algorithm.is_none():
|
if self.spec_algorithm.is_none():
|
||||||
|
|||||||
@@ -333,19 +333,11 @@ class TpModelWorker(BaseTpWorker):
|
|||||||
)
|
)
|
||||||
assert max_req_len > 0, "Memory pool size is too small"
|
assert max_req_len > 0, "Memory pool size is too small"
|
||||||
|
|
||||||
def init_attention_backends(self):
|
def init_backends(self, disable_cuda_graph: bool = False):
|
||||||
"""Initialize attention backends for all model runners."""
|
"""Initialize attention backends and capture cuda graphs."""
|
||||||
self.model_runner.init_attention_backends()
|
self.model_runner.init_backends(disable_cuda_graph=disable_cuda_graph)
|
||||||
for mr in self.model_runner_list[1:]:
|
for mr in self.model_runner_list[1:]:
|
||||||
mr.init_attention_backends()
|
mr.init_backends(disable_cuda_graph=disable_cuda_graph)
|
||||||
|
|
||||||
def init_cuda_graphs(self, capture_decode_cuda_graph: bool = True):
|
|
||||||
"""Capture cuda graphs for all model runners."""
|
|
||||||
self.model_runner.init_cuda_graphs(
|
|
||||||
capture_decode_cuda_graph=capture_decode_cuda_graph
|
|
||||||
)
|
|
||||||
for mr in self.model_runner_list[1:]:
|
|
||||||
mr.init_cuda_graphs(capture_decode_cuda_graph=capture_decode_cuda_graph)
|
|
||||||
|
|
||||||
def _init_model_config(self):
|
def _init_model_config(self):
|
||||||
from sglang.srt.configs.model_config import ModelConfig
|
from sglang.srt.configs.model_config import ModelConfig
|
||||||
|
|||||||
@@ -788,6 +788,18 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
# Deduce KV cache dtype
|
# Deduce KV cache dtype
|
||||||
self.configure_kv_cache_dtype()
|
self.configure_kv_cache_dtype()
|
||||||
|
|
||||||
|
# Snapshot free memory at the end of the weight-load phase. KV-pool
|
||||||
|
# profiling uses this instead of measuring at alloc_memory_pool()
|
||||||
|
# time: draft-model weights load between the two phases and must stay
|
||||||
|
# outside the --mem-fraction-static budget (deployments tune the
|
||||||
|
# fraction assuming draft weights live in the non-static slack).
|
||||||
|
self.post_model_load_memory = get_available_gpu_memory(
|
||||||
|
self.device,
|
||||||
|
self.gpu_id,
|
||||||
|
distributed=get_world_group().world_size > 1,
|
||||||
|
cpu_group=get_world_group().cpu_group,
|
||||||
|
)
|
||||||
|
|
||||||
def get_pp_proxy_topk_size(self) -> Optional[int]:
|
def get_pp_proxy_topk_size(self) -> Optional[int]:
|
||||||
hf_config = self.model_config.hf_text_config
|
hf_config = self.model_config.hf_text_config
|
||||||
if (
|
if (
|
||||||
@@ -850,13 +862,18 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
self.graph_mem_usage = 0
|
self.graph_mem_usage = 0
|
||||||
self.prefill_cuda_graph_runner = None
|
self.prefill_cuda_graph_runner = None
|
||||||
|
|
||||||
def init_attention_backends(self):
|
def init_backends(self, disable_cuda_graph: bool = False):
|
||||||
"""Initialize attention backends only (no cuda graph capture)."""
|
"""Initialize attention backends and capture cuda graphs."""
|
||||||
|
server_args = self.server_args
|
||||||
|
|
||||||
# TODO: Refactor device-specific init branches into platform interface (separate PR).
|
# TODO: Refactor device-specific init branches into platform interface (separate PR).
|
||||||
# Must be called BEFORE init_decode_cuda_graph() so CUDA graph capture
|
# Must be called BEFORE init_decode_cuda_graph() so CUDA graph capture
|
||||||
# runs with aux hidden state capture enabled.
|
# runs with aux hidden state capture enabled.
|
||||||
self.init_aux_hidden_state_capture()
|
self.init_aux_hidden_state_capture()
|
||||||
|
|
||||||
|
# Device-specific attention-backend init. cg_supported gates decode
|
||||||
|
# cuda-graph capture (out-of-tree / unknown platforms may not support it).
|
||||||
|
cg_supported = True
|
||||||
if self.device == "cuda" or self.device == "musa":
|
if self.device == "cuda" or self.device == "musa":
|
||||||
self.init_cublas()
|
self.init_cublas()
|
||||||
self.init_attention_backend()
|
self.init_attention_backend()
|
||||||
@@ -875,15 +892,12 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
get_world_group().world_size,
|
get_world_group().world_size,
|
||||||
get_world_group().cpu_group,
|
get_world_group().cpu_group,
|
||||||
)
|
)
|
||||||
|
elif current_platform.is_out_of_tree():
|
||||||
|
self.init_attention_backend()
|
||||||
|
cg_supported = current_platform.support_cuda_graph()
|
||||||
else:
|
else:
|
||||||
self.init_attention_backend()
|
self.init_attention_backend()
|
||||||
|
cg_supported = False
|
||||||
def init_cuda_graphs(self, capture_decode_cuda_graph: bool = True):
|
|
||||||
"""Capture cuda graphs. Requires init_attention_backends() to have run.
|
|
||||||
|
|
||||||
Spec draft runners pass capture_decode_cuda_graph=False
|
|
||||||
because they capture their own decode-style graphs separately.
|
|
||||||
"""
|
|
||||||
|
|
||||||
# The eager (no-cuda-graph) phase runner, built AFTER the attention
|
# The eager (no-cuda-graph) phase runner, built AFTER the attention
|
||||||
# backend so its __init__ can warm up kernels (run-once) and allocate the
|
# backend so its __init__ can warm up kernels (run-once) and allocate the
|
||||||
@@ -898,27 +912,23 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
# eager buffer allocated above. (init_prefill_cuda_graph routes prefill
|
# eager buffer allocated above. (init_prefill_cuda_graph routes prefill
|
||||||
# to the eager runner when the prefill graph is disabled.)
|
# to the eager runner when the prefill graph is disabled.)
|
||||||
self.init_prefill_cuda_graph()
|
self.init_prefill_cuda_graph()
|
||||||
|
if not disable_cuda_graph and cg_supported:
|
||||||
self.decode_cuda_graph_runner = None
|
|
||||||
self.graph_mem_usage = 0
|
|
||||||
|
|
||||||
if capture_decode_cuda_graph:
|
|
||||||
if self.device in ("cuda", "musa", "cpu", "npu"):
|
|
||||||
self.init_decode_cuda_graph()
|
|
||||||
elif (
|
|
||||||
current_platform.is_out_of_tree()
|
|
||||||
and current_platform.support_cuda_graph()
|
|
||||||
):
|
|
||||||
self.init_decode_cuda_graph()
|
self.init_decode_cuda_graph()
|
||||||
else:
|
else:
|
||||||
|
self.decode_cuda_graph_runner = None
|
||||||
|
self.graph_mem_usage = 0
|
||||||
|
if disable_cuda_graph:
|
||||||
|
# Decode cuda graph disabled: route eager decode through the
|
||||||
|
# EagerRunner (the dispatch gate isinstance(..., EagerRunner) keeps
|
||||||
|
# _forward_raw off any replay branch).
|
||||||
self.decode_cuda_graph_runner = self.eager_runner
|
self.decode_cuda_graph_runner = self.eager_runner
|
||||||
|
|
||||||
# Register forward hooks AFTER cuda-graph capture so their tensor ops are
|
# Register forward hooks AFTER cuda-graph capture so their tensor ops are
|
||||||
# not traced into any captured graph — capture stays hook-free and hooks
|
# not traced into any captured graph — capture stays hook-free and hooks
|
||||||
# fire only on the eager forward path (capture replay never runs Python
|
# fire only on the eager forward path (capture replay never runs Python
|
||||||
# hooks anyway).
|
# hooks anyway).
|
||||||
if self.server_args.forward_hooks:
|
if server_args.forward_hooks:
|
||||||
register_forward_hooks(self.model, self.server_args.forward_hooks)
|
register_forward_hooks(self.model, server_args.forward_hooks)
|
||||||
|
|
||||||
self.prealloc_symmetric_memory_pool()
|
self.prealloc_symmetric_memory_pool()
|
||||||
|
|
||||||
|
|||||||
@@ -82,29 +82,24 @@ _is_hip = is_hip()
|
|||||||
|
|
||||||
class ModelRunnerKVCacheMixin:
|
class ModelRunnerKVCacheMixin:
|
||||||
def _profile_available_bytes(self: ModelRunner, pre_model_load_memory: int) -> int:
|
def _profile_available_bytes(self: ModelRunner, pre_model_load_memory: int) -> int:
|
||||||
# KV pool budget = currently-free GPU memory minus the non-static runtime
|
# Use the snapshot taken at the end of this runner's weight-load phase,
|
||||||
# slack (pre_model_load_memory * (1 - mem_fraction_static)). Whatever is
|
# not the current free memory: draft-model weights loaded after that
|
||||||
# already resident (model weights, etc.) is thus charged against it.
|
# point are charged to the non-static slack, not the static budget.
|
||||||
available_gpu_memory = get_available_gpu_memory(
|
post_model_load_memory = getattr(self, "post_model_load_memory", None)
|
||||||
|
if post_model_load_memory is None:
|
||||||
|
post_model_load_memory = get_available_gpu_memory(
|
||||||
self.device,
|
self.device,
|
||||||
self.gpu_id,
|
self.gpu_id,
|
||||||
distributed=get_world_group().world_size > 1,
|
distributed=get_world_group().world_size > 1,
|
||||||
cpu_group=get_world_group().cpu_group,
|
cpu_group=get_world_group().cpu_group,
|
||||||
)
|
)
|
||||||
|
|
||||||
rest_memory = available_gpu_memory - pre_model_load_memory * (
|
rest_memory = post_model_load_memory - pre_model_load_memory * (
|
||||||
1 - self.mem_fraction_static
|
1 - self.mem_fraction_static
|
||||||
)
|
)
|
||||||
if self.mambaish_config is not None:
|
if self.mambaish_config is not None:
|
||||||
rest_memory = self.handle_max_mamba_cache(rest_memory)
|
rest_memory = self.handle_max_mamba_cache(rest_memory)
|
||||||
|
|
||||||
# Loaded weights (target + draft) can exceed the static budget
|
|
||||||
if rest_memory <= 0:
|
|
||||||
raise ValueError(
|
|
||||||
f"Loaded weights leave no GPU memory for the KV cache under "
|
|
||||||
f"--mem-fraction-static={self.mem_fraction_static}."
|
|
||||||
)
|
|
||||||
|
|
||||||
return int(rest_memory * (1 << 30)) # return in bytes
|
return int(rest_memory * (1 << 30)) # return in bytes
|
||||||
|
|
||||||
def handle_max_mamba_cache(self: ModelRunner, total_rest_memory):
|
def handle_max_mamba_cache(self: ModelRunner, total_rest_memory):
|
||||||
@@ -565,9 +560,9 @@ class ModelRunnerKVCacheMixin:
|
|||||||
self.model_config.hf_text_config.swa_num_key_value_heads
|
self.model_config.hf_text_config.swa_num_key_value_heads
|
||||||
// get_attention_tp_size(),
|
// get_attention_tp_size(),
|
||||||
),
|
),
|
||||||
"swa_head_dim": self.model_config.swa_head_dim,
|
"swa_head_dim": self.model_config.hf_text_config.swa_head_dim,
|
||||||
"swa_v_head_dim": self.model_config.swa_v_head_dim,
|
"swa_v_head_dim": self.model_config.hf_text_config.swa_v_head_dim,
|
||||||
"v_head_dim": self.model_config.v_head_dim,
|
"v_head_dim": self.model_config.hf_text_config.v_head_dim,
|
||||||
}
|
}
|
||||||
self.token_to_kv_pool = SWAKVPool(
|
self.token_to_kv_pool = SWAKVPool(
|
||||||
size=self.full_max_total_num_tokens,
|
size=self.full_max_total_num_tokens,
|
||||||
@@ -688,9 +683,9 @@ class ModelRunnerKVCacheMixin:
|
|||||||
self.model_config.hf_text_config.swa_num_key_value_heads
|
self.model_config.hf_text_config.swa_num_key_value_heads
|
||||||
// get_attention_tp_size(),
|
// get_attention_tp_size(),
|
||||||
),
|
),
|
||||||
"swa_head_dim": self.model_config.swa_head_dim,
|
"swa_head_dim": self.model_config.hf_text_config.swa_head_dim,
|
||||||
"swa_v_head_dim": self.model_config.swa_v_head_dim,
|
"swa_v_head_dim": self.model_config.hf_text_config.swa_v_head_dim,
|
||||||
"v_head_dim": self.model_config.v_head_dim,
|
"v_head_dim": self.model_config.hf_text_config.v_head_dim,
|
||||||
}
|
}
|
||||||
self.token_to_kv_pool = SWAKVPool(
|
self.token_to_kv_pool = SWAKVPool(
|
||||||
size=self.full_max_total_num_tokens,
|
size=self.full_max_total_num_tokens,
|
||||||
|
|||||||
@@ -78,16 +78,15 @@ class EagleDraftWorkerBase(ABC):
|
|||||||
def alloc_memory_pool(self, **kwargs):
|
def alloc_memory_pool(self, **kwargs):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
def init_attention_backends(self):
|
def init_backends(self):
|
||||||
"""Subclasses wrap this with their context managers (draft_tp_context,
|
"""Initialize standard backends (no cuda graphs) then draft-specific backends.
|
||||||
speculative_moe_backend_context, etc.) rather than reimplementing it."""
|
|
||||||
self.draft_worker.init_attention_backends()
|
|
||||||
self.init_attention_backend()
|
|
||||||
|
|
||||||
def init_cuda_graphs(self):
|
Subclasses should wrap this with their context managers (draft_tp_context,
|
||||||
"""Capture draft graphs (decode disabled on the draft TpModelWorker)."""
|
speculative_moe_backend_context, etc.) rather than reimplementing the logic.
|
||||||
self.draft_worker.init_cuda_graphs(capture_decode_cuda_graph=False)
|
"""
|
||||||
self._capture_cuda_graphs()
|
self.draft_worker.init_backends(disable_cuda_graph=True)
|
||||||
|
self.init_attention_backend()
|
||||||
|
self.init_cuda_graphs()
|
||||||
|
|
||||||
def prepare_for_draft_extend(
|
def prepare_for_draft_extend(
|
||||||
self,
|
self,
|
||||||
@@ -294,10 +293,7 @@ class BaseSpecWorker(ABC):
|
|||||||
def alloc_memory_pool(self, **kwargs):
|
def alloc_memory_pool(self, **kwargs):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
def init_attention_backends(self):
|
def init_backends(self):
|
||||||
pass
|
|
||||||
|
|
||||||
def init_cuda_graphs(self):
|
|
||||||
pass
|
pass
|
||||||
|
|
||||||
def on_verify_complete_cpu(
|
def on_verify_complete_cpu(
|
||||||
|
|||||||
@@ -293,23 +293,18 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
token_to_kv_pool_allocator=token_to_kv_pool_allocator,
|
token_to_kv_pool_allocator=token_to_kv_pool_allocator,
|
||||||
)
|
)
|
||||||
|
|
||||||
def init_attention_backends(self):
|
def init_backends(self):
|
||||||
self._draft_worker.init_attention_backends()
|
disable_cuda_graph = False
|
||||||
|
if is_cuda() and not self.server_args.disable_cuda_graph:
|
||||||
def init_cuda_graphs(self):
|
|
||||||
capture_decode_cuda_graph = not self.server_args.disable_cuda_graph
|
|
||||||
if is_cuda() and capture_decode_cuda_graph:
|
|
||||||
available_mem = get_available_gpu_memory(self.device, self.gpu_id)
|
available_mem = get_available_gpu_memory(self.device, self.gpu_id)
|
||||||
if available_mem < 1.0:
|
disable_cuda_graph = available_mem < 1.0
|
||||||
capture_decode_cuda_graph = False
|
if disable_cuda_graph:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Disable DFLASH draft cuda graph because only %.2f GB GPU "
|
"Disable DFLASH draft cuda graph because only %.2f GB GPU "
|
||||||
"memory is available after target backend initialization.",
|
"memory is available after target backend initialization.",
|
||||||
available_mem,
|
available_mem,
|
||||||
)
|
)
|
||||||
self._draft_worker.init_cuda_graphs(
|
self._draft_worker.init_backends(disable_cuda_graph=disable_cuda_graph)
|
||||||
capture_decode_cuda_graph=capture_decode_cuda_graph
|
|
||||||
)
|
|
||||||
|
|
||||||
def _init_fused_kv_helper(self) -> None:
|
def _init_fused_kv_helper(self) -> None:
|
||||||
"""Initialize the fused KV materialization helper with pre-stacked weights."""
|
"""Initialize the fused KV materialization helper with pre-stacked weights."""
|
||||||
|
|||||||
@@ -242,21 +242,15 @@ class EagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
f"({draft_vocab_size}) != target vocab ({target_vocab_size})."
|
f"({draft_vocab_size}) != target vocab ({target_vocab_size})."
|
||||||
)
|
)
|
||||||
|
|
||||||
def init_attention_backends(self):
|
def init_backends(self):
|
||||||
with self.draft_tp_context(
|
with self.draft_tp_context(
|
||||||
self.draft_runner.tp_group
|
self.draft_runner.tp_group
|
||||||
), speculative_moe_backend_context(), speculative_moe_a2a_backend_context():
|
), speculative_moe_backend_context(), speculative_moe_a2a_backend_context():
|
||||||
self.draft_worker.init_attention_backends()
|
self.draft_worker.init_backends(disable_cuda_graph=True)
|
||||||
self.init_attention_backend()
|
self.init_attention_backend()
|
||||||
|
|
||||||
def init_cuda_graphs(self):
|
|
||||||
with self.draft_tp_context(
|
|
||||||
self.draft_runner.tp_group
|
|
||||||
), speculative_moe_backend_context(), speculative_moe_a2a_backend_context():
|
|
||||||
self.draft_worker.init_cuda_graphs(capture_decode_cuda_graph=False)
|
|
||||||
if check_cuda_graph_backend(Phase.PREFILL, Backend.BREAKABLE):
|
if check_cuda_graph_backend(Phase.PREFILL, Backend.BREAKABLE):
|
||||||
self.draft_runner.init_prefill_cuda_graph(force_for_draft_worker=True)
|
self.draft_runner.init_prefill_cuda_graph(force_for_draft_worker=True)
|
||||||
self._capture_cuda_graphs()
|
self.init_cuda_graphs()
|
||||||
|
|
||||||
if (c := self.draft_runner.canary_manager) is not None:
|
if (c := self.draft_runner.canary_manager) is not None:
|
||||||
c.mark_init_finished()
|
c.mark_init_finished()
|
||||||
@@ -364,8 +358,8 @@ class EagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
self.draft_runner.attn_backend = self.draft_extend_attn_backend
|
self.draft_runner.attn_backend = self.draft_extend_attn_backend
|
||||||
self.tree_mask_mode = TreeMaskMode.FULL_MASK
|
self.tree_mask_mode = TreeMaskMode.FULL_MASK
|
||||||
|
|
||||||
def _capture_cuda_graphs(self):
|
def init_cuda_graphs(self):
|
||||||
"""Capture the draft worker's own cuda graphs (decode + draft-extend)."""
|
"""Capture cuda graphs."""
|
||||||
self.cuda_graph_runner = None
|
self.cuda_graph_runner = None
|
||||||
self.cuda_graph_runner_for_draft_extend = None
|
self.cuda_graph_runner_for_draft_extend = None
|
||||||
|
|
||||||
@@ -989,11 +983,8 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
|||||||
self.req_to_token_pool = req_to_token_pool
|
self.req_to_token_pool = req_to_token_pool
|
||||||
self.token_to_kv_pool_allocator = token_to_kv_pool_allocator
|
self.token_to_kv_pool_allocator = token_to_kv_pool_allocator
|
||||||
|
|
||||||
def init_attention_backends(self):
|
def init_backends(self):
|
||||||
self._draft_worker.init_attention_backends()
|
self._draft_worker.init_backends()
|
||||||
|
|
||||||
def init_cuda_graphs(self):
|
|
||||||
self._draft_worker.init_cuda_graphs()
|
|
||||||
# Build adaptive runtime states after target and draft backends exist.
|
# Build adaptive runtime states after target and draft backends exist.
|
||||||
if self.adaptive_controller is not None:
|
if self.adaptive_controller is not None:
|
||||||
with (
|
with (
|
||||||
@@ -1241,7 +1232,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
|||||||
cuda_graph_bs=cuda_graph_bs,
|
cuda_graph_bs=cuda_graph_bs,
|
||||||
):
|
):
|
||||||
self._draft_worker.init_attention_backend()
|
self._draft_worker.init_attention_backend()
|
||||||
self._draft_worker._capture_cuda_graphs()
|
self._draft_worker.init_cuda_graphs()
|
||||||
|
|
||||||
# Build target attention backend and CUDA graph runner
|
# Build target attention backend and CUDA graph runner
|
||||||
target_model_runner = self._target_worker.model_runner
|
target_model_runner = self._target_worker.model_runner
|
||||||
|
|||||||
@@ -197,24 +197,16 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
def init_attention_backends(self):
|
def init_backends(self):
|
||||||
with (
|
with (
|
||||||
self.draft_tp_context(self.draft_model_runner.tp_group),
|
self.draft_tp_context(self.draft_model_runner.tp_group),
|
||||||
speculative_moe_backend_context(),
|
speculative_moe_backend_context(),
|
||||||
speculative_moe_a2a_backend_context(),
|
speculative_moe_a2a_backend_context(),
|
||||||
):
|
):
|
||||||
TpModelWorker.init_attention_backends(self)
|
TpModelWorker.init_backends(self, disable_cuda_graph=True)
|
||||||
self.draft_attn_backend = self._init_draft_attn_backend()
|
self.draft_attn_backend = self._init_draft_attn_backend()
|
||||||
self.draft_model_runner.draft_attn_backend = self.draft_attn_backend
|
self.draft_model_runner.draft_attn_backend = self.draft_attn_backend
|
||||||
|
self.init_cuda_graphs()
|
||||||
def init_cuda_graphs(self):
|
|
||||||
with (
|
|
||||||
self.draft_tp_context(self.draft_model_runner.tp_group),
|
|
||||||
speculative_moe_backend_context(),
|
|
||||||
speculative_moe_a2a_backend_context(),
|
|
||||||
):
|
|
||||||
TpModelWorker.init_cuda_graphs(self, capture_decode_cuda_graph=False)
|
|
||||||
self._capture_cuda_graphs()
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def draft_model_runner(self):
|
def draft_model_runner(self):
|
||||||
@@ -350,7 +342,7 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
|
|||||||
with self._frozen_kv_target_view(forward_batch):
|
with self._frozen_kv_target_view(forward_batch):
|
||||||
self.draft_attn_backend.init_forward_metadata_out_graph(fb_view)
|
self.draft_attn_backend.init_forward_metadata_out_graph(fb_view)
|
||||||
|
|
||||||
def _capture_cuda_graphs(self) -> None:
|
def init_cuda_graphs(self) -> None:
|
||||||
if cuda_graph_fully_disabled() or self.speculative_num_steps <= 1:
|
if cuda_graph_fully_disabled() or self.speculative_num_steps <= 1:
|
||||||
return
|
return
|
||||||
if self.target_worker.device != "cuda":
|
if self.target_worker.device != "cuda":
|
||||||
|
|||||||
@@ -200,17 +200,11 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
device=self.device,
|
device=self.device,
|
||||||
)
|
)
|
||||||
|
|
||||||
def init_attention_backends(self):
|
def init_backends(self):
|
||||||
with self.draft_tp_context(
|
with self.draft_tp_context(
|
||||||
self.draft_runner_list[0].tp_group
|
self.draft_runner_list[0].tp_group
|
||||||
), speculative_moe_backend_context():
|
), speculative_moe_backend_context():
|
||||||
super().init_attention_backends()
|
super().init_backends()
|
||||||
|
|
||||||
def init_cuda_graphs(self):
|
|
||||||
with self.draft_tp_context(
|
|
||||||
self.draft_runner_list[0].tp_group
|
|
||||||
), speculative_moe_backend_context():
|
|
||||||
super().init_cuda_graphs()
|
|
||||||
|
|
||||||
def mtp_model_runner(self, step: int):
|
def mtp_model_runner(self, step: int):
|
||||||
return self.draft_runner_list[step]
|
return self.draft_runner_list[step]
|
||||||
@@ -239,7 +233,8 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
self.draft_extend_attn_backend_list[-1]
|
self.draft_extend_attn_backend_list[-1]
|
||||||
)
|
)
|
||||||
|
|
||||||
def _capture_cuda_graphs(self):
|
def init_cuda_graphs(self):
|
||||||
|
"""Capture cuda graphs."""
|
||||||
self.cuda_graph_runner = None
|
self.cuda_graph_runner = None
|
||||||
self.cuda_graph_runner_for_draft_extend = None
|
self.cuda_graph_runner_for_draft_extend = None
|
||||||
|
|
||||||
@@ -726,11 +721,8 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
|
|||||||
self.req_to_token_pool = req_to_token_pool
|
self.req_to_token_pool = req_to_token_pool
|
||||||
self.token_to_kv_pool_allocator = token_to_kv_pool_allocator
|
self.token_to_kv_pool_allocator = token_to_kv_pool_allocator
|
||||||
|
|
||||||
def init_attention_backends(self):
|
def init_backends(self):
|
||||||
self._draft_worker.init_attention_backends()
|
self._draft_worker.init_backends()
|
||||||
|
|
||||||
def init_cuda_graphs(self):
|
|
||||||
self._draft_worker.init_cuda_graphs()
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def target_worker(self):
|
def target_worker(self):
|
||||||
|
|||||||
@@ -131,17 +131,11 @@ class StandaloneDraftWorker(EagleDraftWorker):
|
|||||||
self.init_token_map()
|
self.init_token_map()
|
||||||
self.init_lm_head()
|
self.init_lm_head()
|
||||||
|
|
||||||
def init_attention_backends(self):
|
def init_backends(self):
|
||||||
with self.draft_tp_context(
|
with self.draft_tp_context(
|
||||||
self.draft_runner.tp_group
|
self.draft_runner.tp_group
|
||||||
), speculative_moe_backend_context():
|
), speculative_moe_backend_context():
|
||||||
super().init_attention_backends()
|
super().init_backends()
|
||||||
|
|
||||||
def init_cuda_graphs(self):
|
|
||||||
with self.draft_tp_context(
|
|
||||||
self.draft_runner.tp_group
|
|
||||||
), speculative_moe_backend_context():
|
|
||||||
super().init_cuda_graphs()
|
|
||||||
|
|
||||||
def init_lm_head(self):
|
def init_lm_head(self):
|
||||||
"""Override to prevent sharing embeddings and lm_head with target model."""
|
"""Override to prevent sharing embeddings and lm_head with target model."""
|
||||||
|
|||||||
@@ -58,7 +58,7 @@ class SpecEagleServerBase(CustomTestCase):
|
|||||||
attention_backend = "flashinfer"
|
attention_backend = "flashinfer"
|
||||||
# Primary axis: False -> overlap scheduler; True -> synchronous (non-overlap).
|
# Primary axis: False -> overlap scheduler; True -> synchronous (non-overlap).
|
||||||
disable_overlap = False
|
disable_overlap = False
|
||||||
mem_fraction_static = 0.85
|
mem_fraction_static = 0.75
|
||||||
max_running_requests = 8
|
max_running_requests = 8
|
||||||
chunked_prefill_size = 128
|
chunked_prefill_size = 128
|
||||||
# bf16 rather than fp16: fp16 activations can overflow (-> Inf -> NaN) on
|
# bf16 rather than fp16: fp16 activations can overflow (-> Inf -> NaN) on
|
||||||
|
|||||||
Reference in New Issue
Block a user