[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,
|
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]:
|
||||||
|
|||||||
Reference in New Issue
Block a user