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