[misc] Extract PP dynamic chunk sizing into a DynamicChunkSizer scheduler component (#37674)
This commit is contained in:
@@ -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(
|
||||||
|
|||||||
Reference in New Issue
Block a user