[Spec] Split the capture width from num_tokens_per_req and gate replay on it (#31255)

This commit is contained in:
Liangsheng Yin
2026-07-16 14:59:11 -07:00
committed by GitHub
parent 059ac7efe8
commit fc1e3797b7
9 changed files with 141 additions and 83 deletions
@@ -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] * (
@@ -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
@@ -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)))
@@ -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
@@ -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
@@ -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(
@@ -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_()
@@ -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]
@@ -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