[misc] Extract PP dynamic chunk sizing into a DynamicChunkSizer scheduler component (#37674)

This commit is contained in:
Liangsheng Yin
2026-09-02 18:35:43 -07:00
committed by GitHub
parent 0dd66def7c
commit a522c8a4b6
5 changed files with 453 additions and 413 deletions
+26 -14
View File
@@ -216,6 +216,9 @@ from sglang.srt.managers.scheduler_components.batch_result_processor import (
SchedulerBatchResultProcessor, SchedulerBatchResultProcessor,
) )
from sglang.srt.managers.scheduler_components.dp_attn import SchedulerDPAttnAdapter from sglang.srt.managers.scheduler_components.dp_attn import SchedulerDPAttnAdapter
from sglang.srt.managers.scheduler_components.dynamic_chunk_sizer import (
DynamicChunkSizer,
)
from sglang.srt.managers.scheduler_components.flush_wrapper import SchedulerFlushWrapper from sglang.srt.managers.scheduler_components.flush_wrapper import SchedulerFlushWrapper
from sglang.srt.managers.scheduler_components.idle_sleeper import ( from sglang.srt.managers.scheduler_components.idle_sleeper import (
IdleSleeper, IdleSleeper,
@@ -627,6 +630,7 @@ class Scheduler(
# Init chunked prefill # Init chunked prefill
self.init_chunked_prefill() self.init_chunked_prefill()
self.maybe_init_dynamic_chunk_sizer()
# Init diffusion LLM # Init diffusion LLM
self.init_diffusion_llm() self.init_diffusion_llm()
@@ -1244,19 +1248,27 @@ class Scheduler(
self.chunked_prefill_size is not None and get_schedule().enable_mixed_chunk self.chunked_prefill_size is not None and get_schedule().enable_mixed_chunk
) )
# Init the dynamic chunking predictor for PP def maybe_init_dynamic_chunk_sizer(self) -> None:
self.enable_dynamic_chunking = ( """Profile a PP prefill latency model that sizes chunks per stage."""
get_schedule().enable_dynamic_chunking and self.ps.pp_size > 1 self.dynamic_chunk_sizer: Optional[DynamicChunkSizer] = None
if not (get_schedule().enable_dynamic_chunking and self.ps.pp_size > 1):
return
sizer = DynamicChunkSizer(
model_runner=self.tp_worker.model_runner,
model_config=self.model_config,
tree_cache=self.tree_cache,
req_to_token_pool=self.req_to_token_pool,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
spec_algorithm=self.spec_algorithm,
chunked_prefill_size=self.chunked_prefill_size,
max_prefill_tokens=self.max_prefill_tokens,
page_size=self.page_size,
device=self.device,
pp_group=self.pp_group,
pp_rank=self.ps.pp_rank,
) )
if self.enable_dynamic_chunking: if sizer.profile_and_fit():
try: self.dynamic_chunk_sizer = sizer
self.profile_and_init_predictor()
except Exception as e:
logger.warning(
f"[PP Dynamic Chunk] Failed to profile prefill latency: {e!r}. "
"Dynamic chunking will be disabled."
)
self.enable_dynamic_chunking = False
def _should_defer_prefill(self) -> bool: def _should_defer_prefill(self) -> bool:
if self._prefill_decode_interval_remaining == 0: if self._prefill_decode_interval_remaining == 0:
@@ -3601,9 +3613,9 @@ class Scheduler(
# Determine chunked_prefill_size for this batch # Determine chunked_prefill_size for this batch
chunked_prefill_size = self.chunked_prefill_size chunked_prefill_size = self.chunked_prefill_size
if self.chunked_req is not None and self.enable_dynamic_chunking: if self.chunked_req is not None and self.dynamic_chunk_sizer is not None:
history_len = len(self.chunked_req.prefix_indices) history_len = len(self.chunked_req.prefix_indices)
dynamic_size = self.predict_next_chunk_size(history_len) dynamic_size = self.dynamic_chunk_sizer.predict(history_len)
if dynamic_size is not None: if dynamic_size is not None:
chunked_prefill_size = dynamic_size chunked_prefill_size = dynamic_size
@@ -0,0 +1,420 @@
from __future__ import annotations
import logging
import math
import time
from array import array
from typing import TYPE_CHECKING, List, Optional, Tuple
import numpy as np
import torch
import torch.distributed
from tqdm import tqdm
from sglang.srt.distributed.communication_op import attn_cp_tp_broadcast_pyobj
from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import (
get_attention_dp_rank,
get_attention_dp_size,
is_dp_attention_enabled,
set_is_extend_in_batch,
)
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
from sglang.srt.mem_cache.common import release_kv_cache
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.srt.utils.common import get_device_module
if TYPE_CHECKING:
from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.distributed.parallel_state import GroupCoordinator
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
logger = logging.getLogger(__name__)
class DynamicChunkSizer:
"""Sizes PP prefill chunks from a profiled quadratic latency model."""
def __init__(
self,
*,
model_runner: ModelRunner,
model_config: ModelConfig,
tree_cache: BasePrefixCache,
req_to_token_pool: ReqToTokenPool,
token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator,
spec_algorithm: SpeculativeAlgorithm,
chunked_prefill_size: int,
max_prefill_tokens: int,
page_size: int,
device: str,
pp_group: GroupCoordinator,
pp_rank: int,
):
self.model_runner = model_runner
self.model_config = model_config
self.tree_cache = tree_cache
self.req_to_token_pool = req_to_token_pool
self.token_to_kv_pool_allocator = token_to_kv_pool_allocator
self.spec_algorithm = spec_algorithm
self.chunked_prefill_size = chunked_prefill_size
self.max_prefill_tokens = max_prefill_tokens
self.page_size = page_size
self.device = device
self.pp_group = pp_group
self.pp_rank = pp_rank
self.predictor = ChunkSizePredictor()
def profile_and_fit(self) -> bool:
"""PP0 profiles synthetic prefills and every rank fits the same samples;
returns whether the predictor is ready."""
try:
seq_lens: List[int] = []
latencies: List[float] = []
if self.pp_group.is_first_rank:
seq_lens, latencies = self._profile_prefill_latency()
seq_lens, latencies = attn_cp_tp_broadcast_pyobj([seq_lens, latencies])
# Broadcast data to all ranks
if torch.distributed.is_available() and torch.distributed.is_initialized():
data_to_sync = [seq_lens, latencies]
self.pp_group.broadcast_object_list(data_to_sync, src=0)
seq_lens, latencies = data_to_sync
# Quadratic model: f(l) = al^2 + bl + c
self.predictor.fit(seq_lens, latencies)
self.predictor.set_target_latency(self.chunked_prefill_size)
self.predictor.is_ready = True
logger.info(
f"[PP Dynamic Chunk] [PP{self.pp_rank}] Predictor ready (quadratic). "
f"Target latency: {self.predictor.target_latency:.2f}ms"
)
except Exception as e:
logger.warning(
f"[PP Dynamic Chunk] Failed to profile prefill latency: {e!r}. "
"Dynamic chunking will be disabled."
)
return False
return True
def predict(self, history_len: int) -> Optional[int]:
"""Chunk size for the next prefill step, or None to keep the static size."""
if not self.predictor.is_ready:
return None
max_chunk_size = self.max_prefill_tokens
predicted_size = self.predictor.predict_next_chunk_size(
history_len=history_len,
base_chunk_size=self.chunked_prefill_size,
page_size=self.page_size,
context_len=self.model_config.context_len,
max_chunk_size=max_chunk_size,
)
if predicted_size is not None:
logger.debug(
f"[PP Dynamic Chunk] [PP{self.pp_rank}] Predicted chunk size: "
f"{predicted_size} (history_len={history_len})"
)
return predicted_size
def _profile_prefill_latency(self) -> Tuple[List[int], List[float]]:
seq_lens: List[int] = []
latencies: List[float] = []
model_runner = self.model_runner
model_config = model_runner.model_config
input_ids_list: List[array[int]] = []
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
input_ids = array(
"q",
np.random.randint(0, 10000, size=chunk_size, dtype=np.int64).tobytes(),
)
input_ids_list.append(input_ids)
sampling_params = SamplingParams(
temperature=0,
max_new_tokens=1,
)
# Create and profile requests
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="",
origin_input_ids=input_ids,
sampling_params=sampling_params,
)
# Walk the same match -> lock -> alloc lifecycle as a scheduled
# request so release_kv_cache can release it symmetrically.
req.init_next_round_input(self.tree_cache)
lock = self.tree_cache.inc_lock_ref(req.last_node)
req.swa_uuid_for_lock = lock.swa_uuid_for_lock
req.set_extend_range(
len(req.prefix_indices), len(req.full_untruncated_fill_ids)
)
# Prepare batch
batch = ScheduleBatch.init_new(
[req],
self.req_to_token_pool,
self.token_to_kv_pool_allocator,
self.tree_cache,
self.model_config,
False,
self.spec_algorithm,
)
current_seq_len = req.extend_range.end
if is_dp_attention_enabled():
# Profiling runs one request on this rank; other DP ranks report 0.
dp_size = get_attention_dp_size()
global_num_tokens = [0] * dp_size
dp_rank = get_attention_dp_rank()
global_num_tokens[dp_rank] = current_seq_len
batch.global_num_tokens = global_num_tokens
batch.global_num_tokens_for_logprob = global_num_tokens
hs = (
getattr(model_config, "hc_hidden_size", None)
or model_config.hidden_size
)
proxy_tensors = {
"hidden_states": torch.zeros(
(current_seq_len, hs),
dtype=model_config.dtype,
device=self.device,
),
"residual": torch.zeros(
(current_seq_len, model_config.hidden_size),
dtype=model_config.dtype,
device=self.device,
),
}
pp_proxy_topk_size = model_runner.get_pp_proxy_topk_size()
if pp_proxy_topk_size is not None:
proxy_tensors["topk_indices"] = torch.zeros(
(current_seq_len, pp_proxy_topk_size),
dtype=torch.int32,
device=self.device,
)
pp_proxy = PPProxyTensors(proxy_tensors)
# Measure latency with device synchronization for accurate timing
device_module = get_device_module()
# Synchronize before starting timing to ensure clean measurement
device_module.synchronize()
start = time.perf_counter()
batch.prepare_for_extend()
# Resolve deferred H2D: prepare_for_extend now leaves input_ids=None
if batch.input_ids is None and batch.prefill_input_ids_cpu is not None:
batch.input_ids = batch.prefill_input_ids_cpu.to(
self.device, non_blocking=True
)
batch.prefill_input_ids_cpu = None
forward_batch = ForwardBatch.init_new(
batch,
model_runner,
return_hidden_states_before_norm=False,
)
set_is_extend_in_batch(batch.forward_mode.is_extend())
_ = model_runner.forward(
forward_batch=forward_batch, pp_proxy_tensors=pp_proxy
)
# Synchronize after forward to ensure GPU operations complete
device_module.synchronize()
latency_seconds = time.perf_counter() - start
latency_ms = latency_seconds * 1e3 # Convert to milliseconds
seq_lens.append(len(input_ids))
latencies.append(latency_ms)
# Release KV and Mamba cache
if req.kv.holds_kv:
release_kv_cache(req, self.tree_cache, is_insert=False)
logger.info(
f"[PP Dynamic Chunk] [PP0] Profiled {len(seq_lens)} samples: "
f"seq_lens={seq_lens}, latencies_ms={latencies}"
)
return seq_lens, latencies
class ChunkSizePredictor:
"""Quadratic latency model f(l) = a*l^2 + b*l + c; predicts the chunk x with
f(L + x) - f(L) = target_latency."""
def __init__(self):
self.quadratic_coeff_a = 0.0
self.linear_coeff_b = 0.0
self.constant_coeff_c = 0.0
self.target_latency: Optional[float] = None
self.is_ready = False
def fit(self, seq_lens: List[int], latencies: List[float]):
"""Fit quadratic coefficients f(l) = al^2 + bl + c from data points."""
# Skip the first data point to reduce fitting bias, as the first run is slower without warmup
L = np.array(seq_lens[1:], dtype=np.float64)
T = np.array(latencies[1:], dtype=np.float64)
if len(L) < 8:
raise ValueError(
f"Not enough data points for quadratic fitting ({len(L)} < 8). "
"Need at least 8 samples with different sequence lengths."
)
# Build design matrix for f(l) = al^2 + bl + c
X = np.column_stack([L * L, L, np.ones_like(L)]) # [l^2, l, 1]
try:
coeffs, residuals, rank, s = np.linalg.lstsq(X, T, rcond=None)
if len(coeffs) >= 3:
fitted_a = float(coeffs[0]) # quadratic coefficient
fitted_b = float(coeffs[1]) # linear coefficient
fitted_c = float(coeffs[2]) # constant coefficient
else:
raise ValueError("Failed to fit coefficients: insufficient rank")
except np.linalg.LinAlgError as e:
raise ValueError(f"Failed to fit f(l) = al^2 + bl + c: {e}")
# Validate coefficients
if fitted_a <= 0:
raise ValueError(
f"Fitted quadratic coefficient a={fitted_a:.2e} is not positive. "
"Attention has O(n^2) complexity, so a must be positive. "
"Check warmup data quality."
)
if fitted_b < 0:
logger.warning(
f"Fitted linear coefficient b={fitted_b:.2e} is negative. Setting b=0."
)
fitted_b = 0.0
self.quadratic_coeff_a = fitted_a
self.linear_coeff_b = fitted_b
self.constant_coeff_c = fitted_c
logger.info(
f"[ChunkSizePredictor] Fitted coefficients: a={fitted_a:.2e}, "
f"b={fitted_b:.2e}, c={fitted_c:.2e}"
)
def set_target_latency(self, base_chunk_size: int):
"""Set target latency based on base chunk size: target = f(base_chunk_size) - f(0)."""
def f(length: float) -> float:
"""Total latency function: f(length) = a*length^2 + b*length + c."""
return (
self.quadratic_coeff_a * length * length
+ self.linear_coeff_b * length
+ self.constant_coeff_c
)
self.target_latency = f(float(base_chunk_size)) - f(0.0)
if self.target_latency <= 0:
raise ValueError(
f"Calculated target_latency={self.target_latency:.2f}ms is not positive. "
"Check warmup data quality."
)
logger.info(
f"[ChunkSizePredictor] Target latency: {self.target_latency:.2f}ms "
f"(base_chunk_size={base_chunk_size})"
)
def predict_next_chunk_size(
self,
history_len: int,
base_chunk_size: int,
page_size: int,
context_len: int,
max_chunk_size: Optional[int] = None,
) -> Optional[int]:
"""Chunk size x with f(L + x) - f(L) = target_latency for L = history_len,
or None when the model cannot say."""
if not self.is_ready or self.target_latency is None:
return None
# Handle quadratic model: f(l) = al^2 + bl + c
if self.quadratic_coeff_a <= 0:
return None
# f(L+x) - f(L) = T expands to a*x^2 + (2aL + b)*x - T = 0.
A = self.quadratic_coeff_a
B = 2 * self.quadratic_coeff_a * history_len + self.linear_coeff_b
C = -self.target_latency
discriminant = B * B - 4 * A * C
if discriminant < 0:
logger.warning(
f"Discriminant is negative ({discriminant:.2e}). "
f"No real solution for chunk size. L={history_len}, T={self.target_latency:.2f}ms."
)
return None
sqrt_discriminant = math.sqrt(discriminant)
calculated_chunk_size_float = (-B + sqrt_discriminant) / (2 * A)
if calculated_chunk_size_float <= 0:
logger.warning(
f"Calculated chunk size is non-positive ({calculated_chunk_size_float:.2f}). "
f"L={history_len}, T={self.target_latency:.2f}ms."
)
return None
# Use a smooth coefficient to reduce the abrupt decrease in chunk size
smooth_coeff = envs.SGLANG_DYNAMIC_CHUNKING_SMOOTH_FACTOR.get()
smoothed_chunk_size = base_chunk_size + smooth_coeff * (
calculated_chunk_size_float - base_chunk_size
)
# Make sure the dynamic chunk size is at least 1/4 of the base chunk size
calculated_chunk_size = max(int(smoothed_chunk_size), base_chunk_size // 4)
# 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
if dynamic_chunk_size < alignment_size:
dynamic_chunk_size = alignment_size
# Apply constraints
max_allowed = context_len - history_len - 100 # Leave 100 tokens margin
if max_chunk_size is not None:
max_allowed = min(max_allowed, max_chunk_size)
dynamic_chunk_size = min(dynamic_chunk_size, max_allowed)
# Align again after min operation
dynamic_chunk_size = (dynamic_chunk_size // alignment_size) * alignment_size
if dynamic_chunk_size < alignment_size:
return None
return dynamic_chunk_size
@@ -1,29 +1,18 @@
from __future__ import annotations from __future__ import annotations
import logging import logging
import math
import time
from array import array
from collections import defaultdict, deque from collections import defaultdict, deque
from dataclasses import dataclass from dataclasses import dataclass
from typing import TYPE_CHECKING, Dict, List, Optional, Tuple from typing import TYPE_CHECKING, Dict, List, Optional, Tuple
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 poll_and_all_reduce_attn_cp_tp_group from sglang.srt.disaggregation.utils import poll_and_all_reduce_attn_cp_tp_group
from sglang.srt.distributed.communication_op import attn_cp_tp_broadcast_pyobj from sglang.srt.distributed.communication_op import attn_cp_tp_broadcast_pyobj
from sglang.srt.distributed.parallel_state import P2PWork from sglang.srt.distributed.parallel_state import P2PWork
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import (
get_attention_dp_rank,
get_attention_dp_size,
is_dp_attention_enabled,
set_is_extend_in_batch,
)
from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.managers.overlap_utils import RelayPayload from sglang.srt.managers.overlap_utils import RelayPayload
from sglang.srt.managers.schedule_batch import FINISH_ABORT, Req, ScheduleBatch from sglang.srt.managers.schedule_batch import FINISH_ABORT, Req, ScheduleBatch
@@ -32,9 +21,7 @@ from sglang.srt.managers.utils import (
get_logprob_dict_from_result, get_logprob_dict_from_result,
get_logprob_from_pp_outputs, get_logprob_from_pp_outputs,
) )
from sglang.srt.mem_cache.common import release_kv_cache
from sglang.srt.model_executor.forward_batch_info import ( from sglang.srt.model_executor.forward_batch_info import (
ForwardBatch,
ForwardMode, ForwardMode,
PPProxyTensors, PPProxyTensors,
) )
@@ -44,9 +31,8 @@ from sglang.srt.sampling.sampling_observer_pp import (
add_auxiliary_output_to_pp_tensors, add_auxiliary_output_to_pp_tensors,
pop_auxiliary_output_from_pp_tensors, pop_auxiliary_output_from_pp_tensors,
) )
from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.srt.utils import DynamicGradMode, point_to_point_pyobj from sglang.srt.utils import DynamicGradMode, point_to_point_pyobj
from sglang.srt.utils.common import get_device_module, is_xpu from sglang.srt.utils.common import is_xpu
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -589,205 +575,6 @@ class SchedulerPPMixin:
defaultdict(deque) defaultdict(deque)
) )
def profile_and_init_predictor(self: Scheduler):
"""
Profile prefill latency for dynamic chunk sizing.
Only runs on PP0 (first rank), then broadcasts data to all ranks.
All ranks fit coefficients using the same data.
"""
seq_lens: List[int] = []
latencies: List[float] = []
if self.pp_group.is_first_rank:
model_runner = self.tp_worker.model_runner
model_config = model_runner.model_config
input_ids_list: List[array[int]] = []
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
input_ids = array(
"q",
np.random.randint(
0, 10000, size=chunk_size, dtype=np.int64
).tobytes(),
)
input_ids_list.append(input_ids)
sampling_params = SamplingParams(
temperature=0,
max_new_tokens=1,
)
# Create and profile requests
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="",
origin_input_ids=input_ids,
sampling_params=sampling_params,
)
# Walk the same match -> lock -> alloc lifecycle as a scheduled
# request so release_kv_cache can release it symmetrically.
req.init_next_round_input(self.tree_cache)
lock = self.tree_cache.inc_lock_ref(req.last_node)
req.swa_uuid_for_lock = lock.swa_uuid_for_lock
req.set_extend_range(
len(req.prefix_indices), len(req.full_untruncated_fill_ids)
)
# Prepare batch
batch = ScheduleBatch.init_new(
[req],
self.req_to_token_pool,
self.token_to_kv_pool_allocator,
self.tree_cache,
self.model_config,
False,
self.spec_algorithm,
)
current_seq_len = req.extend_range.end
if is_dp_attention_enabled():
# For profiling, we only have one request on PP0
# Set global_num_tokens to indicate this rank has tokens, others have 0
dp_size = get_attention_dp_size()
global_num_tokens = [0] * dp_size
dp_rank = get_attention_dp_rank()
global_num_tokens[dp_rank] = current_seq_len
batch.global_num_tokens = global_num_tokens
batch.global_num_tokens_for_logprob = global_num_tokens
hs = (
getattr(model_config, "hc_hidden_size", None)
or model_config.hidden_size
)
proxy_tensors = {
"hidden_states": torch.zeros(
(current_seq_len, hs),
dtype=model_config.dtype,
device=self.device,
),
"residual": torch.zeros(
(current_seq_len, model_config.hidden_size),
dtype=model_config.dtype,
device=self.device,
),
}
pp_proxy_topk_size = model_runner.get_pp_proxy_topk_size()
if pp_proxy_topk_size is not None:
proxy_tensors["topk_indices"] = torch.zeros(
(current_seq_len, pp_proxy_topk_size),
dtype=torch.int32,
device=self.device,
)
pp_proxy = PPProxyTensors(proxy_tensors)
# Measure latency with device synchronization for accurate timing
device_module = get_device_module()
# Synchronize before starting timing to ensure clean measurement
device_module.synchronize()
start = time.perf_counter()
batch.prepare_for_extend()
# Resolve deferred H2D: prepare_for_extend now leaves input_ids=None
if batch.input_ids is None and batch.prefill_input_ids_cpu is not None:
batch.input_ids = batch.prefill_input_ids_cpu.to(
self.device, non_blocking=True
)
batch.prefill_input_ids_cpu = None
forward_batch = ForwardBatch.init_new(
batch,
model_runner,
return_hidden_states_before_norm=False,
)
set_is_extend_in_batch(batch.forward_mode.is_extend())
_ = model_runner.forward(
forward_batch=forward_batch, pp_proxy_tensors=pp_proxy
)
# Synchronize after forward to ensure GPU operations complete
device_module.synchronize()
latency_seconds = time.perf_counter() - start
latency_ms = latency_seconds * 1e3 # Convert to milliseconds
seq_lens.append(len(input_ids))
latencies.append(latency_ms)
# Release KV and Mamba cache
if req.kv.holds_kv:
release_kv_cache(req, self.tree_cache, is_insert=False)
logger.info(
f"[PP Dynamic Chunk] [PP0] Profiled {len(seq_lens)} samples: "
f"seq_lens={seq_lens}, latencies_ms={latencies}"
)
seq_lens, latencies = attn_cp_tp_broadcast_pyobj([seq_lens, latencies])
# Broadcast data to all ranks
if torch.distributed.is_available() and torch.distributed.is_initialized():
data_to_sync = [seq_lens, latencies]
self.pp_group.broadcast_object_list(data_to_sync, src=0)
seq_lens, latencies = data_to_sync
# Quadratic model: f(l) = al^2 + bl + c
self.length_predictor = ChunkSizePredictor()
self.length_predictor.fit(seq_lens, latencies)
self.length_predictor.set_target_latency(self.chunked_prefill_size)
self.length_predictor.is_ready = True
logger.info(
f"[PP Dynamic Chunk] [PP{self.ps.pp_rank}] Predictor ready (quadratic). "
f"Target latency: {self.length_predictor.target_latency:.2f}ms"
)
def predict_next_chunk_size(self: Scheduler, history_len: int) -> Optional[int]:
"""
Predict next chunk size dynamically based on current history length.
Args:
history_len: Current sequence length
Returns:
Predicted chunk size, or None to use default chunked_prefill_size
"""
if (
not self.enable_dynamic_chunking
or self.length_predictor is None
or not self.length_predictor.is_ready
):
return None
max_chunk_size = self.max_prefill_tokens
predicted_size = self.length_predictor.predict_next_chunk_size(
history_len=history_len,
base_chunk_size=self.chunked_prefill_size,
page_size=self.page_size,
context_len=self.model_config.context_len,
max_chunk_size=max_chunk_size,
)
if predicted_size is not None:
logger.debug(
f"[PP Dynamic Chunk] [PP{self.ps.pp_rank}] Predicted chunk size: "
f"{predicted_size} (history_len={history_len})"
)
return predicted_size
def process_bootstrapped_queue( def process_bootstrapped_queue(
self: Scheduler, bootstrapped_rids: Optional[List[str]] self: Scheduler, bootstrapped_rids: Optional[List[str]]
): ):
@@ -1505,177 +1292,3 @@ class SchedulerPPMixin:
self.waiting_queue.extend(released_reqs) self.waiting_queue.extend(released_reqs)
return [req.rid for req in released_reqs] return [req.rid for req in released_reqs]
return None return None
class ChunkSizePredictor:
"""
Predictor for dynamic chunk size based on quadratic latency model.
Models latency as: f(l) = a*l^2 + b*l + c
Predicts next chunk size x such that: f(L+x) - f(L) = target_latency
"""
def __init__(self):
self.quadratic_coeff_a = 0.0
self.linear_coeff_b = 0.0
self.constant_coeff_c = 0.0
self.target_latency: Optional[float] = None
self.is_ready = False
def fit(self, seq_lens: List[int], latencies: List[float]):
"""Fit quadratic coefficients f(l) = al^2 + bl + c from data points."""
# Skip the first data point to reduce fitting bias, as the first run is slower without warmup
L = np.array(seq_lens[1:], dtype=np.float64)
T = np.array(latencies[1:], dtype=np.float64)
if len(L) < 8:
raise ValueError(
f"Not enough data points for quadratic fitting ({len(L)} < 8). "
"Need at least 8 samples with different sequence lengths."
)
# Build design matrix for f(l) = al^2 + bl + c
X = np.column_stack([L * L, L, np.ones_like(L)]) # [l^2, l, 1]
try:
coeffs, residuals, rank, s = np.linalg.lstsq(X, T, rcond=None)
if len(coeffs) >= 3:
fitted_a = float(coeffs[0]) # quadratic coefficient
fitted_b = float(coeffs[1]) # linear coefficient
fitted_c = float(coeffs[2]) # constant coefficient
else:
raise ValueError("Failed to fit coefficients: insufficient rank")
except np.linalg.LinAlgError as e:
raise ValueError(f"Failed to fit f(l) = al^2 + bl + c: {e}")
# Validate coefficients
if fitted_a <= 0:
raise ValueError(
f"Fitted quadratic coefficient a={fitted_a:.2e} is not positive. "
"Attention has O(n^2) complexity, so a must be positive. "
"Check warmup data quality."
)
if fitted_b < 0:
logger.warning(
f"Fitted linear coefficient b={fitted_b:.2e} is negative. Setting b=0."
)
fitted_b = 0.0
self.quadratic_coeff_a = fitted_a
self.linear_coeff_b = fitted_b
self.constant_coeff_c = fitted_c
logger.info(
f"[ChunkSizePredictor] Fitted coefficients: a={fitted_a:.2e}, "
f"b={fitted_b:.2e}, c={fitted_c:.2e}"
)
def set_target_latency(self, base_chunk_size: int):
"""Set target latency based on base chunk size: target = f(base_chunk_size) - f(0)."""
def f(length: float) -> float:
"""Total latency function: f(length) = a*length^2 + b*length + c."""
return (
self.quadratic_coeff_a * length * length
+ self.linear_coeff_b * length
+ self.constant_coeff_c
)
self.target_latency = f(float(base_chunk_size)) - f(0.0)
if self.target_latency <= 0:
raise ValueError(
f"Calculated target_latency={self.target_latency:.2f}ms is not positive. "
"Check warmup data quality."
)
logger.info(
f"[ChunkSizePredictor] Target latency: {self.target_latency:.2f}ms "
f"(base_chunk_size={base_chunk_size})"
)
def predict_next_chunk_size(
self,
history_len: int,
base_chunk_size: int,
page_size: int,
context_len: int,
max_chunk_size: Optional[int] = None,
) -> Optional[int]:
"""
Predict next chunk size x such that f(history_len + x) - f(history_len) = target_latency.
Args:
history_len: Current sequence length (L)
base_chunk_size: Base chunk size
page_size: Page size for alignment
context_len: Maximum context length
max_chunk_size: Maximum allowed chunk size (optional)
Returns:
Predicted chunk size, or None if prediction fails
"""
if not self.is_ready or self.target_latency is None:
return None
# Handle quadratic model: f(l) = al^2 + bl + c
if self.quadratic_coeff_a <= 0:
return None
# Solve f(L+x) - f(L) = T
# where f(L) = a*L^2 + b*L + c
# This expands to: ax^2 + (2aL+b)x - T = 0
# A = a, B = 2aL + b, C = -T
A = self.quadratic_coeff_a
B = 2 * self.quadratic_coeff_a * history_len + self.linear_coeff_b
C = -self.target_latency
discriminant = B * B - 4 * A * C
if discriminant < 0:
logger.warning(
f"Discriminant is negative ({discriminant:.2e}). "
f"No real solution for chunk size. L={history_len}, T={self.target_latency:.2f}ms."
)
return None
sqrt_discriminant = math.sqrt(discriminant)
calculated_chunk_size_float = (-B + sqrt_discriminant) / (2 * A)
if calculated_chunk_size_float <= 0:
logger.warning(
f"Calculated chunk size is non-positive ({calculated_chunk_size_float:.2f}). "
f"L={history_len}, T={self.target_latency:.2f}ms."
)
return None
# Use a smooth coefficient to reduce the abrupt decrease in chunk size
smooth_coeff = envs.SGLANG_DYNAMIC_CHUNKING_SMOOTH_FACTOR.get()
smoothed_chunk_size = base_chunk_size + smooth_coeff * (
calculated_chunk_size_float - base_chunk_size
)
# Make sure the dynamic chunk size is at least 1/4 of the base chunk size
calculated_chunk_size = max(int(smoothed_chunk_size), base_chunk_size // 4)
# 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
if dynamic_chunk_size < alignment_size:
dynamic_chunk_size = alignment_size
# Apply constraints
max_allowed = context_len - history_len - 100 # Leave 100 tokens margin
if max_chunk_size is not None:
max_allowed = min(max_allowed, max_chunk_size)
dynamic_chunk_size = min(dynamic_chunk_size, max_allowed)
# Align again after min operation
dynamic_chunk_size = (dynamic_chunk_size // alignment_size) * alignment_size
if dynamic_chunk_size < alignment_size:
return None
return dynamic_chunk_size
@@ -80,9 +80,7 @@ class TestPPBasic(ScriptedTestCase):
@staticmethod @staticmethod
def _script_pp_static_chunk_size_predictor_returns_none(t: ScriptedContext): def _script_pp_static_chunk_size_predictor_returns_none(t: ScriptedContext):
sched = t.scheduler sched = t.scheduler
assert sched.enable_dynamic_chunking is False assert sched.dynamic_chunk_sizer is None
assert sched.predict_next_chunk_size(0) is None
assert sched.predict_next_chunk_size(VERY_LONG_PROMPT_LEN // 2) is None
r = t.start_req(prompt_len=VERY_LONG_PROMPT_LEN, max_new_tokens=2) r = t.start_req(prompt_len=VERY_LONG_PROMPT_LEN, max_new_tokens=2)
yield from run_until_finished(r, max_steps=800) yield from run_until_finished(r, max_steps=800)
assert r.finished assert r.finished
@@ -211,12 +209,9 @@ class TestPPDynamic(ScriptedTestCase):
@staticmethod @staticmethod
def _script_pp_dynamic_chunk_size_recompute_branch_taken(t: ScriptedContext): def _script_pp_dynamic_chunk_size_recompute_branch_taken(t: ScriptedContext):
sched = t.scheduler sizer = t.scheduler.dynamic_chunk_sizer
assert sched.enable_dynamic_chunking is True assert sizer is not None
assert sched.length_predictor is not None dynamic_size = sizer.predict(0)
assert sched.length_predictor.is_ready is True
dynamic_size = sched.predict_next_chunk_size(0)
assert dynamic_size is not None
assert isinstance(dynamic_size, int) and dynamic_size > 0 assert isinstance(dynamic_size, int) and dynamic_size > 0
r = t.start_req( r = t.start_req(
prompt_len=VERY_LONG_PROMPT_LEN, max_new_tokens=2, prompt_token=4 prompt_len=VERY_LONG_PROMPT_LEN, max_new_tokens=2, prompt_token=4
@@ -885,8 +885,8 @@ class TestSpecialCaseDynamicChunkingPP1(ScriptedTestCase):
@staticmethod @staticmethod
def _script_dynamic_chunking_forced_off_on_pp1(t: ScriptedContext): def _script_dynamic_chunking_forced_off_on_pp1(t: ScriptedContext):
assert t.scheduler.enable_dynamic_chunking is False, ( assert t.scheduler.dynamic_chunk_sizer is None, (
"pp_size==1 must force enable_dynamic_chunking off even when the " "pp_size==1 must leave dynamic chunking off even when the "
"server arg is True (the 'and ps.pp_size > 1' conjunct)" "server arg is True (the 'and ps.pp_size > 1' conjunct)"
) )
r = t.start_req( r = t.start_req(