Optimize get load calls (/v1/loads) using shared-memory load snapshots (#26348)
Co-authored-by: cctry <cctry@meta.com>
This commit is contained in:
@@ -0,0 +1,96 @@
|
||||
"""Integration tests for load snapshot with real servers.
|
||||
|
||||
Tests [no dp, normal dp] x [zmq, shm] by launching real servers
|
||||
and querying /v1/loads.
|
||||
"""
|
||||
|
||||
import json
|
||||
import time
|
||||
import unittest
|
||||
import urllib.request
|
||||
|
||||
from sglang.srt.utils import kill_process_tree
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
from sglang.test.test_utils import (
|
||||
DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
|
||||
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||
DEFAULT_URL_FOR_TEST,
|
||||
CustomTestCase,
|
||||
popen_launch_server,
|
||||
)
|
||||
|
||||
register_cuda_ci(est_time=300, stage="base-b", runner_config="2-gpu-large")
|
||||
|
||||
|
||||
def _query_loads(base_url, retries=5, interval=2.0):
|
||||
url = f"{base_url}/v1/loads"
|
||||
for attempt in range(retries):
|
||||
try:
|
||||
resp = urllib.request.urlopen(url, timeout=5)
|
||||
data = json.loads(resp.read())
|
||||
if data.get("loads"):
|
||||
return data
|
||||
except Exception:
|
||||
pass
|
||||
if attempt < retries - 1:
|
||||
time.sleep(interval)
|
||||
try:
|
||||
resp = urllib.request.urlopen(url, timeout=5)
|
||||
return json.loads(resp.read())
|
||||
except Exception:
|
||||
return {"loads": []}
|
||||
|
||||
|
||||
def _launch_and_check(test_case, other_args=None, env=None, expected_dp_size=1):
|
||||
process = popen_launch_server(
|
||||
DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
|
||||
DEFAULT_URL_FOR_TEST,
|
||||
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||
other_args=other_args or [],
|
||||
env=env,
|
||||
)
|
||||
try:
|
||||
time.sleep(5)
|
||||
data = _query_loads(DEFAULT_URL_FOR_TEST)
|
||||
loads = data.get("loads", [])
|
||||
test_case.assertGreater(len(loads), 0, f"Expected non-empty loads, got: {data}")
|
||||
test_case.assertEqual(len(loads), expected_dp_size)
|
||||
dp_ranks = sorted(l["dp_rank"] for l in loads)
|
||||
test_case.assertEqual(dp_ranks, list(range(expected_dp_size)))
|
||||
for load in loads:
|
||||
test_case.assertGreater(load["max_total_num_tokens"], 0)
|
||||
finally:
|
||||
kill_process_tree(process.pid)
|
||||
|
||||
|
||||
class TestLoadSnapshotNoDP(CustomTestCase):
|
||||
def test_shm_backend(self):
|
||||
_launch_and_check(self, expected_dp_size=1)
|
||||
|
||||
def test_zmq_backend(self):
|
||||
_launch_and_check(
|
||||
self,
|
||||
env={"SGLANG_LOAD_SNAPSHOT_USE_ZMQ": "1"},
|
||||
expected_dp_size=1,
|
||||
)
|
||||
|
||||
|
||||
class TestLoadSnapshotNormalDP(CustomTestCase):
|
||||
def test_shm_backend(self):
|
||||
_launch_and_check(
|
||||
self,
|
||||
other_args=["--dp", "2"],
|
||||
expected_dp_size=2,
|
||||
)
|
||||
|
||||
def test_zmq_backend(self):
|
||||
_launch_and_check(
|
||||
self,
|
||||
other_args=["--dp", "2"],
|
||||
env={"SGLANG_LOAD_SNAPSHOT_USE_ZMQ": "1"},
|
||||
expected_dp_size=2,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -1,75 +1,185 @@
|
||||
"""Unit tests for /v1/loads _compute_aggregate.
|
||||
|
||||
Narrow scope: lock in the semantic of new aggregate keys added by this PR
|
||||
(total_used_tokens vs total_tokens). Trivial helpers (dict filtering,
|
||||
zero-init branch) are not covered — they would just restate Python.
|
||||
"""
|
||||
"""Unit tests for /v1/loads load snapshot response behavior."""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import tempfile
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
|
||||
from sglang.srt.entrypoints.v1_loads import _compute_aggregate
|
||||
import msgspec.msgpack
|
||||
|
||||
from sglang.srt.entrypoints.v1_loads import get_loads
|
||||
from sglang.srt.managers.load_snapshot import (
|
||||
HEADER_STRUCT,
|
||||
MAGIC,
|
||||
SLOT_LEN_STRUCT,
|
||||
SLOT_SIZE,
|
||||
VERSION,
|
||||
LoadSnapshot,
|
||||
ShmLoadSnapshotReader,
|
||||
ShmLoadSnapshotWriter,
|
||||
slot_offset,
|
||||
)
|
||||
from sglang.srt.managers.tokenizer_control_mixin import TokenizerControlMixin
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel
|
||||
|
||||
maybe_stub_sgl_kernel()
|
||||
|
||||
|
||||
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
|
||||
|
||||
|
||||
def _load(
|
||||
*,
|
||||
dp_rank=0,
|
||||
running=0,
|
||||
waiting=0,
|
||||
used=0,
|
||||
total=0,
|
||||
token_usage=0.0,
|
||||
throughput=0.0,
|
||||
utilization=0.0,
|
||||
):
|
||||
return {
|
||||
"dp_rank": dp_rank,
|
||||
"num_running_reqs": running,
|
||||
"num_waiting_reqs": waiting,
|
||||
"num_used_tokens": used,
|
||||
"num_total_tokens": total,
|
||||
"token_usage": token_usage,
|
||||
"gen_throughput": throughput,
|
||||
"utilization": utilization,
|
||||
}
|
||||
def _temp_path() -> str:
|
||||
fd, path = tempfile.mkstemp()
|
||||
os.close(fd)
|
||||
os.unlink(path)
|
||||
return path
|
||||
|
||||
|
||||
class TestComputeAggregate(CustomTestCase):
|
||||
def test_multi_dp_rank_sums(self):
|
||||
agg = _compute_aggregate(
|
||||
class _FakeTokenizerManager(TokenizerControlMixin):
|
||||
def __init__(self, reader, dp_size: int):
|
||||
self.load_snapshot_reader = reader
|
||||
self.server_args = SimpleNamespace(
|
||||
dp_size=dp_size,
|
||||
enable_dp_attention=False,
|
||||
nnodes=1,
|
||||
)
|
||||
|
||||
def auto_create_handle_loop(self):
|
||||
pass
|
||||
|
||||
|
||||
class _FakeHttpTokenizerManager:
|
||||
metrics_collector = None
|
||||
|
||||
def __init__(self, loads):
|
||||
self.loads = loads
|
||||
|
||||
async def get_loads(self, include=None, dp_rank=None):
|
||||
results = []
|
||||
for load in self.loads:
|
||||
if dp_rank is not None and load.dp_rank != dp_rank:
|
||||
continue
|
||||
results.append(load)
|
||||
return results
|
||||
|
||||
|
||||
class TestLoadsResponse(CustomTestCase):
|
||||
def test_response_omits_server_side_aggregate_and_redundant_fields(self):
|
||||
manager = _FakeHttpTokenizerManager(
|
||||
[
|
||||
_load(dp_rank=0, running=3, waiting=1, used=50, total=70),
|
||||
_load(dp_rank=1, running=5, waiting=2, used=80, total=100),
|
||||
_load(dp_rank=2, running=0, waiting=4, used=0, total=40),
|
||||
LoadSnapshot(
|
||||
dp_rank=0,
|
||||
num_running_reqs=3,
|
||||
num_waiting_reqs=2,
|
||||
num_total_tokens=256,
|
||||
)
|
||||
]
|
||||
)
|
||||
self.assertEqual(agg["total_running_reqs"], 8)
|
||||
self.assertEqual(agg["total_waiting_reqs"], 7)
|
||||
self.assertEqual(agg["total_reqs"], 15)
|
||||
self.assertEqual(agg["total_used_tokens"], 130)
|
||||
self.assertEqual(agg["total_tokens"], 210)
|
||||
|
||||
def test_averages_over_dp_count(self):
|
||||
agg = _compute_aggregate(
|
||||
[
|
||||
_load(token_usage=0.6, throughput=100.0, utilization=0.5),
|
||||
_load(token_usage=0.8, throughput=200.0, utilization=0.7),
|
||||
]
|
||||
)
|
||||
self.assertAlmostEqual(agg["avg_token_usage"], 0.7)
|
||||
self.assertAlmostEqual(agg["avg_throughput"], 150.0)
|
||||
self.assertAlmostEqual(agg["avg_utilization"], 0.6)
|
||||
response = asyncio.run(get_loads(tokenizer_manager=manager))
|
||||
|
||||
def test_total_tokens_differs_from_total_used_tokens(self):
|
||||
# Regression: total_tokens sums num_total_tokens, NOT num_used_tokens.
|
||||
# Gateway reads aggregate.total_tokens for DP load estimation, so a
|
||||
# silent swap would under-report load.
|
||||
agg = _compute_aggregate([_load(used=10, total=30), _load(used=20, total=45)])
|
||||
self.assertEqual(agg["total_used_tokens"], 30)
|
||||
self.assertEqual(agg["total_tokens"], 75)
|
||||
self.assertNotIn("dp_rank_count", response)
|
||||
self.assertNotIn("aggregate", response)
|
||||
self.assertEqual(len(response["loads"]), 1)
|
||||
self.assertNotIn("num_total_reqs", response["loads"][0])
|
||||
self.assertEqual(response["loads"][0]["num_running_reqs"], 3)
|
||||
self.assertEqual(response["loads"][0]["num_waiting_reqs"], 2)
|
||||
|
||||
|
||||
class TestGetLoads(CustomTestCase):
|
||||
def test_load_snapshot_wire_format_is_msgpack_slots(self):
|
||||
path = _temp_path()
|
||||
writer = ShmLoadSnapshotWriter(path, dp_size=2, dp_rank=1)
|
||||
try:
|
||||
writer.write(
|
||||
LoadSnapshot(
|
||||
dp_rank=1,
|
||||
num_running_reqs=3,
|
||||
num_waiting_reqs=2,
|
||||
token_usage=0.25,
|
||||
)
|
||||
)
|
||||
|
||||
with open(path, "rb") as f:
|
||||
data = f.read()
|
||||
|
||||
self.assertEqual(len(data), HEADER_STRUCT.size + 2 * SLOT_SIZE)
|
||||
magic, version, dp_size, slot_size = HEADER_STRUCT.unpack_from(data, 0)
|
||||
self.assertEqual(magic, MAGIC)
|
||||
self.assertEqual(version, VERSION)
|
||||
self.assertEqual(dp_size, 2)
|
||||
self.assertEqual(slot_size, SLOT_SIZE)
|
||||
|
||||
offset = slot_offset(1, slot_size)
|
||||
(payload_len,) = SLOT_LEN_STRUCT.unpack_from(data, offset)
|
||||
payload_start = offset + SLOT_LEN_STRUCT.size
|
||||
payload = data[payload_start : payload_start + payload_len]
|
||||
decoded = msgspec.msgpack.decode(payload)
|
||||
|
||||
self.assertEqual(decoded["dp_rank"], 1)
|
||||
self.assertEqual(decoded["num_running_reqs"], 3)
|
||||
self.assertEqual(decoded["num_waiting_reqs"], 2)
|
||||
self.assertEqual(decoded["token_usage"], 0.25)
|
||||
finally:
|
||||
writer.close()
|
||||
if os.path.exists(path):
|
||||
os.unlink(path)
|
||||
|
||||
def test_reads_snapshot_and_filters_sections(self):
|
||||
path = _temp_path()
|
||||
writer = ShmLoadSnapshotWriter(path, dp_size=1, dp_rank=0)
|
||||
reader = ShmLoadSnapshotReader(path, dp_size=1)
|
||||
try:
|
||||
initial_load = reader.read(0)
|
||||
self.assertIsNotNone(initial_load)
|
||||
self.assertEqual(initial_load.num_total_tokens, 0)
|
||||
|
||||
writer.write(
|
||||
LoadSnapshot(
|
||||
dp_rank=0,
|
||||
timestamp=1.25,
|
||||
num_running_reqs=3,
|
||||
num_waiting_reqs=2,
|
||||
num_used_tokens=128,
|
||||
num_total_tokens=256,
|
||||
max_total_num_tokens=4096,
|
||||
token_usage=0.125,
|
||||
gen_throughput=99.5,
|
||||
cache_hit_rate=0.75,
|
||||
utilization=0.5,
|
||||
max_running_requests=128,
|
||||
has_disaggregation=1,
|
||||
disagg_mode=2,
|
||||
decode_transfer_queue_reqs=4,
|
||||
has_queues=1,
|
||||
queue_waiting=2,
|
||||
queue_grammar=1,
|
||||
queue_paused=0,
|
||||
queue_retracted=3,
|
||||
)
|
||||
)
|
||||
|
||||
manager = _FakeTokenizerManager(reader, dp_size=1)
|
||||
loads = asyncio.run(manager.get_loads(include=["core"], dp_rank=0))
|
||||
|
||||
self.assertEqual(len(loads), 1)
|
||||
self.assertEqual(loads[0].num_total_tokens, 256)
|
||||
|
||||
d = loads[0].to_dict({"core"})
|
||||
self.assertNotIn("disaggregation", d)
|
||||
self.assertNotIn("queues", d)
|
||||
|
||||
loads_all = asyncio.run(manager.get_loads(include=["all"], dp_rank=0))
|
||||
d_all = loads_all[0].to_dict()
|
||||
self.assertIn("disaggregation", d_all)
|
||||
self.assertIn("queues", d_all)
|
||||
finally:
|
||||
reader.close()
|
||||
writer.close()
|
||||
if os.path.exists(path):
|
||||
os.unlink(path)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -11,11 +11,12 @@ if a scheduler starts reading another attr. `maybe_external_dp_rank_routing`
|
||||
is exercised as the real method, no mock.
|
||||
"""
|
||||
|
||||
import dataclasses
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import msgspec.structs
|
||||
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel
|
||||
|
||||
@@ -26,29 +27,20 @@ from sglang.srt.managers.data_parallel_controller import (
|
||||
DPBudget,
|
||||
LoadBalanceMethod,
|
||||
)
|
||||
from sglang.srt.managers.io_struct import GetLoadsReqOutput, WatchLoadUpdateReq
|
||||
from sglang.srt.managers.load_snapshot import LoadSnapshot
|
||||
|
||||
register_cpu_ci(est_time=11, suite="base-a-test-cpu")
|
||||
|
||||
|
||||
_BASE_LOAD = GetLoadsReqOutput(
|
||||
dp_rank=0,
|
||||
timestamp=0.0,
|
||||
num_running_reqs=0,
|
||||
num_waiting_reqs=0,
|
||||
num_used_tokens=0,
|
||||
num_total_tokens=0,
|
||||
_BASE_LOAD = msgspec.structs.replace(
|
||||
LoadSnapshot(dp_rank=0),
|
||||
max_total_num_tokens=4096,
|
||||
token_usage=0.0,
|
||||
gen_throughput=0.0,
|
||||
cache_hit_rate=0.0,
|
||||
utilization=0.0,
|
||||
max_running_requests=128,
|
||||
)
|
||||
|
||||
|
||||
def _load(**overrides) -> GetLoadsReqOutput:
|
||||
return dataclasses.replace(_BASE_LOAD, **overrides)
|
||||
def _load(**overrides) -> LoadSnapshot:
|
||||
return msgspec.structs.replace(_BASE_LOAD, **overrides)
|
||||
|
||||
|
||||
def _make_controller(dp_size: int) -> DataParallelController:
|
||||
@@ -74,44 +66,52 @@ class TestDPBudgetUpdateBudget(CustomTestCase):
|
||||
def test_maps_running_plus_waiting_to_total_requests(self):
|
||||
budget = DPBudget(dp_size=2)
|
||||
budget.update_budget(
|
||||
WatchLoadUpdateReq(
|
||||
loads=[
|
||||
_load(dp_rank=0, num_running_reqs=3, num_waiting_reqs=2),
|
||||
_load(dp_rank=1, num_running_reqs=5, num_waiting_reqs=1),
|
||||
]
|
||||
)
|
||||
[
|
||||
_load(dp_rank=0, timestamp=1.0, num_running_reqs=3, num_waiting_reqs=2),
|
||||
_load(dp_rank=1, timestamp=1.0, num_running_reqs=5, num_waiting_reqs=1),
|
||||
]
|
||||
)
|
||||
self.assertEqual(budget.total_requests, [5, 6])
|
||||
|
||||
def test_maps_num_total_tokens_not_num_used_tokens(self):
|
||||
# Reads num_total_tokens (used + pending prefill), NOT num_used_tokens.
|
||||
# A silent swap here would break DP balance for long-prompt workloads.
|
||||
budget = DPBudget(dp_size=2)
|
||||
budget.update_budget(
|
||||
WatchLoadUpdateReq(
|
||||
loads=[
|
||||
_load(dp_rank=0, num_used_tokens=100, num_total_tokens=150),
|
||||
_load(dp_rank=1, num_used_tokens=80, num_total_tokens=80),
|
||||
]
|
||||
)
|
||||
[
|
||||
_load(
|
||||
dp_rank=0, timestamp=1.0, num_used_tokens=100, num_total_tokens=150
|
||||
),
|
||||
_load(
|
||||
dp_rank=1, timestamp=1.0, num_used_tokens=80, num_total_tokens=80
|
||||
),
|
||||
]
|
||||
)
|
||||
self.assertEqual(budget.total_tokens, [150, 80])
|
||||
|
||||
def test_partial_update_only_affects_reported_rank(self):
|
||||
budget = DPBudget(dp_size=3)
|
||||
budget.total_requests = [10, 20, 30]
|
||||
budget.total_tokens = [100, 200, 300]
|
||||
budget.update_budget(
|
||||
WatchLoadUpdateReq(
|
||||
loads=[
|
||||
_load(
|
||||
dp_rank=1,
|
||||
num_running_reqs=1,
|
||||
num_waiting_reqs=1,
|
||||
num_total_tokens=50,
|
||||
)
|
||||
]
|
||||
)
|
||||
[
|
||||
_load(
|
||||
dp_rank=0, timestamp=1.0, num_running_reqs=10, num_total_tokens=100
|
||||
),
|
||||
_load(
|
||||
dp_rank=1, timestamp=1.0, num_running_reqs=20, num_total_tokens=200
|
||||
),
|
||||
_load(
|
||||
dp_rank=2, timestamp=1.0, num_running_reqs=30, num_total_tokens=300
|
||||
),
|
||||
]
|
||||
)
|
||||
budget.update_budget(
|
||||
[
|
||||
_load(
|
||||
dp_rank=1,
|
||||
timestamp=2.0,
|
||||
num_running_reqs=1,
|
||||
num_waiting_reqs=1,
|
||||
num_total_tokens=50,
|
||||
)
|
||||
]
|
||||
)
|
||||
self.assertEqual(budget.total_requests, [10, 2, 30])
|
||||
self.assertEqual(budget.total_tokens, [100, 50, 300])
|
||||
|
||||
@@ -0,0 +1,355 @@
|
||||
"""Unit tests for LoadSnapshot SHM and ZMQ backends."""
|
||||
|
||||
import os
|
||||
import tempfile
|
||||
import time
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
|
||||
from sglang.srt.managers.load_snapshot import (
|
||||
LoadSnapshot,
|
||||
ShmLoadSnapshotReader,
|
||||
ShmLoadSnapshotWriter,
|
||||
ZmqLoadSnapshotWriter,
|
||||
ZmqShmLoadSnapshotReader,
|
||||
_zmq_addr_for,
|
||||
create_load_snapshot_reader,
|
||||
create_load_snapshot_writer,
|
||||
should_use_zmq,
|
||||
)
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel
|
||||
|
||||
maybe_stub_sgl_kernel()
|
||||
|
||||
|
||||
register_cpu_ci(est_time=15, suite="base-a-test-cpu")
|
||||
|
||||
|
||||
def _temp_path() -> str:
|
||||
fd, path = tempfile.mkstemp()
|
||||
os.close(fd)
|
||||
os.unlink(path)
|
||||
return path
|
||||
|
||||
|
||||
def _ipc_addr() -> str:
|
||||
fd, path = tempfile.mkstemp(prefix="sglang_test_zmq_", suffix=".sock")
|
||||
os.close(fd)
|
||||
os.unlink(path)
|
||||
return f"ipc://{path}"
|
||||
|
||||
|
||||
def _warmup_zmq(writers, reader, attempts=20, interval=0.05):
|
||||
"""Send warmup messages until the reader receives from all writers."""
|
||||
expected = {w.dp_rank for w in writers}
|
||||
received = set()
|
||||
for _ in range(attempts):
|
||||
for w in writers:
|
||||
w.write(LoadSnapshot(dp_rank=w.dp_rank, timestamp=-1.0, num_running_reqs=0))
|
||||
time.sleep(interval)
|
||||
for rank in expected:
|
||||
load = reader.read(rank)
|
||||
if load is not None:
|
||||
received.add(rank)
|
||||
if received >= expected:
|
||||
return
|
||||
raise RuntimeError(f"warmup failed: expected {expected}, received {received}")
|
||||
|
||||
|
||||
class TestShmRoundTrip(CustomTestCase):
|
||||
def test_single_rank_write_read(self):
|
||||
path = _temp_path()
|
||||
writer = ShmLoadSnapshotWriter(path, dp_size=1, dp_rank=0)
|
||||
reader = ShmLoadSnapshotReader(path, dp_size=1)
|
||||
try:
|
||||
writer.write(LoadSnapshot(dp_rank=0, num_running_reqs=5, timestamp=1.0))
|
||||
load = reader.read(0)
|
||||
self.assertIsNotNone(load)
|
||||
self.assertEqual(load.num_running_reqs, 5)
|
||||
self.assertEqual(load.timestamp, 1.0)
|
||||
finally:
|
||||
reader.close()
|
||||
writer.close()
|
||||
if os.path.exists(path):
|
||||
os.unlink(path)
|
||||
|
||||
def test_multi_rank_write_read_all(self):
|
||||
path = _temp_path()
|
||||
writers = []
|
||||
try:
|
||||
for rank in range(4):
|
||||
w = ShmLoadSnapshotWriter(path, dp_size=4, dp_rank=rank)
|
||||
w.write(
|
||||
LoadSnapshot(
|
||||
dp_rank=rank,
|
||||
num_running_reqs=rank * 10,
|
||||
timestamp=1.0,
|
||||
)
|
||||
)
|
||||
writers.append(w)
|
||||
|
||||
reader = ShmLoadSnapshotReader(path, dp_size=4)
|
||||
loads = reader.read_all()
|
||||
self.assertEqual(len(loads), 4)
|
||||
for i, load in enumerate(loads):
|
||||
self.assertEqual(load.dp_rank, i)
|
||||
self.assertEqual(load.num_running_reqs, i * 10)
|
||||
reader.close()
|
||||
finally:
|
||||
for w in writers:
|
||||
w.close()
|
||||
if os.path.exists(path):
|
||||
os.unlink(path)
|
||||
|
||||
def test_reader_empty_before_writer(self):
|
||||
path = _temp_path()
|
||||
reader = ShmLoadSnapshotReader(path, dp_size=2)
|
||||
self.assertEqual(reader.read_all(), [])
|
||||
self.assertIsNone(reader.read(0))
|
||||
reader.close()
|
||||
|
||||
|
||||
class TestZmqRoundTrip(CustomTestCase):
|
||||
def test_single_rank_zmq_to_shm(self):
|
||||
shm_path = _temp_path()
|
||||
addr = _ipc_addr()
|
||||
reader = ZmqShmLoadSnapshotReader(addr, shm_path, dp_size=2)
|
||||
writer = ZmqLoadSnapshotWriter(addr, dp_size=2, dp_rank=0)
|
||||
try:
|
||||
_warmup_zmq([writer], reader)
|
||||
|
||||
writer.write(LoadSnapshot(dp_rank=0, num_running_reqs=7, timestamp=2.0))
|
||||
time.sleep(0.05)
|
||||
|
||||
load = reader.read(0)
|
||||
self.assertIsNotNone(load)
|
||||
self.assertEqual(load.num_running_reqs, 7)
|
||||
self.assertEqual(load.timestamp, 2.0)
|
||||
finally:
|
||||
writer.close()
|
||||
reader.close()
|
||||
if os.path.exists(shm_path):
|
||||
os.unlink(shm_path)
|
||||
|
||||
def test_multi_rank_zmq(self):
|
||||
shm_path = _temp_path()
|
||||
addr = _ipc_addr()
|
||||
dp_size = 4
|
||||
reader = ZmqShmLoadSnapshotReader(addr, shm_path, dp_size)
|
||||
writers = []
|
||||
try:
|
||||
for rank in range(dp_size):
|
||||
w = ZmqLoadSnapshotWriter(addr, dp_size, dp_rank=rank)
|
||||
writers.append(w)
|
||||
|
||||
_warmup_zmq(writers, reader)
|
||||
|
||||
for rank, w in enumerate(writers):
|
||||
w.write(
|
||||
LoadSnapshot(dp_rank=rank, num_running_reqs=rank + 1, timestamp=3.0)
|
||||
)
|
||||
time.sleep(0.05)
|
||||
|
||||
loads = reader.read_all()
|
||||
self.assertEqual(len(loads), dp_size)
|
||||
for load in loads:
|
||||
self.assertEqual(load.num_running_reqs, load.dp_rank + 1)
|
||||
finally:
|
||||
for w in writers:
|
||||
w.close()
|
||||
reader.close()
|
||||
if os.path.exists(shm_path):
|
||||
os.unlink(shm_path)
|
||||
|
||||
def test_read_returns_latest(self):
|
||||
shm_path = _temp_path()
|
||||
addr = _ipc_addr()
|
||||
reader = ZmqShmLoadSnapshotReader(addr, shm_path, dp_size=1)
|
||||
writer = ZmqLoadSnapshotWriter(addr, dp_size=1, dp_rank=0)
|
||||
try:
|
||||
_warmup_zmq([writer], reader)
|
||||
|
||||
for i in range(10):
|
||||
writer.write(
|
||||
LoadSnapshot(dp_rank=0, num_running_reqs=i, timestamp=float(i))
|
||||
)
|
||||
time.sleep(0.05)
|
||||
|
||||
load = reader.read(0)
|
||||
self.assertIsNotNone(load)
|
||||
self.assertEqual(load.num_running_reqs, 9)
|
||||
self.assertEqual(load.timestamp, 9.0)
|
||||
finally:
|
||||
writer.close()
|
||||
reader.close()
|
||||
if os.path.exists(shm_path):
|
||||
os.unlink(shm_path)
|
||||
|
||||
def test_zmq_writer_noblock_without_reader(self):
|
||||
addr = _ipc_addr()
|
||||
writer = ZmqLoadSnapshotWriter(addr, dp_size=1, dp_rank=0)
|
||||
try:
|
||||
writer.write(LoadSnapshot(dp_rank=0, num_running_reqs=1, timestamp=1.0))
|
||||
finally:
|
||||
writer.close()
|
||||
ipc_path = addr[len("ipc://") :]
|
||||
if os.path.exists(ipc_path):
|
||||
os.unlink(ipc_path)
|
||||
|
||||
def test_reader_ipc_cleanup(self):
|
||||
addr = _ipc_addr()
|
||||
shm_path = _temp_path()
|
||||
ipc_path = addr[len("ipc://") :]
|
||||
reader = ZmqShmLoadSnapshotReader(addr, shm_path, dp_size=1)
|
||||
self.assertTrue(os.path.exists(ipc_path))
|
||||
reader.close()
|
||||
self.assertFalse(os.path.exists(ipc_path))
|
||||
if os.path.exists(shm_path):
|
||||
os.unlink(shm_path)
|
||||
|
||||
|
||||
class TestFactoryFunctions(CustomTestCase):
|
||||
def test_shm_mode(self):
|
||||
server_args = SimpleNamespace(
|
||||
enable_dp_attention=False,
|
||||
nnodes=1,
|
||||
dp_size=1,
|
||||
load_balance_method="round_robin",
|
||||
node_rank=0,
|
||||
)
|
||||
port_args = SimpleNamespace(instance_id="test_shm_factory")
|
||||
writer = create_load_snapshot_writer(
|
||||
server_args, port_args, dp_size=1, dp_rank=0
|
||||
)
|
||||
self.assertIsInstance(writer, ShmLoadSnapshotWriter)
|
||||
reader = create_load_snapshot_reader(server_args, port_args, caller="tokenizer")
|
||||
self.assertIsInstance(reader, ShmLoadSnapshotReader)
|
||||
reader.close()
|
||||
writer.close()
|
||||
from sglang.srt.managers.load_snapshot import shm_path_for
|
||||
|
||||
path = shm_path_for("test_shm_factory")
|
||||
if os.path.exists(path):
|
||||
os.unlink(path)
|
||||
|
||||
def test_zmq_mode_via_env(self):
|
||||
server_args = SimpleNamespace(
|
||||
enable_dp_attention=False,
|
||||
nnodes=1,
|
||||
dp_size=1,
|
||||
load_balance_method="round_robin",
|
||||
node_rank=0,
|
||||
)
|
||||
port_args = SimpleNamespace(instance_id="test_zmq_factory")
|
||||
os.environ["SGLANG_LOAD_SNAPSHOT_USE_ZMQ"] = "1"
|
||||
try:
|
||||
writer = create_load_snapshot_writer(
|
||||
server_args, port_args, dp_size=1, dp_rank=0
|
||||
)
|
||||
self.assertIsInstance(writer, ZmqLoadSnapshotWriter)
|
||||
reader = create_load_snapshot_reader(
|
||||
server_args, port_args, caller="tokenizer"
|
||||
)
|
||||
self.assertIsInstance(reader, ZmqShmLoadSnapshotReader)
|
||||
reader.close()
|
||||
writer.close()
|
||||
finally:
|
||||
del os.environ["SGLANG_LOAD_SNAPSHOT_USE_ZMQ"]
|
||||
|
||||
def test_should_use_zmq_multinode_dp_attention(self):
|
||||
args = SimpleNamespace(enable_dp_attention=True, nnodes=2)
|
||||
self.assertTrue(should_use_zmq(args))
|
||||
|
||||
def test_should_use_zmq_single_node(self):
|
||||
args = SimpleNamespace(enable_dp_attention=False, nnodes=1)
|
||||
self.assertFalse(should_use_zmq(args))
|
||||
|
||||
def test_should_use_zmq_dp_attention_single_node(self):
|
||||
args = SimpleNamespace(enable_dp_attention=True, nnodes=1)
|
||||
self.assertFalse(should_use_zmq(args))
|
||||
|
||||
|
||||
class TestZmqAddr(CustomTestCase):
|
||||
def test_ipc_for_single_node(self):
|
||||
port_args = SimpleNamespace(instance_id="myinstance")
|
||||
addr = _zmq_addr_for(port_args)
|
||||
self.assertTrue(addr.startswith("ipc://"))
|
||||
self.assertIn("myinstance", addr)
|
||||
|
||||
def test_tcp_from_port_args(self):
|
||||
from sglang.srt.utils.network import NetworkAddress
|
||||
|
||||
port_args = SimpleNamespace(
|
||||
instance_id="myinstance",
|
||||
load_collector_ipc_name=NetworkAddress("10.0.0.1", 29506).to_tcp(),
|
||||
)
|
||||
addr = _zmq_addr_for(port_args)
|
||||
self.assertTrue(addr.startswith("tcp://"))
|
||||
self.assertIn("10.0.0.1", addr)
|
||||
|
||||
|
||||
class TestEndToEndZmqSimulation(CustomTestCase):
|
||||
"""Simulate multi-node DP attention on single machine using IPC."""
|
||||
|
||||
def test_full_flow_dp_size_2(self):
|
||||
shm_path = _temp_path()
|
||||
addr = _ipc_addr()
|
||||
dp_size = 2
|
||||
|
||||
reader = ZmqShmLoadSnapshotReader(addr, shm_path, dp_size)
|
||||
|
||||
writers = []
|
||||
for rank in range(dp_size):
|
||||
w = ZmqLoadSnapshotWriter(addr, dp_size, dp_rank=rank)
|
||||
writers.append(w)
|
||||
|
||||
try:
|
||||
_warmup_zmq(writers, reader)
|
||||
|
||||
for rank, w in enumerate(writers):
|
||||
w.write(
|
||||
LoadSnapshot(
|
||||
dp_rank=rank,
|
||||
timestamp=1.0,
|
||||
num_running_reqs=10 + rank,
|
||||
num_waiting_reqs=5 + rank,
|
||||
num_total_tokens=100 + rank * 50,
|
||||
)
|
||||
)
|
||||
time.sleep(0.05)
|
||||
|
||||
loads = reader.read_all()
|
||||
self.assertEqual(len(loads), dp_size)
|
||||
self.assertEqual(loads[0].num_running_reqs, 10)
|
||||
self.assertEqual(loads[1].num_running_reqs, 11)
|
||||
self.assertEqual(loads[0].num_total_tokens, 100)
|
||||
self.assertEqual(loads[1].num_total_tokens, 150)
|
||||
|
||||
for rank, w in enumerate(writers):
|
||||
w.write(
|
||||
LoadSnapshot(
|
||||
dp_rank=rank,
|
||||
timestamp=2.0,
|
||||
num_running_reqs=20 + rank,
|
||||
num_waiting_reqs=0,
|
||||
num_total_tokens=200 + rank * 50,
|
||||
)
|
||||
)
|
||||
time.sleep(0.05)
|
||||
|
||||
loads = reader.read_all()
|
||||
self.assertEqual(loads[0].num_running_reqs, 20)
|
||||
self.assertEqual(loads[1].num_running_reqs, 21)
|
||||
self.assertEqual(loads[0].num_total_tokens, 200)
|
||||
self.assertEqual(loads[1].num_total_tokens, 250)
|
||||
finally:
|
||||
for w in writers:
|
||||
w.close()
|
||||
reader.close()
|
||||
if os.path.exists(shm_path):
|
||||
os.unlink(shm_path)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user