项目文件夹

文件
wehub-resource-sync 94057c3d3e
PR Test (NPU) / check-changes (push) Has been cancelled
PR Test (NPU) / pr-gate (push) Has been cancelled
PR Test (NPU) / set-image-config (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-4-npu-a3 (push) Has been cancelled
PR Test (NPU) / stage-b-test-16-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-1-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-2-npu-a3 (push) Has been cancelled
PR Test (Arm64) / pr-gate (push) Has been cancelled
PR Test (Arm64) / check-changes (push) Has been cancelled
PR Test (Arm64) / build-test (push) Has been cancelled
PR Test (sgl-router) / gate (push) Has been cancelled
PR Test (sgl-router) / tier-1 — lint (push) Has been cancelled
PR Test (sgl-router) / tier-2 — build + test (push) Has been cancelled
PR Test (sgl-router) / tier-3 — docker (placeholder) (push) Has been cancelled
PR Test (sgl-router) / tier-3 — k8s integration (push) Has been cancelled
PR Test (sgl-router) / tier-3 — e2e (push) Has been cancelled
PR Test (sgl-router) / finish (push) Has been cancelled
PR Test (NPU) / single-node-poc (map[name:qwen3_6_27b_w8a8_1p_in64k_out1k_50ms runner:linux-aarch64-a3-2 test_case:test/registered/ascend/performance/qwen3_6_27b/test_npu_qwen3_6_27b_w8a8_1p_in64k_out1k_50ms.py test_type:perf]) (push) Has been cancelled
PR Test (NPU) / pr-test-npu-finish (push) Has been cancelled
PR Test (Xeon) / pr-gate (push) Has been cancelled
PR Test (Xeon) / check-changes (push) Has been cancelled
PR Test (Xeon) / build-test (, xeon-gnr, base-b-test-cpu) (push) Has been cancelled
PR Test (XPU) / check-changes (push) Has been cancelled
PR Test (XPU) / pr-gate (push) Has been cancelled
PR Test (XPU) / stage-a-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / wait-for-stage-a (push) Has been cancelled
PR Test (XPU) / stage-b-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / finish (push) Has been cancelled
CI Model Inventory / build-inventory (push) Has been cancelled
Lint / lint (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Compilation Check (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Manual Policy (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Request Processing (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Summary (push) Has been cancelled
PR Test (SMG) / build-wheel (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on windows (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (x86_64 - auto) (push) Has been cancelled
PR Test (SMG) / python-unit-tests (push) Has been cancelled
PR Test (SMG) / unit-tests (push) Has been cancelled
PR Test (SMG) / benchmarks (push) Has been cancelled
PR Test (SMG) / chat-completions (push) Has been cancelled
PR Test (SMG) / chat-completions-4gpu (push) Has been cancelled
PR Test (SMG) / e2e (push) Has been cancelled
PR Test (SMG) / docker-build-test (push) Has been cancelled
PR Test (SMG) / k8s-integration (push) Has been cancelled
PR Test (SMG) / finish (push) Has been cancelled
PR Test (SMG) / summarize-benchmarks (push) Has been cancelled
Release SGLang Model Gateway Docker Image / publish (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Build SDist (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Upload to PyPI (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (aarch64, 12.9, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (x86_64, 12.9, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu129 (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (aarch64, 13.0, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (x86_64, 13.0, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu130 (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 700) (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 720) (push) Has been cancelled
Release SGLang Kernels / release-rocm700 (push) Has been cancelled
Release SGLang Kernels / release-rocm720 (push) Has been cancelled
Release SGLang Kernels / build-musa43 (43, 3.10) (push) Has been cancelled
Release SGLang Kernels / release-musa43 (push) Has been cancelled
chore: import upstream snapshot with attribution
2026-07-13 12:38:16 +08:00

820 行
28 KiB
Python

"""Basic CPU unit tests for NIXL disaggregation control paths."""
import struct
import sys
import threading
import types
import unittest
from collections import defaultdict
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import numpy as np
from sglang.srt.disaggregation.base.conn import KVPoll
from sglang.srt.disaggregation.common.conn import CommonKVManager
from sglang.srt.disaggregation.common.staging_handler import PrefillStagingContext
from sglang.srt.disaggregation.common.utils import pack_int_lists
from sglang.srt.disaggregation.nixl.conn import (
KVArgsRegisterInfo,
NixlKVManager,
NixlKVReceiver,
NixlKVSender,
TransferInfo,
TransferKVChunk,
TransferStatus,
)
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=23, suite="base-a-test-cpu")
class NotificationFakeAgent:
def __init__(self, messages):
self.messages = messages
def get_new_notifs(self):
return {"peer": [msg.encode("ascii") for msg in self.messages]}
class StagingFakeAgent:
def __init__(self, register_result=None):
self.register_result = (
register_result if register_result is not None else ["desc"]
)
self.register_memory_calls = []
self.get_xfer_descs_calls = []
self.initialize_xfer_calls = []
self.transfer_calls = []
def register_memory(self, addrs, mem_type):
self.register_memory_calls.append((addrs, mem_type))
return self.register_result
def get_xfer_descs(self, reqs, mem_type):
self.get_xfer_descs_calls.append((reqs, mem_type))
return f"{mem_type}_{len(self.get_xfer_descs_calls)}"
def initialize_xfer(self, *args):
self.initialize_xfer_calls.append(args)
return "handle"
def transfer(self, handle):
self.transfer_calls.append(handle)
return "DONE"
class FakeQueue:
def __init__(self):
self.items = []
def put(self, item):
self.items.append(item)
class FakeTensor:
shape = (1, 1, 8)
def element_size(self):
return 2
class FakeStagingBuffer:
def __init__(self, ptr=0x9000, size=1 << 20):
self.ptr = ptr
self.size = size
def fits(self, required_bytes):
return required_bytes <= self.size
def get_ptr(self):
return self.ptr
class FakeStagingAllocator:
ALLOC_OVERSIZED = -2
def _fake_staging_buffer_module(mock_gather=None):
module = types.ModuleType("sglang.srt.disaggregation.common.staging_buffer")
module.StagingAllocator = FakeStagingAllocator
module.compute_head_slice_params = lambda *args: (0, 1, 0, 1)
module.compute_staging_layout = lambda *args: (2, [256, 256], 512)
module.resolve_total_kv_heads = lambda kv_args, attn_tp_size: 2
module.gather_all_layers_to_staging = mock_gather or MagicMock()
return module
class TestNixlTransferInfo(CustomTestCase):
def test_from_zmq_parses_required_fields(self):
kv_indices = np.array([3, 5, 8], dtype=np.int32)
state_indices = [[1, 2], [], [9]]
msg = [
b"7",
b"127.0.0.1",
b"12345",
b"decode_agent",
kv_indices.tobytes(),
b"4",
b"2",
pack_int_lists(state_indices, "i"),
b"11",
]
info = TransferInfo.from_zmq(msg)
self.assertEqual(info.room, 7)
self.assertEqual(info.endpoint, "127.0.0.1")
self.assertEqual(info.dst_port, 12345)
self.assertEqual(info.agent_name, "decode_agent")
np.testing.assert_array_equal(info.dst_kv_indices, kv_indices)
self.assertEqual(info.dst_aux_index, 4)
self.assertEqual(info.required_dst_info_num, 2)
self.assertEqual(info.dst_state_indices, state_indices)
self.assertEqual(info.decode_prefix_len, 11)
def test_from_zmq_defaults_optional_fields(self):
info = TransferInfo.from_zmq(
[
b"8",
b"127.0.0.1",
b"12346",
b"agent",
np.array([1], dtype=np.int32).tobytes(),
b"0",
b"1",
]
)
self.assertEqual(info.dst_state_indices, [])
self.assertIsNone(info.decode_prefix_len)
def test_decode_radix_full_hit_is_not_dummy(self):
info = TransferInfo.from_zmq(
[
b"9",
b"127.0.0.1",
b"12347",
b"agent",
np.array([], dtype=np.int32).tobytes(),
b"2",
b"1",
b"",
b"128",
]
)
self.assertFalse(info.is_dummy())
def test_empty_indices_without_decode_prefix_is_dummy(self):
info = TransferInfo.from_zmq(
[
b"10",
b"127.0.0.1",
b"12348",
b"agent",
np.array([], dtype=np.int32).tobytes(),
b"2",
b"1",
b"",
b"0",
]
)
self.assertTrue(info.is_dummy())
class TestNixlKVArgsRegisterInfo(CustomTestCase):
def test_from_zmq_preserves_unsigned_pointers_and_optional_fields(self):
high_ptr = 0xFFFF_81AB_54E0_1000
kv_ptrs = [high_ptr, high_ptr + 0x1000]
aux_ptrs = [0x1000, 0x2000]
state_ptrs = [[high_ptr + 0x2000], [high_ptr + 0x3000, high_ptr + 0x4000]]
state_item_lens = [[64], [128, 256]]
state_dims = [[16], [32, 64]]
staging_ptr = high_ptr + 0x5000
msg = [
b"None",
b"10.0.0.2",
b"23456",
b"agent_with_large_ptr",
b"metadata",
b"".join(struct.pack("Q", ptr) for ptr in kv_ptrs),
b"".join(struct.pack("Q", ptr) for ptr in aux_ptrs),
pack_int_lists(state_ptrs, "Q"),
b"3",
b"4",
b"1",
b"1024",
pack_int_lists(state_item_lens, "I"),
pack_int_lists(state_dims, "I"),
struct.pack("Q", staging_ptr),
b"1048576",
b"64",
b"DRAM,DRAM",
b"".join(struct.pack("Q", item_len) for item_len in [1024, 2048]),
]
info = KVArgsRegisterInfo.from_zmq(msg)
self.assertEqual(info.room, "None")
self.assertEqual(info.endpoint, "10.0.0.2")
self.assertEqual(info.dst_port, 23456)
self.assertEqual(info.agent_name, "agent_with_large_ptr")
self.assertEqual(info.agent_metadata, b"metadata")
self.assertEqual(info.dst_kv_ptrs, kv_ptrs)
self.assertEqual(info.dst_aux_ptrs, aux_ptrs)
self.assertEqual(info.dst_state_data_ptrs, state_ptrs)
self.assertEqual(info.gpu_id, 3)
self.assertEqual(info.decode_tp_size, 4)
self.assertEqual(info.decode_tp_rank, 1)
self.assertEqual(info.dst_kv_item_len, 1024)
self.assertEqual(info.dst_kv_item_lens, [1024, 2048])
self.assertEqual(info.dst_num_slots, 64)
self.assertEqual(info.dst_kv_mem_kinds, ["DRAM", "DRAM"])
self.assertEqual(info.dst_state_item_lens, state_item_lens)
self.assertEqual(info.dst_state_dim_per_tensor, state_dims)
self.assertIsNotNone(info.staging)
self.assertEqual(info.staging.base_ptr, staging_ptr)
self.assertEqual(info.staging.total_size, 1048576)
def test_from_zmq_allows_missing_state_and_staging_fields(self):
msg = [
b"None",
b"10.0.0.3",
b"23457",
b"agent",
b"metadata",
struct.pack("Q", 0x1000),
struct.pack("Q", 0x2000),
b"",
b"0",
b"1",
b"0",
b"256",
]
info = KVArgsRegisterInfo.from_zmq(msg)
self.assertEqual(info.dst_state_data_ptrs, [])
self.assertEqual(info.dst_state_item_lens, [])
self.assertEqual(info.dst_state_dim_per_tensor, [])
self.assertEqual(info.dst_kv_item_lens, [256])
self.assertIsNone(info.staging)
class TestNixlTransferStatus(CustomTestCase):
def test_not_done_until_aux_and_expected_count_arrive(self):
status = TransferStatus()
self.assertFalse(status.is_done())
status.received_aux = True
self.assertFalse(status.is_done())
status.num_pp_ranks_expected = 1
self.assertFalse(status.is_done())
status.expected_kvs_per_pp[0] = 1
self.assertFalse(status.is_done())
status.received_kvs_per_pp[0].add(0)
self.assertTrue(status.is_done())
def test_zero_kv_aux_only_completion(self):
status = TransferStatus()
status.received_aux = True
status.num_pp_ranks_expected = 1
status.expected_kvs_per_pp[0] = 0
self.assertTrue(status.is_done())
def test_multi_pp_requires_each_rank_expected_chunks(self):
status = TransferStatus()
status.received_aux = True
status.num_pp_ranks_expected = 2
status.expected_kvs_per_pp[0] = 1
status.received_kvs_per_pp[0].add(0)
self.assertFalse(status.is_done())
status.expected_kvs_per_pp[1] = 2
status.received_kvs_per_pp[1].update({0, 1})
self.assertTrue(status.is_done())
def test_state_required_completion_waits_for_all_pp_ranks(self):
status = TransferStatus()
status.received_aux = True
status.num_pp_ranks_expected = 2
status.expected_kvs_per_pp[0] = 0
status.expected_kvs_per_pp[1] = 0
status.expects_state = True
self.assertFalse(status.is_done())
status.received_state_per_pp.add(0)
self.assertFalse(status.is_done())
status.received_state_per_pp.add(1)
self.assertTrue(status.is_done())
class TestNixlKVSenderChunkPolicy(CustomTestCase):
def test_last_zero_page_chunk_is_sent_for_aux_only_completion(self):
sender = object.__new__(NixlKVSender)
self.assertTrue(sender.should_send_kv_chunk(0, last_chunk=True))
self.assertFalse(sender.should_send_kv_chunk(0, last_chunk=False))
self.assertTrue(sender.should_send_kv_chunk(3, last_chunk=False))
class TestNixlNotifications(CustomTestCase):
def _make_manager(self, messages, required=None):
mgr = object.__new__(NixlKVManager)
mgr.agent = NotificationFakeAgent(messages)
mgr.transfer_statuses = defaultdict(TransferStatus)
mgr.required_prefill_response_num_table = required or {}
mgr.enable_staging = False
mgr._staging_handler = None
mgr._chunk_writer_counts = defaultdict(lambda: defaultdict(list))
return mgr
def test_kv_last_notification_sets_expected_count(self):
mgr = self._make_manager(["5_kv_2_1_0"])
mgr.update_transfer_status()
status = mgr.transfer_statuses[5]
self.assertEqual(status.received_kvs_per_pp[0], {2})
self.assertEqual(status.expected_kvs_per_pp[0], 3)
self.assertEqual(status.num_pp_ranks_expected, 1)
def test_staging_notification_preserves_agent_name_with_underscores(self):
mgr = self._make_manager(["5_stg_0_1_0_2_4_8_agent_with_underscores"])
calls = []
mgr._handle_staging_chunk_arrived = lambda *args: calls.append(args)
mgr.update_transfer_status()
self.assertEqual(calls, [(5, 2, 4, 8, "agent_with_underscores")])
status = mgr.transfer_statuses[5]
self.assertEqual(status.received_kvs_per_pp[0], {0})
self.assertEqual(status.expected_kvs_per_pp[0], 1)
def test_aux_nokv_marks_zero_expected_chunks_for_pp_rank(self):
mgr = self._make_manager(["6_aux_nokv_3"], required={6: 4})
mgr.update_transfer_status()
status = mgr.transfer_statuses[6]
self.assertTrue(status.received_aux)
self.assertEqual(status.expected_kvs_per_pp[3], 0)
self.assertEqual(status.num_pp_ranks_expected, 4)
def test_state_notification_marks_pp_rank(self):
mgr = self._make_manager(["7_state_2"])
mgr.update_transfer_status()
self.assertEqual(mgr.transfer_statuses[7].received_state_per_pp, {2})
def test_aux_nokv_allows_full_hit_completion(self):
mgr = self._make_manager(["8_aux_nokv_0"], required={8: 1})
mgr.update_transfer_status()
self.assertTrue(mgr.transfer_statuses[8].is_done())
class TestNixlReceiverPoll(CustomTestCase):
def _make_receiver(self, status=KVPoll.WaitingForInput):
mgr = MagicMock()
mgr.waiting_timeout = 5
mgr.check_status.return_value = status
mgr.transfer_statuses = {}
mgr.addr_to_rooms_tracker = defaultdict(set)
mgr.addr_to_rooms_tracker["prefill:8998"].add(11)
receiver = object.__new__(NixlKVReceiver)
receiver.kv_mgr = mgr
receiver.bootstrap_room = 11
receiver.bootstrap_addr = "prefill:8998"
receiver.started_transfer = False
receiver.init_time = None
receiver.conclude_state = None
receiver.abort_notified = False
return receiver, mgr
def test_returns_existing_conclude_state_without_polling_manager(self):
receiver, mgr = self._make_receiver()
receiver.conclude_state = KVPoll.Success
self.assertEqual(receiver.poll(), KVPoll.Success)
mgr.check_status.assert_not_called()
def test_returns_bootstrap_status_before_transfer_starts(self):
receiver, mgr = self._make_receiver(status=KVPoll.Bootstrapping)
self.assertEqual(receiver.poll(), KVPoll.Bootstrapping)
mgr.update_transfer_status.assert_not_called()
def test_manager_success_or_failed_status_is_terminal(self):
for terminal_status in (KVPoll.Success, KVPoll.Failed):
receiver, _ = self._make_receiver(status=terminal_status)
self.assertEqual(receiver.poll(), terminal_status)
self.assertEqual(receiver.conclude_state, terminal_status)
@patch("sglang.srt.disaggregation.nixl.conn.time.time")
def test_waiting_timeout_records_failure(self, mock_time):
mock_time.return_value = 20.0
receiver, mgr = self._make_receiver(status=KVPoll.WaitingForInput)
receiver.started_transfer = True
receiver.init_time = 10.0
self.assertEqual(receiver.poll(), KVPoll.Failed)
mgr.record_failure.assert_called_once()
self.assertIn("timed out", mgr.record_failure.call_args[0][1])
mgr.update_status.assert_called_once_with(11, KVPoll.Failed)
@patch("sglang.srt.disaggregation.nixl.conn.time.time")
def test_transfer_done_returns_success_and_cleans_room_state(self, mock_time):
mock_time.return_value = 12.0
receiver, mgr = self._make_receiver(status=KVPoll.WaitingForInput)
receiver.started_transfer = True
receiver.init_time = 10.0
status = TransferStatus()
status.received_aux = True
status.num_pp_ranks_expected = 1
status.expected_kvs_per_pp[0] = 0
mgr.transfer_statuses = {11: status}
mgr.check_transfer_done.return_value = True
self.assertEqual(receiver.poll(), KVPoll.Success)
self.assertNotIn(11, mgr.transfer_statuses)
self.assertNotIn(11, mgr.addr_to_rooms_tracker["prefill:8998"])
self.assertEqual(receiver.conclude_state, KVPoll.Success)
class TestNixlNodeFailure(CustomTestCase):
def _make_manager(self):
mgr = object.__new__(NixlKVManager)
mgr.connection_lock = threading.Lock()
# Connection keys are "{addr}_{dp_rank}_{cp_rank}_{tp_rank}".
mgr.connection_pool = {
"10.0.0.1:8998_0_0_0": [{"rank_ip": "10.0.0.1"}],
"10.0.0.1:8998_0_0_1": [{"rank_ip": "10.0.0.1"}],
"10.0.0.2:8998_0_0_0": [{"rank_ip": "10.0.0.2"}],
}
mgr.prefill_info_table = {
"10.0.0.1:8998": object(),
"10.0.0.2:8998": object(),
}
mgr.addr_to_rooms_tracker = defaultdict(set)
mgr.addr_to_rooms_tracker["10.0.0.1:8998"] = {3, 4, 5}
mgr.request_status = {
3: KVPoll.WaitingForInput,
4: KVPoll.Transferring,
5: KVPoll.Success,
}
mgr.failure_records = {}
mgr.failure_lock = threading.Lock()
mgr.update_status = CommonKVManager.update_status.__get__(mgr, CommonKVManager)
mgr.check_status = CommonKVManager.check_status.__get__(mgr, CommonKVManager)
mgr.record_failure = CommonKVManager.record_failure.__get__(
mgr, CommonKVManager
)
return mgr
def test_handle_node_failure_removes_connections_and_marks_pending_rooms(self):
mgr = self._make_manager()
mgr._handle_node_failure("10.0.0.1:8998")
self.assertNotIn("10.0.0.1:8998_0_0_0", mgr.connection_pool)
self.assertNotIn("10.0.0.1:8998_0_0_1", mgr.connection_pool)
self.assertIn("10.0.0.2:8998_0_0_0", mgr.connection_pool)
self.assertNotIn("10.0.0.1:8998", mgr.prefill_info_table)
self.assertNotIn("10.0.0.1:8998", mgr.addr_to_rooms_tracker)
self.assertEqual(mgr.request_status[3], KVPoll.Failed)
self.assertEqual(mgr.request_status[4], KVPoll.Failed)
self.assertEqual(mgr.request_status[5], KVPoll.Success)
self.assertIn(3, mgr.failure_records)
self.assertIn(4, mgr.failure_records)
self.assertNotIn(5, mgr.failure_records)
def test_late_failed_update_does_not_resurrect_cleared_room(self):
mgr = object.__new__(CommonKVManager)
mgr.request_status = {}
CommonKVManager.update_status(mgr, 9, KVPoll.Failed)
self.assertNotIn(9, mgr.request_status)
class TestNixlStaging(CustomTestCase):
def _make_manager(self, agent=None):
mgr = object.__new__(NixlKVManager)
mgr.agent = agent or StagingFakeAgent()
mgr.attn_tp_size = 2
mgr.is_mla_backend = False
mgr.kv_args = SimpleNamespace(
gpu_id=1,
engine_rank=1,
page_size=2,
total_kv_head_num=2,
kv_head_num=1,
)
mgr.server_args = SimpleNamespace(chunked_prefill_size=4)
return mgr
def test_register_buffer_to_engine_groups_kv_memory_kinds_in_one_pass(self):
agent = StagingFakeAgent(register_result=["desc"])
mgr = self._make_manager(agent)
mgr.kv_args.kv_data_ptrs = [0x1000, 0x2000, 0x3000]
mgr.kv_args.kv_data_lens = [64, 128, 256]
mgr.kv_args.kv_data_mem_kinds = ["VRAM", "DRAM", "VRAM"]
mgr.kv_args.aux_data_ptrs = [0x4000]
mgr.kv_args.aux_data_lens = [32]
mgr.kv_args.state_data_ptrs = []
mgr.kv_args.state_data_lens = []
mgr.register_buffer_to_engine()
self.assertEqual(
agent.register_memory_calls,
[
(
[(0x1000, 64, 1, ""), (0x3000, 256, 1, "")],
"VRAM",
),
([(0x2000, 128, 0, "")], "DRAM"),
([(0x4000, 32, 0, "")], "DRAM"),
],
)
self.assertEqual(mgr.kv_descs, [["desc"], ["desc"]])
self.assertEqual(mgr.aux_descs, ["desc"])
def test_register_staging_memory_uses_vram_and_fails_on_empty_descs(self):
agent = StagingFakeAgent(register_result=["staging"])
mgr = self._make_manager(agent)
mgr._register_staging_memory(0x1000, 4096, 3)
self.assertEqual(
agent.register_memory_calls,
[([(0x1000, 4096, 3, "")], "VRAM")],
)
mgr = self._make_manager(StagingFakeAgent(register_result=[]))
with self.assertRaisesRegex(RuntimeError, "staging buffer"):
mgr._register_staging_memory(0x1000, 4096, 3)
def test_prefetch_staging_reqs_noops_when_disabled_or_missing_kv_buffers(self):
mgr = self._make_manager()
mgr.enable_staging = False
mgr.kv_buffer_tensors = {"k_buffers": [], "v_buffers": [], "page_size": 1}
mgr._prefetch_staging_reqs(3)
mgr.enable_staging = True
mgr.kv_buffer_tensors = None
mgr._prefetch_staging_reqs(3)
def test_prefetch_staging_reqs_marks_room_when_no_peer_needs_staging(self):
mgr = self._make_manager()
mgr.enable_staging = True
mgr.kv_buffer_tensors = {"k_buffers": [], "v_buffers": [], "page_size": 1}
mgr._staging_ctx = PrefillStagingContext()
mgr.transfer_infos = {
3: {
"agent": TransferInfo(
room=3,
endpoint="127.0.0.1",
dst_port=1000,
agent_name="agent",
dst_kv_indices=np.array([1], dtype=np.int32),
dst_aux_index=0,
required_dst_info_num=1,
dst_state_indices=[],
)
}
}
mgr.decode_kv_args_table = {
"agent": SimpleNamespace(decode_tp_size=2),
}
mgr._prefetch_staging_reqs(3)
self.assertIn(3, mgr._staging_ctx.prefetched_rooms)
def test_do_staging_transfer_requeues_when_allocation_not_ready(self):
mgr = self._make_manager()
strategy = MagicMock()
strategy.check_ready.return_value = (False, 0, -1, 0, -1)
kv_chunk = TransferKVChunk(
room=3,
prefill_kv_indices=np.array([10, 11], dtype=np.int32),
index_slice=slice(0, 2),
is_last_chunk=False,
chunk_id=0,
prefill_aux_index=None,
state_indices=None,
)
req = SimpleNamespace(room=3, agent_name="decode_agent")
queue = FakeQueue()
with patch.dict(
sys.modules,
{
"sglang.srt.disaggregation.common.staging_buffer": (
_fake_staging_buffer_module()
)
},
):
handle, deferred = mgr._do_staging_transfer(
strategy,
kv_chunk,
kv_chunk.prefill_kv_indices,
req,
SimpleNamespace(),
queue,
)
self.assertIsNone(handle)
self.assertTrue(deferred)
self.assertEqual(queue.items, [kv_chunk])
def test_do_staging_transfer_raises_for_oversized_allocation(self):
mgr = self._make_manager()
strategy = MagicMock()
strategy.check_ready.return_value = (
False,
0,
FakeStagingAllocator.ALLOC_OVERSIZED,
0,
-1,
)
kv_chunk = TransferKVChunk(
room=3,
prefill_kv_indices=np.array([10], dtype=np.int32),
index_slice=slice(0, 1),
is_last_chunk=False,
chunk_id=0,
prefill_aux_index=None,
state_indices=None,
)
with self.assertRaisesRegex(RuntimeError, "ring buffer total size"):
with patch.dict(
sys.modules,
{
"sglang.srt.disaggregation.common.staging_buffer": (
_fake_staging_buffer_module()
)
},
):
mgr._do_staging_transfer(
strategy,
kv_chunk,
kv_chunk.prefill_kv_indices,
SimpleNamespace(room=3, agent_name="decode_agent"),
SimpleNamespace(),
FakeQueue(),
)
def test_do_staging_transfer_builds_staging_notification(self):
mgr = self._make_manager()
strategy = MagicMock()
strategy.check_ready.return_value = (True, 2, 128, 0, 512)
strategy.staging_buffer = FakeStagingBuffer()
kv_chunk = TransferKVChunk(
room=3,
prefill_kv_indices=np.array([10, 11], dtype=np.int32),
index_slice=slice(4, 6),
is_last_chunk=True,
chunk_id=7,
prefill_aux_index=0,
state_indices=None,
)
dst_info = KVArgsRegisterInfo(
room="None",
endpoint="127.0.0.1",
dst_port=1000,
agent_name="decode_agent",
agent_metadata=b"",
dst_kv_ptrs=[],
dst_kv_mem_kinds=[],
dst_aux_ptrs=[],
dst_state_data_ptrs=[],
gpu_id=5,
decode_tp_size=1,
decode_tp_rank=0,
dst_kv_item_len=128,
dst_kv_item_lens=[],
staging=SimpleNamespace(base_ptr=0x8000, total_size=4096),
)
calls = []
mgr.send_kvcache_staged = (
lambda *args, **kwargs: calls.append((args, kwargs)) or "handle"
)
handle, deferred = mgr._do_staging_transfer(
strategy,
kv_chunk,
kv_chunk.prefill_kv_indices,
SimpleNamespace(room=3, agent_name="decode_agent"),
dst_info,
FakeQueue(),
)
self.assertEqual(handle, "handle")
self.assertFalse(deferred)
self.assertEqual(calls[0][0][8], "3_stg_7_1_1_2_4_2_decode_agent")
def test_send_kvcache_staged_uses_one_bulk_vram_write(self):
mock_gather = MagicMock()
agent = StagingFakeAgent()
mgr = self._make_manager(agent)
mgr.kv_buffer_tensors = {
"k_buffers": [FakeTensor(), FakeTensor()],
"v_buffers": [FakeTensor(), FakeTensor()],
"page_size": 2,
}
with patch.dict(
sys.modules,
{
"sglang.srt.disaggregation.common.staging_buffer": (
_fake_staging_buffer_module(mock_gather)
)
},
):
handle = mgr.send_kvcache_staged(
"peer",
np.array([1, 2], dtype=np.int32),
dst_staging_ptr=0x100000,
dst_staging_size=1 << 20,
dst_gpu_id=4,
dst_tp_rank=0,
dst_attn_tp_size=1,
dst_kv_item_len=128,
notif="3_stg_0_1_1_0_0_2_decode_agent",
staging_buffer=FakeStagingBuffer(ptr=0x9000, size=1 << 20),
)
self.assertEqual(handle, "handle")
mock_gather.assert_called_once()
src_reqs, src_mem = agent.get_xfer_descs_calls[0]
dst_reqs, dst_mem = agent.get_xfer_descs_calls[1]
self.assertEqual(src_mem, "VRAM")
self.assertEqual(dst_mem, "VRAM")
self.assertEqual(src_reqs.shape, (1, 3))
self.assertEqual(dst_reqs.shape, (1, 3))
self.assertTrue(np.issubdtype(src_reqs.dtype, np.integer))
self.assertTrue(np.issubdtype(dst_reqs.dtype, np.integer))
self.assertEqual(int(src_reqs[0, 0]), 0x9000)
self.assertGreaterEqual(int(dst_reqs[0, 0]), 0x100000)
self.assertEqual(agent.initialize_xfer_calls[0][0], "WRITE")
self.assertEqual(
agent.initialize_xfer_calls[0][-1],
b"3_stg_0_1_1_0_0_2_decode_agent",
)
def test_send_kvcache_staged_falls_back_when_prefill_buffer_too_small(self):
mgr = self._make_manager()
mgr.kv_buffer_tensors = {
"k_buffers": [FakeTensor(), FakeTensor()],
"v_buffers": [FakeTensor(), FakeTensor()],
"page_size": 2,
}
with patch.dict(
sys.modules,
{
"sglang.srt.disaggregation.common.staging_buffer": (
_fake_staging_buffer_module()
)
},
):
handle = mgr.send_kvcache_staged(
"peer",
np.array([1, 2], dtype=np.int32),
dst_staging_ptr=0xA000,
dst_staging_size=1 << 20,
dst_gpu_id=4,
dst_tp_rank=0,
dst_attn_tp_size=1,
dst_kv_item_len=128,
notif="notif",
staging_buffer=FakeStagingBuffer(size=1),
)
self.assertIsNone(handle)
if __name__ == "__main__":
unittest.main()