fix(disagg): support pipeline-parallel hybrid-linear transfer (#32270)

This commit is contained in:
YAMY
2026-07-25 13:34:38 -07:00
committed by GitHub
parent 659d349b61
commit 91f386a5b2
12 changed files with 351 additions and 57 deletions
@@ -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,