[Fix] Broadcast PP dynamic-chunk profiling failures so every rank disables together (#37675)
This commit is contained in:
@@ -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():
|
||||
|
||||
@@ -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]:
|
||||
|
||||
Reference in New Issue
Block a user