[PP] Fix dynamic chunking strategy for PP (#15372)

Signed-off-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
Shangming Cai
2025-12-18 14:24:55 +08:00
committed by GitHub
parent ef7c29acd7
commit ee1ca51de8
@@ -10,6 +10,7 @@ from typing import TYPE_CHECKING, Dict, List, Optional, Tuple
import numpy as np import numpy as np
import torch import torch
import torch.distributed import torch.distributed
from tqdm import tqdm
from sglang.srt.disaggregation.base.conn import KVPoll from sglang.srt.disaggregation.base.conn import KVPoll
from sglang.srt.disaggregation.utils import DisaggregationMode, poll_and_all_reduce from sglang.srt.disaggregation.utils import DisaggregationMode, poll_and_all_reduce
@@ -532,13 +533,11 @@ class SchedulerPPMixin:
latencies: List[float] = [] latencies: List[float] = []
if self.pp_group.is_first_rank: 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 = [] input_ids_list = []
for i in range(32): for i in range(128):
chunk_size = self.chunked_prefill_size - i * ( chunk_size = int(
self.chunked_prefill_size // 32 self.chunked_prefill_size * 1.25
- i * (self.chunked_prefill_size * 1.25 // 128)
) )
if chunk_size <= 0: if chunk_size <= 0:
break break
@@ -551,9 +550,13 @@ class SchedulerPPMixin:
temperature=0, temperature=0,
max_new_tokens=1, max_new_tokens=1,
) )
# Create and profile requests # 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( req = Req(
rid=str(i), rid=str(i),
origin_input_text="", origin_input_text="",
@@ -1338,8 +1341,8 @@ class ChunkSizePredictor:
) )
calculated_chunk_size = int(smoothed_chunk_size) calculated_chunk_size = int(smoothed_chunk_size)
# Align to page_size (round down to nearest multiple) # Align to page_size (minimum alignment size is 64)
alignment_size = max(page_size, 1) alignment_size = max(page_size, 64)
dynamic_chunk_size = (calculated_chunk_size // alignment_size) * alignment_size dynamic_chunk_size = (calculated_chunk_size // alignment_size) * alignment_size
# Ensure aligned size is at least alignment_size # Ensure aligned size is at least alignment_size