From 54da74e83e428fd26fdc5858a30e5c951e9db0df Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Wed, 2 Sep 2026 21:09:04 -0700 Subject: [PATCH] [Fix] Broadcast PP dynamic-chunk profiling failures so every rank disables together (#37675) --- python/sglang/srt/managers/scheduler.py | 1 + .../dynamic_chunk_sizer.py | 57 +++++++++++-------- 2 files changed, 34 insertions(+), 24 deletions(-) diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 4d2755acf..b002d9d87 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -1265,6 +1265,7 @@ class Scheduler( page_size=self.page_size, device=self.device, pp_group=self.pp_group, + world_group=self.world_group, pp_rank=self.ps.pp_rank, ) if sizer.profile_and_fit(): diff --git a/python/sglang/srt/managers/scheduler_components/dynamic_chunk_sizer.py b/python/sglang/srt/managers/scheduler_components/dynamic_chunk_sizer.py index 0dc3ab208..ca1f868d4 100644 --- a/python/sglang/srt/managers/scheduler_components/dynamic_chunk_sizer.py +++ b/python/sglang/srt/managers/scheduler_components/dynamic_chunk_sizer.py @@ -8,10 +8,8 @@ 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, @@ -23,6 +21,7 @@ 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 import broadcast_pyobj from sglang.srt.utils.common import get_device_module if TYPE_CHECKING: @@ -54,6 +53,7 @@ class DynamicChunkSizer: page_size: int, device: str, pp_group: GroupCoordinator, + world_group: GroupCoordinator, pp_rank: int, ): self.model_runner = model_runner @@ -67,41 +67,50 @@ class DynamicChunkSizer: self.page_size = page_size self.device = device self.pp_group = pp_group + self.world_group = world_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.""" + samples: Optional[Tuple[List[int], List[float]]] = None + + if self.pp_group.is_first_rank: + try: + samples = self._profile_prefill_latency() + except Exception as e: + logger.warning( + f"[PP Dynamic Chunk] Failed to profile prefill latency: {e!r}. " + "Dynamic chunking will be disabled." + ) + + # The samples are global, so one broadcast from global rank 0 (a PP0 rank) + # reaches every stage and attention rank; a failure travels as None. + samples = broadcast_pyobj( + [samples], self.world_group.rank, self.world_group.cpu_group, src=0 + )[0] + + if samples is None: + return False + + seq_lens, latencies = samples + # Quadratic model: f(l) = al^2 + bl + c 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: + # Every rank fits the same samples, so this fails on all of them alike. logger.warning( - f"[PP Dynamic Chunk] Failed to profile prefill latency: {e!r}. " + f"[PP Dynamic Chunk] Failed to fit the chunk-size predictor: {e!r}. " "Dynamic chunking will be disabled." ) return False + 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" + ) return True def predict(self, history_len: int) -> Optional[int]: