Report accelerator type in /v1/loads (#32348)
Co-authored-by: Yinghai Lu <yinghai@meta.com>
This commit is contained in:
@@ -20,16 +20,24 @@ metrics for load balancing, monitoring, and capacity planning.
|
|||||||
|
|
||||||
import time
|
import time
|
||||||
from datetime import datetime, timezone
|
from datetime import datetime, timezone
|
||||||
|
from functools import lru_cache
|
||||||
from typing import Optional
|
from typing import Optional
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException
|
from fastapi import APIRouter, Depends, HTTPException
|
||||||
from fastapi.responses import Response
|
from fastapi.responses import Response
|
||||||
|
|
||||||
|
from sglang.srt.utils import get_device_name
|
||||||
from sglang.version import __version__
|
from sglang.version import __version__
|
||||||
|
|
||||||
router = APIRouter()
|
router = APIRouter()
|
||||||
|
|
||||||
|
|
||||||
|
@lru_cache(maxsize=1)
|
||||||
|
def _accelerator_name() -> Optional[str]:
|
||||||
|
"""Accelerator marketing name (e.g. "NVIDIA GB300"), None if unavailable."""
|
||||||
|
return get_device_name()
|
||||||
|
|
||||||
|
|
||||||
def _get_tokenizer_manager():
|
def _get_tokenizer_manager():
|
||||||
"""Dependency to get tokenizer_manager from global state."""
|
"""Dependency to get tokenizer_manager from global state."""
|
||||||
from sglang.srt.entrypoints.http_server import get_global_state
|
from sglang.srt.entrypoints.http_server import get_global_state
|
||||||
@@ -87,7 +95,7 @@ async def get_loads(
|
|||||||
format: Response format - 'json' (default) or 'prometheus'
|
format: Response format - 'json' (default) or 'prometheus'
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
JSON response with timestamp, version, and per-DP-rank loads
|
JSON response with timestamp, version, accelerator, and per-DP-rank loads
|
||||||
"""
|
"""
|
||||||
include_list = [s.strip() for s in include.split(",")] if include else None
|
include_list = [s.strip() for s in include.split(",")] if include else None
|
||||||
|
|
||||||
@@ -119,5 +127,6 @@ async def get_loads(
|
|||||||
return {
|
return {
|
||||||
"timestamp": datetime.now(timezone.utc).isoformat(),
|
"timestamp": datetime.now(timezone.utc).isoformat(),
|
||||||
"version": __version__,
|
"version": __version__,
|
||||||
|
"accelerator": _accelerator_name(),
|
||||||
"loads": loads,
|
"loads": loads,
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -5,9 +5,11 @@ import os
|
|||||||
import tempfile
|
import tempfile
|
||||||
import unittest
|
import unittest
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
from unittest import mock
|
||||||
|
|
||||||
import msgspec.msgpack
|
import msgspec.msgpack
|
||||||
|
|
||||||
|
from sglang.srt.entrypoints import v1_loads
|
||||||
from sglang.srt.entrypoints.v1_loads import get_loads
|
from sglang.srt.entrypoints.v1_loads import get_loads
|
||||||
from sglang.srt.managers.load_snapshot import (
|
from sglang.srt.managers.load_snapshot import (
|
||||||
HEADER_STRUCT,
|
HEADER_STRUCT,
|
||||||
@@ -91,6 +93,19 @@ class TestLoadsResponse(CustomTestCase):
|
|||||||
self.assertEqual(response["loads"][0]["num_waiting_reqs"], 2)
|
self.assertEqual(response["loads"][0]["num_waiting_reqs"], 2)
|
||||||
|
|
||||||
|
|
||||||
|
class TestLoadsAcceleratorField(CustomTestCase):
|
||||||
|
def test_accelerator_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)])
|
||||||
|
|
||||||
|
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")
|
||||||
|
|
||||||
|
|
||||||
class TestGetLoads(CustomTestCase):
|
class TestGetLoads(CustomTestCase):
|
||||||
def test_load_snapshot_wire_format_is_msgpack_slots(self):
|
def test_load_snapshot_wire_format_is_msgpack_slots(self):
|
||||||
path = _temp_path()
|
path = _temp_path()
|
||||||
|
|||||||
Reference in New Issue
Block a user