[Spec] Rename num_tokens_per_bs to num_tokens_per_req (#30977)

This commit is contained in:
Liangsheng Yin
2026-07-13 13:47:53 -05:00
committed by GitHub
parent b677babc62
commit c0f1f7e062
38 changed files with 282 additions and 246 deletions
@@ -1052,14 +1052,14 @@ class DeepseekV4AscendAttnBackend(
device = self.device device = self.device
if forward_mode.is_target_verify() or forward_mode.is_draft_extend_v2(): if forward_mode.is_target_verify() or forward_mode.is_draft_extend_v2():
tokens_per_bs = self.speculative_num_draft_tokens tokens_per_req = self.speculative_num_draft_tokens
else: else:
tokens_per_bs = 1 tokens_per_req = 1
metadata.actual_seq_lengths_q_pa = torch.arange( metadata.actual_seq_lengths_q_pa = torch.arange(
0, 0,
bs * tokens_per_bs + tokens_per_bs, bs * tokens_per_req + tokens_per_req,
tokens_per_bs, tokens_per_req,
dtype=torch.int32, dtype=torch.int32,
device=device, device=device,
) )
@@ -1081,7 +1081,7 @@ class DeepseekV4AscendAttnBackend(
:bs, : :bs, :
] ]
n_tok = bs * tokens_per_bs n_tok = bs * tokens_per_req
c4_pad = min(n_tok, n_tok // 4 + bs) c4_pad = min(n_tok, n_tok // 4 + bs)
c128_pad = min(n_tok, n_tok // 128 + bs) c128_pad = min(n_tok, n_tok // 128 + bs)
metadata.swa_loc = torch.zeros(n_tok, dtype=torch.int64, device=device) metadata.swa_loc = torch.zeros(n_tok, dtype=torch.int64, device=device)
@@ -1106,7 +1106,7 @@ class DeepseekV4AscendAttnBackend(
"li_quant_metadata": self.graph_metadata["kernel_metadata_li_quant"], "li_quant_metadata": self.graph_metadata["kernel_metadata_li_quant"],
} }
T = bs * tokens_per_bs T = bs * tokens_per_req
metadata.c4_topk_indices = self.graph_metadata["c4_topk_indices"][:T, :] metadata.c4_topk_indices = self.graph_metadata["c4_topk_indices"][:T, :]
self.forward_metadata = metadata self.forward_metadata = metadata
@@ -1126,16 +1126,16 @@ class DeepseekV4AscendAttnBackend(
device = seq_lens.device device = seq_lens.device
if forward_mode.is_target_verify() or forward_mode.is_draft_extend_v2(): if forward_mode.is_target_verify() or forward_mode.is_draft_extend_v2():
tokens_per_bs = self.speculative_num_draft_tokens tokens_per_req = self.speculative_num_draft_tokens
else: else:
tokens_per_bs = 1 tokens_per_req = 1
seq_lens_cpu = forward_batch.seq_lens_cpu seq_lens_cpu = forward_batch.seq_lens_cpu
assert seq_lens_cpu is not None, "V4 graph replay requires seq_lens_cpu." assert seq_lens_cpu is not None, "V4 graph replay requires seq_lens_cpu."
if forward_mode.is_target_verify(): if forward_mode.is_target_verify():
# In graph replay, buffers.seq_lens already contains the attention KV # In graph replay, buffers.seq_lens already contains the attention KV
# length (live length + draft tokens). Padded rows therefore show up as # length (live length + draft tokens). Padded rows therefore show up as
# tokens_per_bs instead of 0. Use the CPU live lengths as the source of # tokens_per_req instead of 0. Use the CPU live lengths as the source of
# truth so padded rows stay masked out. # truth so padded rows stay masked out.
live_seq_lens = seq_lens_cpu[:bs].to(device=device, dtype=torch.int32) live_seq_lens = seq_lens_cpu[:bs].to(device=device, dtype=torch.int32)
elif seq_lens is not None and seq_lens.device.type != "cpu": elif seq_lens is not None and seq_lens.device.type != "cpu":
@@ -1145,9 +1145,9 @@ class DeepseekV4AscendAttnBackend(
attn_seq_lens = live_seq_lens attn_seq_lens = live_seq_lens
if forward_mode.is_target_verify(): if forward_mode.is_target_verify():
valid_verify_rows = live_seq_lens > 0 valid_verify_rows = live_seq_lens > 0
attn_seq_lens = live_seq_lens + int(tokens_per_bs) attn_seq_lens = live_seq_lens + int(tokens_per_req)
attn_seq_lens = torch.where(valid_verify_rows, attn_seq_lens, live_seq_lens) attn_seq_lens = torch.where(valid_verify_rows, attn_seq_lens, live_seq_lens)
fm.seq_lens_cpu_int = (seq_lens_cpu[:bs] + int(tokens_per_bs)).int() fm.seq_lens_cpu_int = (seq_lens_cpu[:bs] + int(tokens_per_req)).int()
fm.seq_lens_cpu_int = torch.where( fm.seq_lens_cpu_int = torch.where(
seq_lens_cpu[:bs] > 0, seq_lens_cpu[:bs] > 0,
fm.seq_lens_cpu_int, fm.seq_lens_cpu_int,
@@ -1166,8 +1166,8 @@ class DeepseekV4AscendAttnBackend(
_compress_seq_lens = live_seq_lens _compress_seq_lens = live_seq_lens
_compress_seq_lens_max = int(seq_lens_cpu[:bs].max()) if bs > 0 else 0 _compress_seq_lens_max = int(seq_lens_cpu[:bs].max()) if bs > 0 else 0
if _verify_compress: if _verify_compress:
_compress_seq_lens = live_seq_lens + int(tokens_per_bs) _compress_seq_lens = live_seq_lens + int(tokens_per_req)
_compress_seq_lens_max += int(tokens_per_bs) _compress_seq_lens_max += int(tokens_per_req)
result = self._compute_compress_locs( result = self._compute_compress_locs(
pool=pool, pool=pool,
@@ -1218,7 +1218,7 @@ class DeepseekV4AscendAttnBackend(
_copy_1d(getattr(fm, key), result[key]) _copy_1d(getattr(fm, key), result[key])
if _verify_compress: if _verify_compress:
verify_seq_lens_cpu = seq_lens_cpu[:bs] + int(tokens_per_bs) verify_seq_lens_cpu = seq_lens_cpu[:bs] + int(tokens_per_req)
verify_seq_lens_cpu = torch.where( verify_seq_lens_cpu = torch.where(
seq_lens_cpu[:bs] > 0, seq_lens_cpu[:bs] > 0,
verify_seq_lens_cpu, verify_seq_lens_cpu,
@@ -1229,19 +1229,19 @@ class DeepseekV4AscendAttnBackend(
fm.positions_cmp_padding_c4, fm.positions_cmp_padding_c4,
4, 4,
verify_seq_lens_cpu, verify_seq_lens_cpu,
n_draft=tokens_per_bs, n_draft=tokens_per_req,
) )
self._fill_verify_positions_cmp_padding_one( self._fill_verify_positions_cmp_padding_one(
forward_batch.positions, forward_batch.positions,
fm.positions_cmp_padding_c128, fm.positions_cmp_padding_c128,
128, 128,
verify_seq_lens_cpu, verify_seq_lens_cpu,
n_draft=tokens_per_bs, n_draft=tokens_per_req,
) )
fm.start_pos.copy_(live_seq_lens.to(torch.int32)) fm.start_pos.copy_(live_seq_lens.to(torch.int32))
valid = live_seq_lens[:bs] > 0 valid = live_seq_lens[:bs] > 0
fm.seqused.copy_( fm.seqused.copy_(
(valid.to(torch.int32) * int(tokens_per_bs)).to(device=device) (valid.to(torch.int32) * int(tokens_per_req)).to(device=device)
) )
_bundle = getattr(forward_batch, "out_cache_loc_dsv4", None) _bundle = getattr(forward_batch, "out_cache_loc_dsv4", None)
if _bundle is not None: if _bundle is not None:
@@ -1302,7 +1302,7 @@ class DeepseekV4AscendAttnBackend(
actual_seq_lengths_q_pa=fm.actual_seq_lengths_q_pa, actual_seq_lengths_q_pa=fm.actual_seq_lengths_q_pa,
actual_seq_lengths_kv=fm.actual_seq_lengths_kv, actual_seq_lengths_kv=fm.actual_seq_lengths_kv,
block_tables=fm.block_tables, block_tables=fm.block_tables,
max_seqlen_q=tokens_per_bs, max_seqlen_q=tokens_per_req,
is_nextn=False, is_nextn=False,
) )
for key in ( for key in (
@@ -240,7 +240,7 @@ class NPUGraphRunner(DecodeCudaGraphRunner):
or is_deepseek_v4(self.model_runner.model_config.hf_config) or is_deepseek_v4(self.model_runner.model_config.hf_config)
): ):
if forward_batch.forward_mode.is_target_verify(): if forward_batch.forward_mode.is_target_verify():
seq_lens_cpu = forward_batch.seq_lens.cpu() + self.num_tokens_per_bs seq_lens_cpu = forward_batch.seq_lens.cpu() + self.num_tokens_per_req
seq_lens = seq_lens_cpu.tolist() + [0] * (self.bs - self.raw_bs) seq_lens = seq_lens_cpu.tolist() + [0] * (self.bs - self.raw_bs)
else: else:
seq_lens = forward_batch.seq_lens.cpu().tolist() + [0] * ( seq_lens = forward_batch.seq_lens.cpu().tolist() + [0] * (
+3 -3
View File
@@ -87,9 +87,9 @@ class CanaryLaunchCapacities:
f"kv-canary: max_prefill_tokens must be positive, got {max_prefill_tokens}" f"kv-canary: max_prefill_tokens must be positive, got {max_prefill_tokens}"
) )
num_tokens_per_bs = 1 num_tokens_per_req = 1
if spec_num_draft_tokens: if spec_num_draft_tokens:
num_tokens_per_bs = max(num_tokens_per_bs, spec_num_draft_tokens) num_tokens_per_req = max(num_tokens_per_req, spec_num_draft_tokens)
max_bs = max(cuda_graph_max_bs, req_to_token_pool_size) max_bs = max(cuda_graph_max_bs, req_to_token_pool_size)
@@ -102,7 +102,7 @@ class CanaryLaunchCapacities:
max_extend_tokens_per_forward = min(max_prefill_tokens, chunked_limit) max_extend_tokens_per_forward = min(max_prefill_tokens, chunked_limit)
write_entry_capacity = max( write_entry_capacity = max(
max_bs * num_tokens_per_bs, max_extend_tokens_per_forward max_bs * num_tokens_per_req, max_extend_tokens_per_forward
) )
# Radix prefix sharing lets sum_r prefix_lens[r] exceed pool_slot_count; observed up to ~2x # Radix prefix sharing lets sum_r prefix_lens[r] exceed pool_slot_count; observed up to ~2x
@@ -1794,9 +1794,9 @@ class AiterAttnBackend(AttentionBackend):
# EAGLE V2: Fixed num_draft_tokens per batch # EAGLE V2: Fixed num_draft_tokens per batch
self._ensure_spec_v2_topk_supported() self._ensure_spec_v2_topk_supported()
seq_lens = seq_lens[:bs] seq_lens = seq_lens[:bs]
num_tokens_per_bs = self._resolve_v2_num_draft_tokens() num_tokens_per_req = self._resolve_v2_num_draft_tokens()
extend_lens = torch.full( extend_lens = torch.full(
(bs,), num_tokens_per_bs, dtype=torch.int32, device=seq_lens.device (bs,), num_tokens_per_req, dtype=torch.int32, device=seq_lens.device
) )
qo_indptr = self.qo_indptr[: bs + 1] qo_indptr = self.qo_indptr[: bs + 1]
@@ -1815,7 +1815,7 @@ class AiterAttnBackend(AttentionBackend):
) )
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs] kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
max_q_len = num_tokens_per_bs max_q_len = num_tokens_per_req
if self.use_mla and _use_mla_ps_kernel: if self.use_mla and _use_mla_ps_kernel:
num_kv_splits = self.max_split_per_batch num_kv_splits = self.max_split_per_batch
@@ -1043,14 +1043,14 @@ class DeepseekV4AttnBackend(
req_pool_indices: torch.Tensor, req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor, seq_lens: torch.Tensor,
seq_lens_cpu: List[int], seq_lens_cpu: List[int],
num_tokens_per_bs: int, num_tokens_per_req: int,
out_cache_loc: Optional[torch.Tensor] = None, out_cache_loc: Optional[torch.Tensor] = None,
use_prefill_cuda_graph: bool = False, use_prefill_cuda_graph: bool = False,
) -> DSV4Metadata: ) -> DSV4Metadata:
batch_size = len(seq_lens) batch_size = len(seq_lens)
extend_seq_lens_cpu = [num_tokens_per_bs] * batch_size extend_seq_lens_cpu = [num_tokens_per_req] * batch_size
extend_seq_lens = self._move_to_device(extend_seq_lens_cpu) extend_seq_lens = self._move_to_device(extend_seq_lens_cpu)
num_tokens = num_tokens_per_bs * batch_size num_tokens = num_tokens_per_req * batch_size
if out_cache_loc is None: if out_cache_loc is None:
out_cache_loc = seq_lens.new_zeros(num_tokens) out_cache_loc = seq_lens.new_zeros(num_tokens)
return self.init_forward_metadata_prefill( return self.init_forward_metadata_prefill(
@@ -1287,13 +1287,13 @@ class DeepseekV4AttnBackend(
req_pool_indices, req_pool_indices,
seq_lens, seq_lens,
) )
num_tokens_per_bs = self.draft_extend_num_tokens_per_bs num_tokens_per_req = self.draft_extend_num_tokens_per_req
if out_cache_loc is not None: if out_cache_loc is not None:
# Pad the real write locations to the captured token count so # Pad the real write locations to the captured token count so
# raw_out_loc reflects the actual replay out_cache_loc. # raw_out_loc reflects the actual replay out_cache_loc.
out_cache_loc = torch.nn.functional.pad( out_cache_loc = torch.nn.functional.pad(
out_cache_loc, out_cache_loc,
pad=(0, num_tokens_per_bs * bs - len(out_cache_loc)), pad=(0, num_tokens_per_req * bs - len(out_cache_loc)),
mode="constant", mode="constant",
value=0, value=0,
) )
@@ -1305,7 +1305,7 @@ class DeepseekV4AttnBackend(
req_pool_indices=req_pool_indices, req_pool_indices=req_pool_indices,
seq_lens=seq_lens, seq_lens=seq_lens,
seq_lens_cpu=draft_extend_seq_lens_cpu, seq_lens_cpu=draft_extend_seq_lens_cpu,
num_tokens_per_bs=num_tokens_per_bs, num_tokens_per_req=num_tokens_per_req,
out_cache_loc=out_cache_loc, out_cache_loc=out_cache_loc,
use_prefill_cuda_graph=True, use_prefill_cuda_graph=True,
) )
@@ -1477,7 +1477,7 @@ class DeepseekV4AttnBackend(
], ],
], ],
] = {bucket: {} for bucket in _GraphBucket} ] = {bucket: {} for bucket in _GraphBucket}
self.draft_extend_num_tokens_per_bs = ( self.draft_extend_num_tokens_per_req = (
max_num_tokens // max_bs if max_bs > 0 else 1 max_num_tokens // max_bs if max_bs > 0 else 1
) )
@@ -727,14 +727,14 @@ class DeepseekV4HipRadixBackend(
req_pool_indices: torch.Tensor, req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor, seq_lens: torch.Tensor,
seq_lens_cpu: List[int], seq_lens_cpu: List[int],
num_tokens_per_bs: int, num_tokens_per_req: int,
out_cache_loc: Optional[torch.Tensor] = None, out_cache_loc: Optional[torch.Tensor] = None,
use_prefill_cuda_graph: bool = False, use_prefill_cuda_graph: bool = False,
) -> DSV4Metadata: ) -> DSV4Metadata:
batch_size = len(seq_lens) batch_size = len(seq_lens)
extend_seq_lens_cpu = [num_tokens_per_bs] * batch_size extend_seq_lens_cpu = [num_tokens_per_req] * batch_size
extend_seq_lens = self._move_to_device(extend_seq_lens_cpu) extend_seq_lens = self._move_to_device(extend_seq_lens_cpu)
num_tokens = num_tokens_per_bs * batch_size num_tokens = num_tokens_per_req * batch_size
if out_cache_loc is None: if out_cache_loc is None:
out_cache_loc = seq_lens.new_zeros(num_tokens) out_cache_loc = seq_lens.new_zeros(num_tokens)
return self.init_forward_metadata_prefill( return self.init_forward_metadata_prefill(
@@ -889,13 +889,13 @@ class DeepseekV4HipRadixBackend(
seq_lens_cpu=seq_lens_cpu.tolist(), seq_lens_cpu=seq_lens_cpu.tolist(),
) )
elif bucket == _GraphBucket.DRAFT_EXTEND: elif bucket == _GraphBucket.DRAFT_EXTEND:
num_tokens_per_bs = self.draft_extend_num_tokens_per_bs num_tokens_per_req = self.draft_extend_num_tokens_per_req
if out_cache_loc is not None: if out_cache_loc is not None:
# Pad the real write locations to the captured token count so # Pad the real write locations to the captured token count so
# raw_out_loc reflects the actual replay out_cache_loc. # raw_out_loc reflects the actual replay out_cache_loc.
out_cache_loc = torch.nn.functional.pad( out_cache_loc = torch.nn.functional.pad(
out_cache_loc, out_cache_loc,
pad=(0, num_tokens_per_bs * bs - len(out_cache_loc)), pad=(0, num_tokens_per_req * bs - len(out_cache_loc)),
mode="constant", mode="constant",
value=0, value=0,
) )
@@ -904,7 +904,7 @@ class DeepseekV4HipRadixBackend(
req_pool_indices=req_pool_indices, req_pool_indices=req_pool_indices,
seq_lens=seq_lens, seq_lens=seq_lens,
seq_lens_cpu=seq_lens_cpu.tolist(), seq_lens_cpu=seq_lens_cpu.tolist(),
num_tokens_per_bs=num_tokens_per_bs, num_tokens_per_req=num_tokens_per_req,
out_cache_loc=out_cache_loc, out_cache_loc=out_cache_loc,
use_prefill_cuda_graph=True, use_prefill_cuda_graph=True,
) )
@@ -1012,7 +1012,7 @@ class DeepseekV4HipRadixBackend(
], ],
], ],
] = {bucket: {} for bucket in _GraphBucket} ] = {bucket: {} for bucket in _GraphBucket}
self.draft_extend_num_tokens_per_bs = ( self.draft_extend_num_tokens_per_req = (
max_num_tokens // max_bs if max_bs > 0 else 1 max_num_tokens // max_bs if max_bs > 0 else 1
) )
@@ -282,7 +282,7 @@ class FlashAttentionBackend(AttentionBackend):
): ):
self.speculative_num_draft_tokens = SpeculativeAlgorithm.from_string( self.speculative_num_draft_tokens = SpeculativeAlgorithm.from_string(
model_runner.server_args.speculative_algorithm model_runner.server_args.speculative_algorithm
).get_num_tokens_per_bs_for_target_verify( ).get_num_tokens_per_req_for_target_verify(
int(self.speculative_num_draft_tokens), is_draft_worker=True int(self.speculative_num_draft_tokens), is_draft_worker=True
) )
self.speculative_step_id = speculative_step_id self.speculative_step_id = speculative_step_id
@@ -513,7 +513,7 @@ class FlashAttentionBackend(AttentionBackend):
# CUDA graph bakes max_seq_len_q as a constant. replay() sets it to # CUDA graph bakes max_seq_len_q as a constant. replay() sets it to
# max(num_accept_tokens_cpu) which is None/empty at capture time, # max(num_accept_tokens_cpu) which is None/empty at capture time,
# falling back to 1. Restore the correct upper bound so the kernel # falling back to 1. Restore the correct upper bound so the kernel
# sees num_tokens_per_bs (not 1) for all replays of this graph. # sees num_tokens_per_req (not 1) for all replays of this graph.
self.forward_metadata.max_seq_len_q = num_tokens // bs self.forward_metadata.max_seq_len_q = num_tokens // bs
else: else:
self._apply_cuda_graph_metadata( self._apply_cuda_graph_metadata(
@@ -2353,11 +2353,11 @@ class FlashAttentionBackend(AttentionBackend):
metadata.swa_spec_metadata = metadata_swa metadata.swa_spec_metadata = metadata_swa
elif forward_mode.is_draft_extend_v2(): elif forward_mode.is_draft_extend_v2():
num_tokens_per_bs = num_tokens // bs num_tokens_per_req = num_tokens // bs
metadata.cache_seqlens_int32 = self.draft_extend_metadata["cache_seqlens"][ metadata.cache_seqlens_int32 = self.draft_extend_metadata["cache_seqlens"][
:bs :bs
] ]
metadata.max_seq_len_q = num_tokens_per_bs metadata.max_seq_len_q = num_tokens_per_req
metadata.cu_seqlens_q = self.draft_extend_metadata["cu_seqlens_q"][: bs + 1] metadata.cu_seqlens_q = self.draft_extend_metadata["cu_seqlens_q"][: bs + 1]
metadata.cu_seqlens_k = self.draft_extend_metadata["cu_seqlens_k"][ metadata.cu_seqlens_k = self.draft_extend_metadata["cu_seqlens_k"][
: (bs + 1) : (bs + 1)
@@ -507,7 +507,7 @@ class TritonAttnBackend(AttentionBackend):
seq_lens = seq_lens[:bs] seq_lens = seq_lens[:bs]
# V2 draft-extend fills num_draft_tokens per req; num_steps+1 only equals # V2 draft-extend fills num_draft_tokens per req; num_steps+1 only equals
# that when topk == 1. # that when topk == 1.
num_tokens_per_bs = ( num_tokens_per_req = (
self.num_draft_tokens self.num_draft_tokens
if forward_mode.is_draft_extend_v2() if forward_mode.is_draft_extend_v2()
else self.speculative_num_steps + 1 else self.speculative_num_steps + 1
@@ -515,8 +515,8 @@ class TritonAttnBackend(AttentionBackend):
qo_indptr = self.qo_indptr[: bs + 1] qo_indptr = self.qo_indptr[: bs + 1]
qo_indptr[: bs + 1] = torch.arange( qo_indptr[: bs + 1] = torch.arange(
0, 0,
bs * num_tokens_per_bs + 1, bs * num_tokens_per_req + 1,
step=num_tokens_per_bs, step=num_tokens_per_req,
dtype=torch.int32, dtype=torch.int32,
device=self.device, device=self.device,
) )
@@ -534,7 +534,7 @@ class TritonAttnBackend(AttentionBackend):
kv_indptr = self._fill_kv_indptr_and_indices( kv_indptr = self._fill_kv_indptr_and_indices(
bs, kv_lens, req_pool_indices, self.cuda_graph_kv_indices bs, kv_lens, req_pool_indices, self.cuda_graph_kv_indices
) )
return qo_indptr, kv_indptr, num_tokens_per_bs return qo_indptr, kv_indptr, num_tokens_per_req
def init_forward_metadata_out_graph( def init_forward_metadata_out_graph(
self, self,
@@ -1084,7 +1084,7 @@ class TritonAttnBackend(AttentionBackend):
return ForwardMetadata( return ForwardMetadata(
attn_logits=None, attn_logits=None,
attn_lse=None, attn_lse=None,
# Must match the per-req query count (num_tokens_per_bs) used to # Must match the per-req query count (num_tokens_per_req) used to
# build qo_indptr above, else the extend kernel grid is too small # build qo_indptr above, else the extend kernel grid is too small
# for topk > 1 (num_draft_tokens > num_steps+1) and drops query # for topk > 1 (num_draft_tokens > num_steps+1) and drops query
# blocks. # blocks.
@@ -487,13 +487,13 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
) )
self.target_verify_metadata[bs] = metadata self.target_verify_metadata[bs] = metadata
elif forward_mode.is_draft_extend_v2(): elif forward_mode.is_draft_extend_v2():
num_tokens_per_bs = num_tokens // bs num_tokens_per_req = num_tokens // bs
metadata.cache_seqlens_int32 = self.draft_extend_metadata["cache_seqlens"][ metadata.cache_seqlens_int32 = self.draft_extend_metadata["cache_seqlens"][
:bs :bs
] ]
metadata.cu_seqlens_q = self.draft_extend_metadata["cu_seqlens_q"][: bs + 1] metadata.cu_seqlens_q = self.draft_extend_metadata["cu_seqlens_q"][: bs + 1]
metadata.cu_seqlens_k = self.draft_extend_metadata["cu_seqlens_k"][: bs + 1] metadata.cu_seqlens_k = self.draft_extend_metadata["cu_seqlens_k"][: bs + 1]
metadata.max_seq_len_q = num_tokens_per_bs metadata.max_seq_len_q = num_tokens_per_req
metadata.page_table = self.draft_extend_metadata["page_table"][:bs, :] metadata.page_table = self.draft_extend_metadata["page_table"][:bs, :]
self._bind_swa_page_table( self._bind_swa_page_table(
metadata, metadata,
@@ -567,9 +567,9 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
# Static per-request query width, fixed by the captured graph shape. # Static per-request query width, fixed by the captured graph shape.
# Do not inspect replay-time tensors here; this body is recorded into # Do not inspect replay-time tensors here; this body is recorded into
# the CUDA graph. # the CUDA graph.
num_tokens_per_bs = metadata.max_seq_len_q num_tokens_per_req = metadata.max_seq_len_q
cu_seqlens_q = metadata.cu_seqlens_q cu_seqlens_q = metadata.cu_seqlens_q
q_stride = num_tokens_per_bs q_stride = num_tokens_per_req
q_mode = Q_MODE_STRIDED q_mode = Q_MODE_STRIDED
else: else:
raise ValueError( raise ValueError(
@@ -283,13 +283,13 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
self.decode_cuda_graph_kv_indices = torch.full( self.decode_cuda_graph_kv_indices = torch.full(
(max_bs, max_blocks_per_seq), -1, dtype=torch.int32, device=self.device (max_bs, max_blocks_per_seq), -1, dtype=torch.int32, device=self.device
) )
num_tokens_per_bs = max_num_tokens // max_bs num_tokens_per_req = max_num_tokens // max_bs
if is_float4_e2m1fn_x2(self.data_type): if is_float4_e2m1fn_x2(self.data_type):
# Buffer for padded query: (max_bs, max_draft_tokens, num_q_heads, v_head_dim) # Buffer for padded query: (max_bs, max_draft_tokens, num_q_heads, v_head_dim)
self.store_dtype = torch.uint8 self.store_dtype = torch.uint8
self.padded_q_buffer = torch.zeros( self.padded_q_buffer = torch.zeros(
(max_bs, num_tokens_per_bs // 2, self.num_q_heads, self.kv_cache_dim), (max_bs, num_tokens_per_req // 2, self.num_q_heads, self.kv_cache_dim),
dtype=self.store_dtype, dtype=self.store_dtype,
device=self.device, device=self.device,
) )
@@ -303,7 +303,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
else: else:
# Buffer for padded query: (max_bs, max_draft_tokens, num_q_heads, v_head_dim) # Buffer for padded query: (max_bs, max_draft_tokens, num_q_heads, v_head_dim)
self.padded_q_buffer = torch.zeros( self.padded_q_buffer = torch.zeros(
(max_bs, num_tokens_per_bs, self.num_q_heads, self.kv_cache_dim), (max_bs, num_tokens_per_req, self.num_q_heads, self.kv_cache_dim),
dtype=self.data_type, dtype=self.data_type,
device=self.device, device=self.device,
) )
@@ -343,18 +343,18 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
if forward_mode.is_target_verify(): if forward_mode.is_target_verify():
metadata.seq_lens_k = torch.zeros((bs,), dtype=torch.int32, device=device) metadata.seq_lens_k = torch.zeros((bs,), dtype=torch.int32, device=device)
elif forward_mode.is_draft_extend_v2(): elif forward_mode.is_draft_extend_v2():
num_tokens_per_bs = self.num_draft_tokens num_tokens_per_req = self.num_draft_tokens
metadata.max_seq_len_q = num_tokens_per_bs metadata.max_seq_len_q = num_tokens_per_req
metadata.sum_seq_lens_q = num_tokens_per_bs * bs metadata.sum_seq_lens_q = num_tokens_per_req * bs
metadata.cu_seqlens_q = torch.arange( metadata.cu_seqlens_q = torch.arange(
0, 0,
bs * num_tokens_per_bs + 1, bs * num_tokens_per_req + 1,
num_tokens_per_bs, num_tokens_per_req,
dtype=torch.int32, dtype=torch.int32,
device=device, device=device,
) )
metadata.seq_lens_q = torch.full( metadata.seq_lens_q = torch.full(
(bs,), num_tokens_per_bs, dtype=torch.int32, device=device (bs,), num_tokens_per_req, dtype=torch.int32, device=device
) )
metadata.seq_lens_k = torch.zeros((bs,), dtype=torch.int32, device=device) metadata.seq_lens_k = torch.zeros((bs,), dtype=torch.int32, device=device)
@@ -385,9 +385,9 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
seq_lens = seq_lens[:bs] + self.num_draft_tokens seq_lens = seq_lens[:bs] + self.num_draft_tokens
metadata.seq_lens_k.copy_(seq_lens) metadata.seq_lens_k.copy_(seq_lens)
elif forward_mode.is_draft_extend_v2(): elif forward_mode.is_draft_extend_v2():
num_tokens_per_bs = self.num_draft_tokens num_tokens_per_req = self.num_draft_tokens
metadata.max_seq_len_q = num_tokens_per_bs metadata.max_seq_len_q = num_tokens_per_req
metadata.sum_seq_lens_q = num_tokens_per_bs * bs metadata.sum_seq_lens_q = num_tokens_per_req * bs
seq_lens = seq_lens[:bs] seq_lens = seq_lens[:bs]
metadata.seq_lens_k.copy_(seq_lens) metadata.seq_lens_k.copy_(seq_lens)
@@ -204,7 +204,7 @@ class AscendLoRABackend(BaseLoRABackend):
def init_cuda_graph_batch_info( def init_cuda_graph_batch_info(
self, self,
max_bs_in_cuda_graph: int, max_bs_in_cuda_graph: int,
num_tokens_per_bs: int, num_tokens_per_req: int,
): ):
with torch.device("npu"): with torch.device("npu"):
self.npu_graph_batch_info = LoRABatchInfo( self.npu_graph_batch_info = LoRABatchInfo(
@@ -212,10 +212,10 @@ class AscendLoRABackend(BaseLoRABackend):
use_cuda_graph=True, use_cuda_graph=True,
num_segments=None, num_segments=None,
seg_lens=torch.full( seg_lens=torch.full(
(max_bs_in_cuda_graph,), num_tokens_per_bs, dtype=torch.int32 (max_bs_in_cuda_graph,), num_tokens_per_req, dtype=torch.int32
), ),
seg_indptr=torch.empty(max_bs_in_cuda_graph + 1, dtype=torch.int32), seg_indptr=torch.empty(max_bs_in_cuda_graph + 1, dtype=torch.int32),
max_len=num_tokens_per_bs, max_len=num_tokens_per_req,
weight_indices=torch.zeros(max_bs_in_cuda_graph, dtype=torch.int32), weight_indices=torch.zeros(max_bs_in_cuda_graph, dtype=torch.int32),
lora_ranks=torch.zeros(self.max_loras_per_batch, dtype=torch.int32), lora_ranks=torch.zeros(self.max_loras_per_batch, dtype=torch.int32),
scalings=torch.zeros(self.max_loras_per_batch, dtype=torch.float), scalings=torch.zeros(self.max_loras_per_batch, dtype=torch.float),
@@ -149,7 +149,7 @@ class BaseLoRABackend(LoRABackendLmHeadMixing):
def init_cuda_graph_batch_info( def init_cuda_graph_batch_info(
self, self,
max_bs_in_cuda_graph: int, max_bs_in_cuda_graph: int,
num_tokens_per_bs: int, num_tokens_per_req: int,
): ):
"""Phase 2 of LoRA CUDA graph init: dense LoRA batch metadata. """Phase 2 of LoRA CUDA graph init: dense LoRA batch metadata.
@@ -157,7 +157,7 @@ class BaseLoRABackend(LoRABackendLmHeadMixing):
Args: Args:
max_bs_in_cuda_graph: maximum batch size for CUDA Graph mode max_bs_in_cuda_graph: maximum batch size for CUDA Graph mode
num_tokens_per_bs: number of tokens per sequence (1 for decoding, >1 for target_verify) num_tokens_per_req: number of tokens per sequence (1 for decoding, >1 for target_verify)
""" """
pass pass
@@ -218,12 +218,12 @@ class ChunkedSgmvLoRABackend(BaseLoRABackend):
def init_cuda_graph_batch_info( def init_cuda_graph_batch_info(
self, self,
max_bs_in_cuda_graph: int, max_bs_in_cuda_graph: int,
num_tokens_per_bs: int, num_tokens_per_req: int,
): ):
max_num_segments = ( max_num_segments = (
(num_tokens_per_bs + MIN_CHUNK_SIZE - 1) // MIN_CHUNK_SIZE (num_tokens_per_req + MIN_CHUNK_SIZE - 1) // MIN_CHUNK_SIZE
) * max_bs_in_cuda_graph ) * max_bs_in_cuda_graph
max_num_tokens = max_bs_in_cuda_graph * num_tokens_per_bs max_num_tokens = max_bs_in_cuda_graph * num_tokens_per_req
with torch.device("cuda"): with torch.device("cuda"):
self.cuda_graph_batch_info = LoRABatchInfo( self.cuda_graph_batch_info = LoRABatchInfo(
bs=max_bs_in_cuda_graph, bs=max_bs_in_cuda_graph,
@@ -164,7 +164,7 @@ class TorchNativeLoRABackend(BaseLoRABackend):
def init_cuda_graph_batch_info( def init_cuda_graph_batch_info(
self, self,
max_bs_in_cuda_graph: int, max_bs_in_cuda_graph: int,
num_tokens_per_bs: int, num_tokens_per_req: int,
): ):
with torch.device("cuda"): with torch.device("cuda"):
self.cuda_graph_batch_info = TorchNativeLoRABatchInfo( self.cuda_graph_batch_info = TorchNativeLoRABatchInfo(
@@ -172,14 +172,14 @@ class TorchNativeLoRABackend(BaseLoRABackend):
bs=max_bs_in_cuda_graph, bs=max_bs_in_cuda_graph,
num_segments=self.max_loras_per_batch, num_segments=self.max_loras_per_batch,
seg_lens=torch.full( seg_lens=torch.full(
(max_bs_in_cuda_graph,), num_tokens_per_bs, dtype=torch.int32 (max_bs_in_cuda_graph,), num_tokens_per_req, dtype=torch.int32
), ),
seg_indptr=torch.zeros(max_bs_in_cuda_graph + 1, dtype=torch.int32), seg_indptr=torch.zeros(max_bs_in_cuda_graph + 1, dtype=torch.int32),
weight_indices=torch.zeros(max_bs_in_cuda_graph, dtype=torch.int32), weight_indices=torch.zeros(max_bs_in_cuda_graph, dtype=torch.int32),
lora_ranks=torch.zeros(self.max_loras_per_batch, dtype=torch.int32), lora_ranks=torch.zeros(self.max_loras_per_batch, dtype=torch.int32),
scalings=torch.zeros(self.max_loras_per_batch, dtype=torch.float), scalings=torch.zeros(self.max_loras_per_batch, dtype=torch.float),
permutation=None, permutation=None,
max_len=num_tokens_per_bs, max_len=num_tokens_per_req,
) )
# Initialize seg_indptr for CUDA graph as they remain constant # Initialize seg_indptr for CUDA graph as they remain constant
@@ -140,9 +140,9 @@ class TritonLoRABackend(BaseLoRABackend):
def init_cuda_graph_batch_info( def init_cuda_graph_batch_info(
self, self,
max_bs_in_cuda_graph: int, max_bs_in_cuda_graph: int,
num_tokens_per_bs: int, num_tokens_per_req: int,
): ):
max_tokens = max_bs_in_cuda_graph * num_tokens_per_bs max_tokens = max_bs_in_cuda_graph * num_tokens_per_req
mlpb = self.max_loras_per_batch mlpb = self.max_loras_per_batch
with torch.device("cuda"): with torch.device("cuda"):
self.cuda_graph_batch_info = LoRABatchInfo( self.cuda_graph_batch_info = LoRABatchInfo(
@@ -150,10 +150,10 @@ class TritonLoRABackend(BaseLoRABackend):
use_cuda_graph=True, use_cuda_graph=True,
num_segments=None, num_segments=None,
seg_lens=torch.full( seg_lens=torch.full(
(max_bs_in_cuda_graph,), num_tokens_per_bs, dtype=torch.int32 (max_bs_in_cuda_graph,), num_tokens_per_req, dtype=torch.int32
), ),
seg_indptr=torch.zeros(max_bs_in_cuda_graph + 1, dtype=torch.int32), seg_indptr=torch.zeros(max_bs_in_cuda_graph + 1, dtype=torch.int32),
max_len=num_tokens_per_bs, max_len=num_tokens_per_req,
weight_indices=torch.zeros(max_bs_in_cuda_graph, dtype=torch.int32), weight_indices=torch.zeros(max_bs_in_cuda_graph, dtype=torch.int32),
lora_ranks=torch.zeros(mlpb, dtype=torch.int32), lora_ranks=torch.zeros(mlpb, dtype=torch.int32),
scalings=torch.zeros(mlpb, dtype=torch.float), scalings=torch.zeros(mlpb, dtype=torch.float),
+2 -2
View File
@@ -112,7 +112,7 @@ class LoRAManager:
) )
def init_cuda_graph_batch_info( def init_cuda_graph_batch_info(
self, max_bs_in_cuda_graph: int, num_tokens_per_bs: int self, max_bs_in_cuda_graph: int, num_tokens_per_req: int
): ):
"""Phase 2 of LoRA CUDA graph init: dense LoRA batch metadata. """Phase 2 of LoRA CUDA graph init: dense LoRA batch metadata.
@@ -122,7 +122,7 @@ class LoRAManager:
self.max_bs_in_cuda_graph = max_bs_in_cuda_graph self.max_bs_in_cuda_graph = max_bs_in_cuda_graph
self.lora_backend.init_cuda_graph_batch_info( self.lora_backend.init_cuda_graph_batch_info(
max_bs_in_cuda_graph=max_bs_in_cuda_graph, max_bs_in_cuda_graph=max_bs_in_cuda_graph,
num_tokens_per_bs=num_tokens_per_bs, num_tokens_per_req=num_tokens_per_req,
) )
# ===== TO BE REFACTORED ==== # ===== TO BE REFACTORED ====
@@ -581,7 +581,7 @@ class CPUGraphRunner:
self.capture_forward_mode = ForwardMode.DECODE self.capture_forward_mode = ForwardMode.DECODE
self.capture_hidden_mode = CaptureHiddenMode.NULL self.capture_hidden_mode = CaptureHiddenMode.NULL
self.num_tokens_per_bs = 1 self.num_tokens_per_req = 1
# If returning hidden states is enabled, set initial capture hidden mode to full to avoid double-capture on startup # 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: if self.enable_return_hidden_states:
@@ -618,7 +618,7 @@ class CPUGraphRunner:
self.captured_forward_batches_cross = {} self.captured_forward_batches_cross = {}
# Attention backend # Attention backend
self.max_bs = max(self.capture_bs) self.max_bs = max(self.capture_bs)
self.max_num_token = self.max_bs * self.num_tokens_per_bs self.max_num_token = self.max_bs * self.num_tokens_per_req
self.model_runner.attn_backend.init_cpu_graph_state( self.model_runner.attn_backend.init_cpu_graph_state(
self.max_bs, self.max_num_token self.max_bs, self.max_num_token
) )
@@ -646,7 +646,7 @@ class CPUGraphRunner:
self.custom_mask = torch.ones( self.custom_mask = torch.ones(
( (
(self.seq_lens.sum().item() + self.max_num_token) (self.seq_lens.sum().item() + self.max_num_token)
* self.num_tokens_per_bs * self.num_tokens_per_req
), ),
dtype=torch.bool, dtype=torch.bool,
device=self.device, device=self.device,
@@ -725,7 +725,7 @@ class CPUGraphRunner:
with patch_model( with patch_model(
self.model_runner.model, self.model_runner.model,
bs in self.capture_bs, bs in self.capture_bs,
num_tokens=bs * self.num_tokens_per_bs, num_tokens=bs * self.num_tokens_per_req,
tp_group=self.model_runner.tp_group, tp_group=self.model_runner.tp_group,
) as forward: ) as forward:
graph, output_buffers = self.capture_one_batch_size( graph, output_buffers = self.capture_one_batch_size(
@@ -767,7 +767,7 @@ class CPUGraphRunner:
def capture_one_batch_size( def capture_one_batch_size(
self, bs: int, forward: Callable, skip_cross_attention: bool = False self, bs: int, forward: Callable, skip_cross_attention: bool = False
): ):
num_tokens = bs * self.num_tokens_per_bs num_tokens = bs * self.num_tokens_per_req
# Graph inputs # Graph inputs
input_ids = self.input_ids[:num_tokens] input_ids = self.input_ids[:num_tokens]
@@ -916,7 +916,7 @@ class CPUGraphRunner:
self.model_runner.attn_backend.init_forward_metadata(forward_batch) self.model_runner.attn_backend.init_forward_metadata(forward_batch)
return forward_batch return forward_batch
raw_num_token = raw_bs * self.num_tokens_per_bs raw_num_token = raw_bs * self.num_tokens_per_req
index = bisect.bisect_left(self.capture_bs, raw_bs) index = bisect.bisect_left(self.capture_bs, raw_bs)
bs = self.capture_bs[index] bs = self.capture_bs[index]
assert bs > raw_bs assert bs > raw_bs
@@ -100,7 +100,7 @@ class FillContext:
Carries both the bs-axis and tokens-axis raw/padded counts so a hook can Carries both the bs-axis and tokens-axis raw/padded counts so a hook can
derive values regardless of its own slot's axis — e.g. the padded token derive values regardless of its own slot's axis — e.g. the padded token
count (``padded_num_tokens`` == padded_bs * num_tokens_per_bs), which the count (``padded_num_tokens`` == padded_bs * num_tokens_per_req), which the
global-num-tokens fill and the local-num-token-non-padded transform need. global-num-tokens fill and the local-num-token-non-padded transform need.
""" """
@@ -842,14 +842,14 @@ class ModelRunner(ModelRunnerKVCacheMixin):
return None return None
return getattr(hf_config, "index_topk", None) return getattr(hf_config, "index_topk", None)
def decode_num_tokens_per_bs( def decode_num_tokens_per_req(
self, *, num_draft_tokens: Optional[int] = None self, *, num_draft_tokens: Optional[int] = None
) -> int: ) -> int:
"""Logits rows per decode batch slot.""" """Logits rows per decode batch slot."""
if self.spec_algorithm.is_speculative(): if self.spec_algorithm.is_speculative():
if num_draft_tokens is None: if num_draft_tokens is None:
num_draft_tokens = self.server_args.speculative_num_draft_tokens num_draft_tokens = self.server_args.speculative_num_draft_tokens
return self.spec_algorithm.get_num_tokens_per_bs_for_target_verify( return self.spec_algorithm.get_num_tokens_per_req_for_target_verify(
num_draft_tokens, self.is_draft_worker num_draft_tokens, self.is_draft_worker
) )
dllm_config = DllmConfig.from_server_args(self.server_args) dllm_config = DllmConfig.from_server_args(self.server_args)
@@ -857,9 +857,9 @@ class ModelRunner(ModelRunnerKVCacheMixin):
def max_decode_logits_rows(self) -> int: def max_decode_logits_rows(self) -> int:
"""Rows the shared logits buffer needs.""" """Rows the shared logits buffer needs."""
num_tokens_per_bs = self.decode_num_tokens_per_bs() num_tokens_per_req = self.decode_num_tokens_per_req()
capture_bs, _ = get_batch_sizes_to_capture(self, num_tokens_per_bs) capture_bs, _ = get_batch_sizes_to_capture(self, num_tokens_per_req)
return max(capture_bs) * num_tokens_per_bs return max(capture_bs) * num_tokens_per_req
def alloc_memory_pool(self, memory_pool_config: Optional[MemoryPoolConfig] = None): def alloc_memory_pool(self, memory_pool_config: Optional[MemoryPoolConfig] = None):
"""Allocate KV cache memory pools only (no backends or cuda graphs).""" """Allocate KV cache memory pools only (no backends or cuda graphs)."""
@@ -2234,7 +2234,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
Phase 2 (dense LoRA batch metadata) is handled later in Phase 2 (dense LoRA batch metadata) is handled later in
CudaGraphRunner.__init__() via lora_manager.init_cuda_graph_batch_info(), CudaGraphRunner.__init__() via lora_manager.init_cuda_graph_batch_info(),
because it needs capture-time parameters (max_bs, num_tokens_per_bs) because it needs capture-time parameters (max_bs, num_tokens_per_req)
that are only available at that stage. that are only available at that stage.
""" """
from sglang.srt.lora.layers import FusedMoEWithLoRA from sglang.srt.lora.layers import FusedMoEWithLoRA
@@ -2640,20 +2640,20 @@ class ModelRunner(ModelRunnerKVCacheMixin):
role = "draft" if self.is_draft_worker else "target" role = "draft" if self.is_draft_worker else "target"
if self.spec_algorithm.is_speculative(): if self.spec_algorithm.is_speculative():
capture_name = f"{role} verify" capture_name = f"{role} verify"
num_tokens_per_bs = ( num_tokens_per_req = (
self.spec_algorithm.get_num_tokens_per_bs_for_target_verify( self.spec_algorithm.get_num_tokens_per_req_for_target_verify(
self.server_args.speculative_num_draft_tokens, self.server_args.speculative_num_draft_tokens,
self.is_draft_worker, self.is_draft_worker,
) )
) )
else: else:
capture_name = f"{role} decode" capture_name = f"{role} decode"
num_tokens_per_bs = 1 num_tokens_per_req = 1
capture_bs, _ = get_batch_sizes_to_capture(self, num_tokens_per_bs) capture_bs, _ = get_batch_sizes_to_capture(self, num_tokens_per_req)
decode_backend = self.server_args.cuda_graph_config.decode.backend decode_backend = self.server_args.cuda_graph_config.decode.backend
logger.info( logger.info(
f"Capture {capture_name} {graph_backend[self.device]} begin. " f"Capture {capture_name} {graph_backend[self.device]} begin. "
f"backend={decode_backend}, num_tokens_per_bs={num_tokens_per_bs}, " f"backend={decode_backend}, num_tokens_per_req={num_tokens_per_req}, "
f"bs={capture_bs}, avail mem={before_mem:.2f} GB" f"bs={capture_bs}, avail mem={before_mem:.2f} GB"
) )
@@ -56,7 +56,7 @@ def freeze_gc(enable_cudagraph_gc: bool):
def get_batch_sizes_to_capture( def get_batch_sizes_to_capture(
model_runner: ModelRunner, num_tokens_per_bs: int = 1 model_runner: ModelRunner, num_tokens_per_req: int = 1
) -> Tuple[List[int], List[int]]: ) -> Tuple[List[int], List[int]]:
"""Build the (capture_bs, compile_bs) lists for the decode runner. """Build the (capture_bs, compile_bs) lists for the decode runner.
@@ -71,7 +71,7 @@ def get_batch_sizes_to_capture(
mul_base = 1 mul_base = 1
if server_args.enable_two_batch_overlap: if server_args.enable_two_batch_overlap:
mul_base *= 2 mul_base *= 2
num_tokens_per_bs = 1 num_tokens_per_req = 1
if require_gathered_buffer(server_args): if require_gathered_buffer(server_args):
mul_base *= get_parallel().attn_tp_size mul_base *= get_parallel().attn_tp_size
@@ -86,8 +86,8 @@ def get_batch_sizes_to_capture(
# is very small. We add more values here to make sure we capture the maximum bs. # is very small. We add more values here to make sure we capture the maximum bs.
capture_bs += [num_max_requests] capture_bs += [num_max_requests]
# Model input token count = bs * num_tokens_per_bs; must be a multiple of attn_tp_size. # 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_bs % mul_base == 0] capture_bs = [bs for bs in capture_bs if bs * num_tokens_per_req % mul_base == 0]
capture_bs = [bs for bs in capture_bs if bs <= num_max_requests] capture_bs = [bs for bs in capture_bs if bs <= num_max_requests]
capture_bs = list(sorted(set(capture_bs))) capture_bs = list(sorted(set(capture_bs)))
@@ -74,7 +74,7 @@ def _allocate_decode_buffers(
require_mlp_tp_gather: bool, require_mlp_tp_gather: bool,
seq_len_fill_value: int, seq_len_fill_value: int,
encoder_len_fill_value: int, encoder_len_fill_value: int,
num_tokens_per_bs: int, num_tokens_per_req: int,
cache_loc_dtype: torch.dtype, cache_loc_dtype: torch.dtype,
enable_mamba_track: bool, enable_mamba_track: bool,
ne_token_table: Optional[torch.Tensor] = None, ne_token_table: Optional[torch.Tensor] = None,
@@ -92,7 +92,7 @@ def _allocate_decode_buffers(
mrope_positions = torch.zeros((3, max_num_token), dtype=torch.int64) mrope_positions = torch.zeros((3, max_num_token), dtype=torch.int64)
num_token_non_padded = torch.zeros((1,), dtype=torch.int32) num_token_non_padded = torch.zeros((1,), dtype=torch.int32)
custom_mask = torch.ones( custom_mask = torch.ones(
(max_bs * seq_len_fill_value + max_num_token) * num_tokens_per_bs, (max_bs * seq_len_fill_value + max_num_token) * num_tokens_per_req,
dtype=torch.bool, dtype=torch.bool,
) )
next_token_logits_buffer = torch.zeros( next_token_logits_buffer = torch.zeros(
@@ -280,9 +280,9 @@ class BaseRunner(ABC):
run_flashinfer_autotune_forward(self.model_runner, forward_fn, skip_logits=True) run_flashinfer_autotune_forward(self.model_runner, forward_fn, skip_logits=True)
def _alloc_dummy_decode_buffers(self, max_bs: int, *, num_tokens_per_bs: int = 1): def _alloc_dummy_decode_buffers(self, max_bs: int, *, num_tokens_per_req: int = 1):
"""Allocate one static decode-buffer set for a dummy forward, sized to """Allocate one static decode-buffer set for a dummy forward, sized to
(max_bs, max_bs * num_tokens_per_bs). (max_bs, max_bs * num_tokens_per_req).
The PP-parallel DeepGEMM warmup sweeps batch sizes far larger than any The PP-parallel DeepGEMM warmup sweeps batch sizes far larger than any
runner's max_bs (up to ~n_sms*block_m), so no pre-allocated runner buffer runner's max_bs (up to ~n_sms*block_m), so no pre-allocated runner buffer
@@ -295,7 +295,7 @@ class BaseRunner(ABC):
return _allocate_decode_buffers( return _allocate_decode_buffers(
device=mr.device, device=mr.device,
max_bs=max_bs, max_bs=max_bs,
max_num_token=max_bs * num_tokens_per_bs, max_num_token=max_bs * num_tokens_per_req,
hidden_size=mr.model_config.hidden_size, hidden_size=mr.model_config.hidden_size,
vocab_size=mr.model_config.vocab_size, vocab_size=mr.model_config.vocab_size,
dtype=mr.model_config.dtype, dtype=mr.model_config.dtype,
@@ -309,7 +309,7 @@ class BaseRunner(ABC):
if mr.model_config.is_encoder_decoder if mr.model_config.is_encoder_decoder
else 0 else 0
), ),
num_tokens_per_bs=num_tokens_per_bs, num_tokens_per_req=num_tokens_per_req,
cache_loc_dtype=torch.int64, cache_loc_dtype=torch.int64,
enable_mamba_track=False, enable_mamba_track=False,
ne_token_table=mr.token_table if mr.use_ngram_embedding else None, ne_token_table=mr.token_table if mr.use_ngram_embedding else None,
@@ -350,14 +350,14 @@ class BaseRunner(ABC):
else: else:
capture_forward_mode = ForwardMode.EXTEND capture_forward_mode = ForwardMode.EXTEND
capture_hidden_mode = CaptureHiddenMode.NULL capture_hidden_mode = CaptureHiddenMode.NULL
num_tokens_per_bs = 1 num_tokens_per_req = 1
if mr.spec_algorithm.is_speculative(): if mr.spec_algorithm.is_speculative():
if mr.is_draft_worker: if mr.is_draft_worker:
if not mr.spec_algorithm.supports_target_verify_for_draft(): if not mr.spec_algorithm.supports_target_verify_for_draft():
raise RuntimeError("This should not happen") raise RuntimeError("This should not happen")
capture_forward_mode = ForwardMode.TARGET_VERIFY capture_forward_mode = ForwardMode.TARGET_VERIFY
num_tokens_per_bs = ( num_tokens_per_req = (
mr.spec_algorithm.get_num_tokens_per_bs_for_target_verify( mr.spec_algorithm.get_num_tokens_per_req_for_target_verify(
mr.server_args.speculative_num_draft_tokens, mr.is_draft_worker mr.server_args.speculative_num_draft_tokens, mr.is_draft_worker
) )
) )
@@ -365,7 +365,7 @@ class BaseRunner(ABC):
if mr.server_args.enable_return_hidden_states: if mr.server_args.enable_return_hidden_states:
capture_hidden_mode = CaptureHiddenMode.FULL capture_hidden_mode = CaptureHiddenMode.FULL
num_tokens = batch_size * num_tokens_per_bs num_tokens = batch_size * num_tokens_per_req
# Caller owns the shape: passes a static buffer >= the dummy shape; no # Caller owns the shape: passes a static buffer >= the dummy shape; no
# allocation, no re-padding (would overflow the reused buffers). # allocation, no re-padding (would overflow the reused buffers).
@@ -439,7 +439,7 @@ class BaseRunner(ABC):
(batch_size,), dtype=torch.int32, device=mr.device (batch_size,), dtype=torch.int32, device=mr.device
) )
extend_start_loc = torch.arange( extend_start_loc = torch.arange(
0, num_tokens, num_tokens_per_bs, dtype=torch.int32, device=mr.device 0, num_tokens, num_tokens_per_req, dtype=torch.int32, device=mr.device
) )
else: else:
extend_prefix_lens_cpu = None extend_prefix_lens_cpu = None
@@ -484,7 +484,7 @@ class BaseRunner(ABC):
mr.spec_algorithm, mr.spec_algorithm,
mr.server_args, mr.server_args,
buffers.custom_mask, buffers.custom_mask,
num_tokens_per_bs, num_tokens_per_req,
mr.is_draft_worker, mr.is_draft_worker,
) )
if spec_info is not None and ( if spec_info is not None and (
@@ -254,7 +254,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
# --- capture mode + tokens-per-bs ------------------------------ # --- capture mode + tokens-per-bs ------------------------------
self.capture_forward_mode = ForwardMode.DECODE self.capture_forward_mode = ForwardMode.DECODE
self.capture_hidden_mode = CaptureHiddenMode.NULL self.capture_hidden_mode = CaptureHiddenMode.NULL
self.num_tokens_per_bs = model_runner.decode_num_tokens_per_bs( self.num_tokens_per_req = model_runner.decode_num_tokens_per_req(
num_draft_tokens=self.speculative_num_draft_tokens num_draft_tokens=self.speculative_num_draft_tokens
) )
if model_runner.spec_algorithm.is_speculative(): if model_runner.spec_algorithm.is_speculative():
@@ -270,7 +270,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
# --- bucket sizes --------------------------------------------- # --- bucket sizes ---------------------------------------------
self.capture_bs, self.compile_bs = get_batch_sizes_to_capture( self.capture_bs, self.compile_bs = get_batch_sizes_to_capture(
model_runner, self.num_tokens_per_bs model_runner, self.num_tokens_per_req
) )
if KTRANSFORMERS_AVAILABLE: if KTRANSFORMERS_AVAILABLE:
KTMoEWrapper.set_capture_batch_sizes(self.capture_bs) KTMoEWrapper.set_capture_batch_sizes(self.capture_bs)
@@ -304,7 +304,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
# Attention backend # Attention backend
self.max_bs = max(self.capture_bs) self.max_bs = max(self.capture_bs)
self.max_num_token = self.max_bs * self.num_tokens_per_bs self.max_num_token = self.max_bs * self.num_tokens_per_req
self.attn_backend.init_cuda_graph_state(self.max_bs, self.max_num_token) self.attn_backend.init_cuda_graph_state(self.max_bs, self.max_num_token)
# Init PDMux if needed # Init PDMux if needed
@@ -331,7 +331,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
# lora_manager.init_cuda_graph_moe_buffers(). # lora_manager.init_cuda_graph_moe_buffers().
self.model_runner.lora_manager.init_cuda_graph_batch_info( self.model_runner.lora_manager.init_cuda_graph_batch_info(
max_bs_in_cuda_graph=self.max_bs, max_bs_in_cuda_graph=self.max_bs,
num_tokens_per_bs=self.num_tokens_per_bs, num_tokens_per_req=self.num_tokens_per_req,
) )
enable_mamba_track = ( enable_mamba_track = (
@@ -358,7 +358,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
require_mlp_tp_gather=self.require_mlp_tp_gather, require_mlp_tp_gather=self.require_mlp_tp_gather,
seq_len_fill_value=self.seq_len_fill_value, seq_len_fill_value=self.seq_len_fill_value,
encoder_len_fill_value=self.encoder_len_fill_value, encoder_len_fill_value=self.encoder_len_fill_value,
num_tokens_per_bs=self.num_tokens_per_bs, num_tokens_per_req=self.num_tokens_per_req,
cache_loc_dtype=self._cache_loc_dtype(), cache_loc_dtype=self._cache_loc_dtype(),
enable_mamba_track=enable_mamba_track, enable_mamba_track=enable_mamba_track,
ne_token_table=( ne_token_table=(
@@ -404,7 +404,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
) )
def _build_ragged_verify_token_buckets(self) -> list[int]: def _build_ragged_verify_token_buckets(self) -> list[int]:
buckets = sorted({bs * self.num_tokens_per_bs for bs in self.capture_bs}) buckets = sorted({bs * self.num_tokens_per_req for bs in self.capture_bs})
assert buckets and buckets[0] > 0, f"{buckets=}" assert buckets and buckets[0] > 0, f"{buckets=}"
return buckets return buckets
@@ -468,7 +468,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
def _ragged_capture_slots(self, num_tokens: int) -> int: def _ragged_capture_slots(self, num_tokens: int) -> int:
if envs.SGLANG_TEST_RAGGED_VERIFY_FORCE_UNIFORM_CAPTURE.get(): if envs.SGLANG_TEST_RAGGED_VERIFY_FORCE_UNIFORM_CAPTURE.get():
return num_tokens // self.num_tokens_per_bs return num_tokens // self.num_tokens_per_req
return min(num_tokens, self.max_bs) return min(num_tokens, self.max_bs)
def _capture_ragged_verify_layout(self, num_tokens: int): def _capture_ragged_verify_layout(self, num_tokens: int):
@@ -484,7 +484,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
verify_lens_cpu = build_capture_verify_lens( verify_lens_cpu = build_capture_verify_lens(
num_tokens=num_tokens, num_tokens=num_tokens,
num_slots=self._ragged_capture_slots(num_tokens), num_slots=self._ragged_capture_slots(num_tokens),
num_draft_tokens=self.num_tokens_per_bs, num_draft_tokens=self.num_tokens_per_req,
) )
return RaggedVerifyLayout.from_verify_lens( return RaggedVerifyLayout.from_verify_lens(
verify_lens_cpu=verify_lens_cpu, verify_lens_cpu=verify_lens_cpu,
@@ -509,7 +509,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
if self.require_mlp_tp_gather: if self.require_mlp_tp_gather:
cuda_graph_bs = ( cuda_graph_bs = (
max(forward_batch.global_num_tokens_cpu) // self.num_tokens_per_bs max(forward_batch.global_num_tokens_cpu) // self.num_tokens_per_req
if self.model_runner.spec_algorithm.is_eagle() if self.model_runner.spec_algorithm.is_eagle()
or self.model_runner.spec_algorithm.is_standalone() or self.model_runner.spec_algorithm.is_standalone()
or self.model_runner.spec_algorithm.is_dflash_family() or self.model_runner.spec_algorithm.is_dflash_family()
@@ -559,7 +559,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
is_ngram_supported = ( is_ngram_supported = (
( (
forward_batch.batch_size * self.num_tokens_per_bs forward_batch.batch_size * self.num_tokens_per_req
== forward_batch.input_ids.numel() == forward_batch.input_ids.numel()
) )
if self.model_runner.spec_algorithm.is_ngram() if self.model_runner.spec_algorithm.is_ngram()
@@ -658,7 +658,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
populate static input buffers, choose the active attn backend, and populate static input buffers, choose the active attn backend, and
optionally build pp_proxy_tensors. optionally build pp_proxy_tensors.
num_tokens defaults to the uniform bs * num_tokens_per_bs; ragged num_tokens defaults to the uniform bs * num_tokens_per_req; ragged
verify capture passes the decoupled (slots, tier tokens) pair. verify capture passes the decoupled (slots, tier tokens) pair.
Returns (forward_batch, attn_backend, pp_proxy_tensors); Returns (forward_batch, attn_backend, pp_proxy_tensors);
@@ -667,7 +667,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
bs = size bs = size
buffers: DecodeInputBuffers = self.buffers buffers: DecodeInputBuffers = self.buffers
if num_tokens is None: if num_tokens is None:
num_tokens = bs * self.num_tokens_per_bs num_tokens = bs * self.num_tokens_per_req
# Registry-owned FB-shared slots come through the registry (which # Registry-owned FB-shared slots come through the registry (which
# shares physical storage with self.buffers via source=...); the rest # shares physical storage with self.buffers via source=...); the rest
@@ -815,7 +815,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
if self.enable_torch_compile and not (get_flags().capture.enable_torch_compile): if self.enable_torch_compile and not (get_flags().capture.enable_torch_compile):
self.enable_torch_compile = False self.enable_torch_compile = False
_, self.compile_bs = get_batch_sizes_to_capture( _, self.compile_bs = get_batch_sizes_to_capture(
self.model_runner, self.num_tokens_per_bs self.model_runner, self.num_tokens_per_req
) )
profile_context = empty_context() profile_context = empty_context()
if self.enable_profile_cuda_graph: if self.enable_profile_cuda_graph:
@@ -889,7 +889,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
with torch_compile_decoration.patch_model( with torch_compile_decoration.patch_model(
self.model_runner.model, self.model_runner.model,
bs in self.compile_bs, bs in self.compile_bs,
num_tokens=bs * self.num_tokens_per_bs, num_tokens=bs * self.num_tokens_per_req,
tp_group=self.model_runner.tp_group, tp_group=self.model_runner.tp_group,
) as forward: ) as forward:
self.capture_one_shape(bs, forward, stream_idx, variant_label) self.capture_one_shape(bs, forward, stream_idx, variant_label)
@@ -901,7 +901,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
stream_idx: Optional[int] = None, stream_idx: Optional[int] = None,
variant_label: Optional[str] = None, variant_label: Optional[str] = None,
): ):
num_tokens = size * self.num_tokens_per_bs num_tokens = size * self.num_tokens_per_req
bs = self._ragged_capture_slots(num_tokens) if self.ragged_verify_mode else size bs = self._ragged_capture_slots(num_tokens) if self.ragged_verify_mode else size
# Sanity-check: --debug-cuda-graph requires breakable backend. # Sanity-check: --debug-cuda-graph requires breakable backend.
@@ -1060,7 +1060,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
self._ragged_graph_size self._ragged_graph_size
if is_ragged if is_ragged
else self._capture_graph_size( else self._capture_graph_size(
bs=self.bs, num_tokens=self.bs * self.num_tokens_per_bs bs=self.bs, num_tokens=self.bs * self.num_tokens_per_req
) )
) )
if is_ragged: if is_ragged:
@@ -1105,11 +1105,11 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
) )
padded_num_tokens = graph_size_key padded_num_tokens = graph_size_key
else: else:
raw_num_token = raw_bs * self.num_tokens_per_bs raw_num_token = raw_bs * self.num_tokens_per_req
if self.require_mlp_tp_gather: if self.require_mlp_tp_gather:
max_num_tokens = max(forward_batch.global_num_tokens_cpu) max_num_tokens = max(forward_batch.global_num_tokens_cpu)
max_batch_size = ( max_batch_size = (
max_num_tokens / self.num_tokens_per_bs max_num_tokens / self.num_tokens_per_req
if self.model_runner.spec_algorithm.is_eagle() if self.model_runner.spec_algorithm.is_eagle()
or self.model_runner.spec_algorithm.is_standalone() or self.model_runner.spec_algorithm.is_standalone()
or self.model_runner.spec_algorithm.is_dflash_family() or self.model_runner.spec_algorithm.is_dflash_family()
@@ -1118,7 +1118,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
bs = self._pad_to_bucket(int(max_batch_size), self.capture_bs) bs = self._pad_to_bucket(int(max_batch_size), self.capture_bs)
else: else:
bs = self._pad_to_bucket(raw_bs, self.capture_bs) bs = self._pad_to_bucket(raw_bs, self.capture_bs)
padded_num_tokens = bs * self.num_tokens_per_bs padded_num_tokens = bs * self.num_tokens_per_req
graph_size_key = self._capture_graph_size( graph_size_key = self._capture_graph_size(
bs=bs, num_tokens=padded_num_tokens bs=bs, num_tokens=padded_num_tokens
) )
@@ -1326,7 +1326,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
spec_info = DFlashVerifyInput( spec_info = DFlashVerifyInput(
draft_token=None, draft_token=None,
positions=None, positions=None,
draft_token_num=self.num_tokens_per_bs, draft_token_num=self.num_tokens_per_req,
custom_mask=( custom_mask=(
None None
if (self.model_runner.is_draft_worker or not build_custom_mask) if (self.model_runner.is_draft_worker or not build_custom_mask)
@@ -1350,7 +1350,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
retrieve_index=None, retrieve_index=None,
retrieve_next_token=None, retrieve_next_token=None,
retrieve_next_sibling=None, retrieve_next_sibling=None,
draft_token_num=self.num_tokens_per_bs, draft_token_num=self.num_tokens_per_req,
) )
spec_info.capture_hidden_mode = CaptureHiddenMode.NULL spec_info.capture_hidden_mode = CaptureHiddenMode.NULL
@@ -68,12 +68,12 @@ class EagerRunner(BaseRunner):
sa = mr.server_args sa = mr.server_args
# Built first so the cg runners coalesce onto its buffers via the shared # Built first so the cg runners coalesce onto its buffers via the shared
# input pool; size to the largest tokens/req across modes the worker hits. # input pool; size to the largest tokens/req across modes the worker hits.
num_tokens_per_bs = 1 num_tokens_per_req = 1
if mr.spec_algorithm.is_speculative(): if mr.spec_algorithm.is_speculative():
# speculative_adaptive can grow draft tokens at runtime; size to the max. # speculative_adaptive can grow draft tokens at runtime; size to the max.
num_draft_tokens = sa.max_speculative_num_draft_tokens or 1 num_draft_tokens = sa.max_speculative_num_draft_tokens or 1
if mr.is_draft_worker: if mr.is_draft_worker:
num_tokens_per_bs = max( num_tokens_per_req = max(
sa.speculative_eagle_topk or 1, sa.speculative_eagle_topk or 1,
num_draft_tokens, num_draft_tokens,
( (
@@ -83,8 +83,8 @@ class EagerRunner(BaseRunner):
), ),
) )
else: else:
num_tokens_per_bs = ( num_tokens_per_req = (
mr.spec_algorithm.get_num_tokens_per_bs_for_target_verify( mr.spec_algorithm.get_num_tokens_per_req_for_target_verify(
num_draft_tokens, mr.is_draft_worker num_draft_tokens, mr.is_draft_worker
) )
) )
@@ -92,7 +92,7 @@ class EagerRunner(BaseRunner):
dllm_config = DllmConfig.from_server_args(sa) dllm_config = DllmConfig.from_server_args(sa)
if dllm_config is not None: if dllm_config is not None:
# dLLM runs block_size tokens/request (DLLM_EXTEND). # dLLM runs block_size tokens/request (DLLM_EXTEND).
num_tokens_per_bs = dllm_config.block_size num_tokens_per_req = dllm_config.block_size
max_bs = mr.max_running_requests max_bs = mr.max_running_requests
if ( if (
mr.is_draft_worker mr.is_draft_worker
@@ -109,12 +109,12 @@ class EagerRunner(BaseRunner):
max_bs = ceil_align(max_bs, self.attn_tp_size) max_bs = ceil_align(max_bs, self.attn_tp_size)
max_bs = ceil_align(max_bs, get_cp_padding_align_size()) max_bs = ceil_align(max_bs, get_cp_padding_align_size())
prefill_ceiling = max(mr.max_total_num_tokens, sa.max_prefill_buffer_tokens()) prefill_ceiling = max(mr.max_total_num_tokens, sa.max_prefill_buffer_tokens())
max_num_token = max(prefill_ceiling, max_bs * num_tokens_per_bs) max_num_token = max(prefill_ceiling, max_bs * num_tokens_per_req)
if require_mlp_sync(sa): if require_mlp_sync(sa):
max_num_token = ceil_align(max_num_token, self.attn_tp_size) max_num_token = ceil_align(max_num_token, self.attn_tp_size)
max_num_token = ceil_align(max_num_token, get_cp_padding_align_size()) max_num_token = ceil_align(max_num_token, get_cp_padding_align_size())
self._eager_max_bs = max_bs self._eager_max_bs = max_bs
self._eager_num_tokens_per_bs = num_tokens_per_bs self._eager_num_tokens_per_req = num_tokens_per_req
is_encoder_decoder = mr.model_config.is_encoder_decoder is_encoder_decoder = mr.model_config.is_encoder_decoder
self._eager_registry = build_eager_registry( self._eager_registry = build_eager_registry(
device=mr.device, device=mr.device,
@@ -139,7 +139,7 @@ class EagerRunner(BaseRunner):
self.warmup() self.warmup()
def _autotune_buffers(self) -> Tuple[Any, int]: def _autotune_buffers(self) -> Tuple[Any, int]:
"""Decode-shaped dummy buffers (bs * num_tokens_per_bs) for the warmup """Decode-shaped dummy buffers (bs * num_tokens_per_req) for the warmup
flashinfer-autotune forward. flashinfer-autotune forward.
flashinfer's MoE autotuner times candidate tactics against the buffer it flashinfer's MoE autotuner times candidate tactics against the buffer it
@@ -148,16 +148,16 @@ class EagerRunner(BaseRunner):
ceiling; the dummy run only needs the decode-sized slice. ceiling; the dummy run only needs the decode-sized slice.
""" """
mr = self.model_runner mr = self.model_runner
num_tokens_per_bs = 1 num_tokens_per_req = 1
if mr.spec_algorithm.is_speculative(): if mr.spec_algorithm.is_speculative():
num_tokens_per_bs = ( num_tokens_per_req = (
mr.spec_algorithm.get_num_tokens_per_bs_for_target_verify( mr.spec_algorithm.get_num_tokens_per_req_for_target_verify(
mr.server_args.speculative_num_draft_tokens, mr.is_draft_worker mr.server_args.speculative_num_draft_tokens, mr.is_draft_worker
) )
) )
return ( return (
self._alloc_dummy_decode_buffers( self._alloc_dummy_decode_buffers(
self._eager_max_bs, num_tokens_per_bs=num_tokens_per_bs self._eager_max_bs, num_tokens_per_req=num_tokens_per_req
), ),
self._eager_max_bs, self._eager_max_bs,
) )
@@ -99,7 +99,7 @@ class DecodeInputBuffers(ForwardInputBuffers):
require_mlp_tp_gather: bool, require_mlp_tp_gather: bool,
seq_len_fill_value: int, seq_len_fill_value: int,
encoder_len_fill_value: int, encoder_len_fill_value: int,
num_tokens_per_bs: int, num_tokens_per_req: int,
cache_loc_dtype: torch.dtype, cache_loc_dtype: torch.dtype,
enable_mamba_track: bool, enable_mamba_track: bool,
ne_token_table: Optional[torch.Tensor] = None, ne_token_table: Optional[torch.Tensor] = None,
@@ -116,7 +116,7 @@ class DecodeInputBuffers(ForwardInputBuffers):
mrope_positions = torch.zeros((3, max_num_token), dtype=torch.int64) mrope_positions = torch.zeros((3, max_num_token), dtype=torch.int64)
num_token_non_padded = torch.zeros((1,), dtype=torch.int32) num_token_non_padded = torch.zeros((1,), dtype=torch.int32)
custom_mask = torch.ones( custom_mask = torch.ones(
(max_bs * seq_len_fill_value + max_num_token) * num_tokens_per_bs, (max_bs * seq_len_fill_value + max_num_token) * num_tokens_per_req,
dtype=torch.bool, dtype=torch.bool,
) )
mamba_track_indices = ( mamba_track_indices = (
@@ -220,7 +220,7 @@ class DecodeInputBuffers(ForwardInputBuffers):
bs: int, bs: int,
seq_len_fill_value: int, seq_len_fill_value: int,
require_gathered_buffer: bool, require_gathered_buffer: bool,
num_tokens_per_bs: int, num_tokens_per_req: int,
dsa_enable_prefill_cp: bool, dsa_enable_prefill_cp: bool,
enable_num_token_non_padded_flag: bool, enable_num_token_non_padded_flag: bool,
pp_proxy_tensors: Optional[PPProxyTensors] = None, pp_proxy_tensors: Optional[PPProxyTensors] = None,
@@ -290,12 +290,12 @@ class DecodeInputBuffers(ForwardInputBuffers):
srcs.append(forward_batch.bootstrap_room_ids_int) srcs.append(forward_batch.bootstrap_room_ids_int)
if require_gathered_buffer: if require_gathered_buffer:
self.global_num_tokens_gpu.fill_(bs * num_tokens_per_bs) self.global_num_tokens_gpu.fill_(bs * num_tokens_per_req)
self.global_num_tokens_for_logprob_gpu.fill_(bs * num_tokens_per_bs) self.global_num_tokens_for_logprob_gpu.fill_(bs * num_tokens_per_req)
if enable_num_token_non_padded_flag: if enable_num_token_non_padded_flag:
if require_gathered_buffer and not dsa_enable_prefill_cp: if require_gathered_buffer and not dsa_enable_prefill_cp:
num_tokens_per_dp = bs * num_tokens_per_bs num_tokens_per_dp = bs * num_tokens_per_req
local = compute_local_num_token_non_padded( local = compute_local_num_token_non_padded(
global_num_token_non_padded=forward_batch.num_token_non_padded, global_num_token_non_padded=forward_batch.num_token_non_padded,
num_tokens_per_dp=num_tokens_per_dp, num_tokens_per_dp=num_tokens_per_dp,
+6 -6
View File
@@ -2926,7 +2926,7 @@ class ServerArgs:
handle_speculative_decoding(self) handle_speculative_decoding(self)
# Validate the CuteDSL A2A token budget now that num_tokens_per_bs is final. # Validate the CuteDSL A2A token budget now that num_tokens_per_req is final.
self._validate_cutedsl_a2a_token_budget() self._validate_cutedsl_a2a_token_budget()
# Handle model loading format. # Handle model loading format.
@@ -5419,19 +5419,19 @@ class ServerArgs:
MoE layer on one (DP) rank. Single source of truth for both the MoE layer on one (DP) rank. Single source of truth for both the
standard-allgather wrapper buffers and the FlashInfer A2A dispatcher standard-allgather wrapper buffers and the FlashInfer A2A dispatcher
budget. Max over the prefill (max_prefill_tokens), piecewise-prefill budget. Max over the prefill (max_prefill_tokens), piecewise-prefill
capture, and decode/verify bounds; num_tokens_per_bs is capture, and decode/verify bounds; num_tokens_per_req is
speculative_num_draft_tokens under speculative decoding, else 1. speculative_num_draft_tokens under speculative decoding, else 1.
""" """
if self.speculative_algorithm: if self.speculative_algorithm:
num_tokens_per_bs = self.speculative_num_draft_tokens or 1 num_tokens_per_req = self.speculative_num_draft_tokens or 1
else: else:
num_tokens_per_bs = 1 num_tokens_per_req = 1
prefill_tokens = self.max_prefill_tokens prefill_tokens = self.max_prefill_tokens
cg_config = self.cuda_graph_config cg_config = self.cuda_graph_config
if cg_config is not None and cg_config.prefill.backend == Backend.TC_PIECEWISE: if cg_config is not None and cg_config.prefill.backend == Backend.TC_PIECEWISE:
prefill_tokens = max(prefill_tokens, cg_config.prefill.max_bs or 0) prefill_tokens = max(prefill_tokens, cg_config.prefill.max_bs or 0)
decode_max_bs = (cg_config.decode.max_bs if cg_config is not None else 0) or 0 decode_max_bs = (cg_config.decode.max_bs if cg_config is not None else 0) or 0
decode_tokens = decode_max_bs * num_tokens_per_bs decode_tokens = decode_max_bs * num_tokens_per_req
return max(prefill_tokens, decode_tokens) return max(prefill_tokens, decode_tokens)
def max_prefill_buffer_tokens(self) -> int: def max_prefill_buffer_tokens(self) -> int:
@@ -5452,7 +5452,7 @@ class ServerArgs:
def _validate_cutedsl_a2a_token_budget(self): def _validate_cutedsl_a2a_token_budget(self):
"""Fail fast if the FlashInfer A2A dispatcher workspace cannot cover the """Fail fast if the FlashInfer A2A dispatcher workspace cannot cover the
largest CuteDSL MoE forward. Runs after speculative decoding is resolved largest CuteDSL MoE forward. Runs after speculative decoding is resolved
so cutedsl_moe_max_num_tokens() sees the final num_tokens_per_bs.""" so cutedsl_moe_max_num_tokens() sees the final num_tokens_per_req."""
from sglang.srt.arg_groups.overrides import resolved_view from sglang.srt.arg_groups.overrides import resolved_view
view = resolved_view(self) view = resolved_view(self)
@@ -145,9 +145,9 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
# Bucket sizes # Bucket sizes
self.capture_bs, _ = get_batch_sizes_to_capture(model_runner) self.capture_bs, _ = get_batch_sizes_to_capture(model_runner)
self.num_tokens_per_bs = self.topk self.num_tokens_per_req = self.topk
self.max_bs = max(self.capture_bs) self.max_bs = max(self.capture_bs)
self.max_num_token = self.max_bs * self.num_tokens_per_bs self.max_num_token = self.max_bs * self.num_tokens_per_req
# Attention backend init # Attention backend init
self.draft_attn_backend.init_cuda_graph_state(self.max_bs, self.max_num_token) self.draft_attn_backend.init_cuda_graph_state(self.max_bs, self.max_num_token)
@@ -290,7 +290,7 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
def can_run_graph(self, forward_batch: ForwardBatch): def can_run_graph(self, forward_batch: ForwardBatch):
if self.require_mlp_tp_gather: if self.require_mlp_tp_gather:
cuda_graph_bs = ( cuda_graph_bs = (
max(forward_batch.global_num_tokens_cpu) // self.num_tokens_per_bs max(forward_batch.global_num_tokens_cpu) // self.num_tokens_per_req
if self.model_runner.spec_algorithm.is_eagle() if self.model_runner.spec_algorithm.is_eagle()
or self.model_runner.spec_algorithm.is_standalone() or self.model_runner.spec_algorithm.is_standalone()
else max(forward_batch.global_num_tokens_cpu) else max(forward_batch.global_num_tokens_cpu)
@@ -321,7 +321,7 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
): ):
num_seqs = size # EAGLE legacy name num_seqs = size # EAGLE legacy name
buffers = self.buffers buffers = self.buffers
num_tokens = num_seqs * self.num_tokens_per_bs num_tokens = num_seqs * self.num_tokens_per_req
# Graph inputs # Graph inputs
req_pool_indices = buffers.req_pool_indices[:num_seqs] req_pool_indices = buffers.req_pool_indices[:num_seqs]
@@ -490,13 +490,13 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
buffers = self.buffers buffers = self.buffers
raw_bs = forward_batch.batch_size raw_bs = forward_batch.batch_size
raw_num_token = raw_bs * self.num_tokens_per_bs raw_num_token = raw_bs * self.num_tokens_per_req
# Pad to nearest captured shape # Pad to nearest captured shape
if self.require_mlp_tp_gather: if self.require_mlp_tp_gather:
max_num_tokens = max(forward_batch.global_num_tokens_cpu) max_num_tokens = max(forward_batch.global_num_tokens_cpu)
max_batch_size = ( max_batch_size = (
max_num_tokens // self.num_tokens_per_bs max_num_tokens // self.num_tokens_per_req
if self.model_runner.spec_algorithm.is_eagle() if self.model_runner.spec_algorithm.is_eagle()
or self.model_runner.spec_algorithm.is_standalone() or self.model_runner.spec_algorithm.is_standalone()
else max_num_tokens else max_num_tokens
@@ -523,7 +523,7 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
buffers.dsa_seed_topk.zero_() buffers.dsa_seed_topk.zero_()
buffers.req_pool_indices.zero_() buffers.req_pool_indices.zero_()
num_tokens = bs * self.num_tokens_per_bs num_tokens = bs * self.num_tokens_per_req
maybe_detect_nan( maybe_detect_nan(
forward_batch.spec_info.topk_p, forward_batch.spec_info.topk_p,
@@ -598,8 +598,10 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
# TODO(ch-wan): support num_token_non_padded # TODO(ch-wan): support num_token_non_padded
if self.require_gathered_buffer: if self.require_gathered_buffer:
buffers.global_num_tokens_gpu.fill_(bs * self.num_tokens_per_bs) buffers.global_num_tokens_gpu.fill_(bs * self.num_tokens_per_req)
buffers.global_num_tokens_for_logprob_gpu.fill_(bs * self.num_tokens_per_bs) buffers.global_num_tokens_for_logprob_gpu.fill_(
bs * self.num_tokens_per_req
)
# Save the raw seq_lens_sum; it is restored after replay. While the graph # Save the raw seq_lens_sum; it is restored after replay. While the graph
# runs it must reflect the padded fake rows (set below), since draft decode # runs it must reflect the padded fake rows (set below), since draft decode
@@ -134,9 +134,9 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
# Size cuda-graph buffers by num_draft_tokens (full tree width), not # Size cuda-graph buffers by num_draft_tokens (full tree width), not
# num_steps + 1, or topk > 1 draft-extend overflows them. # num_steps + 1, or topk > 1 draft-extend overflows them.
self.num_tokens_per_bs = model_runner.server_args.speculative_num_draft_tokens self.num_tokens_per_req = model_runner.server_args.speculative_num_draft_tokens
self.max_bs = max(self.capture_bs) self.max_bs = max(self.capture_bs)
self.max_num_token = self.max_bs * self.num_tokens_per_bs self.max_num_token = self.max_bs * self.num_tokens_per_req
self.draft_extend_attn_backend.init_cuda_graph_state( self.draft_extend_attn_backend.init_cuda_graph_state(
self.max_bs, self.max_num_token self.max_bs, self.max_num_token
@@ -144,7 +144,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
self.seq_len_fill_value = ( self.seq_len_fill_value = (
self.draft_extend_attn_backend.get_cuda_graph_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_bs] * self.max_bs self.extend_seq_lens_cpu = [self.num_tokens_per_req] * self.max_bs
if self.enable_torch_compile: if self.enable_torch_compile:
set_torch_compile_config() set_torch_compile_config()
@@ -182,13 +182,13 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
(self.max_bs,), self.seq_len_fill_value, dtype=torch.int64 (self.max_bs,), self.seq_len_fill_value, dtype=torch.int64
) )
extend_seq_lens = torch.full( extend_seq_lens = torch.full(
(self.max_bs,), self.num_tokens_per_bs, dtype=torch.int32 (self.max_bs,), self.num_tokens_per_req, dtype=torch.int32
) )
num_correct_drafts = torch.full( num_correct_drafts = torch.full(
(self.max_bs,), self.num_tokens_per_bs, dtype=torch.int32 (self.max_bs,), self.num_tokens_per_req, dtype=torch.int32
) )
num_accept_tokens = torch.full( num_accept_tokens = torch.full(
(self.max_bs,), self.num_tokens_per_bs, dtype=torch.int32 (self.max_bs,), self.num_tokens_per_req, dtype=torch.int32
) )
if self.require_gathered_buffer: if self.require_gathered_buffer:
@@ -227,7 +227,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
next_token_logits_buffer = ( next_token_logits_buffer = (
self.model_runner.graph_shared_output.get_logits_buffer( self.model_runner.graph_shared_output.get_logits_buffer(
vocab_size, rows=self.max_bs * self.num_tokens_per_bs vocab_size, rows=self.max_bs * self.num_tokens_per_req
) )
) )
@@ -287,7 +287,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
def can_run_graph(self, forward_batch: ForwardBatch): def can_run_graph(self, forward_batch: ForwardBatch):
if self.require_mlp_tp_gather: if self.require_mlp_tp_gather:
cuda_graph_bs = ( cuda_graph_bs = (
max(forward_batch.global_num_tokens_cpu) // self.num_tokens_per_bs max(forward_batch.global_num_tokens_cpu) // self.num_tokens_per_req
if self.model_runner.spec_algorithm.is_eagle() if self.model_runner.spec_algorithm.is_eagle()
or self.model_runner.spec_algorithm.is_standalone() or self.model_runner.spec_algorithm.is_standalone()
else max(forward_batch.global_num_tokens_cpu) else max(forward_batch.global_num_tokens_cpu)
@@ -315,7 +315,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
): ):
bs = size bs = size
buffers = self.buffers buffers = self.buffers
num_tokens = bs * self.num_tokens_per_bs num_tokens = bs * self.num_tokens_per_req
# Graph inputs # Graph inputs
input_ids = buffers.input_ids[:num_tokens] input_ids = buffers.input_ids[:num_tokens]
@@ -370,7 +370,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
num_correct_drafts=num_correct_drafts, num_correct_drafts=num_correct_drafts,
num_accept_tokens=num_accept_tokens, num_accept_tokens=num_accept_tokens,
# Padded tree width per req; drives the constant qo layout. # Padded tree width per req; drives the constant qo layout.
num_tokens_per_req=self.num_tokens_per_bs, num_tokens_per_req=self.num_tokens_per_req,
) )
forward_batch = ForwardBatch( forward_batch = ForwardBatch(
@@ -482,7 +482,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
if self.require_mlp_tp_gather: if self.require_mlp_tp_gather:
max_num_tokens = max(forward_batch.global_num_tokens_cpu) max_num_tokens = max(forward_batch.global_num_tokens_cpu)
max_batch_size = ( max_batch_size = (
max_num_tokens // self.num_tokens_per_bs max_num_tokens // self.num_tokens_per_req
if self.model_runner.spec_algorithm.is_eagle() if self.model_runner.spec_algorithm.is_eagle()
else max_num_tokens else max_num_tokens
) )
@@ -490,16 +490,16 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
else: else:
bs = self._pad_to_bucket(raw_bs, self.capture_bs) bs = self._pad_to_bucket(raw_bs, self.capture_bs)
if bs * self.num_tokens_per_bs != num_tokens: if bs * self.num_tokens_per_req != num_tokens:
buffers.seq_lens.fill_(self.seq_len_fill_value) buffers.seq_lens.fill_(self.seq_len_fill_value)
buffers.out_cache_loc.zero_() buffers.out_cache_loc.zero_()
buffers.positions.zero_() buffers.positions.zero_()
# Pair with seq_lens fill: padded rows must point at reserved # Pair with seq_lens fill: padded rows must point at reserved
# req_pool slot 0 (req_to_token[0, :] is all zeros from init). # req_pool slot 0 (req_to_token[0, :] is all zeros from init).
buffers.req_pool_indices.zero_() buffers.req_pool_indices.zero_()
buffers.num_correct_drafts.fill_(self.num_tokens_per_bs) buffers.num_correct_drafts.fill_(self.num_tokens_per_req)
buffers.num_accept_tokens.fill_(self.num_tokens_per_bs) buffers.num_accept_tokens.fill_(self.num_tokens_per_req)
buffers.extend_seq_lens.fill_(self.num_tokens_per_bs) buffers.extend_seq_lens.fill_(self.num_tokens_per_req)
# Batch the small per-field device copies into a grouped foreach copy # Batch the small per-field device copies into a grouped foreach copy
# (one foreach call per dtype pair) to cut launch overhead. hidden_states # (one foreach call per dtype pair) to cut launch overhead. hidden_states
@@ -523,7 +523,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
copy_dsts.append(buffers.extend_seq_lens[:raw_bs]) copy_dsts.append(buffers.extend_seq_lens[:raw_bs])
copy_srcs.append(forward_batch.extend_seq_lens) copy_srcs.append(forward_batch.extend_seq_lens)
else: else:
buffers.extend_seq_lens[:raw_bs].fill_(self.num_tokens_per_bs) buffers.extend_seq_lens[:raw_bs].fill_(self.num_tokens_per_req)
if forward_batch.spec_info.num_correct_drafts is not None: if forward_batch.spec_info.num_correct_drafts is not None:
copy_dsts.append(buffers.num_correct_drafts[:raw_bs]) copy_dsts.append(buffers.num_correct_drafts[:raw_bs])
copy_srcs.append(forward_batch.spec_info.num_correct_drafts) copy_srcs.append(forward_batch.spec_info.num_correct_drafts)
@@ -545,8 +545,10 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
# TODO(ch-wan): support num_token_non_padded # TODO(ch-wan): support num_token_non_padded
if self.require_gathered_buffer: if self.require_gathered_buffer:
buffers.global_num_tokens_gpu.fill_(bs * self.num_tokens_per_bs) buffers.global_num_tokens_gpu.fill_(bs * self.num_tokens_per_req)
buffers.global_num_tokens_for_logprob_gpu.fill_(bs * self.num_tokens_per_bs) buffers.global_num_tokens_for_logprob_gpu.fill_(
bs * self.num_tokens_per_req
)
if forward_batch.seq_lens_cpu is not None: if forward_batch.seq_lens_cpu is not None:
if bs != raw_bs: if bs != raw_bs:
@@ -556,9 +558,9 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
if forward_batch.extend_seq_lens_cpu is not None: if forward_batch.extend_seq_lens_cpu is not None:
self.extend_seq_lens_cpu[:raw_bs] = forward_batch.extend_seq_lens_cpu self.extend_seq_lens_cpu[:raw_bs] = forward_batch.extend_seq_lens_cpu
else: else:
self.extend_seq_lens_cpu[:raw_bs] = [self.num_tokens_per_bs] * raw_bs self.extend_seq_lens_cpu[:raw_bs] = [self.num_tokens_per_req] * raw_bs
if bs > raw_bs: if bs > raw_bs:
self.extend_seq_lens_cpu[raw_bs:bs] = [self.num_tokens_per_bs] * ( self.extend_seq_lens_cpu[raw_bs:bs] = [self.num_tokens_per_req] * (
bs - raw_bs bs - raw_bs
) )
forward_batch.spec_info.extend_seq_lens_cpu = list( forward_batch.spec_info.extend_seq_lens_cpu = list(
@@ -424,7 +424,7 @@ class EagleDraftWorker(EagleDraftWorkerBase):
log_info_on_rank0( log_info_on_rank0(
logger, logger,
f"Capture draft decode CUDA graph begin. backend={decode_backend}, " f"Capture draft decode CUDA graph begin. backend={decode_backend}, "
f"num_tokens_per_bs={self.topk}, bs={capture_bs}, " f"num_tokens_per_req={self.topk}, bs={capture_bs}, "
f"avail mem={before_mem:.2f} GB", f"avail mem={before_mem:.2f} GB",
) )
self.cuda_graph_runner = Device2DraftCudaGraphRunner[ self.cuda_graph_runner = Device2DraftCudaGraphRunner[
@@ -501,7 +501,7 @@ class EagleDraftWorker(EagleDraftWorkerBase):
log_info_on_rank0( log_info_on_rank0(
logger, logger,
f"Capture draft extend CUDA graph begin. backend={decode_backend}, " f"Capture draft extend CUDA graph begin. backend={decode_backend}, "
f"num_tokens_per_bs={self.speculative_num_draft_tokens}, " f"num_tokens_per_req={self.speculative_num_draft_tokens}, "
f"bs={capture_bs}, avail mem={before_mem:.2f} GB", f"bs={capture_bs}, avail mem={before_mem:.2f} GB",
) )
self.cuda_graph_runner_for_draft_extend = Device2ExtendCudaGraphRunner[ self.cuda_graph_runner_for_draft_extend = Device2ExtendCudaGraphRunner[
@@ -113,12 +113,12 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
self.capture_forward_mode = ForwardMode.DECODE self.capture_forward_mode = ForwardMode.DECODE
self.capture_hidden_mode = CaptureHiddenMode.LAST self.capture_hidden_mode = CaptureHiddenMode.LAST
self.num_tokens_per_bs = self.topk self.num_tokens_per_req = self.topk
self.capture_bs, _ = get_batch_sizes_to_capture( self.capture_bs, _ = get_batch_sizes_to_capture(
model_runner, self.num_tokens_per_bs model_runner, self.num_tokens_per_req
) )
self.max_bs = max(self.capture_bs) self.max_bs = max(self.capture_bs)
self.max_num_token = self.max_bs * self.num_tokens_per_bs self.max_num_token = self.max_bs * self.num_tokens_per_req
self.draft_attn_backend.init_cuda_graph_state(self.max_bs, self.max_num_token) self.draft_attn_backend.init_cuda_graph_state(self.max_bs, self.max_num_token)
self.seq_len_fill_value = ( self.seq_len_fill_value = (
@@ -227,7 +227,7 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
del forward, stream_idx, variant_label del forward, stream_idx, variant_label
buffers = self.buffers buffers = self.buffers
request_bs = size request_bs = size
expanded_bs = request_bs * self.num_tokens_per_bs expanded_bs = request_bs * self.num_tokens_per_req
req_pool_indices = buffers.req_pool_indices[:expanded_bs] req_pool_indices = buffers.req_pool_indices[:expanded_bs]
positions = buffers.positions[:expanded_bs] positions = buffers.positions[:expanded_bs]
@@ -362,7 +362,7 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
raw_expanded_bs = forward_batch.batch_size raw_expanded_bs = forward_batch.batch_size
raw_bs = ( raw_bs = (
raw_expanded_bs // self.num_tokens_per_bs raw_expanded_bs // self.num_tokens_per_req
if self.topk > 1 if self.topk > 1
else raw_expanded_bs else raw_expanded_bs
) )
@@ -371,13 +371,13 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
if self.require_mlp_tp_gather: if self.require_mlp_tp_gather:
max_num_tokens = max(forward_batch.global_num_tokens_cpu) max_num_tokens = max(forward_batch.global_num_tokens_cpu)
max_batch_size = max_num_tokens // ( max_batch_size = max_num_tokens // (
self.num_tokens_per_bs * self.num_tokens_per_bs self.num_tokens_per_req * self.num_tokens_per_req
) )
bs = self._pad_to_bucket(int(max_batch_size), self.capture_bs) bs = self._pad_to_bucket(int(max_batch_size), self.capture_bs)
else: else:
bs = self._pad_to_bucket(raw_bs, self.capture_bs) bs = self._pad_to_bucket(raw_bs, self.capture_bs)
expanded_bs = bs * self.num_tokens_per_bs expanded_bs = bs * self.num_tokens_per_req
if bs != raw_bs: if bs != raw_bs:
buffers.seq_lens.fill_(self.seq_len_fill_value) buffers.seq_lens.fill_(self.seq_len_fill_value)
buffers.positions.zero_() buffers.positions.zero_()
@@ -160,10 +160,10 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
# Fixed window: every step extends each request by the same number of # Fixed window: every step extends each request by the same number of
# tokens, which lets all steps share one buffer set. # tokens, which lets all steps share one buffer set.
self.num_tokens_per_bs = self.speculative_num_draft_tokens self.num_tokens_per_req = self.speculative_num_draft_tokens
self.max_bs = max(self.capture_bs) self.max_bs = max(self.capture_bs)
self.max_num_token = self.max_bs * self.num_tokens_per_bs self.max_num_token = self.max_bs * self.num_tokens_per_req
self.extend_seq_lens_cpu = [self.num_tokens_per_bs] * self.max_bs self.extend_seq_lens_cpu = [self.num_tokens_per_req] * self.max_bs
self.eagle_worker.draft_extend_attn_backend_list[ self.eagle_worker.draft_extend_attn_backend_list[
self.step self.step
@@ -198,7 +198,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
def can_run_graph(self, forward_batch: ForwardBatch): def can_run_graph(self, forward_batch: ForwardBatch):
if self.require_mlp_tp_gather: if self.require_mlp_tp_gather:
cuda_graph_bs = ( cuda_graph_bs = (
max(forward_batch.global_num_tokens_cpu) // self.num_tokens_per_bs max(forward_batch.global_num_tokens_cpu) // self.num_tokens_per_req
if self.model_runner.spec_algorithm.is_eagle() if self.model_runner.spec_algorithm.is_eagle()
else max(forward_batch.global_num_tokens_cpu) else max(forward_batch.global_num_tokens_cpu)
) )
@@ -218,7 +218,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
def get_forward_batch(self, bs: int) -> ForwardBatch: def get_forward_batch(self, bs: int) -> ForwardBatch:
buffers = self.buffers buffers = self.buffers
num_tokens = bs * self.num_tokens_per_bs num_tokens = bs * self.num_tokens_per_req
input_ids = buffers.input_ids[:num_tokens] input_ids = buffers.input_ids[:num_tokens]
req_pool_indices = buffers.req_pool_indices[:bs] req_pool_indices = buffers.req_pool_indices[:bs]
@@ -303,8 +303,8 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
extend_seq_lens_cpu=extend_seq_lens_cpu, extend_seq_lens_cpu=extend_seq_lens_cpu,
padded_static_len=self.padded_static_len, padded_static_len=self.padded_static_len,
extend_start_loc=extend_start_loc, extend_start_loc=extend_start_loc,
extend_num_tokens=self.num_tokens_per_bs * bs, extend_num_tokens=self.num_tokens_per_req * bs,
num_token_non_padded_cpu=self.num_tokens_per_bs * bs, num_token_non_padded_cpu=self.num_tokens_per_req * bs,
return_hidden_states_before_norm=True, return_hidden_states_before_norm=True,
) )
return forward_batch return forward_batch
@@ -332,7 +332,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
bs = size bs = size
buffers = self.buffers buffers = self.buffers
num_tokens = bs * self.num_tokens_per_bs num_tokens = bs * self.num_tokens_per_req
forward_batch = self.get_forward_batch(bs) forward_batch = self.get_forward_batch(bs)
forward_batch = self._postprocess_forward_batch(forward_batch, bs) forward_batch = self._postprocess_forward_batch(forward_batch, bs)
attn_backend = self.eagle_worker.draft_extend_attn_backend_list[self.step] attn_backend = self.eagle_worker.draft_extend_attn_backend_list[self.step]
@@ -398,7 +398,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
write + worker-side rotation (steps > 0).""" write + worker-side rotation (steps > 0)."""
self.deepep_adapter.replay() self.deepep_adapter.replay()
buffers = self.buffers buffers = self.buffers
num_tokens = bs * self.num_tokens_per_bs num_tokens = bs * self.num_tokens_per_req
if self.require_gathered_buffer: if self.require_gathered_buffer:
buffers.global_num_tokens_gpu.fill_(num_tokens) buffers.global_num_tokens_gpu.fill_(num_tokens)
@@ -453,7 +453,7 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner:
self.runners: List[Optional[MultiLayerEagleDraftExtendCudaGraphRunner]] = [] self.runners: List[Optional[MultiLayerEagleDraftExtendCudaGraphRunner]] = []
self.seq_len_fill_value = 1 self.seq_len_fill_value = 1
self.max_bs = 1 self.max_bs = 1
self.num_tokens_per_bs = 1 self.num_tokens_per_req = 1
self._init_and_capture() self._init_and_capture()
@@ -487,7 +487,7 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner:
self.runners.append(runner) self.runners.append(runner)
self.seq_len_fill_value = runner.seq_len_fill_value self.seq_len_fill_value = runner.seq_len_fill_value
self.max_bs = runner.max_bs self.max_bs = runner.max_bs
self.num_tokens_per_bs = runner.num_tokens_per_bs self.num_tokens_per_req = runner.num_tokens_per_req
self.capture_bs = runner.capture_bs self.capture_bs = runner.capture_bs
self.require_gathered_buffer = runner.require_gathered_buffer self.require_gathered_buffer = runner.require_gathered_buffer
self.require_mlp_tp_gather = runner.require_mlp_tp_gather self.require_mlp_tp_gather = runner.require_mlp_tp_gather
@@ -533,8 +533,8 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner:
runner = next(r for r in self.runners if r is not None) runner = next(r for r in self.runners if r is not None)
model_runner = runner.model_runner model_runner = runner.model_runner
max_bs = self.max_bs max_bs = self.max_bs
num_tokens_per_bs = self.num_tokens_per_bs num_tokens_per_req = self.num_tokens_per_req
max_num_token = max_bs * num_tokens_per_bs max_num_token = max_bs * num_tokens_per_req
hidden_size = get_draft_input_from_target_hidden_dim(model_runner) hidden_size = get_draft_input_from_target_hidden_dim(model_runner)
dtype = model_runner.model_config.dtype dtype = model_runner.model_config.dtype
vocab_size = self._vocab_size() vocab_size = self._vocab_size()
@@ -553,13 +553,13 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner:
num_correct_drafts = torch.full((max_bs,), 1, dtype=torch.int32) num_correct_drafts = torch.full((max_bs,), 1, dtype=torch.int32)
num_accept_tokens = torch.full((max_bs,), 1, dtype=torch.int32) num_accept_tokens = torch.full((max_bs,), 1, dtype=torch.int32)
# Fixed window: every request extends by exactly num_tokens_per_bs # Fixed window: every request extends by exactly num_tokens_per_req
# tokens, and start locs are a constant arange. # tokens, and start locs are a constant arange.
extend_seq_lens = torch.full( extend_seq_lens = torch.full(
(max_bs,), num_tokens_per_bs, dtype=torch.int32 (max_bs,), num_tokens_per_req, dtype=torch.int32
) )
extend_start_loc = torch.arange( extend_start_loc = torch.arange(
0, max_num_token, step=num_tokens_per_bs, dtype=torch.int32 0, max_num_token, step=num_tokens_per_req, dtype=torch.int32
) )
select_index = torch.zeros((max_bs,), dtype=torch.int64) select_index = torch.zeros((max_bs,), dtype=torch.int64)
@@ -610,7 +610,7 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner:
the batch size. Subsequent ``replay(step)`` calls reuse this state.""" the batch size. Subsequent ``replay(step)`` calls reuse this state."""
buffers = self.buffers buffers = self.buffers
raw_bs = forward_batch.batch_size raw_bs = forward_batch.batch_size
num_tokens = raw_bs * self.num_tokens_per_bs num_tokens = raw_bs * self.num_tokens_per_req
# Bucketize to a captured batch size (padding the tail). # Bucketize to a captured batch size (padding the tail).
if self.require_mlp_tp_gather: if self.require_mlp_tp_gather:
@@ -656,21 +656,23 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner:
# and by the worker's rotation. # and by the worker's rotation.
arange = torch.arange(bs, device=self.device, dtype=torch.int64) arange = torch.arange(bs, device=self.device, dtype=torch.int64)
buffers.select_index[:bs].copy_( buffers.select_index[:bs].copy_(
arange * self.num_tokens_per_bs + buffers.num_correct_drafts[:bs] arange * self.num_tokens_per_req + buffers.num_correct_drafts[:bs]
) )
if self.require_gathered_buffer: if self.require_gathered_buffer:
buffers.global_num_tokens_gpu.fill_(bs * self.num_tokens_per_bs) buffers.global_num_tokens_gpu.fill_(bs * self.num_tokens_per_req)
buffers.global_num_tokens_for_logprob_gpu.fill_(bs * self.num_tokens_per_bs) buffers.global_num_tokens_for_logprob_gpu.fill_(
bs * self.num_tokens_per_req
)
# Reusable spec_info for per-step attention metadata. # Reusable spec_info for per-step attention metadata.
padded_num_tokens = bs * self.num_tokens_per_bs padded_num_tokens = bs * self.num_tokens_per_req
spec_info = EagleDraftExtendInput( spec_info = EagleDraftExtendInput(
hidden_states=buffers.hidden_states[:padded_num_tokens], hidden_states=buffers.hidden_states[:padded_num_tokens],
num_correct_drafts=buffers.num_correct_drafts[:bs], num_correct_drafts=buffers.num_correct_drafts[:bs],
num_accept_tokens=buffers.num_accept_tokens[:bs], num_accept_tokens=buffers.num_accept_tokens[:bs],
) )
spec_info.num_tokens_per_req = self.num_tokens_per_bs spec_info.num_tokens_per_req = self.num_tokens_per_req
spec_info.num_tokens_for_logprob_per_req = 1 spec_info.num_tokens_for_logprob_per_req = 1
spec_info.positions = buffers.positions[:padded_num_tokens] spec_info.positions = buffers.positions[:padded_num_tokens]
spec_info.extend_seq_lens_tensor = buffers.extend_seq_lens[:bs] spec_info.extend_seq_lens_tensor = buffers.extend_seq_lens[:bs]
+18 -3
View File
@@ -1,5 +1,6 @@
from __future__ import annotations from __future__ import annotations
import warnings
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from enum import Enum, IntEnum, auto from enum import Enum, IntEnum, auto
from typing import TYPE_CHECKING, Callable, List, Optional, Tuple, Type, Union from typing import TYPE_CHECKING, Callable, List, Optional, Tuple, Type, Union
@@ -210,7 +211,7 @@ class SpeculativeAlgorithm(Enum):
elif self.is_ngram(): elif self.is_ngram():
_handle_ngram(server_args) _handle_ngram(server_args)
def get_num_tokens_per_bs_for_target_verify( def get_num_tokens_per_req_for_target_verify(
self, num_draft_tokens: int, is_draft_worker: bool self, num_draft_tokens: int, is_draft_worker: bool
) -> int: ) -> int:
# FIXME: Remove this after the forward mode refactor. Target verify is # FIXME: Remove this after the forward mode refactor. Target verify is
@@ -222,6 +223,20 @@ class SpeculativeAlgorithm(Enum):
return num_draft_tokens - 1 return num_draft_tokens - 1
return num_draft_tokens return num_draft_tokens
def get_num_tokens_per_bs_for_target_verify(
self, num_draft_tokens: int, is_draft_worker: bool
) -> int:
# Deprecated alias; remove together with the FIXME above.
warnings.warn(
"get_num_tokens_per_bs_for_target_verify is deprecated; use "
"get_num_tokens_per_req_for_target_verify instead.",
DeprecationWarning,
stacklevel=2,
)
return self.get_num_tokens_per_req_for_target_verify(
num_draft_tokens, is_draft_worker
)
def create_worker( def create_worker(
self, server_args: ServerArgs self, server_args: ServerArgs
) -> Optional[Union[Type[BaseSpecWorker], Type[TpModelWorker], Type[NGRAMWorker]]]: ) -> Optional[Union[Type[BaseSpecWorker], Type[TpModelWorker], Type[NGRAMWorker]]]:
@@ -343,7 +358,7 @@ def create_dummy_verify_input(
spec_algorithm: SpeculativeAlgorithm, spec_algorithm: SpeculativeAlgorithm,
server_args: ServerArgs, server_args: ServerArgs,
custom_mask: torch.Tensor, custom_mask: torch.Tensor,
num_tokens_per_bs: int, num_tokens_per_req: int,
is_draft_worker: bool, is_draft_worker: bool,
) -> Optional[SpecInput]: ) -> Optional[SpecInput]:
"""Dummy verify ``SpecInput`` for CUDA-graph capture (per-algorithm dispatch).""" """Dummy verify ``SpecInput`` for CUDA-graph capture (per-algorithm dispatch)."""
@@ -395,7 +410,7 @@ def create_dummy_verify_input(
retrieve_index=None, retrieve_index=None,
retrieve_next_token=None, retrieve_next_token=None,
retrieve_next_sibling=None, retrieve_next_sibling=None,
draft_token_num=num_tokens_per_bs, draft_token_num=num_tokens_per_req,
) )
spec_info.capture_hidden_mode = CaptureHiddenMode.NULL spec_info.capture_hidden_mode = CaptureHiddenMode.NULL
+16 -1
View File
@@ -5,6 +5,7 @@ should use that classmethod API; do not import from this module directly.
from __future__ import annotations from __future__ import annotations
import logging import logging
import warnings
from typing import TYPE_CHECKING, Callable, Dict, Optional, Type from typing import TYPE_CHECKING, Callable, Dict, Optional, Type
import torch import torch
@@ -119,7 +120,7 @@ class CustomSpecAlgo:
) )
return self.factory(server_args) return self.factory(server_args)
def get_num_tokens_per_bs_for_target_verify( def get_num_tokens_per_req_for_target_verify(
self, num_draft_tokens: int, is_draft_worker: bool self, num_draft_tokens: int, is_draft_worker: bool
) -> int: ) -> int:
# FIXME: Remove this after the forward mode refactor. Target verify is # FIXME: Remove this after the forward mode refactor. Target verify is
@@ -129,6 +130,20 @@ class CustomSpecAlgo:
# Here, we expose this interface to allow the other use cases. # Here, we expose this interface to allow the other use cases.
return num_draft_tokens return num_draft_tokens
def get_num_tokens_per_bs_for_target_verify(
self, num_draft_tokens: int, is_draft_worker: bool
) -> int:
# Deprecated alias; remove together with the FIXME above.
warnings.warn(
"get_num_tokens_per_bs_for_target_verify is deprecated; use "
"get_num_tokens_per_req_for_target_verify instead.",
DeprecationWarning,
stacklevel=2,
)
return self.get_num_tokens_per_req_for_target_verify(
num_draft_tokens, is_draft_worker
)
def build_disagg_draft_input( def build_disagg_draft_input(
self, self,
batch: ScheduleBatch, batch: ScheduleBatch,
@@ -21,7 +21,7 @@ from .cuda_graph_decode_runner import (
# "prod_fill": mirrors `eagle_draft_extend_cuda_graph_runner.py:466-474` # "prod_fill": mirrors `eagle_draft_extend_cuda_graph_runner.py:466-474`
# (and similar in `multi_layer_eagle_draft_extend_cuda_graph_runner.py`): # (and similar in `multi_layer_eagle_draft_extend_cuda_graph_runner.py`):
# padded rows are pure scratch — `seq_lens[padded] = seq_len_fill_value`, # padded rows are pure scratch — `seq_lens[padded] = seq_len_fill_value`,
# `extend_seq_lens[padded] = num_tokens_per_bs`, `req_pool_indices[padded] = 0`, # `extend_seq_lens[padded] = num_tokens_per_req`, `req_pool_indices[padded] = 0`,
# `out_cache_loc[padded] = 0`, `positions[padded] = 0`. seq_lens and # `out_cache_loc[padded] = 0`, `positions[padded] = 0`. seq_lens and
# extend_seq_lens are intentionally inconsistent for padded rows (their # extend_seq_lens are intentionally inconsistent for padded rows (their
# subtraction goes negative), so backends must defend against that — the # subtraction goes negative), so backends must defend against that — the
@@ -55,7 +55,7 @@ class SpeculativeCudaGraphAdapter:
pad_style: PadStyle = "small_real" pad_style: PadStyle = "small_real"
# Required when pad_style == "prod_fill": draft tokens per request, # Required when pad_style == "prod_fill": draft tokens per request,
# used to fill the padded slots of extend_seq_lens / spec_info. # used to fill the padded slots of extend_seq_lens / spec_info.
pad_num_tokens_per_bs: Optional[int] = None pad_num_tokens_per_req: Optional[int] = None
def _apply_prod_fill_padding( def _apply_prod_fill_padding(
@@ -64,7 +64,7 @@ def _apply_prod_fill_padding(
real_bs: int, real_bs: int,
capture_bs: int, capture_bs: int,
seq_len_fill_value: int, seq_len_fill_value: int,
num_tokens_per_bs: int, num_tokens_per_req: int,
) -> None: ) -> None:
"""Overwrite padded slots of `batch` to match the production CG runner. """Overwrite padded slots of `batch` to match the production CG runner.
@@ -85,20 +85,20 @@ def _apply_prod_fill_padding(
batch.seq_lens_sum = int(batch.seq_lens_cpu.sum()) batch.seq_lens_sum = int(batch.seq_lens_cpu.sum())
if getattr(batch, "extend_seq_lens", None) is not None: if getattr(batch, "extend_seq_lens", None) is not None:
batch.extend_seq_lens[pad_lo:pad_hi] = num_tokens_per_bs batch.extend_seq_lens[pad_lo:pad_hi] = num_tokens_per_req
if getattr(batch, "extend_seq_lens_cpu", None) is not None: if getattr(batch, "extend_seq_lens_cpu", None) is not None:
ext = list(batch.extend_seq_lens_cpu) ext = list(batch.extend_seq_lens_cpu)
for i in range(pad_lo, min(pad_hi, len(ext))): for i in range(pad_lo, min(pad_hi, len(ext))):
ext[i] = num_tokens_per_bs ext[i] = num_tokens_per_req
batch.extend_seq_lens_cpu = ext batch.extend_seq_lens_cpu = ext
# Per-request slot tensors. # Per-request slot tensors.
batch.req_pool_indices[pad_lo:pad_hi] = 0 batch.req_pool_indices[pad_lo:pad_hi] = 0
# Per-token tensors: padded rows occupy slots # Per-token tensors: padded rows occupy slots
# [real_bs * num_tokens_per_bs, capture_bs * num_tokens_per_bs). # [real_bs * num_tokens_per_req, capture_bs * num_tokens_per_req).
tok_lo = pad_lo * num_tokens_per_bs tok_lo = pad_lo * num_tokens_per_req
tok_hi = pad_hi * num_tokens_per_bs tok_hi = pad_hi * num_tokens_per_req
for field in ("out_cache_loc", "positions", "input_ids"): for field in ("out_cache_loc", "positions", "input_ids"):
t = getattr(batch, field, None) t = getattr(batch, field, None)
if t is not None and t.numel() >= tok_hi: if t is not None and t.numel() >= tok_hi:
@@ -111,11 +111,11 @@ def _apply_prod_fill_padding(
if spec_info is not None: if spec_info is not None:
eslt = getattr(spec_info, "extend_seq_lens_tensor", None) eslt = getattr(spec_info, "extend_seq_lens_tensor", None)
if isinstance(eslt, torch.Tensor) and eslt.numel() >= pad_hi: if isinstance(eslt, torch.Tensor) and eslt.numel() >= pad_hi:
eslt[pad_lo:pad_hi] = num_tokens_per_bs eslt[pad_lo:pad_hi] = num_tokens_per_req
eslc = getattr(spec_info, "extend_seq_lens_cpu", None) eslc = getattr(spec_info, "extend_seq_lens_cpu", None)
if isinstance(eslc, list): if isinstance(eslc, list):
for i in range(pad_lo, min(pad_hi, len(eslc))): for i in range(pad_lo, min(pad_hi, len(eslc))):
eslc[i] = num_tokens_per_bs eslc[i] = num_tokens_per_req
def _check_speculative_cuda_graph_case( def _check_speculative_cuda_graph_case(
@@ -296,9 +296,9 @@ def run_speculative_cuda_graph_case(
and adapter.allow_padding and adapter.allow_padding
and real_bs < capture_batch_size and real_bs < capture_batch_size
): ):
if adapter.pad_num_tokens_per_bs is None: if adapter.pad_num_tokens_per_req is None:
raise ValueError( raise ValueError(
"SpeculativeCudaGraphAdapter.pad_num_tokens_per_bs must be set " "SpeculativeCudaGraphAdapter.pad_num_tokens_per_req must be set "
"when pad_style='prod_fill'." "when pad_style='prod_fill'."
) )
_apply_prod_fill_padding( _apply_prod_fill_padding(
@@ -306,7 +306,7 @@ def run_speculative_cuda_graph_case(
real_bs=real_bs, real_bs=real_bs,
capture_bs=capture_batch_size, capture_bs=capture_batch_size,
seq_len_fill_value=capture_prefix_len, seq_len_fill_value=capture_prefix_len,
num_tokens_per_bs=adapter.pad_num_tokens_per_bs, num_tokens_per_req=adapter.pad_num_tokens_per_req,
) )
with torch.no_grad(), forward_context(ForwardContext(attn_backend=backend)): with torch.no_grad(), forward_context(ForwardContext(attn_backend=backend)):
@@ -192,7 +192,7 @@ def _run_draft_extend_cuda_graph_case(
run_graph_eager: bool = True, run_graph_eager: bool = True,
compare_replay_to_graph_eager: bool = True, compare_replay_to_graph_eager: bool = True,
pad_style: str = "small_real", pad_style: str = "small_real",
pad_num_tokens_per_bs: int | None = None, pad_num_tokens_per_req: int | None = None,
): ):
adapter = SpeculativeCudaGraphAdapter( adapter = SpeculativeCudaGraphAdapter(
build_fixture=build_fixture, build_fixture=build_fixture,
@@ -217,7 +217,7 @@ def _run_draft_extend_cuda_graph_case(
atol=atol, atol=atol,
rtol=rtol, rtol=rtol,
pad_style=pad_style, pad_style=pad_style,
pad_num_tokens_per_bs=pad_num_tokens_per_bs, pad_num_tokens_per_req=pad_num_tokens_per_req,
) )
run_speculative_cuda_graph_case( run_speculative_cuda_graph_case(
testcase, testcase,
@@ -299,7 +299,7 @@ def run_dense_draft_extend_v2_cuda_graph_case(
run_graph_eager=False, run_graph_eager=False,
compare_replay_to_graph_eager=False, compare_replay_to_graph_eager=False,
pad_style=pad_style, pad_style=pad_style,
pad_num_tokens_per_bs=num_tokens_per_req, pad_num_tokens_per_req=num_tokens_per_req,
) )
@@ -367,7 +367,7 @@ def run_mla_draft_extend_v2_cuda_graph_case(
run_graph_eager=False, run_graph_eager=False,
compare_replay_to_graph_eager=False, compare_replay_to_graph_eager=False,
pad_style=pad_style, pad_style=pad_style,
pad_num_tokens_per_bs=num_tokens_per_req, pad_num_tokens_per_req=num_tokens_per_req,
) )
@@ -170,8 +170,8 @@ class TestTboAttnDenseAttentionBackendCorrectness(CustomTestCase):
) )
capture_bs = case.batch_size capture_bs = case.batch_size
num_tokens_per_bs = sum(case.extend_lens) // capture_bs num_tokens_per_req = sum(case.extend_lens) // capture_bs
num_tokens = capture_bs * num_tokens_per_bs num_tokens = capture_bs * num_tokens_per_req
split_seq_index, split_token_index = ( split_seq_index, split_token_index = (
compute_split_indices_for_cuda_graph_replay( compute_split_indices_for_cuda_graph_replay(
forward_mode=batch.forward_mode, forward_mode=batch.forward_mode,
@@ -59,7 +59,7 @@ class TestComputeLaunchCapacities(CustomTestCase):
) )
def test_from_args_treats_missing_speculative_draft_tokens_as_zero(self) -> None: def test_from_args_treats_missing_speculative_draft_tokens_as_zero(self) -> None:
"""per_forward_write_entry_capacity is floored by max_prefill_tokens when batch * tokens_per_bs is smaller.""" """per_forward_write_entry_capacity is floored by max_prefill_tokens when batch * tokens_per_req is smaller."""
server_args = self._make_server_args(max_bs=2) server_args = self._make_server_args(max_bs=2)
server_args.speculative_num_draft_tokens = None server_args.speculative_num_draft_tokens = None
@@ -1032,7 +1032,7 @@ class TestChunkedSGMV(unittest.TestCase):
backend = ChunkedSgmvLoRABackend( backend = ChunkedSgmvLoRABackend(
max_loras_per_batch=5, device=self.device, server_args=mock_server_args max_loras_per_batch=5, device=self.device, server_args=mock_server_args
) )
backend.init_cuda_graph_batch_info(max_bs_in_cuda_graph=8, num_tokens_per_bs=1) backend.init_cuda_graph_batch_info(max_bs_in_cuda_graph=8, num_tokens_per_req=1)
lora_ranks = [8] * 5 lora_ranks = [8] * 5
scalings = [1.0] * 5 scalings = [1.0] * 5
@@ -77,7 +77,7 @@ class TestEagleDraftCudaGraphRunner(CustomTestCase):
dsa_seed_topk=None, dsa_seed_topk=None,
) )
runner.capture_bs = [1, CAPTURE_BS] runner.capture_bs = [1, CAPTURE_BS]
runner.num_tokens_per_bs = 1 runner.num_tokens_per_req = 1
runner.speculative_num_steps = NUM_STEPS runner.speculative_num_steps = NUM_STEPS
runner.seq_len_fill_value = SEQ_LEN_FILL_VALUE runner.seq_len_fill_value = SEQ_LEN_FILL_VALUE
runner.require_mlp_tp_gather = False runner.require_mlp_tp_gather = False