[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
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:
tokens_per_bs = 1
tokens_per_req = 1
metadata.actual_seq_lengths_q_pa = torch.arange(
0,
bs * tokens_per_bs + tokens_per_bs,
tokens_per_bs,
bs * tokens_per_req + tokens_per_req,
tokens_per_req,
dtype=torch.int32,
device=device,
)
@@ -1081,7 +1081,7 @@ class DeepseekV4AscendAttnBackend(
:bs, :
]
n_tok = bs * tokens_per_bs
n_tok = bs * tokens_per_req
c4_pad = min(n_tok, n_tok // 4 + bs)
c128_pad = min(n_tok, n_tok // 128 + bs)
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"],
}
T = bs * tokens_per_bs
T = bs * tokens_per_req
metadata.c4_topk_indices = self.graph_metadata["c4_topk_indices"][:T, :]
self.forward_metadata = metadata
@@ -1126,16 +1126,16 @@ class DeepseekV4AscendAttnBackend(
device = seq_lens.device
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:
tokens_per_bs = 1
tokens_per_req = 1
seq_lens_cpu = forward_batch.seq_lens_cpu
assert seq_lens_cpu is not None, "V4 graph replay requires seq_lens_cpu."
if forward_mode.is_target_verify():
# In graph replay, buffers.seq_lens already contains the attention KV
# 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.
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":
@@ -1145,9 +1145,9 @@ class DeepseekV4AscendAttnBackend(
attn_seq_lens = live_seq_lens
if forward_mode.is_target_verify():
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)
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(
seq_lens_cpu[:bs] > 0,
fm.seq_lens_cpu_int,
@@ -1166,8 +1166,8 @@ class DeepseekV4AscendAttnBackend(
_compress_seq_lens = live_seq_lens
_compress_seq_lens_max = int(seq_lens_cpu[:bs].max()) if bs > 0 else 0
if _verify_compress:
_compress_seq_lens = live_seq_lens + int(tokens_per_bs)
_compress_seq_lens_max += int(tokens_per_bs)
_compress_seq_lens = live_seq_lens + int(tokens_per_req)
_compress_seq_lens_max += int(tokens_per_req)
result = self._compute_compress_locs(
pool=pool,
@@ -1218,7 +1218,7 @@ class DeepseekV4AscendAttnBackend(
_copy_1d(getattr(fm, key), result[key])
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(
seq_lens_cpu[:bs] > 0,
verify_seq_lens_cpu,
@@ -1229,19 +1229,19 @@ class DeepseekV4AscendAttnBackend(
fm.positions_cmp_padding_c4,
4,
verify_seq_lens_cpu,
n_draft=tokens_per_bs,
n_draft=tokens_per_req,
)
self._fill_verify_positions_cmp_padding_one(
forward_batch.positions,
fm.positions_cmp_padding_c128,
128,
verify_seq_lens_cpu,
n_draft=tokens_per_bs,
n_draft=tokens_per_req,
)
fm.start_pos.copy_(live_seq_lens.to(torch.int32))
valid = live_seq_lens[:bs] > 0
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)
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_kv=fm.actual_seq_lengths_kv,
block_tables=fm.block_tables,
max_seqlen_q=tokens_per_bs,
max_seqlen_q=tokens_per_req,
is_nextn=False,
)
for key in (
@@ -240,7 +240,7 @@ class NPUGraphRunner(DecodeCudaGraphRunner):
or is_deepseek_v4(self.model_runner.model_config.hf_config)
):
if forward_batch.forward_mode.is_target_verify():
seq_lens_cpu = forward_batch.seq_lens.cpu() + self.num_tokens_per_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)
else:
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}"
)
num_tokens_per_bs = 1
num_tokens_per_req = 1
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)
@@ -102,7 +102,7 @@ class CanaryLaunchCapacities:
max_extend_tokens_per_forward = min(max_prefill_tokens, chunked_limit)
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
@@ -1794,9 +1794,9 @@ class AiterAttnBackend(AttentionBackend):
# EAGLE V2: Fixed num_draft_tokens per batch
self._ensure_spec_v2_topk_supported()
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(
(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]
@@ -1815,7 +1815,7 @@ class AiterAttnBackend(AttentionBackend):
)
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:
num_kv_splits = self.max_split_per_batch
@@ -1043,14 +1043,14 @@ class DeepseekV4AttnBackend(
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_cpu: List[int],
num_tokens_per_bs: int,
num_tokens_per_req: int,
out_cache_loc: Optional[torch.Tensor] = None,
use_prefill_cuda_graph: bool = False,
) -> DSV4Metadata:
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)
num_tokens = num_tokens_per_bs * batch_size
num_tokens = num_tokens_per_req * batch_size
if out_cache_loc is None:
out_cache_loc = seq_lens.new_zeros(num_tokens)
return self.init_forward_metadata_prefill(
@@ -1287,13 +1287,13 @@ class DeepseekV4AttnBackend(
req_pool_indices,
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:
# Pad the real write locations to the captured token count so
# raw_out_loc reflects the actual replay out_cache_loc.
out_cache_loc = torch.nn.functional.pad(
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",
value=0,
)
@@ -1305,7 +1305,7 @@ class DeepseekV4AttnBackend(
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
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,
use_prefill_cuda_graph=True,
)
@@ -1477,7 +1477,7 @@ class DeepseekV4AttnBackend(
],
],
] = {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
)
@@ -727,14 +727,14 @@ class DeepseekV4HipRadixBackend(
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_cpu: List[int],
num_tokens_per_bs: int,
num_tokens_per_req: int,
out_cache_loc: Optional[torch.Tensor] = None,
use_prefill_cuda_graph: bool = False,
) -> DSV4Metadata:
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)
num_tokens = num_tokens_per_bs * batch_size
num_tokens = num_tokens_per_req * batch_size
if out_cache_loc is None:
out_cache_loc = seq_lens.new_zeros(num_tokens)
return self.init_forward_metadata_prefill(
@@ -889,13 +889,13 @@ class DeepseekV4HipRadixBackend(
seq_lens_cpu=seq_lens_cpu.tolist(),
)
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:
# Pad the real write locations to the captured token count so
# raw_out_loc reflects the actual replay out_cache_loc.
out_cache_loc = torch.nn.functional.pad(
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",
value=0,
)
@@ -904,7 +904,7 @@ class DeepseekV4HipRadixBackend(
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
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,
use_prefill_cuda_graph=True,
)
@@ -1012,7 +1012,7 @@ class DeepseekV4HipRadixBackend(
],
],
] = {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
)
@@ -282,7 +282,7 @@ class FlashAttentionBackend(AttentionBackend):
):
self.speculative_num_draft_tokens = SpeculativeAlgorithm.from_string(
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
)
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
# max(num_accept_tokens_cpu) which is None/empty at capture time,
# 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
else:
self._apply_cuda_graph_metadata(
@@ -2353,11 +2353,11 @@ class FlashAttentionBackend(AttentionBackend):
metadata.swa_spec_metadata = metadata_swa
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"][
: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_k = self.draft_extend_metadata["cu_seqlens_k"][
: (bs + 1)
@@ -507,7 +507,7 @@ class TritonAttnBackend(AttentionBackend):
seq_lens = seq_lens[:bs]
# V2 draft-extend fills num_draft_tokens per req; num_steps+1 only equals
# that when topk == 1.
num_tokens_per_bs = (
num_tokens_per_req = (
self.num_draft_tokens
if forward_mode.is_draft_extend_v2()
else self.speculative_num_steps + 1
@@ -515,8 +515,8 @@ class TritonAttnBackend(AttentionBackend):
qo_indptr = self.qo_indptr[: bs + 1]
qo_indptr[: bs + 1] = torch.arange(
0,
bs * num_tokens_per_bs + 1,
step=num_tokens_per_bs,
bs * num_tokens_per_req + 1,
step=num_tokens_per_req,
dtype=torch.int32,
device=self.device,
)
@@ -534,7 +534,7 @@ class TritonAttnBackend(AttentionBackend):
kv_indptr = self._fill_kv_indptr_and_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(
self,
@@ -1084,7 +1084,7 @@ class TritonAttnBackend(AttentionBackend):
return ForwardMetadata(
attn_logits=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
# for topk > 1 (num_draft_tokens > num_steps+1) and drops query
# blocks.
@@ -487,13 +487,13 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
)
self.target_verify_metadata[bs] = metadata
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"][
:bs
]
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.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, :]
self._bind_swa_page_table(
metadata,
@@ -567,9 +567,9 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
# Static per-request query width, fixed by the captured graph shape.
# Do not inspect replay-time tensors here; this body is recorded into
# 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
q_stride = num_tokens_per_bs
q_stride = num_tokens_per_req
q_mode = Q_MODE_STRIDED
else:
raise ValueError(
@@ -283,13 +283,13 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
self.decode_cuda_graph_kv_indices = torch.full(
(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):
# Buffer for padded query: (max_bs, max_draft_tokens, num_q_heads, v_head_dim)
self.store_dtype = torch.uint8
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,
device=self.device,
)
@@ -303,7 +303,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
else:
# Buffer for padded query: (max_bs, max_draft_tokens, num_q_heads, v_head_dim)
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,
device=self.device,
)
@@ -343,18 +343,18 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
if forward_mode.is_target_verify():
metadata.seq_lens_k = torch.zeros((bs,), dtype=torch.int32, device=device)
elif forward_mode.is_draft_extend_v2():
num_tokens_per_bs = self.num_draft_tokens
metadata.max_seq_len_q = num_tokens_per_bs
metadata.sum_seq_lens_q = num_tokens_per_bs * bs
num_tokens_per_req = self.num_draft_tokens
metadata.max_seq_len_q = num_tokens_per_req
metadata.sum_seq_lens_q = num_tokens_per_req * bs
metadata.cu_seqlens_q = torch.arange(
0,
bs * num_tokens_per_bs + 1,
num_tokens_per_bs,
bs * num_tokens_per_req + 1,
num_tokens_per_req,
dtype=torch.int32,
device=device,
)
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)
@@ -385,9 +385,9 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
seq_lens = seq_lens[:bs] + self.num_draft_tokens
metadata.seq_lens_k.copy_(seq_lens)
elif forward_mode.is_draft_extend_v2():
num_tokens_per_bs = self.num_draft_tokens
metadata.max_seq_len_q = num_tokens_per_bs
metadata.sum_seq_lens_q = num_tokens_per_bs * bs
num_tokens_per_req = self.num_draft_tokens
metadata.max_seq_len_q = num_tokens_per_req
metadata.sum_seq_lens_q = num_tokens_per_req * bs
seq_lens = seq_lens[:bs]
metadata.seq_lens_k.copy_(seq_lens)
@@ -204,7 +204,7 @@ class AscendLoRABackend(BaseLoRABackend):
def init_cuda_graph_batch_info(
self,
max_bs_in_cuda_graph: int,
num_tokens_per_bs: int,
num_tokens_per_req: int,
):
with torch.device("npu"):
self.npu_graph_batch_info = LoRABatchInfo(
@@ -212,10 +212,10 @@ class AscendLoRABackend(BaseLoRABackend):
use_cuda_graph=True,
num_segments=None,
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),
max_len=num_tokens_per_bs,
max_len=num_tokens_per_req,
weight_indices=torch.zeros(max_bs_in_cuda_graph, 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),
@@ -149,7 +149,7 @@ class BaseLoRABackend(LoRABackendLmHeadMixing):
def init_cuda_graph_batch_info(
self,
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.
@@ -157,7 +157,7 @@ class BaseLoRABackend(LoRABackendLmHeadMixing):
Args:
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
@@ -218,12 +218,12 @@ class ChunkedSgmvLoRABackend(BaseLoRABackend):
def init_cuda_graph_batch_info(
self,
max_bs_in_cuda_graph: int,
num_tokens_per_bs: int,
num_tokens_per_req: int,
):
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_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"):
self.cuda_graph_batch_info = LoRABatchInfo(
bs=max_bs_in_cuda_graph,
@@ -164,7 +164,7 @@ class TorchNativeLoRABackend(BaseLoRABackend):
def init_cuda_graph_batch_info(
self,
max_bs_in_cuda_graph: int,
num_tokens_per_bs: int,
num_tokens_per_req: int,
):
with torch.device("cuda"):
self.cuda_graph_batch_info = TorchNativeLoRABatchInfo(
@@ -172,14 +172,14 @@ class TorchNativeLoRABackend(BaseLoRABackend):
bs=max_bs_in_cuda_graph,
num_segments=self.max_loras_per_batch,
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),
weight_indices=torch.zeros(max_bs_in_cuda_graph, 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),
permutation=None,
max_len=num_tokens_per_bs,
max_len=num_tokens_per_req,
)
# Initialize seg_indptr for CUDA graph as they remain constant
@@ -140,9 +140,9 @@ class TritonLoRABackend(BaseLoRABackend):
def init_cuda_graph_batch_info(
self,
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
with torch.device("cuda"):
self.cuda_graph_batch_info = LoRABatchInfo(
@@ -150,10 +150,10 @@ class TritonLoRABackend(BaseLoRABackend):
use_cuda_graph=True,
num_segments=None,
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),
max_len=num_tokens_per_bs,
max_len=num_tokens_per_req,
weight_indices=torch.zeros(max_bs_in_cuda_graph, dtype=torch.int32),
lora_ranks=torch.zeros(mlpb, dtype=torch.int32),
scalings=torch.zeros(mlpb, dtype=torch.float),
+2 -2
View File
@@ -112,7 +112,7 @@ class LoRAManager:
)
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.
@@ -122,7 +122,7 @@ class LoRAManager:
self.max_bs_in_cuda_graph = max_bs_in_cuda_graph
self.lora_backend.init_cuda_graph_batch_info(
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 ====
@@ -581,7 +581,7 @@ class CPUGraphRunner:
self.capture_forward_mode = ForwardMode.DECODE
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 self.enable_return_hidden_states:
@@ -618,7 +618,7 @@ class CPUGraphRunner:
self.captured_forward_batches_cross = {}
# Attention backend
self.max_bs = max(self.capture_bs)
self.max_num_token = self.max_bs * self.num_tokens_per_bs
self.max_num_token = self.max_bs * self.num_tokens_per_req
self.model_runner.attn_backend.init_cpu_graph_state(
self.max_bs, self.max_num_token
)
@@ -646,7 +646,7 @@ class CPUGraphRunner:
self.custom_mask = torch.ones(
(
(self.seq_lens.sum().item() + self.max_num_token)
* self.num_tokens_per_bs
* self.num_tokens_per_req
),
dtype=torch.bool,
device=self.device,
@@ -725,7 +725,7 @@ class CPUGraphRunner:
with patch_model(
self.model_runner.model,
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,
) as forward:
graph, output_buffers = self.capture_one_batch_size(
@@ -767,7 +767,7 @@ class CPUGraphRunner:
def capture_one_batch_size(
self, bs: int, forward: Callable, skip_cross_attention: bool = False
):
num_tokens = bs * self.num_tokens_per_bs
num_tokens = bs * self.num_tokens_per_req
# Graph inputs
input_ids = self.input_ids[:num_tokens]
@@ -916,7 +916,7 @@ class CPUGraphRunner:
self.model_runner.attn_backend.init_forward_metadata(forward_batch)
return forward_batch
raw_num_token = raw_bs * self.num_tokens_per_bs
raw_num_token = raw_bs * self.num_tokens_per_req
index = bisect.bisect_left(self.capture_bs, raw_bs)
bs = self.capture_bs[index]
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
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.
"""
@@ -842,14 +842,14 @@ class ModelRunner(ModelRunnerKVCacheMixin):
return 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
) -> int:
"""Logits rows per decode batch slot."""
if self.spec_algorithm.is_speculative():
if num_draft_tokens is None:
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
)
dllm_config = DllmConfig.from_server_args(self.server_args)
@@ -857,9 +857,9 @@ class ModelRunner(ModelRunnerKVCacheMixin):
def max_decode_logits_rows(self) -> int:
"""Rows the shared logits buffer needs."""
num_tokens_per_bs = self.decode_num_tokens_per_bs()
capture_bs, _ = get_batch_sizes_to_capture(self, num_tokens_per_bs)
return max(capture_bs) * 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_req)
return max(capture_bs) * num_tokens_per_req
def alloc_memory_pool(self, memory_pool_config: Optional[MemoryPoolConfig] = None):
"""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
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.
"""
from sglang.srt.lora.layers import FusedMoEWithLoRA
@@ -2640,20 +2640,20 @@ class ModelRunner(ModelRunnerKVCacheMixin):
role = "draft" if self.is_draft_worker else "target"
if self.spec_algorithm.is_speculative():
capture_name = f"{role} verify"
num_tokens_per_bs = (
self.spec_algorithm.get_num_tokens_per_bs_for_target_verify(
num_tokens_per_req = (
self.spec_algorithm.get_num_tokens_per_req_for_target_verify(
self.server_args.speculative_num_draft_tokens,
self.is_draft_worker,
)
)
else:
capture_name = f"{role} decode"
num_tokens_per_bs = 1
capture_bs, _ = get_batch_sizes_to_capture(self, num_tokens_per_bs)
num_tokens_per_req = 1
capture_bs, _ = get_batch_sizes_to_capture(self, num_tokens_per_req)
decode_backend = self.server_args.cuda_graph_config.decode.backend
logger.info(
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"
)
@@ -56,7 +56,7 @@ def freeze_gc(enable_cudagraph_gc: bool):
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]]:
"""Build the (capture_bs, compile_bs) lists for the decode runner.
@@ -71,7 +71,7 @@ def get_batch_sizes_to_capture(
mul_base = 1
if server_args.enable_two_batch_overlap:
mul_base *= 2
num_tokens_per_bs = 1
num_tokens_per_req = 1
if require_gathered_buffer(server_args):
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.
capture_bs += [num_max_requests]
# Model input token count = bs * num_tokens_per_bs; 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]
# Model input token count = bs * num_tokens_per_req; must be a multiple of attn_tp_size.
capture_bs = [bs for bs in capture_bs if bs * num_tokens_per_req % mul_base == 0]
capture_bs = [bs for bs in capture_bs if bs <= num_max_requests]
capture_bs = list(sorted(set(capture_bs)))
@@ -74,7 +74,7 @@ def _allocate_decode_buffers(
require_mlp_tp_gather: bool,
seq_len_fill_value: int,
encoder_len_fill_value: int,
num_tokens_per_bs: int,
num_tokens_per_req: int,
cache_loc_dtype: torch.dtype,
enable_mamba_track: bool,
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)
num_token_non_padded = torch.zeros((1,), dtype=torch.int32)
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,
)
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)
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
(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
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(
device=mr.device,
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,
vocab_size=mr.model_config.vocab_size,
dtype=mr.model_config.dtype,
@@ -309,7 +309,7 @@ class BaseRunner(ABC):
if mr.model_config.is_encoder_decoder
else 0
),
num_tokens_per_bs=num_tokens_per_bs,
num_tokens_per_req=num_tokens_per_req,
cache_loc_dtype=torch.int64,
enable_mamba_track=False,
ne_token_table=mr.token_table if mr.use_ngram_embedding else None,
@@ -350,14 +350,14 @@ class BaseRunner(ABC):
else:
capture_forward_mode = ForwardMode.EXTEND
capture_hidden_mode = CaptureHiddenMode.NULL
num_tokens_per_bs = 1
num_tokens_per_req = 1
if mr.spec_algorithm.is_speculative():
if mr.is_draft_worker:
if not mr.spec_algorithm.supports_target_verify_for_draft():
raise RuntimeError("This should not happen")
capture_forward_mode = ForwardMode.TARGET_VERIFY
num_tokens_per_bs = (
mr.spec_algorithm.get_num_tokens_per_bs_for_target_verify(
num_tokens_per_req = (
mr.spec_algorithm.get_num_tokens_per_req_for_target_verify(
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:
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
# allocation, no re-padding (would overflow the reused buffers).
@@ -439,7 +439,7 @@ class BaseRunner(ABC):
(batch_size,), dtype=torch.int32, device=mr.device
)
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:
extend_prefix_lens_cpu = None
@@ -484,7 +484,7 @@ class BaseRunner(ABC):
mr.spec_algorithm,
mr.server_args,
buffers.custom_mask,
num_tokens_per_bs,
num_tokens_per_req,
mr.is_draft_worker,
)
if spec_info is not None and (
@@ -254,7 +254,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
# --- capture mode + tokens-per-bs ------------------------------
self.capture_forward_mode = ForwardMode.DECODE
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
)
if model_runner.spec_algorithm.is_speculative():
@@ -270,7 +270,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
# --- bucket sizes ---------------------------------------------
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:
KTMoEWrapper.set_capture_batch_sizes(self.capture_bs)
@@ -304,7 +304,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
# Attention backend
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)
# Init PDMux if needed
@@ -331,7 +331,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
# lora_manager.init_cuda_graph_moe_buffers().
self.model_runner.lora_manager.init_cuda_graph_batch_info(
max_bs_in_cuda_graph=self.max_bs,
num_tokens_per_bs=self.num_tokens_per_bs,
num_tokens_per_req=self.num_tokens_per_req,
)
enable_mamba_track = (
@@ -358,7 +358,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
require_mlp_tp_gather=self.require_mlp_tp_gather,
seq_len_fill_value=self.seq_len_fill_value,
encoder_len_fill_value=self.encoder_len_fill_value,
num_tokens_per_bs=self.num_tokens_per_bs,
num_tokens_per_req=self.num_tokens_per_req,
cache_loc_dtype=self._cache_loc_dtype(),
enable_mamba_track=enable_mamba_track,
ne_token_table=(
@@ -404,7 +404,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
)
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=}"
return buckets
@@ -468,7 +468,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
def _ragged_capture_slots(self, num_tokens: int) -> int:
if envs.SGLANG_TEST_RAGGED_VERIFY_FORCE_UNIFORM_CAPTURE.get():
return num_tokens // self.num_tokens_per_bs
return num_tokens // self.num_tokens_per_req
return min(num_tokens, self.max_bs)
def _capture_ragged_verify_layout(self, num_tokens: int):
@@ -484,7 +484,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
verify_lens_cpu = build_capture_verify_lens(
num_tokens=num_tokens,
num_slots=self._ragged_capture_slots(num_tokens),
num_draft_tokens=self.num_tokens_per_bs,
num_draft_tokens=self.num_tokens_per_req,
)
return RaggedVerifyLayout.from_verify_lens(
verify_lens_cpu=verify_lens_cpu,
@@ -509,7 +509,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
if self.require_mlp_tp_gather:
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()
or self.model_runner.spec_algorithm.is_standalone()
or self.model_runner.spec_algorithm.is_dflash_family()
@@ -559,7 +559,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
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()
)
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
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.
Returns (forward_batch, attn_backend, pp_proxy_tensors);
@@ -667,7 +667,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
bs = size
buffers: DecodeInputBuffers = self.buffers
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
# 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):
self.enable_torch_compile = False
_, 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()
if self.enable_profile_cuda_graph:
@@ -889,7 +889,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
with torch_compile_decoration.patch_model(
self.model_runner.model,
bs in self.compile_bs,
num_tokens=bs * self.num_tokens_per_bs,
num_tokens=bs * self.num_tokens_per_req,
tp_group=self.model_runner.tp_group,
) as forward:
self.capture_one_shape(bs, forward, stream_idx, variant_label)
@@ -901,7 +901,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
stream_idx: Optional[int] = 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
# Sanity-check: --debug-cuda-graph requires breakable backend.
@@ -1060,7 +1060,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
self._ragged_graph_size
if is_ragged
else self._capture_graph_size(
bs=self.bs, num_tokens=self.bs * self.num_tokens_per_bs
bs=self.bs, num_tokens=self.bs * self.num_tokens_per_req
)
)
if is_ragged:
@@ -1105,11 +1105,11 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
)
padded_num_tokens = graph_size_key
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:
max_num_tokens = max(forward_batch.global_num_tokens_cpu)
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()
or self.model_runner.spec_algorithm.is_standalone()
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)
else:
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(
bs=bs, num_tokens=padded_num_tokens
)
@@ -1326,7 +1326,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
spec_info = DFlashVerifyInput(
draft_token=None,
positions=None,
draft_token_num=self.num_tokens_per_bs,
draft_token_num=self.num_tokens_per_req,
custom_mask=(
None
if (self.model_runner.is_draft_worker or not build_custom_mask)
@@ -1350,7 +1350,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
retrieve_index=None,
retrieve_next_token=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
@@ -68,12 +68,12 @@ class EagerRunner(BaseRunner):
sa = mr.server_args
# 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.
num_tokens_per_bs = 1
num_tokens_per_req = 1
if mr.spec_algorithm.is_speculative():
# speculative_adaptive can grow draft tokens at runtime; size to the max.
num_draft_tokens = sa.max_speculative_num_draft_tokens or 1
if mr.is_draft_worker:
num_tokens_per_bs = max(
num_tokens_per_req = max(
sa.speculative_eagle_topk or 1,
num_draft_tokens,
(
@@ -83,8 +83,8 @@ class EagerRunner(BaseRunner):
),
)
else:
num_tokens_per_bs = (
mr.spec_algorithm.get_num_tokens_per_bs_for_target_verify(
num_tokens_per_req = (
mr.spec_algorithm.get_num_tokens_per_req_for_target_verify(
num_draft_tokens, mr.is_draft_worker
)
)
@@ -92,7 +92,7 @@ class EagerRunner(BaseRunner):
dllm_config = DllmConfig.from_server_args(sa)
if dllm_config is not None:
# 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
if (
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, get_cp_padding_align_size())
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):
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())
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
self._eager_registry = build_eager_registry(
device=mr.device,
@@ -139,7 +139,7 @@ class EagerRunner(BaseRunner):
self.warmup()
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'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.
"""
mr = self.model_runner
num_tokens_per_bs = 1
num_tokens_per_req = 1
if mr.spec_algorithm.is_speculative():
num_tokens_per_bs = (
mr.spec_algorithm.get_num_tokens_per_bs_for_target_verify(
num_tokens_per_req = (
mr.spec_algorithm.get_num_tokens_per_req_for_target_verify(
mr.server_args.speculative_num_draft_tokens, mr.is_draft_worker
)
)
return (
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,
)
@@ -99,7 +99,7 @@ class DecodeInputBuffers(ForwardInputBuffers):
require_mlp_tp_gather: bool,
seq_len_fill_value: int,
encoder_len_fill_value: int,
num_tokens_per_bs: int,
num_tokens_per_req: int,
cache_loc_dtype: torch.dtype,
enable_mamba_track: bool,
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)
num_token_non_padded = torch.zeros((1,), dtype=torch.int32)
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,
)
mamba_track_indices = (
@@ -220,7 +220,7 @@ class DecodeInputBuffers(ForwardInputBuffers):
bs: int,
seq_len_fill_value: int,
require_gathered_buffer: bool,
num_tokens_per_bs: int,
num_tokens_per_req: int,
dsa_enable_prefill_cp: bool,
enable_num_token_non_padded_flag: bool,
pp_proxy_tensors: Optional[PPProxyTensors] = None,
@@ -290,12 +290,12 @@ class DecodeInputBuffers(ForwardInputBuffers):
srcs.append(forward_batch.bootstrap_room_ids_int)
if require_gathered_buffer:
self.global_num_tokens_gpu.fill_(bs * num_tokens_per_bs)
self.global_num_tokens_for_logprob_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_req)
if enable_num_token_non_padded_flag:
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(
global_num_token_non_padded=forward_batch.num_token_non_padded,
num_tokens_per_dp=num_tokens_per_dp,
+6 -6
View File
@@ -2926,7 +2926,7 @@ class ServerArgs:
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()
# Handle model loading format.
@@ -5419,19 +5419,19 @@ class ServerArgs:
MoE layer on one (DP) rank. Single source of truth for both the
standard-allgather wrapper buffers and the FlashInfer A2A dispatcher
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.
"""
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:
num_tokens_per_bs = 1
num_tokens_per_req = 1
prefill_tokens = self.max_prefill_tokens
cg_config = self.cuda_graph_config
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)
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)
def max_prefill_buffer_tokens(self) -> int:
@@ -5452,7 +5452,7 @@ class ServerArgs:
def _validate_cutedsl_a2a_token_budget(self):
"""Fail fast if the FlashInfer A2A dispatcher workspace cannot cover the
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
view = resolved_view(self)
@@ -145,9 +145,9 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
# Bucket sizes
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_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
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):
if self.require_mlp_tp_gather:
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()
or self.model_runner.spec_algorithm.is_standalone()
else max(forward_batch.global_num_tokens_cpu)
@@ -321,7 +321,7 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
):
num_seqs = size # EAGLE legacy name
buffers = self.buffers
num_tokens = num_seqs * self.num_tokens_per_bs
num_tokens = num_seqs * self.num_tokens_per_req
# Graph inputs
req_pool_indices = buffers.req_pool_indices[:num_seqs]
@@ -490,13 +490,13 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
buffers = self.buffers
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
if self.require_mlp_tp_gather:
max_num_tokens = max(forward_batch.global_num_tokens_cpu)
max_batch_size = (
max_num_tokens // self.num_tokens_per_bs
max_num_tokens // self.num_tokens_per_req
if self.model_runner.spec_algorithm.is_eagle()
or self.model_runner.spec_algorithm.is_standalone()
else max_num_tokens
@@ -523,7 +523,7 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
buffers.dsa_seed_topk.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(
forward_batch.spec_info.topk_p,
@@ -598,8 +598,10 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
# TODO(ch-wan): support num_token_non_padded
if self.require_gathered_buffer:
buffers.global_num_tokens_gpu.fill_(bs * self.num_tokens_per_bs)
buffers.global_num_tokens_for_logprob_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_req
)
# 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
@@ -134,9 +134,9 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
# Size cuda-graph buffers by num_draft_tokens (full tree width), not
# 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_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.max_bs, self.max_num_token
@@ -144,7 +144,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
self.seq_len_fill_value = (
self.draft_extend_attn_backend.get_cuda_graph_seq_len_fill_value()
)
self.extend_seq_lens_cpu = [self.num_tokens_per_bs] * self.max_bs
self.extend_seq_lens_cpu = [self.num_tokens_per_req] * self.max_bs
if self.enable_torch_compile:
set_torch_compile_config()
@@ -182,13 +182,13 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
(self.max_bs,), self.seq_len_fill_value, dtype=torch.int64
)
extend_seq_lens = torch.full(
(self.max_bs,), self.num_tokens_per_bs, dtype=torch.int32
(self.max_bs,), self.num_tokens_per_req, dtype=torch.int32
)
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(
(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:
@@ -227,7 +227,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
next_token_logits_buffer = (
self.model_runner.graph_shared_output.get_logits_buffer(
vocab_size, rows=self.max_bs * self.num_tokens_per_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):
if self.require_mlp_tp_gather:
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()
or self.model_runner.spec_algorithm.is_standalone()
else max(forward_batch.global_num_tokens_cpu)
@@ -315,7 +315,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
):
bs = size
buffers = self.buffers
num_tokens = bs * self.num_tokens_per_bs
num_tokens = bs * self.num_tokens_per_req
# Graph inputs
input_ids = buffers.input_ids[:num_tokens]
@@ -370,7 +370,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
num_correct_drafts=num_correct_drafts,
num_accept_tokens=num_accept_tokens,
# Padded tree width per req; drives the constant qo layout.
num_tokens_per_req=self.num_tokens_per_bs,
num_tokens_per_req=self.num_tokens_per_req,
)
forward_batch = ForwardBatch(
@@ -482,7 +482,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
if self.require_mlp_tp_gather:
max_num_tokens = max(forward_batch.global_num_tokens_cpu)
max_batch_size = (
max_num_tokens // self.num_tokens_per_bs
max_num_tokens // self.num_tokens_per_req
if self.model_runner.spec_algorithm.is_eagle()
else max_num_tokens
)
@@ -490,16 +490,16 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
else:
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.out_cache_loc.zero_()
buffers.positions.zero_()
# Pair with seq_lens fill: padded rows must point at reserved
# req_pool slot 0 (req_to_token[0, :] is all zeros from init).
buffers.req_pool_indices.zero_()
buffers.num_correct_drafts.fill_(self.num_tokens_per_bs)
buffers.num_accept_tokens.fill_(self.num_tokens_per_bs)
buffers.extend_seq_lens.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_req)
buffers.extend_seq_lens.fill_(self.num_tokens_per_req)
# Batch the small per-field device copies into a grouped foreach copy
# (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_srcs.append(forward_batch.extend_seq_lens)
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:
copy_dsts.append(buffers.num_correct_drafts[:raw_bs])
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
if self.require_gathered_buffer:
buffers.global_num_tokens_gpu.fill_(bs * self.num_tokens_per_bs)
buffers.global_num_tokens_for_logprob_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_req
)
if forward_batch.seq_lens_cpu is not None:
if bs != raw_bs:
@@ -556,9 +558,9 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
if forward_batch.extend_seq_lens_cpu is not None:
self.extend_seq_lens_cpu[:raw_bs] = forward_batch.extend_seq_lens_cpu
else:
self.extend_seq_lens_cpu[:raw_bs] = [self.num_tokens_per_bs] * raw_bs
self.extend_seq_lens_cpu[:raw_bs] = [self.num_tokens_per_req] * 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
)
forward_batch.spec_info.extend_seq_lens_cpu = list(
@@ -424,7 +424,7 @@ class EagleDraftWorker(EagleDraftWorkerBase):
log_info_on_rank0(
logger,
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",
)
self.cuda_graph_runner = Device2DraftCudaGraphRunner[
@@ -501,7 +501,7 @@ class EagleDraftWorker(EagleDraftWorkerBase):
log_info_on_rank0(
logger,
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",
)
self.cuda_graph_runner_for_draft_extend = Device2ExtendCudaGraphRunner[
@@ -113,12 +113,12 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
self.capture_forward_mode = ForwardMode.DECODE
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(
model_runner, self.num_tokens_per_bs
model_runner, self.num_tokens_per_req
)
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.seq_len_fill_value = (
@@ -227,7 +227,7 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
del forward, stream_idx, variant_label
buffers = self.buffers
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]
positions = buffers.positions[:expanded_bs]
@@ -362,7 +362,7 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
raw_expanded_bs = forward_batch.batch_size
raw_bs = (
raw_expanded_bs // self.num_tokens_per_bs
raw_expanded_bs // self.num_tokens_per_req
if self.topk > 1
else raw_expanded_bs
)
@@ -371,13 +371,13 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
if self.require_mlp_tp_gather:
max_num_tokens = max(forward_batch.global_num_tokens_cpu)
max_batch_size = max_num_tokens // (
self.num_tokens_per_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)
else:
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:
buffers.seq_lens.fill_(self.seq_len_fill_value)
buffers.positions.zero_()
@@ -160,10 +160,10 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
# Fixed window: every step extends each request by the same number of
# tokens, which lets all steps share one buffer set.
self.num_tokens_per_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_num_token = self.max_bs * self.num_tokens_per_bs
self.extend_seq_lens_cpu = [self.num_tokens_per_bs] * self.max_bs
self.max_num_token = self.max_bs * self.num_tokens_per_req
self.extend_seq_lens_cpu = [self.num_tokens_per_req] * self.max_bs
self.eagle_worker.draft_extend_attn_backend_list[
self.step
@@ -198,7 +198,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
def can_run_graph(self, forward_batch: ForwardBatch):
if self.require_mlp_tp_gather:
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()
else max(forward_batch.global_num_tokens_cpu)
)
@@ -218,7 +218,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
def get_forward_batch(self, bs: int) -> ForwardBatch:
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]
req_pool_indices = buffers.req_pool_indices[:bs]
@@ -303,8 +303,8 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
extend_seq_lens_cpu=extend_seq_lens_cpu,
padded_static_len=self.padded_static_len,
extend_start_loc=extend_start_loc,
extend_num_tokens=self.num_tokens_per_bs * bs,
num_token_non_padded_cpu=self.num_tokens_per_bs * bs,
extend_num_tokens=self.num_tokens_per_req * bs,
num_token_non_padded_cpu=self.num_tokens_per_req * bs,
return_hidden_states_before_norm=True,
)
return forward_batch
@@ -332,7 +332,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
bs = size
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._postprocess_forward_batch(forward_batch, bs)
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)."""
self.deepep_adapter.replay()
buffers = self.buffers
num_tokens = bs * self.num_tokens_per_bs
num_tokens = bs * self.num_tokens_per_req
if self.require_gathered_buffer:
buffers.global_num_tokens_gpu.fill_(num_tokens)
@@ -453,7 +453,7 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner:
self.runners: List[Optional[MultiLayerEagleDraftExtendCudaGraphRunner]] = []
self.seq_len_fill_value = 1
self.max_bs = 1
self.num_tokens_per_bs = 1
self.num_tokens_per_req = 1
self._init_and_capture()
@@ -487,7 +487,7 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner:
self.runners.append(runner)
self.seq_len_fill_value = runner.seq_len_fill_value
self.max_bs = runner.max_bs
self.num_tokens_per_bs = runner.num_tokens_per_bs
self.num_tokens_per_req = runner.num_tokens_per_req
self.capture_bs = runner.capture_bs
self.require_gathered_buffer = runner.require_gathered_buffer
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)
model_runner = runner.model_runner
max_bs = self.max_bs
num_tokens_per_bs = self.num_tokens_per_bs
max_num_token = max_bs * num_tokens_per_bs
num_tokens_per_req = self.num_tokens_per_req
max_num_token = max_bs * num_tokens_per_req
hidden_size = get_draft_input_from_target_hidden_dim(model_runner)
dtype = model_runner.model_config.dtype
vocab_size = self._vocab_size()
@@ -553,13 +553,13 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner:
num_correct_drafts = 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.
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(
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)
@@ -610,7 +610,7 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner:
the batch size. Subsequent ``replay(step)`` calls reuse this state."""
buffers = self.buffers
raw_bs = forward_batch.batch_size
num_tokens = raw_bs * self.num_tokens_per_bs
num_tokens = raw_bs * self.num_tokens_per_req
# Bucketize to a captured batch size (padding the tail).
if self.require_mlp_tp_gather:
@@ -656,21 +656,23 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner:
# and by the worker's rotation.
arange = torch.arange(bs, device=self.device, dtype=torch.int64)
buffers.select_index[:bs].copy_(
arange * self.num_tokens_per_bs + buffers.num_correct_drafts[:bs]
arange * self.num_tokens_per_req + buffers.num_correct_drafts[:bs]
)
if self.require_gathered_buffer:
buffers.global_num_tokens_gpu.fill_(bs * self.num_tokens_per_bs)
buffers.global_num_tokens_for_logprob_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_req
)
# 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(
hidden_states=buffers.hidden_states[:padded_num_tokens],
num_correct_drafts=buffers.num_correct_drafts[: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.positions = buffers.positions[:padded_num_tokens]
spec_info.extend_seq_lens_tensor = buffers.extend_seq_lens[:bs]
+18 -3
View File
@@ -1,5 +1,6 @@
from __future__ import annotations
import warnings
from abc import ABC, abstractmethod
from enum import Enum, IntEnum, auto
from typing import TYPE_CHECKING, Callable, List, Optional, Tuple, Type, Union
@@ -210,7 +211,7 @@ class SpeculativeAlgorithm(Enum):
elif self.is_ngram():
_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
) -> int:
# 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
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(
self, server_args: ServerArgs
) -> Optional[Union[Type[BaseSpecWorker], Type[TpModelWorker], Type[NGRAMWorker]]]:
@@ -343,7 +358,7 @@ def create_dummy_verify_input(
spec_algorithm: SpeculativeAlgorithm,
server_args: ServerArgs,
custom_mask: torch.Tensor,
num_tokens_per_bs: int,
num_tokens_per_req: int,
is_draft_worker: bool,
) -> Optional[SpecInput]:
"""Dummy verify ``SpecInput`` for CUDA-graph capture (per-algorithm dispatch)."""
@@ -395,7 +410,7 @@ def create_dummy_verify_input(
retrieve_index=None,
retrieve_next_token=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
+16 -1
View File
@@ -5,6 +5,7 @@ should use that classmethod API; do not import from this module directly.
from __future__ import annotations
import logging
import warnings
from typing import TYPE_CHECKING, Callable, Dict, Optional, Type
import torch
@@ -119,7 +120,7 @@ class CustomSpecAlgo:
)
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
) -> int:
# 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.
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(
self,
batch: ScheduleBatch,
@@ -21,7 +21,7 @@ from .cuda_graph_decode_runner import (
# "prod_fill": mirrors `eagle_draft_extend_cuda_graph_runner.py:466-474`
# (and similar in `multi_layer_eagle_draft_extend_cuda_graph_runner.py`):
# 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
# extend_seq_lens are intentionally inconsistent for padded rows (their
# subtraction goes negative), so backends must defend against that — the
@@ -55,7 +55,7 @@ class SpeculativeCudaGraphAdapter:
pad_style: PadStyle = "small_real"
# Required when pad_style == "prod_fill": draft tokens per request,
# 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(
@@ -64,7 +64,7 @@ def _apply_prod_fill_padding(
real_bs: int,
capture_bs: int,
seq_len_fill_value: int,
num_tokens_per_bs: int,
num_tokens_per_req: int,
) -> None:
"""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())
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:
ext = list(batch.extend_seq_lens_cpu)
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
# Per-request slot tensors.
batch.req_pool_indices[pad_lo:pad_hi] = 0
# Per-token tensors: padded rows occupy slots
# [real_bs * num_tokens_per_bs, capture_bs * num_tokens_per_bs).
tok_lo = pad_lo * num_tokens_per_bs
tok_hi = pad_hi * num_tokens_per_bs
# [real_bs * num_tokens_per_req, capture_bs * num_tokens_per_req).
tok_lo = pad_lo * num_tokens_per_req
tok_hi = pad_hi * num_tokens_per_req
for field in ("out_cache_loc", "positions", "input_ids"):
t = getattr(batch, field, None)
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:
eslt = getattr(spec_info, "extend_seq_lens_tensor", None)
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)
if isinstance(eslc, list):
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(
@@ -296,9 +296,9 @@ def run_speculative_cuda_graph_case(
and adapter.allow_padding
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(
"SpeculativeCudaGraphAdapter.pad_num_tokens_per_bs must be set "
"SpeculativeCudaGraphAdapter.pad_num_tokens_per_req must be set "
"when pad_style='prod_fill'."
)
_apply_prod_fill_padding(
@@ -306,7 +306,7 @@ def run_speculative_cuda_graph_case(
real_bs=real_bs,
capture_bs=capture_batch_size,
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)):
@@ -192,7 +192,7 @@ def _run_draft_extend_cuda_graph_case(
run_graph_eager: bool = True,
compare_replay_to_graph_eager: bool = True,
pad_style: str = "small_real",
pad_num_tokens_per_bs: int | None = None,
pad_num_tokens_per_req: int | None = None,
):
adapter = SpeculativeCudaGraphAdapter(
build_fixture=build_fixture,
@@ -217,7 +217,7 @@ def _run_draft_extend_cuda_graph_case(
atol=atol,
rtol=rtol,
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(
testcase,
@@ -299,7 +299,7 @@ def run_dense_draft_extend_v2_cuda_graph_case(
run_graph_eager=False,
compare_replay_to_graph_eager=False,
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,
compare_replay_to_graph_eager=False,
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
num_tokens_per_bs = sum(case.extend_lens) // capture_bs
num_tokens = capture_bs * num_tokens_per_bs
num_tokens_per_req = sum(case.extend_lens) // capture_bs
num_tokens = capture_bs * num_tokens_per_req
split_seq_index, split_token_index = (
compute_split_indices_for_cuda_graph_replay(
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:
"""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.speculative_num_draft_tokens = None
@@ -1032,7 +1032,7 @@ class TestChunkedSGMV(unittest.TestCase):
backend = ChunkedSgmvLoRABackend(
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
scalings = [1.0] * 5
@@ -77,7 +77,7 @@ class TestEagleDraftCudaGraphRunner(CustomTestCase):
dsa_seed_topk=None,
)
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.seq_len_fill_value = SEQ_LEN_FILL_VALUE
runner.require_mlp_tp_gather = False