[Spec] Rename num_tokens_per_bs to num_tokens_per_req (#30977)
This commit is contained in:
@@ -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] * (
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
+13
-13
@@ -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)):
|
||||
|
||||
+4
-4
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user