[Spec] Split the capture width from num_tokens_per_req and gate replay on it (#31255)
This commit is contained in:
@@ -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]
|
||||
|
||||
Reference in New Issue
Block a user