[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

View File

@@ -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