diff --git a/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py b/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py index 35bf0bc4a..2ee4a53b0 100644 --- a/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py @@ -1,15 +1,11 @@ from __future__ import annotations -import bisect from dataclasses import dataclass from typing import TYPE_CHECKING, Callable, Optional import torch -import tqdm from sglang.srt.compilation.torch_compile_decoration import set_torch_compile_config -from sglang.srt.distributed import get_tensor_model_parallel_rank -from sglang.srt.distributed.parallel_state import graph_capture from sglang.srt.layers.dp_attention import ( DpPaddingMode, set_dp_buffer_len, @@ -23,13 +19,13 @@ from sglang.srt.model_executor.forward_batch_info import ( from sglang.srt.model_executor.forward_context import ForwardContext, forward_context from sglang.srt.model_executor.input_buffers import ForwardInputBuffers from sglang.srt.model_executor.runner import ( + DecodeCudaGraphRunner, DeepEPCudaGraphRunnerAdapter, - freeze_gc, + ShapeKey, get_batch_sizes_to_capture, - get_global_graph_memory_pool, model_capture_mode, - set_global_graph_memory_pool, ) +from sglang.srt.model_executor.runner_backend.utils import resolve_decode_backend from sglang.srt.model_executor.runner_backend_utils import ( CUDA_GRAPH_CAPTURE_FAILED_MSG, ) @@ -61,14 +57,30 @@ class FrozenKVMTPInputBuffers(ForwardInputBuffers): global_num_tokens_for_logprob_gpu: Optional[torch.Tensor] -class FrozenKVMTPCudaGraphRunner: - """CUDA graph runner for the Frozen-KV MTP recurrent draft-loop step.""" +class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner): + """CUDA graph runner for the Frozen-KV MTP recurrent draft-loop step. + + Subclasses DecodeCudaGraphRunner to inherit the outer capture loop + (capture() / _capture_one_stream()), the bucket-padding helper + (_pad_to_bucket), and the backend-driven capture/replay scaffolding. + Frozen-KV-MTP-specific bits — the buffer dataclass, the dummy + ForwardBatch + FrozenKVMTPDraftInput built in capture_one_shape, the + target-KV-pool swap during capture, the worker's frozen-KV metadata + helpers, the topk*topk bucket math, the expanded-bs bookkeeping, and + the 3-tuple replay output — are overridden. + + Like the EAGLE draft runner, it does NOT call + DecodeCudaGraphRunner.__init__ (that init sets up decode-only state); + it sets up its own fields directly while satisfying the parent's + capture() / backend contract. + """ def __init__(self, frozen_kv_mtp_worker: FrozenKVMTPDraftWorker): self.frozen_kv_mtp_worker = frozen_kv_mtp_worker self.model_runner = model_runner = frozen_kv_mtp_worker.draft_model_runner - self.graphs = {} - self.output_buffers = {} + + self.device = model_runner.device + self.device_module = torch.get_device_module(self.device) self.enable_torch_compile = model_runner.server_args.enable_torch_compile self.disable_padding = model_runner.server_args.disable_cuda_graph_padding self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args) @@ -77,17 +89,28 @@ class FrozenKVMTPCudaGraphRunner: self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args) self.tp_size = self.model_runner.tp_size self.dp_size = self.model_runner.dp_size + self.pp_size = model_runner.server_args.pp_size self.speculative_num_steps = model_runner.server_args.speculative_num_steps self.topk = model_runner.server_args.speculative_eagle_topk self.draft_attn_backend = frozen_kv_mtp_worker.draft_attn_backend self.enable_profile_cuda_graph = ( model_runner.server_args.enable_profile_cuda_graph ) + + self.attn_backend = self.draft_attn_backend + + self.compile_bs = [] self.enable_pdmux = False + self.record_nolora_graph = False + self.is_dllm = False + self.deepep_adapter = DeepEPCudaGraphRunnerAdapter() + self.capture_forward_mode = ForwardMode.DECODE + self.capture_hidden_mode = CaptureHiddenMode.LAST + self.num_tokens_per_bs = self.topk - self.capture_bs, self.compile_bs = get_batch_sizes_to_capture( + self.capture_bs, _ = get_batch_sizes_to_capture( model_runner, self.num_tokens_per_bs ) self.max_bs = max(self.capture_bs) @@ -152,6 +175,8 @@ class FrozenKVMTPCudaGraphRunner: ) self.buffers.share_buffers() + self.backend = resolve_decode_backend(self) + try: with model_capture_mode(): self.capture() @@ -161,6 +186,12 @@ class FrozenKVMTPCudaGraphRunner: f"{CUDA_GRAPH_CAPTURE_FAILED_MSG}" ) + def _make_graph_key(self, bs, stream_idx=None, variant_label=None): + return ShapeKey(size=bs) + + def _replay_graph(self, shape_key, forward_batch): + return self.backend.replay(shape_key, forward_batch) + def can_run(self, forward_batch: ForwardBatch): if self.require_mlp_tp_gather: cuda_graph_bs = max(forward_batch.global_num_tokens_cpu) // ( @@ -174,7 +205,7 @@ class FrozenKVMTPCudaGraphRunner: ) is_bs_supported = ( - cuda_graph_bs in self.graphs + self.backend.can_run(forward_batch, self._make_graph_key(cuda_graph_bs)) if self.disable_padding else cuda_graph_bs <= self.max_bs ) @@ -182,45 +213,16 @@ class FrozenKVMTPCudaGraphRunner: is_bs_supported = is_bs_supported and forward_batch.can_run_dp_cuda_graph return is_bs_supported - def _create_graph(self): - return torch.cuda.CUDAGraph() - - def _capture_init(self, run_once_fn): - for _ in range(2): - torch.cuda.synchronize() - self.model_runner.tp_group.barrier() - run_once_fn() - - def _capture_graph(self, graph, pool, stream, run_once_fn): - with torch.cuda.graph(graph, pool=pool, stream=stream): - out = run_once_fn() - return out - - def _replay(self): - self.graphs[self.bs].replay() - - def capture(self): - with freeze_gc(self.model_runner.server_args.enable_cudagraph_gc): - with graph_capture() as graph_capture_context: - self.stream = graph_capture_context.stream - capture_range = ( - tqdm.tqdm(list(reversed(self.capture_bs))) - if get_tensor_model_parallel_rank() == 0 - else reversed(self.capture_bs) - ) - for bs in capture_range: - graph, output_buffers = self.capture_one_batch_size(bs, None) - self.graphs[bs] = graph - self.output_buffers[bs] = output_buffers - - def capture_one_batch_size( - self, num_seqs: int, forward: Callable, stream_idx: int = 0 + def capture_one_shape( + self, + size: int, + forward: Callable, + stream_idx: Optional[int] = None, + variant_label: Optional[str] = None, ): - del forward, stream_idx + del forward, stream_idx, variant_label buffers = self.buffers - graph = self._create_graph() - stream = self.stream - request_bs = num_seqs + request_bs = size expanded_bs = request_bs * self.num_tokens_per_bs req_pool_indices = buffers.req_pool_indices[:expanded_bs] @@ -318,14 +320,17 @@ class FrozenKVMTPCudaGraphRunner: forward_batch ) self.deepep_adapter.capture(is_extend_in_batch=False) - self._capture_init(run_once) - out = self._capture_graph( - graph, get_global_graph_memory_pool(), stream, run_once + shape_key = self._make_graph_key(request_bs) + self.backend.capture_one( + shape_key, + run_once, + dummies=None, + post_warmup_hook=getattr( + self.draft_attn_backend, "on_after_cuda_graph_warmup", None + ), ) finally: self.draft_attn_backend.token_to_kv_pool = saved_backend_pool - set_global_graph_memory_pool(graph.pool()) - return graph, out def _postprocess_output_to_raw_bs(self, out, raw_bs): parent_list, top_scores_index, draft_tokens = (t[:raw_bs] for t in out) @@ -348,11 +353,10 @@ class FrozenKVMTPCudaGraphRunner: max_batch_size = max_num_tokens // ( self.num_tokens_per_bs * self.num_tokens_per_bs ) - index = bisect.bisect_left(self.capture_bs, max_batch_size) + bs = self._pad_to_bucket(int(max_batch_size), self.capture_bs) else: - index = bisect.bisect_left(self.capture_bs, raw_bs) + bs = self._pad_to_bucket(raw_bs, self.capture_bs) - bs = self.capture_bs[index] expanded_bs = bs * self.num_tokens_per_bs if bs != raw_bs: buffers.seq_lens.fill_(self.seq_len_fill_value) @@ -400,14 +404,14 @@ class FrozenKVMTPCudaGraphRunner: self.raw_bs = raw_bs self.bs = bs + shape_key = self._make_graph_key(bs) # NVTX span: the graph bypasses `model_runner.forward`'s record_function. span_name = f"step[DRAFT_LOOP raw_bs={raw_bs} bs={bs} topk={self.topk}]" if torch.autograd._profiler_enabled(): with torch.profiler.record_function(span_name): - self._replay() + out = self._replay_graph(shape_key, forward_batch) else: - self._replay() - out = self.output_buffers[bs] + out = self._replay_graph(shape_key, forward_batch) if bs != raw_bs: out = self._postprocess_output_to_raw_bs(out, raw_bs) diff --git a/python/sglang/test/kits/attention_unittest/runner_modes/speculative_draft_runner.py b/python/sglang/test/kits/attention_unittest/runner_modes/speculative_draft_runner.py index 52bd0ec54..61e997b58 100644 --- a/python/sglang/test/kits/attention_unittest/runner_modes/speculative_draft_runner.py +++ b/python/sglang/test/kits/attention_unittest/runner_modes/speculative_draft_runner.py @@ -459,13 +459,17 @@ def _capture_frozen_kv_mtp_graph_runner( ) -> FrozenKVMTPCudaGraphRunner: with ( patch( - "sglang.srt.speculative.frozen_kv_mtp_cuda_graph_runner.graph_capture", + "sglang.srt.model_executor.runner.decode_cuda_graph_runner.graph_capture", _single_rank_graph_capture, ), patch( - "sglang.srt.speculative.frozen_kv_mtp_cuda_graph_runner.get_tensor_model_parallel_rank", + "sglang.srt.model_executor.runner.decode_cuda_graph_runner.get_tensor_model_parallel_rank", lambda: 0, ), + patch( + "sglang.srt.model_executor.runner.decode_cuda_graph_runner.get_available_gpu_memory", + lambda *args, **kwargs: 0.0, + ), patch( "sglang.srt.model_executor.runner.base_cuda_graph_runner.get_attention_cp_size", lambda: 1,