From fc1e3797b77b746345636177130d15c978b454ca Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Thu, 16 Jul 2026 14:59:11 -0700 Subject: [PATCH] [Spec] Split the capture width from `num_tokens_per_req` and gate replay on it (#31255) --- .../npu/graph_runner/npu_graph_runner.py | 2 +- .../srt/model_executor/cpu_graph_runner.py | 12 ++--- .../runner/base_cuda_graph_runner.py | 11 +++-- .../runner/decode_cuda_graph_runner.py | 49 ++++++++++++------- .../eagle_draft_cuda_graph_runner.py | 27 +++++++--- .../eagle_draft_extend_cuda_graph_runner.py | 49 ++++++++++++------- .../frozen_kv_mtp_cuda_graph_runner.py | 25 +++++++--- ...er_eagle_draft_extend_cuda_graph_runner.py | 47 +++++++++++------- .../test_eagle_draft_cuda_graph_runner.py | 2 +- 9 files changed, 141 insertions(+), 83 deletions(-) diff --git a/python/sglang/srt/hardware_backend/npu/graph_runner/npu_graph_runner.py b/python/sglang/srt/hardware_backend/npu/graph_runner/npu_graph_runner.py index cfae98079..b82c973bc 100644 --- a/python/sglang/srt/hardware_backend/npu/graph_runner/npu_graph_runner.py +++ b/python/sglang/srt/hardware_backend/npu/graph_runner/npu_graph_runner.py @@ -240,7 +240,7 @@ class NPUGraphRunner(DecodeCudaGraphRunner): or is_deepseek_v4(self.model_runner.model_config.hf_config) ): if forward_batch.forward_mode.is_target_verify(): - seq_lens_cpu = forward_batch.seq_lens.cpu() + self.num_tokens_per_req + seq_lens_cpu = forward_batch.seq_lens.cpu() + self.captured_req_width seq_lens = seq_lens_cpu.tolist() + [0] * (self.bs - self.raw_bs) else: seq_lens = forward_batch.seq_lens.cpu().tolist() + [0] * ( diff --git a/python/sglang/srt/model_executor/cpu_graph_runner.py b/python/sglang/srt/model_executor/cpu_graph_runner.py index 76a990fb3..28d91fb1c 100644 --- a/python/sglang/srt/model_executor/cpu_graph_runner.py +++ b/python/sglang/srt/model_executor/cpu_graph_runner.py @@ -591,7 +591,7 @@ class CPUGraphRunner: self.capture_forward_mode = ForwardMode.DECODE self.capture_hidden_mode = CaptureHiddenMode.NULL # Static capture width: CPU graphs are decode-only. - self.num_tokens_per_req = 1 + self.captured_req_width = 1 # If returning hidden states is enabled, set initial capture hidden mode to full to avoid double-capture on startup if self.enable_return_hidden_states: @@ -628,7 +628,7 @@ class CPUGraphRunner: self.captured_forward_batches_cross = {} # Attention backend self.max_bs = max(self.capture_bs) - self.max_num_token = self.max_bs * self.num_tokens_per_req + self.max_num_token = self.max_bs * self.captured_req_width self.model_runner.attn_backend.init_cpu_graph_state( self.max_bs, self.max_num_token ) @@ -656,7 +656,7 @@ class CPUGraphRunner: self.custom_mask = torch.ones( ( (self.seq_lens.sum().item() + self.max_num_token) - * self.num_tokens_per_req + * self.captured_req_width ), dtype=torch.bool, device=self.device, @@ -735,7 +735,7 @@ class CPUGraphRunner: with patch_model( self.model_runner.model, bs in self.capture_bs, - num_tokens=bs * self.num_tokens_per_req, + num_tokens=bs * self.captured_req_width, tp_group=self.model_runner.tp_group, ) as forward: graph, output_buffers = self.capture_one_batch_size( @@ -777,7 +777,7 @@ class CPUGraphRunner: def capture_one_batch_size( self, bs: int, forward: Callable, skip_cross_attention: bool = False ): - num_tokens = bs * self.num_tokens_per_req + num_tokens = bs * self.captured_req_width # Graph inputs input_ids = self.input_ids[:num_tokens] @@ -926,7 +926,7 @@ class CPUGraphRunner: self.model_runner.attn_backend.init_forward_metadata(forward_batch) return forward_batch - raw_num_token = raw_bs * self.num_tokens_per_req + raw_num_token = raw_bs * self.captured_req_width index = bisect.bisect_left(self.capture_bs, raw_bs) bs = self.capture_bs[index] assert bs > raw_bs diff --git a/python/sglang/srt/model_executor/runner/base_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/base_cuda_graph_runner.py index a241844f6..55f6fac78 100644 --- a/python/sglang/srt/model_executor/runner/base_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/base_cuda_graph_runner.py @@ -56,7 +56,7 @@ def freeze_gc(enable_cudagraph_gc: bool): def get_batch_sizes_to_capture( - model_runner: ModelRunner, num_tokens_per_req: int = 1 + model_runner: ModelRunner, captured_req_width: int = 1 ) -> Tuple[List[int], List[int]]: """Build the (capture_bs, compile_bs) lists for the decode runner. @@ -69,9 +69,12 @@ def get_batch_sizes_to_capture( num_max_requests = model_runner.req_to_token_pool.size mul_base = 1 + # TBO splits each request's rows across two micro-batches, so the + # alignment constraint applies per request rather than per token row. + alignment_width = captured_req_width if server_args.enable_two_batch_overlap: mul_base *= 2 - num_tokens_per_req = 1 + alignment_width = 1 if require_gathered_buffer(server_args): mul_base *= get_parallel().attn_tp_size @@ -86,8 +89,8 @@ def get_batch_sizes_to_capture( # is very small. We add more values here to make sure we capture the maximum bs. capture_bs += [num_max_requests] - # Model input token count = bs * num_tokens_per_req; must be a multiple of attn_tp_size. - capture_bs = [bs for bs in capture_bs if bs * num_tokens_per_req % mul_base == 0] + # Model input token count = bs * alignment_width; must be a multiple of attn_tp_size. + capture_bs = [bs for bs in capture_bs if bs * alignment_width % mul_base == 0] capture_bs = [bs for bs in capture_bs if bs <= num_max_requests] capture_bs = list(sorted(set(capture_bs))) diff --git a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py index e9c91b0d0..68c5ee820 100644 --- a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py @@ -259,7 +259,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): self.capture_forward_mode = ForwardMode.DECODE self.capture_hidden_mode = CaptureHiddenMode.NULL # Static capture width. - self.num_tokens_per_req = model_runner.decode_num_tokens_per_req( + self.captured_req_width = model_runner.decode_num_tokens_per_req( num_draft_tokens=self.speculative_num_draft_tokens ) if model_runner.spec_algorithm.is_speculative(): @@ -275,7 +275,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): # --- bucket sizes --------------------------------------------- self.capture_bs, self.compile_bs = get_batch_sizes_to_capture( - model_runner, self.num_tokens_per_req + model_runner, self.captured_req_width ) if KTRANSFORMERS_AVAILABLE: KTMoEWrapper.set_capture_batch_sizes(self.capture_bs) @@ -309,7 +309,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): # Attention backend self.max_bs = max(self.capture_bs) - self.max_num_token = self.max_bs * self.num_tokens_per_req + self.max_num_token = self.max_bs * self.captured_req_width self.attn_backend.init_cuda_graph_state(self.max_bs, self.max_num_token) # Init PDMux if needed @@ -336,7 +336,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): # lora_manager.init_cuda_graph_moe_buffers(). self.model_runner.lora_manager.init_cuda_graph_batch_info( max_bs_in_cuda_graph=self.max_bs, - num_tokens_per_req=self.num_tokens_per_req, + num_tokens_per_req=self.captured_req_width, ) enable_mamba_track = ( @@ -363,7 +363,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): require_mlp_tp_gather=self.require_mlp_tp_gather, seq_len_fill_value=self.seq_len_fill_value, encoder_len_fill_value=self.encoder_len_fill_value, - num_tokens_per_req=self.num_tokens_per_req, + num_tokens_per_req=self.captured_req_width, cache_loc_dtype=self._cache_loc_dtype(), enable_mamba_track=enable_mamba_track, ne_token_table=( @@ -411,7 +411,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): ) def _build_ragged_verify_token_buckets(self) -> list[int]: - buckets = sorted({bs * self.num_tokens_per_req for bs in self.capture_bs}) + buckets = sorted({bs * self.captured_req_width for bs in self.capture_bs}) assert buckets and buckets[0] > 0, f"{buckets=}" return buckets @@ -475,7 +475,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): def _ragged_capture_slots(self, num_tokens: int) -> int: if envs.SGLANG_TEST_RAGGED_VERIFY_FORCE_UNIFORM_CAPTURE.get(): - return num_tokens // self.num_tokens_per_req + return num_tokens // self.captured_req_width return min(num_tokens, self.max_bs) def _capture_ragged_verify_layout(self, num_tokens: int): @@ -491,7 +491,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): verify_lens_cpu = build_capture_verify_lens( num_tokens=num_tokens, num_slots=self._ragged_capture_slots(num_tokens), - num_draft_tokens=self.num_tokens_per_req, + num_draft_tokens=self.captured_req_width, ) return RaggedVerifyLayout.from_verify_lens( verify_lens_cpu=verify_lens_cpu, @@ -514,6 +514,17 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): if self.ragged_verify_mode and forward_batch.forward_mode.is_target_verify(): return False + # Uniform-width replay invariant: the batch's actual per-request width + # must match this runner's capture width; anything else falls back to + # eager. (Unset widths pass: not every path fills the field yet.) + spec_info = forward_batch.spec_info + if ( + spec_info is not None + and spec_info.num_tokens_per_req > 0 + and spec_info.num_tokens_per_req != self.captured_req_width + ): + return False + if self.require_mlp_tp_gather: # Raw sync values are per-rank request counts on decode-family # rounds -- no width division, no per-algorithm enumeration. @@ -562,7 +573,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): is_ngram_supported = ( ( - forward_batch.batch_size * self.num_tokens_per_req + forward_batch.batch_size * self.captured_req_width == forward_batch.input_ids.numel() ) if self.model_runner.spec_algorithm.is_ngram() @@ -670,7 +681,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): bs = size buffers: DecodeInputBuffers = self.buffers if num_tokens is None: - num_tokens = bs * self.num_tokens_per_req + num_tokens = bs * self.captured_req_width # Registry-owned FB-shared slots come through the registry (which # shares physical storage with self.buffers via source=...); the rest @@ -818,7 +829,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): if self.enable_torch_compile and not (get_flags().capture.enable_torch_compile): self.enable_torch_compile = False _, self.compile_bs = get_batch_sizes_to_capture( - self.model_runner, self.num_tokens_per_req + self.model_runner, self.captured_req_width ) profile_context = empty_context() if self.enable_profile_cuda_graph: @@ -892,7 +903,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): with torch_compile_decoration.patch_model( self.model_runner.model, bs in self.compile_bs, - num_tokens=bs * self.num_tokens_per_req, + num_tokens=bs * self.captured_req_width, tp_group=self.model_runner.tp_group, ) as forward: self.capture_one_shape(bs, forward, stream_idx, variant_label) @@ -904,7 +915,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): stream_idx: Optional[int] = None, variant_label: Optional[str] = None, ): - num_tokens = size * self.num_tokens_per_req + num_tokens = size * self.captured_req_width bs = self._ragged_capture_slots(num_tokens) if self.ragged_verify_mode else size # Sanity-check: --debug-cuda-graph requires breakable backend. @@ -1063,7 +1074,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): self._ragged_graph_size if is_ragged else self._capture_graph_size( - bs=self.bs, num_tokens=self.bs * self.num_tokens_per_req + bs=self.bs, num_tokens=self.bs * self.captured_req_width ) ) if is_ragged: @@ -1108,11 +1119,11 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): ) padded_num_tokens = graph_size_key else: - raw_num_token = raw_bs * self.num_tokens_per_req + raw_num_token = raw_bs * self.captured_req_width if self.require_mlp_tp_gather: max_num_tokens = max(forward_batch.global_num_tokens_cpu) max_batch_size = ( - max_num_tokens / self.num_tokens_per_req + max_num_tokens / self.captured_req_width if self.model_runner.spec_algorithm.is_eagle() or self.model_runner.spec_algorithm.is_standalone() or self.model_runner.spec_algorithm.is_dflash_family() @@ -1121,7 +1132,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): bs = self._pad_to_bucket(int(max_batch_size), self.capture_bs) else: bs = self._pad_to_bucket(raw_bs, self.capture_bs) - padded_num_tokens = bs * self.num_tokens_per_req + padded_num_tokens = bs * self.captured_req_width graph_size_key = self._capture_graph_size( bs=bs, num_tokens=padded_num_tokens ) @@ -1329,7 +1340,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): spec_info = DFlashVerifyInput( draft_token=None, positions=None, - draft_token_num=self.num_tokens_per_req, + draft_token_num=self.captured_req_width, custom_mask=( None if (self.model_runner.is_draft_worker or not build_custom_mask) @@ -1353,7 +1364,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): retrieve_index=None, retrieve_next_token=None, retrieve_next_sibling=None, - draft_token_num=self.num_tokens_per_req, + draft_token_num=self.captured_req_width, ) spec_info.capture_hidden_mode = CaptureHiddenMode.NULL diff --git a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py index 2cda3f115..3d09b2ef6 100644 --- a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py @@ -147,11 +147,11 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner): # Bucket sizes self.capture_bs, _ = get_batch_sizes_to_capture(model_runner) # Static capture width. - self.num_tokens_per_req = resolve_num_tokens_per_req( + self.captured_req_width = resolve_num_tokens_per_req( phase="draft_decode", server_args=model_runner.server_args ) self.max_bs = max(self.capture_bs) - self.max_num_token = self.max_bs * self.num_tokens_per_req + self.max_num_token = self.max_bs * self.captured_req_width # Attention backend init self.draft_attn_backend.init_cuda_graph_state(self.max_bs, self.max_num_token) @@ -292,6 +292,17 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner): # can_run_graph # ----------------------------------------------------------------- def can_run_graph(self, forward_batch: ForwardBatch): + # Uniform-width replay invariant: the batch's actual per-request width + # must match this runner's capture width; anything else falls back to + # eager. (Unset widths pass: not every path fills the field yet.) + spec_info = forward_batch.spec_info + if ( + spec_info is not None + and spec_info.num_tokens_per_req > 0 + and spec_info.num_tokens_per_req != self.captured_req_width + ): + return False + if self.require_mlp_tp_gather: # Raw sync values are per-rank request counts on decode-family rounds. cuda_graph_bs = max(forward_batch.original_global_num_tokens_cpu) @@ -321,7 +332,7 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner): ): num_seqs = size # EAGLE legacy name buffers = self.buffers - num_tokens = num_seqs * self.num_tokens_per_req + num_tokens = num_seqs * self.captured_req_width # Graph inputs req_pool_indices = buffers.req_pool_indices[:num_seqs] @@ -492,13 +503,13 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner): buffers = self.buffers raw_bs = forward_batch.batch_size - raw_num_token = raw_bs * self.num_tokens_per_req + raw_num_token = raw_bs * self.captured_req_width # Pad to nearest captured shape if self.require_mlp_tp_gather: max_num_tokens = max(forward_batch.global_num_tokens_cpu) max_batch_size = ( - max_num_tokens // self.num_tokens_per_req + max_num_tokens // self.captured_req_width if self.model_runner.spec_algorithm.is_eagle() or self.model_runner.spec_algorithm.is_standalone() else max_num_tokens @@ -525,7 +536,7 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner): buffers.dsa_seed_topk.zero_() buffers.req_pool_indices.zero_() - num_tokens = bs * self.num_tokens_per_req + num_tokens = bs * self.captured_req_width maybe_detect_nan( forward_batch.spec_info.topk_p, @@ -600,9 +611,9 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner): # TODO(ch-wan): support num_token_non_padded if self.require_gathered_buffer: - buffers.global_num_tokens_gpu.fill_(bs * self.num_tokens_per_req) + buffers.global_num_tokens_gpu.fill_(bs * self.captured_req_width) buffers.global_num_tokens_for_logprob_gpu.fill_( - bs * self.num_tokens_per_req + bs * self.captured_req_width ) # Save the raw seq_lens_sum; it is restored after replay. While the graph diff --git a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py index 71844c60d..f95ae9b1e 100644 --- a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py @@ -136,11 +136,11 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): # Static capture width: full tree width (num_draft_tokens), not # num_steps + 1 -- topk > 1 draft-extend overflows the buffers. - self.num_tokens_per_req = resolve_num_tokens_per_req( + self.captured_req_width = resolve_num_tokens_per_req( phase="draft_extend", server_args=model_runner.server_args ) self.max_bs = max(self.capture_bs) - self.max_num_token = self.max_bs * self.num_tokens_per_req + self.max_num_token = self.max_bs * self.captured_req_width self.draft_extend_attn_backend.init_cuda_graph_state( self.max_bs, self.max_num_token @@ -148,7 +148,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): self.seq_len_fill_value = ( self.draft_extend_attn_backend.get_cuda_graph_seq_len_fill_value() ) - self.extend_seq_lens_cpu = [self.num_tokens_per_req] * self.max_bs + self.extend_seq_lens_cpu = [self.captured_req_width] * self.max_bs if self.enable_torch_compile: set_torch_compile_config() @@ -186,13 +186,13 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): (self.max_bs,), self.seq_len_fill_value, dtype=torch.int64 ) extend_seq_lens = torch.full( - (self.max_bs,), self.num_tokens_per_req, dtype=torch.int32 + (self.max_bs,), self.captured_req_width, dtype=torch.int32 ) num_correct_drafts = torch.full( - (self.max_bs,), self.num_tokens_per_req, dtype=torch.int32 + (self.max_bs,), self.captured_req_width, dtype=torch.int32 ) num_accept_tokens = torch.full( - (self.max_bs,), self.num_tokens_per_req, dtype=torch.int32 + (self.max_bs,), self.captured_req_width, dtype=torch.int32 ) if self.require_gathered_buffer: @@ -231,7 +231,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): next_token_logits_buffer = ( self.model_runner.graph_shared_output.get_logits_buffer( - vocab_size, rows=self.max_bs * self.num_tokens_per_req + vocab_size, rows=self.max_bs * self.captured_req_width ) ) @@ -289,6 +289,17 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): return ShapeKey(size=bs) def can_run_graph(self, forward_batch: ForwardBatch): + # Uniform-width replay invariant: the batch's actual per-request width + # must match this runner's capture width; anything else falls back to + # eager. (Unset widths pass: not every path fills the field yet.) + spec_info = forward_batch.spec_info + if ( + spec_info is not None + and spec_info.num_tokens_per_req > 0 + and spec_info.num_tokens_per_req != self.captured_req_width + ): + return False + if self.require_mlp_tp_gather: # Raw sync values are per-rank request counts on decode-family rounds. cuda_graph_bs = max(forward_batch.original_global_num_tokens_cpu) @@ -315,7 +326,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): ): bs = size buffers = self.buffers - num_tokens = bs * self.num_tokens_per_req + num_tokens = bs * self.captured_req_width # Graph inputs input_ids = buffers.input_ids[:num_tokens] @@ -370,7 +381,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): num_correct_drafts=num_correct_drafts, num_accept_tokens=num_accept_tokens, # Padded tree width per req; drives the constant qo layout. - num_tokens_per_req=self.num_tokens_per_req, + num_tokens_per_req=self.captured_req_width, ) forward_batch = ForwardBatch( @@ -481,7 +492,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): if self.require_mlp_tp_gather: max_num_tokens = max(forward_batch.global_num_tokens_cpu) max_batch_size = ( - max_num_tokens // self.num_tokens_per_req + max_num_tokens // self.captured_req_width if self.model_runner.spec_algorithm.is_eagle() else max_num_tokens ) @@ -489,16 +500,16 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): else: bs = self._pad_to_bucket(raw_bs, self.capture_bs) - if bs * self.num_tokens_per_req != num_tokens: + if bs * self.captured_req_width != num_tokens: buffers.seq_lens.fill_(self.seq_len_fill_value) buffers.out_cache_loc.zero_() buffers.positions.zero_() # Pair with seq_lens fill: padded rows must point at reserved # req_pool slot 0 (req_to_token[0, :] is all zeros from init). buffers.req_pool_indices.zero_() - buffers.num_correct_drafts.fill_(self.num_tokens_per_req) - buffers.num_accept_tokens.fill_(self.num_tokens_per_req) - buffers.extend_seq_lens.fill_(self.num_tokens_per_req) + buffers.num_correct_drafts.fill_(self.captured_req_width) + buffers.num_accept_tokens.fill_(self.captured_req_width) + buffers.extend_seq_lens.fill_(self.captured_req_width) # Batch the small per-field device copies into a grouped foreach copy # (one foreach call per dtype pair) to cut launch overhead. hidden_states @@ -522,7 +533,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): copy_dsts.append(buffers.extend_seq_lens[:raw_bs]) copy_srcs.append(forward_batch.extend_seq_lens) else: - buffers.extend_seq_lens[:raw_bs].fill_(self.num_tokens_per_req) + buffers.extend_seq_lens[:raw_bs].fill_(self.captured_req_width) if forward_batch.spec_info.num_correct_drafts is not None: copy_dsts.append(buffers.num_correct_drafts[:raw_bs]) copy_srcs.append(forward_batch.spec_info.num_correct_drafts) @@ -544,9 +555,9 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): # TODO(ch-wan): support num_token_non_padded if self.require_gathered_buffer: - buffers.global_num_tokens_gpu.fill_(bs * self.num_tokens_per_req) + buffers.global_num_tokens_gpu.fill_(bs * self.captured_req_width) buffers.global_num_tokens_for_logprob_gpu.fill_( - bs * self.num_tokens_per_req + bs * self.captured_req_width ) if forward_batch.seq_lens_cpu is not None: @@ -557,9 +568,9 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): if forward_batch.extend_seq_lens_cpu is not None: self.extend_seq_lens_cpu[:raw_bs] = forward_batch.extend_seq_lens_cpu else: - self.extend_seq_lens_cpu[:raw_bs] = [self.num_tokens_per_req] * raw_bs + self.extend_seq_lens_cpu[:raw_bs] = [self.captured_req_width] * raw_bs if bs > raw_bs: - self.extend_seq_lens_cpu[raw_bs:bs] = [self.num_tokens_per_req] * ( + self.extend_seq_lens_cpu[raw_bs:bs] = [self.captured_req_width] * ( bs - raw_bs ) forward_batch.spec_info.extend_seq_lens_cpu = list( 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 bed3369f5..05486883b 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 @@ -115,14 +115,14 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner): self.capture_hidden_mode = CaptureHiddenMode.LAST # Static capture width. - self.num_tokens_per_req = resolve_num_tokens_per_req( + self.captured_req_width = resolve_num_tokens_per_req( phase="draft_decode", server_args=model_runner.server_args ) self.capture_bs, _ = get_batch_sizes_to_capture( - model_runner, self.num_tokens_per_req + model_runner, self.captured_req_width ) self.max_bs = max(self.capture_bs) - self.max_num_token = self.max_bs * self.num_tokens_per_req + self.max_num_token = self.max_bs * self.captured_req_width self.draft_attn_backend.init_cuda_graph_state(self.max_bs, self.max_num_token) self.seq_len_fill_value = ( @@ -201,6 +201,17 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner): return self.backend.replay(shape_key, forward_batch) def can_run_graph(self, forward_batch: ForwardBatch): + # Uniform-width replay invariant: the batch's actual per-request width + # must match this runner's capture width; anything else falls back to + # eager. (Unset widths pass: not every path fills the field yet.) + spec_info = forward_batch.spec_info + if ( + spec_info is not None + and spec_info.num_tokens_per_req > 0 + and spec_info.num_tokens_per_req != self.captured_req_width + ): + return False + if self.require_mlp_tp_gather: # Raw sync values are per-rank request counts on decode-family # rounds; / topk maps the expanded batch to graph-key units @@ -234,7 +245,7 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner): del forward, stream_idx, variant_label buffers = self.buffers request_bs = size - expanded_bs = request_bs * self.num_tokens_per_req + expanded_bs = request_bs * self.captured_req_width req_pool_indices = buffers.req_pool_indices[:expanded_bs] positions = buffers.positions[:expanded_bs] @@ -370,7 +381,7 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner): raw_expanded_bs = forward_batch.batch_size raw_bs = ( - raw_expanded_bs // self.num_tokens_per_req + raw_expanded_bs // self.captured_req_width if self.topk > 1 else raw_expanded_bs ) @@ -379,13 +390,13 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner): if self.require_mlp_tp_gather: max_num_tokens = max(forward_batch.global_num_tokens_cpu) max_batch_size = max_num_tokens // ( - self.num_tokens_per_req * self.num_tokens_per_req + self.captured_req_width * self.captured_req_width ) bs = self._pad_to_bucket(int(max_batch_size), self.capture_bs) else: bs = self._pad_to_bucket(raw_bs, self.capture_bs) - expanded_bs = bs * self.num_tokens_per_req + expanded_bs = bs * self.captured_req_width if bs != raw_bs: buffers.seq_lens.fill_(self.seq_len_fill_value) buffers.positions.zero_() diff --git a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py index 9179902cd..07aa50dec 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py @@ -162,12 +162,12 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): # Fixed window: every step extends each request by the same number of # tokens, which lets all steps share one buffer set. - self.num_tokens_per_req = resolve_num_tokens_per_req( + self.captured_req_width = resolve_num_tokens_per_req( phase="draft_extend", server_args=model_runner.server_args ) self.max_bs = max(self.capture_bs) - self.max_num_token = self.max_bs * self.num_tokens_per_req - self.extend_seq_lens_cpu = [self.num_tokens_per_req] * self.max_bs + self.max_num_token = self.max_bs * self.captured_req_width + self.extend_seq_lens_cpu = [self.captured_req_width] * self.max_bs self.eagle_worker.draft_extend_attn_backend_list[ self.step @@ -200,6 +200,17 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): return ShapeKey(size=bs) def can_run_graph(self, forward_batch: ForwardBatch): + # Uniform-width replay invariant: the batch's actual per-request width + # must match this runner's capture width; anything else falls back to + # eager. (Unset widths pass: not every path fills the field yet.) + spec_info = forward_batch.spec_info + if ( + spec_info is not None + and spec_info.num_tokens_per_req > 0 + and spec_info.num_tokens_per_req != self.captured_req_width + ): + return False + if self.require_mlp_tp_gather: # Raw sync values are per-rank request counts on decode-family rounds. cuda_graph_bs = max(forward_batch.original_global_num_tokens_cpu) @@ -219,7 +230,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): def get_forward_batch(self, bs: int) -> ForwardBatch: buffers = self.buffers - num_tokens = bs * self.num_tokens_per_req + num_tokens = bs * self.captured_req_width input_ids = buffers.input_ids[:num_tokens] req_pool_indices = buffers.req_pool_indices[:bs] @@ -270,7 +281,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): num_correct_drafts=num_correct_drafts, num_accept_tokens=num_accept_tokens, ) - spec_info.num_tokens_per_req = self.num_tokens_per_req + spec_info.num_tokens_per_req = self.captured_req_width spec_info.positions = None capture_mode = ( @@ -304,8 +315,8 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): extend_seq_lens=extend_seq_lens, extend_seq_lens_cpu=extend_seq_lens_cpu, extend_start_loc=extend_start_loc, - extend_num_tokens=self.num_tokens_per_req * bs, - num_token_non_padded_cpu=self.num_tokens_per_req * bs, + extend_num_tokens=self.captured_req_width * bs, + num_token_non_padded_cpu=self.captured_req_width * bs, return_hidden_states_before_norm=True, ) return forward_batch @@ -333,7 +344,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): bs = size buffers = self.buffers - num_tokens = bs * self.num_tokens_per_req + num_tokens = bs * self.captured_req_width forward_batch = self.get_forward_batch(bs) forward_batch = self._postprocess_forward_batch(forward_batch, bs) attn_backend = self.eagle_worker.draft_extend_attn_backend_list[self.step] @@ -405,7 +416,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): write + worker-side rotation (steps > 0).""" self.deepep_adapter.replay() buffers = self.buffers - num_tokens = bs * self.num_tokens_per_req + num_tokens = bs * self.captured_req_width if self.require_gathered_buffer: buffers.global_num_tokens_gpu.fill_(num_tokens) @@ -460,7 +471,7 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner: self.runners: List[Optional[MultiLayerEagleDraftExtendCudaGraphRunner]] = [] self.seq_len_fill_value = 1 self.max_bs = 1 - self.num_tokens_per_req = 1 + self.captured_req_width = 1 self._init_and_capture() @@ -494,7 +505,7 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner: self.runners.append(runner) self.seq_len_fill_value = runner.seq_len_fill_value self.max_bs = runner.max_bs - self.num_tokens_per_req = runner.num_tokens_per_req + self.captured_req_width = runner.captured_req_width self.capture_bs = runner.capture_bs self.require_gathered_buffer = runner.require_gathered_buffer self.require_mlp_tp_gather = runner.require_mlp_tp_gather @@ -540,7 +551,7 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner: runner = next(r for r in self.runners if r is not None) model_runner = runner.model_runner max_bs = self.max_bs - num_tokens_per_req = self.num_tokens_per_req + num_tokens_per_req = self.captured_req_width max_num_token = max_bs * num_tokens_per_req hidden_size = get_draft_input_from_target_hidden_dim(model_runner) dtype = model_runner.model_config.dtype @@ -617,7 +628,7 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner: the batch size. Subsequent ``replay(step)`` calls reuse this state.""" buffers = self.buffers raw_bs = forward_batch.batch_size - num_tokens = raw_bs * self.num_tokens_per_req + num_tokens = raw_bs * self.captured_req_width # Bucketize to a captured batch size (padding the tail). if self.require_mlp_tp_gather: @@ -668,24 +679,24 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner: # and by the worker's rotation. arange = torch.arange(bs, device=self.device, dtype=torch.int64) buffers.select_index[:bs].copy_( - arange * self.num_tokens_per_req + buffers.num_correct_drafts[:bs] + arange * self.captured_req_width + buffers.num_correct_drafts[:bs] ) if self.require_gathered_buffer: - buffers.global_num_tokens_gpu.fill_(bs * self.num_tokens_per_req) + buffers.global_num_tokens_gpu.fill_(bs * self.captured_req_width) buffers.global_num_tokens_for_logprob_gpu.fill_( - bs * self.num_tokens_per_req + bs * self.captured_req_width ) # Reusable spec_info for per-step attention metadata. - padded_num_tokens = bs * self.num_tokens_per_req + padded_num_tokens = bs * self.captured_req_width spec_info = EagleDraftExtendInput( hidden_states=buffers.hidden_states[:padded_num_tokens], num_correct_drafts=buffers.num_correct_drafts[:bs], num_accept_tokens=buffers.num_accept_tokens[:bs], ) # Actual width of the captured forward == static width by construction. - spec_info.num_tokens_per_req = self.num_tokens_per_req + spec_info.num_tokens_per_req = self.captured_req_width spec_info.num_tokens_for_logprob_per_req = 1 spec_info.positions = buffers.positions[:padded_num_tokens] spec_info.extend_seq_lens_tensor = buffers.extend_seq_lens[:bs] diff --git a/test/registered/unit/spec/test_eagle_draft_cuda_graph_runner.py b/test/registered/unit/spec/test_eagle_draft_cuda_graph_runner.py index 2af5bae32..51e4e9722 100644 --- a/test/registered/unit/spec/test_eagle_draft_cuda_graph_runner.py +++ b/test/registered/unit/spec/test_eagle_draft_cuda_graph_runner.py @@ -77,7 +77,7 @@ class TestEagleDraftCudaGraphRunner(CustomTestCase): dsa_seed_topk=None, ) runner.capture_bs = [1, CAPTURE_BS] - runner.num_tokens_per_req = 1 + runner.captured_req_width = 1 runner.speculative_num_steps = NUM_STEPS runner.seq_len_fill_value = SEQ_LEN_FILL_VALUE runner.require_mlp_tp_gather = False