fix(disagg): support pipeline-parallel hybrid-linear transfer (#32270)
This commit is contained in:
@@ -14,7 +14,7 @@ from sglang.test.test_utils import (
|
||||
popen_launch_server,
|
||||
)
|
||||
|
||||
register_cuda_ci(est_time=240, stage="base-c", runner_config="4-gpu-h100")
|
||||
register_cuda_ci(est_time=480, stage="base-c", runner_config="4-gpu-h100")
|
||||
|
||||
KIMI_LINEAR_MODEL = "yujiepan/kimi-linear-tiny-random"
|
||||
SERVER_ENV = {"SGLANG_BATCH_INVARIANT_OPS_ENABLE_MM_DEEPGEMM": "0"}
|
||||
@@ -38,6 +38,7 @@ class TestKimiLinearHeterogeneousTPDisaggregation(PDDisaggregationServerBase):
|
||||
prefill_tp_size = 2
|
||||
decode_tp_size = 1
|
||||
decode_base_gpu_id = 2
|
||||
reference_parallel_args = ["--tp-size", "2"]
|
||||
extra_prefill_args = SERVER_ARGS
|
||||
extra_decode_args = SERVER_ARGS
|
||||
extra_prefill_env = SERVER_ENV
|
||||
@@ -72,7 +73,9 @@ class TestKimiLinearHeterogeneousTPDisaggregation(PDDisaggregationServerBase):
|
||||
self.model,
|
||||
self.lb_url,
|
||||
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||
other_args=["--tp-size", "2", "--trust-remote-code"] + SERVER_ARGS,
|
||||
other_args=self.reference_parallel_args
|
||||
+ ["--trust-remote-code"]
|
||||
+ SERVER_ARGS,
|
||||
env=SERVER_ENV,
|
||||
)
|
||||
try:
|
||||
@@ -101,5 +104,13 @@ class TestKimiLinearHeterogeneousTPDisaggregation(PDDisaggregationServerBase):
|
||||
assert_process_healthy(self, "decode", self.process_decode, self.decode_url)
|
||||
|
||||
|
||||
class TestKimiLinearPipelineDisaggregation(TestKimiLinearHeterogeneousTPDisaggregation):
|
||||
prefill_tp_size = 1
|
||||
decode_tp_size = 1
|
||||
decode_base_gpu_id = 2
|
||||
reference_parallel_args = ["--tp-size", "1", "--pp-size", "2"]
|
||||
extra_prefill_args = SERVER_ARGS + ["--pp-size", "2"]
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -451,8 +451,10 @@ class TestNixlTransferWorker(CustomTestCase):
|
||||
mgr.enable_staging = False
|
||||
mgr._staging_ctx = None
|
||||
mgr.is_mla_backend = False
|
||||
mgr.is_hybrid_mla_backend = False
|
||||
mgr.attn_tp_size = 1
|
||||
mgr.kv_args = SimpleNamespace(engine_rank=0)
|
||||
mgr.transfer_source_rank = 0
|
||||
mgr.kv_args = SimpleNamespace(engine_rank=0, kv_data_ptrs=[0])
|
||||
mgr.exceptions = {}
|
||||
mgr.failure_lock = threading.Lock()
|
||||
mgr.failure_records = {}
|
||||
@@ -493,6 +495,7 @@ class TestNixlTransferWorker(CustomTestCase):
|
||||
self.assertEqual(mgr.request_status[room], KVPoll.Failed)
|
||||
self.assertNotIn(room, mgr.transfer_infos)
|
||||
self.assertNotIn(room, mgr.req_to_decode_prefix_len)
|
||||
mgr.send_aux.assert_called_once()
|
||||
|
||||
def test_given_non_last_chunk_aborts_mid_transfer_when_worker_finishes_then_failed_status_is_preserved(
|
||||
self,
|
||||
@@ -507,6 +510,7 @@ class TestNixlTransferWorker(CustomTestCase):
|
||||
self.assertEqual(mgr.request_status[room], KVPoll.Failed)
|
||||
self.assertIn(room, mgr.transfer_infos)
|
||||
self.assertIn(room, mgr.req_to_decode_prefix_len)
|
||||
mgr.send_kvcache.assert_called_once()
|
||||
|
||||
|
||||
class TestNixlNotifications(CustomTestCase):
|
||||
@@ -699,6 +703,7 @@ class TestNixlStaging(CustomTestCase):
|
||||
mgr.agent = agent or StagingFakeAgent()
|
||||
mgr.attn_tp_size = 2
|
||||
mgr.is_mla_backend = False
|
||||
mgr.transfer_source_rank = 1
|
||||
mgr.kv_args = SimpleNamespace(
|
||||
gpu_id=1,
|
||||
engine_rank=1,
|
||||
|
||||
Reference in New Issue
Block a user