[Router] Publish cache-aware load state (#38139)
This commit is contained in:
@@ -1,10 +1,9 @@
|
|||||||
"""Per-scheduler load reporting for load-aware routers.
|
"""Per-scheduler load reporting for load-aware routers.
|
||||||
|
|
||||||
Each scheduler publishes a periodic `LoadStat` gauge on its own ZMQ PUB
|
Each scheduler publishes a periodic `LoadStat` gauge on its own ZMQ PUB
|
||||||
socket so out-of-process load-aware routers (e.g. sgl-router's
|
socket so out-of-process load-aware routers can route on real queue depth
|
||||||
`cache_aware_zmq` policy) can route on real queue depth instead of a
|
instead of a router-side in-flight counter. The in-deployment counterpart
|
||||||
router-side in-flight counter. The in-deployment counterpart lives in
|
lives in `sglang.srt.managers.load_snapshot` (SHM / PUSH to node 0), which a router
|
||||||
`sglang.srt.managers.load_snapshot` (SHM / PUSH to node 0), which a router
|
|
||||||
that only knows the worker URL cannot subscribe to; the port is instead
|
that only knows the worker URL cannot subscribe to; the port is instead
|
||||||
advertised via `/server_info` (`runtime_context.describe_kv_events_publisher`).
|
advertised via `/server_info` (`runtime_context.describe_kv_events_publisher`).
|
||||||
The payload is a compact tagged subset of `LoadSnapshot` so the wire
|
The payload is a compact tagged subset of `LoadSnapshot` so the wire
|
||||||
@@ -80,9 +79,12 @@ class LoadStat(
|
|||||||
"""Per-scheduler runtime load snapshot.
|
"""Per-scheduler runtime load snapshot.
|
||||||
|
|
||||||
Wire shape (tag + array_like): ``["LoadStat", num_running_reqs,
|
Wire shape (tag + array_like): ``["LoadStat", num_running_reqs,
|
||||||
num_waiting_reqs, num_tokens, max_total_num_tokens, attn_dp_rank]``. The
|
num_waiting_reqs, num_tokens, max_total_num_tokens, attn_dp_rank,
|
||||||
router reads the four counts; array_like always emits the trailing field
|
num_waiting_uncached_tokens, num_total_tokens, max_running_requests,
|
||||||
(null when unset), so a decoder must tolerate it.
|
total_prefill_uncached_tokens, total_prefill_busy_us]``.
|
||||||
|
|
||||||
|
The first six positions are the #34608-compatible prefix. Current routers
|
||||||
|
ignore appended fields, so extending the payload preserves compatibility.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
num_running_reqs: int
|
num_running_reqs: int
|
||||||
@@ -92,6 +94,11 @@ class LoadStat(
|
|||||||
# attn_dp_rank under DP attention, else the plain dp_rank; informational
|
# attn_dp_rank under DP attention, else the plain dp_rank; informational
|
||||||
# only (the router keys by socket rank). Name follows EventBatch's.
|
# only (the router keys by socket rank). Name follows EventBatch's.
|
||||||
attn_dp_rank: Optional[int] = None
|
attn_dp_rank: Optional[int] = None
|
||||||
|
num_waiting_uncached_tokens: int = 0 # Uncached prompt tokens awaiting prefill
|
||||||
|
num_total_tokens: int = 0 # KV tokens in use plus queued request tokens
|
||||||
|
max_running_requests: int = 0 # Scheduler running-request limit
|
||||||
|
total_prefill_uncached_tokens: int = 0 # Cumulative uncached prefill tokens
|
||||||
|
total_prefill_busy_us: int = 0 # Cumulative prefill step time in microseconds
|
||||||
|
|
||||||
|
|
||||||
def _open_pub_socket(endpoint: str) -> zmq.Socket:
|
def _open_pub_socket(endpoint: str) -> zmq.Socket:
|
||||||
@@ -228,6 +235,11 @@ class SchedulerLoadPublisher:
|
|||||||
load.num_waiting_reqs,
|
load.num_waiting_reqs,
|
||||||
load.num_used_tokens,
|
load.num_used_tokens,
|
||||||
load.max_total_num_tokens,
|
load.max_total_num_tokens,
|
||||||
|
load.num_waiting_uncached_tokens,
|
||||||
|
load.num_total_tokens,
|
||||||
|
load.max_running_requests,
|
||||||
|
load.total_prefill_uncached_tokens,
|
||||||
|
load.total_prefill_busy_us,
|
||||||
)
|
)
|
||||||
if (
|
if (
|
||||||
counts == self._last_counts
|
counts == self._last_counts
|
||||||
@@ -241,6 +253,11 @@ class SchedulerLoadPublisher:
|
|||||||
num_tokens=counts[2],
|
num_tokens=counts[2],
|
||||||
max_total_num_tokens=counts[3],
|
max_total_num_tokens=counts[3],
|
||||||
attn_dp_rank=self._rank,
|
attn_dp_rank=self._rank,
|
||||||
|
num_waiting_uncached_tokens=counts[4],
|
||||||
|
num_total_tokens=counts[5],
|
||||||
|
max_running_requests=counts[6],
|
||||||
|
total_prefill_uncached_tokens=counts[7],
|
||||||
|
total_prefill_busy_us=counts[8],
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
seq = next(self._seq).to_bytes(8, "big")
|
seq = next(self._seq).to_bytes(8, "big")
|
||||||
|
|||||||
@@ -1,11 +1,13 @@
|
|||||||
"""Wire contract and port/rank gating for the LoadStat load snapshot.
|
"""Wire contract and port/rank gating for the LoadStat load snapshot.
|
||||||
|
|
||||||
Locks the msgpack array shape the sgl-router `cache_aware_zmq` policy will
|
Locks the msgpack array shape the sgl-router engine load subscriber in #38108
|
||||||
decode positionally (that consumer lands with the router PR; it is not yet
|
will decode positionally (that consumer is not yet in this tree, so this pins
|
||||||
in this tree, so this pins only the Python side):
|
only the Python side):
|
||||||
|
|
||||||
["LoadStat", num_running_reqs, num_waiting_reqs, num_tokens,
|
["LoadStat", num_running_reqs, num_waiting_reqs, num_tokens,
|
||||||
max_total_num_tokens, attn_dp_rank]
|
max_total_num_tokens, attn_dp_rank,
|
||||||
|
num_waiting_uncached_tokens, num_total_tokens, max_running_requests,
|
||||||
|
total_prefill_uncached_tokens, total_prefill_busy_us]
|
||||||
|
|
||||||
carried as the payload of a three-frame message ``[b"load", BE-i64 seq,
|
carried as the payload of a three-frame message ``[b"load", BE-i64 seq,
|
||||||
payload]``. A field reorder or rename is a silent cross-language break, so
|
payload]``. A field reorder or rename is a silent cross-language break, so
|
||||||
@@ -46,7 +48,7 @@ class TestLoadStatWire(CustomTestCase):
|
|||||||
attn_dp_rank=2,
|
attn_dp_rank=2,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
self.assertEqual(raw.hex(), "96a84c6f6164537461740703cd0400cd200002")
|
self.assertEqual(raw.hex(), "9ba84c6f6164537461740703cd0400cd2000020000000000")
|
||||||
|
|
||||||
def test_loadstat_msgpack_array_shape(self):
|
def test_loadstat_msgpack_array_shape(self):
|
||||||
raw = msgspec.msgpack.Encoder().encode(
|
raw = msgspec.msgpack.Encoder().encode(
|
||||||
@@ -59,10 +61,10 @@ class TestLoadStatWire(CustomTestCase):
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
# tag=True + array_like → [tag, *fields] in declaration order; the
|
# tag=True + array_like → [tag, *fields] in declaration order; the
|
||||||
# router reads the tag + four counts and ignores the trailing field.
|
# router reads the tag and four base fields and ignores the suffix.
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
msgspec.msgpack.Decoder().decode(raw),
|
msgspec.msgpack.Decoder().decode(raw),
|
||||||
["LoadStat", 7, 3, 1024, 8192, 2],
|
["LoadStat", 7, 3, 1024, 8192, 2, 0, 0, 0, 0, 0],
|
||||||
)
|
)
|
||||||
|
|
||||||
def test_loadstat_tag_is_class_name(self):
|
def test_loadstat_tag_is_class_name(self):
|
||||||
@@ -78,8 +80,9 @@ class TestLoadStatWire(CustomTestCase):
|
|||||||
)
|
)
|
||||||
decoded = msgspec.msgpack.Decoder().decode(raw)
|
decoded = msgspec.msgpack.Decoder().decode(raw)
|
||||||
# LoadStat sets no omit_defaults, so the trailing field is always
|
# LoadStat sets no omit_defaults, so the trailing field is always
|
||||||
# emitted (null when unset); a decoder must tolerate it.
|
# emitted (null when unset); a decoder must tolerate it. The new suffix
|
||||||
self.assertEqual(decoded, ["LoadStat", 0, 0, 0, 0, None])
|
# fields are always emitted as well.
|
||||||
|
self.assertEqual(decoded, ["LoadStat", 0, 0, 0, 0, None, 0, 0, 0, 0, 0])
|
||||||
|
|
||||||
|
|
||||||
ZMQ_ENDPOINT = '{"publisher": "zmq", "endpoint": "tcp://*:5557"}'
|
ZMQ_ENDPOINT = '{"publisher": "zmq", "endpoint": "tcp://*:5557"}'
|
||||||
@@ -336,6 +339,11 @@ class TestLoadPublisherGating(CustomTestCase):
|
|||||||
num_waiting_reqs=2,
|
num_waiting_reqs=2,
|
||||||
num_used_tokens=3,
|
num_used_tokens=3,
|
||||||
max_total_num_tokens=4,
|
max_total_num_tokens=4,
|
||||||
|
num_waiting_uncached_tokens=5,
|
||||||
|
num_total_tokens=6,
|
||||||
|
max_running_requests=7,
|
||||||
|
total_prefill_uncached_tokens=8,
|
||||||
|
total_prefill_busy_us=9,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -356,6 +364,11 @@ class TestLoadPublisherGating(CustomTestCase):
|
|||||||
num_waiting_reqs=2,
|
num_waiting_reqs=2,
|
||||||
num_used_tokens=3,
|
num_used_tokens=3,
|
||||||
max_total_num_tokens=4,
|
max_total_num_tokens=4,
|
||||||
|
num_waiting_uncached_tokens=5,
|
||||||
|
num_total_tokens=6,
|
||||||
|
max_running_requests=7,
|
||||||
|
total_prefill_uncached_tokens=8,
|
||||||
|
total_prefill_busy_us=9,
|
||||||
)
|
)
|
||||||
pub.publish_load_stat(provider, force=True, snapshot=snap)
|
pub.publish_load_stat(provider, force=True, snapshot=snap)
|
||||||
provider.assert_not_called()
|
provider.assert_not_called()
|
||||||
@@ -372,7 +385,7 @@ class TestLoadPublisherGating(CustomTestCase):
|
|||||||
self.assertEqual(seq, (0).to_bytes(8, "big"))
|
self.assertEqual(seq, (0).to_bytes(8, "big"))
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
msgspec.msgpack.Decoder().decode(payload),
|
msgspec.msgpack.Decoder().decode(payload),
|
||||||
["LoadStat", 1, 2, 3, 4, 0],
|
["LoadStat", 1, 2, 3, 4, 0, 5, 6, 7, 8, 9],
|
||||||
)
|
)
|
||||||
|
|
||||||
def test_unchanged_stat_is_deduped_to_the_heartbeat(self):
|
def test_unchanged_stat_is_deduped_to_the_heartbeat(self):
|
||||||
@@ -475,6 +488,11 @@ class TestLoadStatIntegration(CustomTestCase):
|
|||||||
num_waiting_reqs=3,
|
num_waiting_reqs=3,
|
||||||
num_used_tokens=1024,
|
num_used_tokens=1024,
|
||||||
max_total_num_tokens=8192,
|
max_total_num_tokens=8192,
|
||||||
|
num_waiting_uncached_tokens=512,
|
||||||
|
num_total_tokens=4096,
|
||||||
|
max_running_requests=64,
|
||||||
|
total_prefill_uncached_tokens=20_000,
|
||||||
|
total_prefill_busy_us=2_000_000,
|
||||||
)
|
)
|
||||||
# PUB/SUB drops messages sent before the subscription propagates, so
|
# PUB/SUB drops messages sent before the subscription propagates, so
|
||||||
# re-publish until one lands (heartbeat reset each pass).
|
# re-publish until one lands (heartbeat reset each pass).
|
||||||
@@ -492,7 +510,7 @@ class TestLoadStatIntegration(CustomTestCase):
|
|||||||
self.assertEqual(len(seq), 8)
|
self.assertEqual(len(seq), 8)
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
msgspec.msgpack.Decoder().decode(payload),
|
msgspec.msgpack.Decoder().decode(payload),
|
||||||
["LoadStat", 7, 3, 1024, 8192, 0],
|
["LoadStat", 7, 3, 1024, 8192, 0, 512, 4096, 64, 20_000, 2_000_000],
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user