[PD][MoRI] Align hybrid state transfer with per-component schema (#26539)

This commit is contained in:
Niko Ma
2026-05-29 00:54:46 -07:00
committed by GitHub
parent a42a7654a2
commit 4d1163e6a9
2 changed files with 184 additions and 88 deletions
@@ -8,13 +8,14 @@ from sglang.test.server_fixtures.disaggregation_fixture import (
PDDisaggregationServerBase,
)
from sglang.test.test_utils import (
DEFAULT_HYBRID_MAMBA_MODEL_NAME_FOR_TEST,
DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
popen_launch_pd_server,
try_cached_model,
)
register_amd_ci(est_time=300, suite="stage-b-test-large-8-gpu-mi35x-disaggregation-amd")
register_amd_ci(est_time=900, suite="stage-b-test-large-8-gpu-mi35x-disaggregation-amd")
class MoriTransferEngineBase(PDDisaggregationServerBase):
@@ -24,6 +25,12 @@ class MoriTransferEngineBase(PDDisaggregationServerBase):
decode_base_gpu_id = 1
required_gpus = 2
# Subclasses can override to pick a different model or pass extra args.
model_default = DEFAULT_SMALL_MODEL_NAME_FOR_TEST
model_env_var = "SGLANG_MORI_E2E_TEST_MODEL"
extra_prefill_args: list = []
extra_decode_args: list = []
@classmethod
def setUpClass(cls):
try:
@@ -56,10 +63,7 @@ class MoriTransferEngineBase(PDDisaggregationServerBase):
cls._shift_ports()
cls.model = try_cached_model(
os.environ.get(
"SGLANG_MORI_E2E_TEST_MODEL",
DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
)
os.environ.get(cls.model_env_var, cls.model_default)
)
cls.start_prefill()
@@ -111,7 +115,7 @@ class MoriTransferEngineBase(PDDisaggregationServerBase):
str(cls.prefill_tp),
"--attention-backend",
"aiter",
]
] + list(cls.extra_prefill_args)
prefill_args += cls.transfer_backend + cls.rdma_devices
cls.process_prefill = popen_launch_pd_server(
cls.model,
@@ -134,7 +138,7 @@ class MoriTransferEngineBase(PDDisaggregationServerBase):
str(cls.decode_base_gpu_id),
"--attention-backend",
"aiter",
]
] + list(cls.extra_decode_args)
decode_args += cls.transfer_backend + cls.rdma_devices
cls.process_decode = popen_launch_pd_server(
cls.model,
@@ -175,5 +179,19 @@ class TestMoriTransferEngineTPMismatchE2E(MoriTransferEngineBase):
self._assert_generate_smoke()
class TestMoriTransferEngineHybridMambaE2E(MoriTransferEngineBase):
port_delta = 20
prefill_tp = 4
decode_tp = 4
decode_base_gpu_id = 4
required_gpus = 8
model_default = DEFAULT_HYBRID_MAMBA_MODEL_NAME_FOR_TEST
model_env_var = "SGLANG_MORI_HYBRID_E2E_TEST_MODEL"
def test_generate_smoke_hybrid_mamba(self):
self._assert_generate_smoke()
if __name__ == "__main__":
unittest.main()