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