[Fix] Broadcast PP dynamic-chunk profiling failures so every rank disables together (#37675)

This commit is contained in:
Liangsheng Yin
2026-09-02 21:09:04 -07:00
committed by GitHub
parent 98ef7d8ae6
commit 54da74e83e
2 changed files with 34 additions and 24 deletions
+1
View File
@@ -1265,6 +1265,7 @@ class Scheduler(
page_size=self.page_size, page_size=self.page_size,
device=self.device, device=self.device,
pp_group=self.pp_group, pp_group=self.pp_group,
world_group=self.world_group,
pp_rank=self.ps.pp_rank, pp_rank=self.ps.pp_rank,
) )
if sizer.profile_and_fit(): if sizer.profile_and_fit():
@@ -8,10 +8,8 @@ from typing import TYPE_CHECKING, List, Optional, Tuple
import numpy as np import numpy as np
import torch import torch
import torch.distributed
from tqdm import tqdm 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.environ import envs
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import (
get_attention_dp_rank, 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.mem_cache.common import release_kv_cache
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.sampling.sampling_params import SamplingParams from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.srt.utils import broadcast_pyobj
from sglang.srt.utils.common import get_device_module from sglang.srt.utils.common import get_device_module
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -54,6 +53,7 @@ class DynamicChunkSizer:
page_size: int, page_size: int,
device: str, device: str,
pp_group: GroupCoordinator, pp_group: GroupCoordinator,
world_group: GroupCoordinator,
pp_rank: int, pp_rank: int,
): ):
self.model_runner = model_runner self.model_runner = model_runner
@@ -67,41 +67,50 @@ class DynamicChunkSizer:
self.page_size = page_size self.page_size = page_size
self.device = device self.device = device
self.pp_group = pp_group self.pp_group = pp_group
self.world_group = world_group
self.pp_rank = pp_rank self.pp_rank = pp_rank
self.predictor = ChunkSizePredictor() self.predictor = ChunkSizePredictor()
def profile_and_fit(self) -> bool: def profile_and_fit(self) -> bool:
"""PP0 profiles synthetic prefills and every rank fits the same samples; """PP0 profiles synthetic prefills and every rank fits the same samples;
returns whether the predictor is ready.""" 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: 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.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: except Exception as e:
# Every rank fits the same samples, so this fails on all of them alike.
logger.warning( 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." "Dynamic chunking will be disabled."
) )
return False 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 return True
def predict(self, history_len: int) -> Optional[int]: def predict(self, history_len: int) -> Optional[int]: