hisparse: support NIXL DRAM KV destinations for HiSparse (#27563)

Co-authored-by: Zhangheng <hzh0425@apache.org>
Co-authored-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
ishandhanani
2026-06-27 22:32:29 +08:00
committed by GitHub
co-authored by Zhangheng Shangming Cai
parent c1b5c7e499
commit b030b1a5f3
12 changed files with 726 additions and 167 deletions
@@ -17,10 +17,6 @@ register_cuda_ci(est_time=500, stage="base-c", runner_config="deepep-8-gpu-h200"
DSV4_FLASH_MODEL = "sgl-project/DeepSeek-V4-Flash-FP8"
DEEPEP_CONFIG = '{"normal_dispatch":{"num_sms":96},"normal_combine":{"num_sms":96}}'
DSV4_FLASH_LOADER_CONFIG = '{"enable_multithread_load": true, "num_threads": 64}'
DSV4_HISPARSE_CONFIG = (
'{"top_k":512,"device_buffer_size":4096,"host_to_device_ratio":2}'
)
DSV4_FLASH_ENV = {
"SGLANG_DSV4_FP4_EXPERTS": "0",
@@ -129,104 +125,5 @@ class TestDisaggregationDSV4(SpecDecodingMixin, PDDisaggregationServerBase, GSM8
)
class TestDisaggregationDSV4HiSparseMooncake(PDDisaggregationServerBase, GSM8KMixin):
gsm8k_accuracy_thres = 0.93
gsm8k_num_questions = 200
gsm8k_num_shots = 20
@classmethod
def setUpClass(cls):
super().setUpClass()
cls.model = try_cached_model(DSV4_FLASH_MODEL)
cls.start_prefill()
cls.start_decode()
cls.wait_server_ready(cls.prefill_url + "/health", process=cls.process_prefill)
cls.wait_server_ready(cls.decode_url + "/health", process=cls.process_decode)
cls.launch_lb()
@classmethod
def start_prefill(cls):
prefill_args = [
"--trust-remote-code",
"--disaggregation-mode",
"prefill",
"--disaggregation-bootstrap-port",
cls.bootstrap_port,
"--tp",
4,
"--page-size",
256,
"--chunked-prefill-size",
8192,
"--max-running-requests",
16,
"--mem-fraction-static",
0.9,
"--skip-server-warmup",
"--reasoning-parser",
"deepseek-v4",
"--tool-call-parser",
"deepseekv4",
"--model-loader-extra-config",
DSV4_FLASH_LOADER_CONFIG,
"--watchdog-timeout",
"900",
]
prefill_args += cls.transfer_backend + cls.rdma_devices
cls.process_prefill = popen_launch_pd_server(
cls.model,
cls.prefill_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=prefill_args,
env=DSV4_FLASH_ENV,
)
@classmethod
def start_decode(cls):
decode_args = [
"--trust-remote-code",
"--disaggregation-mode",
"decode",
"--disaggregation-bootstrap-port",
cls.bootstrap_port,
"--tp",
4,
"--base-gpu-id",
4,
"--page-size",
256,
"--chunked-prefill-size",
8192,
"--max-running-requests",
16,
"--mem-fraction-static",
0.9,
"--skip-server-warmup",
"--reasoning-parser",
"deepseek-v4",
"--tool-call-parser",
"deepseekv4",
"--model-loader-extra-config",
DSV4_FLASH_LOADER_CONFIG,
"--enable-hisparse",
"--hisparse-config",
DSV4_HISPARSE_CONFIG,
"--watchdog-timeout",
"900",
]
decode_args += cls.transfer_backend + cls.rdma_devices
cls.process_decode = popen_launch_pd_server(
cls.model,
cls.decode_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=decode_args,
env=DSV4_FLASH_ENV,
)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,165 @@
import unittest
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.eval_accuracy_kit import GSM8KMixin
from sglang.test.server_fixtures.disaggregation_fixture import (
PDDisaggregationServerBase,
)
from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
is_in_ci,
popen_launch_pd_server,
try_cached_model,
)
register_cuda_ci(est_time=1000, stage="extra-b", runner_config="deepep-8-gpu-h200")
DSV4_FLASH_MODEL = "sgl-project/DeepSeek-V4-Flash-FP8"
DSV4_FLASH_LOADER_CONFIG = '{"enable_multithread_load": true, "num_threads": 64}'
DSV4_HISPARSE_CONFIG = (
'{"top_k":512,"device_buffer_size":4096,"host_to_device_ratio":2}'
)
DSV4_FLASH_ENV = {
"SGLANG_DSV4_FP4_EXPERTS": "0",
"SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK": "256",
}
DSV4_NIXL_SERVER_LAUNCH_TIMEOUT = 1800
def _has_nixl():
try:
import nixl._api # noqa: F401
except Exception:
return False
return True
class TestDisaggregationDSV4HiSparseBase(PDDisaggregationServerBase, GSM8KMixin):
gsm8k_accuracy_thres = 0.93
gsm8k_num_questions = 200
gsm8k_num_shots = 20
@classmethod
def setUpClass(cls):
super().setUpClass()
cls.model = try_cached_model(DSV4_FLASH_MODEL)
cls.start_prefill()
cls.start_decode()
cls.wait_server_ready(cls.prefill_url + "/health", process=cls.process_prefill)
cls.wait_server_ready(cls.decode_url + "/health", process=cls.process_decode)
cls.launch_lb()
@classmethod
def start_prefill(cls):
prefill_args = [
"--trust-remote-code",
"--disaggregation-mode",
"prefill",
"--disaggregation-bootstrap-port",
cls.bootstrap_port,
"--tp",
4,
"--page-size",
256,
"--chunked-prefill-size",
8192,
"--max-running-requests",
16,
"--mem-fraction-static",
0.9,
"--skip-server-warmup",
"--reasoning-parser",
"deepseek-v4",
"--tool-call-parser",
"deepseekv4",
"--model-loader-extra-config",
DSV4_FLASH_LOADER_CONFIG,
"--watchdog-timeout",
"900",
]
prefill_args += cls.transfer_backend + cls.rdma_devices
cls.process_prefill = popen_launch_pd_server(
cls.model,
cls.prefill_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=prefill_args,
env=DSV4_FLASH_ENV,
)
@classmethod
def start_decode(cls):
decode_args = [
"--trust-remote-code",
"--disaggregation-mode",
"decode",
"--disaggregation-bootstrap-port",
cls.bootstrap_port,
"--tp",
4,
"--base-gpu-id",
4,
"--page-size",
256,
"--chunked-prefill-size",
8192,
"--max-running-requests",
16,
"--mem-fraction-static",
0.9,
"--skip-server-warmup",
"--reasoning-parser",
"deepseek-v4",
"--tool-call-parser",
"deepseekv4",
"--model-loader-extra-config",
DSV4_FLASH_LOADER_CONFIG,
"--enable-hisparse",
"--hisparse-config",
DSV4_HISPARSE_CONFIG,
"--watchdog-timeout",
"900",
]
decode_args += cls.transfer_backend + cls.rdma_devices
cls.process_decode = popen_launch_pd_server(
cls.model,
cls.decode_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=decode_args,
env=DSV4_FLASH_ENV,
)
@unittest.skipUnless(
is_in_ci() or _has_nixl(),
"NIXL is required for DSV4 HiSparse disaggregation coverage.",
)
class TestDisaggregationDSV4HiSparseNixl(TestDisaggregationDSV4HiSparseBase):
@classmethod
def setUpClass(cls):
PDDisaggregationServerBase.setUpClass.__func__(cls)
cls.transfer_backend = ["--disaggregation-transfer-backend", "nixl"]
cls.model = try_cached_model(DSV4_FLASH_MODEL)
cls.start_prefill()
cls.start_decode()
cls.wait_server_ready(
cls.prefill_url + "/health",
timeout=DSV4_NIXL_SERVER_LAUNCH_TIMEOUT,
process=cls.process_prefill,
)
cls.wait_server_ready(
cls.decode_url + "/health",
timeout=DSV4_NIXL_SERVER_LAUNCH_TIMEOUT,
process=cls.process_decode,
)
cls.launch_lb()
if __name__ == "__main__":
unittest.main()
@@ -212,6 +212,9 @@ class TestNixlKVArgsRegisterInfo(CustomTestCase):
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)
@@ -228,6 +231,9 @@ class TestNixlKVArgsRegisterInfo(CustomTestCase):
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)
@@ -255,6 +261,7 @@ class TestNixlKVArgsRegisterInfo(CustomTestCase):
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)
@@ -523,6 +530,33 @@ class TestNixlStaging(CustomTestCase):
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)
@@ -601,7 +635,12 @@ class TestNixlStaging(CustomTestCase):
},
):
handle, deferred = mgr._do_staging_transfer(
strategy, kv_chunk, req, SimpleNamespace(), queue
strategy,
kv_chunk,
kv_chunk.prefill_kv_indices,
req,
SimpleNamespace(),
queue,
)
self.assertIsNone(handle)
@@ -640,6 +679,7 @@ class TestNixlStaging(CustomTestCase):
mgr._do_staging_transfer(
strategy,
kv_chunk,
kv_chunk.prefill_kv_indices,
SimpleNamespace(room=3, agent_name="decode_agent"),
SimpleNamespace(),
FakeQueue(),
@@ -666,12 +706,14 @@ class TestNixlStaging(CustomTestCase):
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 = []
@@ -682,6 +724,7 @@ class TestNixlStaging(CustomTestCase):
handle, deferred = mgr._do_staging_transfer(
strategy,
kv_chunk,
kv_chunk.prefill_kv_indices,
SimpleNamespace(room=3, agent_name="decode_agent"),
dst_info,
FakeQueue(),