Support EPLB balancedness prometheus metric without GPU->CPU synchronize (#15401)

This commit is contained in:
fzyzcjy
2025-12-18 22:24:23 +08:00
committed by GitHub
parent 602fe3b296
commit 88a405cc10
8 changed files with 111 additions and 33 deletions
+58 -29
View File
@@ -20,6 +20,7 @@ import time
from abc import ABC
from collections import deque
from contextlib import contextmanager
from dataclasses import dataclass
from pathlib import Path
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Type
@@ -43,6 +44,14 @@ logger = logging.getLogger(__name__)
_OutputMode = Literal["file", "object"]
@dataclass
class ExpertDistributionMetrics:
eplb_balancedness: torch.Tensor
def copy_to_cpu(self):
self.eplb_balancedness = self.eplb_balancedness.to("cpu", non_blocking=True)
class ExpertDistributionRecorder(ABC):
"""Global expert distribution recording"""
@@ -78,7 +87,7 @@ class ExpertDistributionRecorder(ABC):
@contextmanager
def with_forward_pass(self, forward_pass_id: int, forward_batch: ForwardBatch):
yield
yield {}
def on_select_experts(self, topk_ids: torch.Tensor):
pass
@@ -157,12 +166,13 @@ class _ExpertDistributionRecorderReal(ExpertDistributionRecorder):
@contextmanager
def with_forward_pass(self, forward_pass_id: int, forward_batch: ForwardBatch):
outputs = {}
with self._current_forward_pass_id.with_value(forward_pass_id):
self._on_forward_pass_start(forward_batch)
try:
yield
yield outputs
finally:
self._on_forward_pass_end(forward_pass_id)
self._on_forward_pass_end(forward_pass_id, outputs)
@contextmanager
def disable_this_region(self):
@@ -181,12 +191,14 @@ class _ExpertDistributionRecorderReal(ExpertDistributionRecorder):
gatherer.reset()
gatherer.on_forward_pass_start(forward_batch)
def _on_forward_pass_end(self, forward_pass_id: int):
def _on_forward_pass_end(self, forward_pass_id: int, outputs: Dict[str, Any]):
if not self._recording:
return
for gatherer_key, gatherer in self._single_pass_gatherers.items():
single_pass_data = gatherer.collect()
self._accumulator.append(forward_pass_id, gatherer_key, single_pass_data)
self._accumulator.append(
forward_pass_id, gatherer_key, single_pass_data, outputs
)
def on_select_experts(self, topk_ids: torch.Tensor):
self._on_hook("on_select_experts", topk_ids=topk_ids)
@@ -636,6 +648,7 @@ class _Accumulator(ABC):
forward_pass_id: int,
gatherer_key: str,
single_pass_data: Dict,
outputs: Dict[str, Any],
):
pass
@@ -659,18 +672,19 @@ class _UtilizationRateAccumulatorMixin(_Accumulator):
self._expert_dispatch_collector = ExpertDispatchCollector(
self._expert_location_metadata.ep_size
)
self._collection_counter = 0
self._metric_heatmap_collection_counter = 0
def append(
self,
forward_pass_id: int,
gatherer_key: str,
single_pass_data: Dict,
outputs: Dict[str, Any],
):
super().append(forward_pass_id, gatherer_key, single_pass_data)
super().append(forward_pass_id, gatherer_key, single_pass_data, outputs)
if self._enable:
self._append_utilization_rate(
forward_pass_id, single_pass_data["global_physical_count"]
return self._append_utilization_rate(
forward_pass_id, single_pass_data["global_physical_count"], outputs
)
def reset(self):
@@ -679,7 +693,10 @@ class _UtilizationRateAccumulatorMixin(_Accumulator):
self._history.clear()
def _append_utilization_rate(
self, forward_pass_id: int, single_pass_global_physical_count: torch.Tensor
self,
forward_pass_id: int,
single_pass_global_physical_count: torch.Tensor,
outputs: Dict[str, Any],
):
gpu_physical_count = compute_gpu_physical_count(
single_pass_global_physical_count,
@@ -691,27 +708,37 @@ class _UtilizationRateAccumulatorMixin(_Accumulator):
)
if self._rank == 0:
self._collect_metrics_if_needed(gpu_physical_count)
self._handle_metric_eplb_heatmap(gpu_physical_count)
utilization_rate_tensor = compute_utilization_rate(gpu_physical_count)
utilization_rate = torch.mean(utilization_rate_tensor).item()
self._history.append(utilization_rate)
gpu_physical_count_sum = gpu_physical_count.sum().item()
logger.info(
f"[Expert Balancedness] "
f"forward_pass_id={forward_pass_id} "
f"current_pass_balancedness={utilization_rate:.03f} "
f"{''.join(f'last_{size}_average_balancedness={value:.03f} ' for size, value in self._history.mean().items())} "
f"gpu_physical_count_sum={gpu_physical_count_sum}"
# f"current_pass_per_layer={[round(x, 2) for x in utilization_rate_tensor.cpu().tolist()]}"
utilization_rate_gpu = torch.mean(
compute_utilization_rate(gpu_physical_count)
)
if envs.SGLANG_ENABLE_EPLB_BALANCEDNESS_METRIC.get():
print(f"hi {self._rank=} {utilization_rate_gpu=}")
outputs["metrics"] = ExpertDistributionMetrics(
eplb_balancedness=utilization_rate_gpu,
)
else:
# TODO maybe refactor this part to also avoid a `.item()` gpu->cpu sync
utilization_rate_cpu = utilization_rate_gpu.item()
self._history.append(utilization_rate_cpu)
def _collect_metrics_if_needed(self, gpu_physical_count: torch.Tensor):
gpu_physical_count_sum = gpu_physical_count.sum().item()
logger.info(
f"[Expert Balancedness] "
f"forward_pass_id={forward_pass_id} "
f"current_pass_balancedness={utilization_rate_cpu:.03f} "
f"{''.join(f'last_{size}_average_balancedness={value:.03f} ' for size, value in self._history.mean().items())} "
f"gpu_physical_count_sum={gpu_physical_count_sum}"
# f"current_pass_per_layer={[round(x, 2) for x in utilization_rate_tensor.cpu().tolist()]}"
)
# TODO refactor
def _handle_metric_eplb_heatmap(self, gpu_physical_count: torch.Tensor):
# sglang:eplb_gpu_physical_count metric is disabled if SGLANG_EPLB_HEATMAP_COLLECTION_INTERVAL <= 0
interval = get_int_env_var("SGLANG_EPLB_HEATMAP_COLLECTION_INTERVAL", 0)
if interval > 0 and self._collection_counter % interval == 0:
if interval > 0 and self._metric_heatmap_collection_counter % interval == 0:
for layer_idx in range(self._expert_location_metadata.num_layers):
count_of_layer = (
self._expert_dispatch_collector.eplb_gpu_physical_count.labels(
@@ -728,7 +755,7 @@ class _UtilizationRateAccumulatorMixin(_Accumulator):
if count > 0:
count_of_layer._sum.inc(count * gpu_rank)
count_of_layer._buckets[gpu_rank].inc(count)
self._collection_counter += 1
self._metric_heatmap_collection_counter += 1
class _DequeCollection:
@@ -767,8 +794,9 @@ class _DetailAccumulator(_UtilizationRateAccumulatorMixin):
forward_pass_id: int,
gatherer_key: str,
single_pass_data: Dict,
outputs: Dict[str, Any],
):
super().append(forward_pass_id, gatherer_key, single_pass_data)
super().append(forward_pass_id, gatherer_key, single_pass_data, outputs)
def _process_object(obj):
if isinstance(obj, torch.Tensor):
@@ -824,8 +852,9 @@ class _StatAccumulator(_UtilizationRateAccumulatorMixin):
forward_pass_id: int,
gatherer_key: str,
single_pass_data: Dict,
outputs: Dict[str, Any],
):
super().append(forward_pass_id, gatherer_key, single_pass_data)
super().append(forward_pass_id, gatherer_key, single_pass_data, outputs)
# Can optimize if overhead here is large
self._global_physical_count_of_buffered_step.append(
single_pass_data["global_physical_count"]