diff --git a/tests/distributed/test_eplb_algo.py b/tests/distributed/test_eplb_algo.py index 721132d15b1d..20360c8c61b9 100644 --- a/tests/distributed/test_eplb_algo.py +++ b/tests/distributed/test_eplb_algo.py @@ -5,10 +5,29 @@ import pytest import torch -from vllm.distributed.eplb.eplb_state import compute_logical_maps +from vllm.distributed.eplb.eplb_state import ( + _compute_eplb_load_stats, # pyright: ignore[reportPrivateUsage] + compute_logical_maps, +) from vllm.distributed.eplb.policy.default import DefaultEplbPolicy +def test_eplb_load_stats_reduce_across_ranks(): + num_tokens_per_rank = torch.tensor( + [ + [100, 0, 0, 0], + [100, 0, 0, 0], + ], + dtype=torch.float32, + ) + + avg_tokens, max_tokens = _compute_eplb_load_stats(num_tokens_per_rank) + + assert avg_tokens.item() == 50 + assert max_tokens.item() == 200 + assert (avg_tokens / max_tokens).item() == 0.25 + + def test_basic_rebalance(): """Test basic rebalancing functionality""" # Example from https://github.com/deepseek-ai/eplb diff --git a/vllm/distributed/eplb/eplb_state.py b/vllm/distributed/eplb/eplb_state.py index d34b844bd3d2..267beda9f188 100644 --- a/vllm/distributed/eplb/eplb_state.py +++ b/vllm/distributed/eplb/eplb_state.py @@ -61,6 +61,14 @@ logger = init_logger(__name__) +def _compute_eplb_load_stats( + num_tokens_per_rank: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor]: + avg_tokens = num_tokens_per_rank.mean(dim=1).sum() + max_tokens = num_tokens_per_rank.max(dim=1).values.sum() + return avg_tokens, max_tokens + + @dataclass class EplbStats: """ @@ -584,8 +592,9 @@ def step( # Compute balancedness ratio: # for each layer: # (mean load across ranks) / (max load across ranks) - avg_tokens_tensor = num_tokens_per_rank.mean(dim=0).sum(dim=0) - max_tokens_tensor = num_tokens_per_rank.max(dim=0).values.sum(dim=0) + avg_tokens_tensor, max_tokens_tensor = _compute_eplb_load_stats( + num_tokens_per_rank + ) # Just to make type checker happy tokens_tensors: list[float] = torch.stack(