From ee1ca51de891bd6d9c913e0613a0f1aa8c54b1c3 Mon Sep 17 00:00:00 2001 From: Shangming Cai Date: Thu, 18 Dec 2025 14:24:55 +0800 Subject: [PATCH] [PP] Fix dynamic chunking strategy for PP (#15372) Signed-off-by: Shangming Cai --- .../sglang/srt/managers/scheduler_pp_mixin.py | 23 +++++++++++-------- 1 file changed, 13 insertions(+), 10 deletions(-) diff --git a/python/sglang/srt/managers/scheduler_pp_mixin.py b/python/sglang/srt/managers/scheduler_pp_mixin.py index dc13fb840..ea15cfd8b 100644 --- a/python/sglang/srt/managers/scheduler_pp_mixin.py +++ b/python/sglang/srt/managers/scheduler_pp_mixin.py @@ -10,6 +10,7 @@ from typing import TYPE_CHECKING, Dict, List, Optional, Tuple import numpy as np import torch import torch.distributed +from tqdm import tqdm from sglang.srt.disaggregation.base.conn import KVPoll from sglang.srt.disaggregation.utils import DisaggregationMode, poll_and_all_reduce @@ -532,13 +533,11 @@ class SchedulerPPMixin: latencies: List[float] = [] if self.pp_group.is_first_rank: - logger.info("Profiling prefill latency for dynamic chunk sizing...") - - # Create requests with different lengths: base_chunk_size // (2**i) for i in range(10) input_ids_list = [] - for i in range(32): - chunk_size = self.chunked_prefill_size - i * ( - self.chunked_prefill_size // 32 + for i in range(128): + chunk_size = int( + self.chunked_prefill_size * 1.25 + - i * (self.chunked_prefill_size * 1.25 // 128) ) if chunk_size <= 0: break @@ -551,9 +550,13 @@ class SchedulerPPMixin: temperature=0, max_new_tokens=1, ) - # Create and profile requests - for i, input_ids in enumerate(input_ids_list): + for i, input_ids in enumerate( + tqdm( + input_ids_list, + desc="Profiling prefill latency for dynamic chunking", + ) + ): req = Req( rid=str(i), origin_input_text="", @@ -1338,8 +1341,8 @@ class ChunkSizePredictor: ) calculated_chunk_size = int(smoothed_chunk_size) - # Align to page_size (round down to nearest multiple) - alignment_size = max(page_size, 1) + # Align to page_size (minimum alignment size is 64) + alignment_size = max(page_size, 64) dynamic_chunk_size = (calculated_chunk_size // alignment_size) * alignment_size # Ensure aligned size is at least alignment_size