[Router] Report accelerator count in /v1/loads (#33548)
Co-authored-by: Yinghai Lu <yinghai@meta.com>
This commit is contained in:
co-authored by
Yinghai Lu
parent
8f2a3ad6d7
commit
d257b58e67
@@ -38,6 +38,19 @@ def _accelerator_name() -> Optional[str]:
|
||||
return get_device_name()
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def _num_accelerators_per_dp_rank(
|
||||
tp_size: int,
|
||||
pp_size: int,
|
||||
dp_size: int,
|
||||
enable_dp_attention: bool,
|
||||
) -> int:
|
||||
num_accelerators = tp_size * pp_size
|
||||
if enable_dp_attention:
|
||||
num_accelerators //= dp_size
|
||||
return num_accelerators
|
||||
|
||||
|
||||
def _get_tokenizer_manager():
|
||||
"""Dependency to get tokenizer_manager from global state."""
|
||||
from sglang.srt.entrypoints.http_server import get_global_state
|
||||
@@ -95,7 +108,8 @@ async def get_loads(
|
||||
format: Response format - 'json' (default) or 'prometheus'
|
||||
|
||||
Returns:
|
||||
JSON response with timestamp, version, accelerator, and per-DP-rank loads
|
||||
JSON response with timestamp, version, accelerator metadata, and
|
||||
per-DP-rank loads
|
||||
"""
|
||||
include_list = [s.strip() for s in include.split(",")] if include else None
|
||||
|
||||
@@ -128,5 +142,11 @@ async def get_loads(
|
||||
"timestamp": datetime.now(timezone.utc).isoformat(),
|
||||
"version": __version__,
|
||||
"accelerator": _accelerator_name(),
|
||||
"num_accelerators": _num_accelerators_per_dp_rank(
|
||||
tokenizer_manager.server_args.tp_size,
|
||||
tokenizer_manager.server_args.pp_size,
|
||||
tokenizer_manager.server_args.dp_size,
|
||||
tokenizer_manager.server_args.enable_dp_attention,
|
||||
),
|
||||
"loads": loads,
|
||||
}
|
||||
|
||||
@@ -58,8 +58,21 @@ class _FakeTokenizerManager(TokenizerControlMixin):
|
||||
class _FakeHttpTokenizerManager:
|
||||
metrics_collector = None
|
||||
|
||||
def __init__(self, loads):
|
||||
def __init__(
|
||||
self,
|
||||
loads,
|
||||
tp_size=1,
|
||||
dp_size=1,
|
||||
pp_size=1,
|
||||
enable_dp_attention=False,
|
||||
):
|
||||
self.loads = loads
|
||||
self.server_args = SimpleNamespace(
|
||||
tp_size=tp_size,
|
||||
dp_size=dp_size,
|
||||
pp_size=pp_size,
|
||||
enable_dp_attention=enable_dp_attention,
|
||||
)
|
||||
|
||||
async def get_loads(self, include=None, dp_rank=None):
|
||||
results = []
|
||||
@@ -94,16 +107,29 @@ class TestLoadsResponse(CustomTestCase):
|
||||
|
||||
|
||||
class TestLoadsAcceleratorField(CustomTestCase):
|
||||
def test_accelerator_reported_in_json(self):
|
||||
def test_accelerator_metadata_reported_in_json(self):
|
||||
"""Guards the response contract: the JSON envelope carries an
|
||||
"accelerator" field with the detected device name."""
|
||||
manager = _FakeHttpTokenizerManager([LoadSnapshot(dp_rank=0)])
|
||||
accelerator name and the accelerator count for each DP rank."""
|
||||
manager = _FakeHttpTokenizerManager([LoadSnapshot(dp_rank=0)], tp_size=16)
|
||||
|
||||
with mock.patch.object(
|
||||
v1_loads, "_accelerator_name", return_value="NVIDIA GB300"
|
||||
):
|
||||
response = asyncio.run(get_loads(tokenizer_manager=manager))
|
||||
self.assertEqual(response["accelerator"], "NVIDIA GB300")
|
||||
self.assertEqual(response["num_accelerators"], 16)
|
||||
|
||||
def test_dp_attention_reports_accelerators_per_dp_rank(self):
|
||||
manager = _FakeHttpTokenizerManager(
|
||||
[LoadSnapshot(dp_rank=rank) for rank in range(8)],
|
||||
tp_size=8,
|
||||
dp_size=8,
|
||||
enable_dp_attention=True,
|
||||
)
|
||||
|
||||
response = asyncio.run(get_loads(tokenizer_manager=manager))
|
||||
|
||||
self.assertEqual(response["num_accelerators"], 1)
|
||||
|
||||
|
||||
class TestGetLoads(CustomTestCase):
|
||||
|
||||
Reference in New Issue
Block a user