[refactor] Fold FrozenKVMTPCudaGraphRunner onto the shared DecodeCudaGraphRunner base (#28081)
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user