[Fix][PD] Give prefill and decode their own RDMA NICs in disaggregation tests (#39584)

This commit is contained in:
Liangsheng Yin
2026-09-15 15:07:42 -07:00
committed by GitHub
parent 406c9c71d8
commit 87c9f78e99
12 changed files with 103 additions and 120 deletions
@@ -58,7 +58,7 @@ class TestDisaggregationMooncakeAARCH64Accuracy(PDDisaggregationServerBase):
"--nccl-port",
str(NCCL_PORT_BASE),
]
prefill_args += cls.transfer_backend + cls.rdma_devices
prefill_args += cls.transfer_backend + cls.rdma_devices_for(range(2))
cls.process_prefill = popen_launch_pd_server(
cls.model,
cls.prefill_url,
@@ -81,7 +81,7 @@ class TestDisaggregationMooncakeAARCH64Accuracy(PDDisaggregationServerBase):
"--nccl-port",
str(NCCL_PORT_BASE + 1),
]
decode_args += cls.transfer_backend + cls.rdma_devices
decode_args += cls.transfer_backend + cls.rdma_devices_for(range(2, 4))
cls.process_decode = popen_launch_pd_server(
cls.model,
cls.decode_url,
@@ -89,7 +89,7 @@ class TestDisaggregationDecodeOffload(PDDisaggregationServerBase):
"--hicache-ratio",
"2",
]
prefill_args += cls.transfer_backend + cls.rdma_devices
prefill_args += cls.transfer_backend + cls.rdma_devices_for(range(1))
cls.process_prefill = popen_launch_pd_server(
cls.model,
cls.prefill_url,
@@ -119,7 +119,7 @@ class TestDisaggregationDecodeOffload(PDDisaggregationServerBase):
"--hicache-storage-backend",
"file",
]
decode_args += cls.transfer_backend + cls.rdma_devices
decode_args += cls.transfer_backend + cls.rdma_devices_for(range(1, 2))
cls.process_decode = popen_launch_pd_server(
cls.model,
cls.decode_url,
@@ -52,7 +52,7 @@ class TestDisaggregationMooncakePrefillLargerTP(PDDisaggregationServerBase):
"--enable-metrics",
"--enable-request-time-stats-logging",
]
prefill_args += cls.transfer_backend + cls.rdma_devices
prefill_args += cls.transfer_backend + cls.rdma_devices_for(range(4))
cls.process_prefill = popen_launch_pd_server(
cls.model,
cls.prefill_url,
@@ -75,7 +75,7 @@ class TestDisaggregationMooncakePrefillLargerTP(PDDisaggregationServerBase):
"--enable-metrics",
"--enable-request-time-stats-logging",
]
decode_args += cls.transfer_backend + cls.rdma_devices
decode_args += cls.transfer_backend + cls.rdma_devices_for(range(4, 6))
cls.process_decode = popen_launch_pd_server(
cls.model,
cls.decode_url,
@@ -131,7 +131,7 @@ class TestDisaggregationMooncakeDecodeLargerTP(PDDisaggregationServerBase):
"--enable-metrics",
"--enable-request-time-stats-logging",
]
prefill_args += cls.transfer_backend + cls.rdma_devices
prefill_args += cls.transfer_backend + cls.rdma_devices_for(range(2))
cls.process_prefill = popen_launch_pd_server(
cls.model,
cls.prefill_url,
@@ -154,7 +154,7 @@ class TestDisaggregationMooncakeDecodeLargerTP(PDDisaggregationServerBase):
"--enable-metrics",
"--enable-request-time-stats-logging",
]
decode_args += cls.transfer_backend + cls.rdma_devices
decode_args += cls.transfer_backend + cls.rdma_devices_for(range(4, 8))
cls.process_decode = popen_launch_pd_server(
cls.model,
cls.decode_url,
@@ -210,7 +210,7 @@ class TestDisaggregationMooncakeMHAPrefillLargerTP(PDDisaggregationServerBase):
"--enable-metrics",
"--enable-request-time-stats-logging",
]
prefill_args += cls.transfer_backend + cls.rdma_devices
prefill_args += cls.transfer_backend + cls.rdma_devices_for(range(4))
cls.process_prefill = popen_launch_pd_server(
cls.model,
cls.prefill_url,
@@ -233,7 +233,7 @@ class TestDisaggregationMooncakeMHAPrefillLargerTP(PDDisaggregationServerBase):
"--enable-metrics",
"--enable-request-time-stats-logging",
]
decode_args += cls.transfer_backend + cls.rdma_devices
decode_args += cls.transfer_backend + cls.rdma_devices_for(range(4, 6))
cls.process_decode = popen_launch_pd_server(
cls.model,
cls.decode_url,
@@ -289,7 +289,7 @@ class TestDisaggregationMooncakeMHADecodeLargerTP(PDDisaggregationServerBase):
"--enable-metrics",
"--enable-request-time-stats-logging",
]
prefill_args += cls.transfer_backend + cls.rdma_devices
prefill_args += cls.transfer_backend + cls.rdma_devices_for(range(2))
cls.process_prefill = popen_launch_pd_server(
cls.model,
cls.prefill_url,
@@ -312,7 +312,7 @@ class TestDisaggregationMooncakeMHADecodeLargerTP(PDDisaggregationServerBase):
"--enable-metrics",
"--enable-request-time-stats-logging",
]
decode_args += cls.transfer_backend + cls.rdma_devices
decode_args += cls.transfer_backend + cls.rdma_devices_for(range(4, 8))
cls.process_decode = popen_launch_pd_server(
cls.model,
cls.decode_url,
@@ -375,7 +375,7 @@ class TestDisaggregationStagingPrefillLargerTP(PDDisaggregationServerBase):
"--enable-metrics",
"--enable-request-time-stats-logging",
]
prefill_args += cls.transfer_backend + cls.rdma_devices
prefill_args += cls.transfer_backend + cls.rdma_devices_for(range(4))
env = {**os.environ, **STAGING_ENV}
cls.process_prefill = popen_launch_pd_server(
cls.model,
@@ -400,7 +400,7 @@ class TestDisaggregationStagingPrefillLargerTP(PDDisaggregationServerBase):
"--enable-metrics",
"--enable-request-time-stats-logging",
]
decode_args += cls.transfer_backend + cls.rdma_devices
decode_args += cls.transfer_backend + cls.rdma_devices_for(range(4, 6))
env = {**os.environ, **STAGING_ENV}
cls.process_decode = popen_launch_pd_server(
cls.model,
@@ -456,7 +456,7 @@ class TestDisaggregationStagingDecodeLargerTP(PDDisaggregationServerBase):
"--enable-metrics",
"--enable-request-time-stats-logging",
]
prefill_args += cls.transfer_backend + cls.rdma_devices
prefill_args += cls.transfer_backend + cls.rdma_devices_for(range(2))
env = {**os.environ, **STAGING_ENV}
cls.process_prefill = popen_launch_pd_server(
cls.model,
@@ -481,7 +481,7 @@ class TestDisaggregationStagingDecodeLargerTP(PDDisaggregationServerBase):
"--enable-metrics",
"--enable-request-time-stats-logging",
]
decode_args += cls.transfer_backend + cls.rdma_devices
decode_args += cls.transfer_backend + cls.rdma_devices_for(range(4, 8))
env = {**os.environ, **STAGING_ENV}
cls.process_decode = popen_launch_pd_server(
cls.model,
@@ -549,7 +549,7 @@ class TestDisaggregationGDNHybridHeteroTP(PDDisaggregationServerBase):
"--enable-metrics",
"--enable-request-time-stats-logging",
]
prefill_args += cls.transfer_backend + cls.rdma_devices
prefill_args += cls.transfer_backend + cls.rdma_devices_for(range(1))
cls.process_prefill = popen_launch_pd_server(
cls.model,
cls.prefill_url,
@@ -572,7 +572,7 @@ class TestDisaggregationGDNHybridHeteroTP(PDDisaggregationServerBase):
"--enable-metrics",
"--enable-request-time-stats-logging",
]
decode_args += cls.transfer_backend + cls.rdma_devices
decode_args += cls.transfer_backend + cls.rdma_devices_for(range(4, 8))
cls.process_decode = popen_launch_pd_server(
cls.model,
cls.decode_url,
@@ -641,7 +641,7 @@ class TestDisaggregationStagingRadixPrefillLargerTP(PDDisaggregationServerBase):
"--enable-metrics",
"--enable-request-time-stats-logging",
]
prefill_args += cls.transfer_backend + cls.rdma_devices
prefill_args += cls.transfer_backend + cls.rdma_devices_for(range(4))
env = {**os.environ, **STAGING_ENV}
cls.process_prefill = popen_launch_pd_server(
cls.model,
@@ -669,7 +669,7 @@ class TestDisaggregationStagingRadixPrefillLargerTP(PDDisaggregationServerBase):
"--enable-metrics",
"--enable-request-time-stats-logging",
]
decode_args += cls.transfer_backend + cls.rdma_devices
decode_args += cls.transfer_backend + cls.rdma_devices_for(range(4, 6))
env = {**os.environ, **STAGING_ENV}
cls.process_decode = popen_launch_pd_server(
cls.model,
@@ -60,7 +60,9 @@ class TestDisaggregationDPAttention(PDDisaggregationServerBase):
"--load-balance-method",
cls.LOAD_BALANCE_METHOD,
]
prefill_args += cls.transfer_backend + cls.rdma_devices
prefill_args += cls.transfer_backend + cls.rdma_devices_for(
range(cls.PREFILL_DP_SIZE)
)
cls.process_prefill = popen_launch_pd_server(
cls.model,
cls.prefill_url,
@@ -86,7 +88,9 @@ class TestDisaggregationDPAttention(PDDisaggregationServerBase):
"--load-balance-method",
cls.LOAD_BALANCE_METHOD,
]
decode_args += cls.transfer_backend + cls.rdma_devices
decode_args += cls.transfer_backend + cls.rdma_devices_for(
range(cls.PREFILL_DP_SIZE, cls.PREFILL_DP_SIZE + cls.DECODE_DP_SIZE)
)
cls.process_decode = popen_launch_pd_server(
cls.model,
cls.decode_url,
@@ -81,7 +81,7 @@ class TestDisaggregationDSV4(SpecDecodingMixin, PDDisaggregationServerBase, GSM8
"--watchdog-timeout",
"900",
]
prefill_args += cls.transfer_backend + cls.rdma_devices
prefill_args += cls.transfer_backend + cls.rdma_devices_for(range(4))
cls.process_prefill = popen_launch_pd_server(
cls.model,
cls.prefill_url,
@@ -119,7 +119,7 @@ class TestDisaggregationDSV4(SpecDecodingMixin, PDDisaggregationServerBase, GSM8
"--watchdog-timeout",
"900",
]
decode_args += cls.transfer_backend + cls.rdma_devices
decode_args += cls.transfer_backend + cls.rdma_devices_for(range(4, 8))
cls.process_decode = popen_launch_pd_server(
cls.model,
cls.decode_url,
@@ -80,7 +80,7 @@ class TestDisaggregationDSV4HiSparseBase(PDDisaggregationServerBase, GSM8KMixin)
"--watchdog-timeout",
"900",
]
prefill_args += cls.transfer_backend + cls.rdma_devices
prefill_args += cls.transfer_backend + cls.rdma_devices_for(range(4))
cls.process_prefill = popen_launch_pd_server(
cls.model,
cls.prefill_url,
@@ -122,7 +122,7 @@ class TestDisaggregationDSV4HiSparseBase(PDDisaggregationServerBase, GSM8KMixin)
"--watchdog-timeout",
"900",
]
decode_args += cls.transfer_backend + cls.rdma_devices
decode_args += cls.transfer_backend + cls.rdma_devices_for(range(4, 8))
cls.process_decode = popen_launch_pd_server(
cls.model,
cls.decode_url,
@@ -43,7 +43,7 @@ class TestDisaggregationHybridAttentionGDN(PDDisaggregationServerBase):
"--tp",
"4",
]
prefill_args += cls.transfer_backend + cls.rdma_devices
prefill_args += cls.transfer_backend + cls.rdma_devices_for(range(4))
cls.process_prefill = popen_launch_pd_server(
cls.model,
cls.prefill_url,
@@ -64,7 +64,7 @@ class TestDisaggregationHybridAttentionGDN(PDDisaggregationServerBase):
"--base-gpu-id",
"4",
]
decode_args += cls.transfer_backend + cls.rdma_devices
decode_args += cls.transfer_backend + cls.rdma_devices_for(range(4, 8))
cls.process_decode = popen_launch_pd_server(
cls.model,
cls.decode_url,
@@ -117,7 +117,7 @@ class TestDisaggregationHybridAttentionGDNExtraBuffer(PDDisaggregationServerBase
"--mamba-radix-cache-strategy",
"extra_buffer",
]
prefill_args += cls.transfer_backend + cls.rdma_devices
prefill_args += cls.transfer_backend + cls.rdma_devices_for(range(4))
cls.process_prefill = popen_launch_pd_server(
cls.model,
cls.prefill_url,
@@ -140,7 +140,7 @@ class TestDisaggregationHybridAttentionGDNExtraBuffer(PDDisaggregationServerBase
"--mamba-radix-cache-strategy",
"extra_buffer",
]
decode_args += cls.transfer_backend + cls.rdma_devices
decode_args += cls.transfer_backend + cls.rdma_devices_for(range(4, 8))
cls.process_decode = popen_launch_pd_server(
cls.model,
cls.decode_url,
@@ -194,7 +194,7 @@ class TestDisaggregationHybridAttentionGDNDPDecode(PDDisaggregationServerBase):
"--tp",
"2",
]
prefill_args += cls.transfer_backend + cls.rdma_devices
prefill_args += cls.transfer_backend + cls.rdma_devices_for(range(2))
cls.process_prefill = popen_launch_pd_server(
cls.model,
cls.prefill_url,
@@ -219,7 +219,7 @@ class TestDisaggregationHybridAttentionGDNDPDecode(PDDisaggregationServerBase):
"--base-gpu-id",
"2",
]
decode_args += cls.transfer_backend + cls.rdma_devices
decode_args += cls.transfer_backend + cls.rdma_devices_for(range(2, 4))
cls.process_decode = popen_launch_pd_server(
cls.model,
cls.decode_url,
@@ -271,7 +271,7 @@ class TestDisaggregationHybridAttentionMamba(PDDisaggregationServerBase):
"--tp",
"4",
]
prefill_args += cls.transfer_backend + cls.rdma_devices
prefill_args += cls.transfer_backend + cls.rdma_devices_for(range(4))
cls.process_prefill = popen_launch_pd_server(
cls.model,
cls.prefill_url,
@@ -292,7 +292,7 @@ class TestDisaggregationHybridAttentionMamba(PDDisaggregationServerBase):
"--base-gpu-id",
"4",
]
decode_args += cls.transfer_backend + cls.rdma_devices
decode_args += cls.transfer_backend + cls.rdma_devices_for(range(4, 8))
cls.process_decode = popen_launch_pd_server(
cls.model,
cls.decode_url,
@@ -345,7 +345,7 @@ class TestDisaggregationHybridAttentionMambaExtraBuffer(PDDisaggregationServerBa
"--mamba-radix-cache-strategy",
"extra_buffer",
]
prefill_args += cls.transfer_backend + cls.rdma_devices
prefill_args += cls.transfer_backend + cls.rdma_devices_for(range(4))
cls.process_prefill = popen_launch_pd_server(
cls.model,
cls.prefill_url,
@@ -368,7 +368,7 @@ class TestDisaggregationHybridAttentionMambaExtraBuffer(PDDisaggregationServerBa
"--mamba-radix-cache-strategy",
"extra_buffer",
]
decode_args += cls.transfer_backend + cls.rdma_devices
decode_args += cls.transfer_backend + cls.rdma_devices_for(range(4, 8))
cls.process_decode = popen_launch_pd_server(
cls.model,
cls.decode_url,
@@ -97,7 +97,8 @@ class TestDisaggregationInklingMXFP8(PDDisaggregationServerBase, GSM8KMixin):
*COMMON_ARGS,
"--enable-hierarchical-cache",
]
prefill_args += cls.transfer_backend + cls.rdma_devices
# COMMON_ARGS carries --tp 2, so each side spans two GPUs.
prefill_args += cls.transfer_backend + cls.rdma_devices_for(range(2))
cls.process_prefill = popen_launch_pd_server(
cls.model,
cls.prefill_url,
@@ -118,7 +119,7 @@ class TestDisaggregationInklingMXFP8(PDDisaggregationServerBase, GSM8KMixin):
"--base-gpu-id",
"2",
]
decode_args += cls.transfer_backend + cls.rdma_devices
decode_args += cls.transfer_backend + cls.rdma_devices_for(range(2, 4))
cls.process_decode = popen_launch_pd_server(
cls.model,
cls.decode_url,
@@ -151,7 +151,9 @@ class NixlPDDisaggregationServerBase(PDDisaggregationServerBase):
"--tp",
str(cls.prefill_tp_size),
] + list(cls.extra_prefill_args)
prefill_args += cls.transfer_backend + cls.rdma_devices
prefill_args += cls.transfer_backend + cls.rdma_devices_for(
range(cls.prefill_tp_size)
)
cls.process_prefill = popen_launch_pd_server(
cls.model,
cls.prefill_url,
@@ -178,7 +180,9 @@ class NixlPDDisaggregationServerBase(PDDisaggregationServerBase):
"--base-gpu-id",
str(cls.decode_base_gpu_id),
] + list(cls.extra_decode_args)
decode_args += cls.transfer_backend + cls.rdma_devices
decode_args += cls.transfer_backend + cls.rdma_devices_for(
range(cls.decode_base_gpu_id, cls.decode_base_gpu_id + cls.decode_tp_size)
)
cls.process_decode = popen_launch_pd_server(
cls.model,
cls.decode_url,
@@ -48,7 +48,7 @@ class TestDisaggregationPrefillPPDynamicChunkAccuracy(PDDisaggregationServerBase
"--disable-overlap-schedule",
"--enable-dynamic-chunking",
]
prefill_args += cls.transfer_backend + cls.rdma_devices
prefill_args += cls.transfer_backend + cls.rdma_devices_for(range(4))
cls.process_prefill = popen_launch_pd_server(
cls.model,
cls.prefill_url,
@@ -69,7 +69,7 @@ class TestDisaggregationPrefillPPDynamicChunkAccuracy(PDDisaggregationServerBase
"--base-gpu-id",
"4",
]
decode_args += cls.transfer_backend + cls.rdma_devices
decode_args += cls.transfer_backend + cls.rdma_devices_for(range(4, 6))
cls.process_decode = popen_launch_pd_server(
cls.model,
cls.decode_url,
@@ -125,7 +125,7 @@ class TestDisaggregationDecodePPAccuracy(PDDisaggregationServerBase):
"2",
"--disable-overlap-schedule",
]
prefill_args += cls.transfer_backend + cls.rdma_devices
prefill_args += cls.transfer_backend + cls.rdma_devices_for(range(4))
cls.process_prefill = popen_launch_pd_server(
cls.model,
cls.prefill_url,
@@ -148,7 +148,7 @@ class TestDisaggregationDecodePPAccuracy(PDDisaggregationServerBase):
"--base-gpu-id",
"4",
]
decode_args += cls.transfer_backend + cls.rdma_devices
decode_args += cls.transfer_backend + cls.rdma_devices_for(range(4, 8))
cls.process_decode = popen_launch_pd_server(
cls.model,
cls.decode_url,
@@ -187,7 +187,7 @@ class TestEPDDisaggregationOmni(PDDisaggregationServerBase):
"--port",
cls.prefill_port,
]
prefill_args += cls.transfer_backend + cls.rdma_devices
prefill_args += cls.transfer_backend + cls.rdma_devices_for(range(1, 2))
prefill_env = os.environ.copy()
if cls.server_type == "grpc":
prefill_env["SGLANG_ENCODER_MM_RECEIVER_MODE"] = "grpc"
@@ -214,7 +214,7 @@ class TestEPDDisaggregationOmni(PDDisaggregationServerBase):
"--port",
cls.decode_port,
]
decode_args += cls.transfer_backend + cls.rdma_devices
decode_args += cls.transfer_backend + cls.rdma_devices_for(range(2, 3))
cls.process_decode = popen_launch_server(
cls.model,
base_url=cls.decode_url,
@@ -703,7 +703,7 @@ class TestEPDDisaggregationOneEncoder(MMMUMixin, PDDisaggregationServerBase):
"--port",
cls.prefill_port,
]
prefill_args += cls.transfer_backend + cls.rdma_devices
prefill_args += cls.transfer_backend + cls.rdma_devices_for(range(1, 2))
cls.process_prefill = popen_launch_server(
cls.model,
base_url=cls.prefill_url,
@@ -727,7 +727,7 @@ class TestEPDDisaggregationOneEncoder(MMMUMixin, PDDisaggregationServerBase):
"--port",
cls.decode_port,
]
decode_args += cls.transfer_backend + cls.rdma_devices
decode_args += cls.transfer_backend + cls.rdma_devices_for(range(2, 3))
cls.process_decode = popen_launch_server(
cls.model,
base_url=cls.decode_url,
@@ -824,7 +824,7 @@ class TestEPDDisaggregationKimiVL(PDDisaggregationServerBase):
cls.prefill_port,
*cls.model_args,
]
prefill_args += cls.transfer_backend + cls.rdma_devices
prefill_args += cls.transfer_backend + cls.rdma_devices_for(range(1, 2))
cls.process_prefill = popen_launch_server(
cls.model,
base_url=cls.prefill_url,
@@ -848,7 +848,7 @@ class TestEPDDisaggregationKimiVL(PDDisaggregationServerBase):
cls.decode_port,
*cls.model_args,
]
decode_args += cls.transfer_backend + cls.rdma_devices
decode_args += cls.transfer_backend + cls.rdma_devices_for(range(2, 3))
cls.process_decode = popen_launch_server(
cls.model,
base_url=cls.decode_url,
@@ -1219,7 +1219,7 @@ class TestEPDDisaggregationMultiEncoders(MMMUMixin, PDDisaggregationServerBase):
"--port",
cls.prefill_port,
]
prefill_args += cls.transfer_backend + cls.rdma_devices
prefill_args += cls.transfer_backend + cls.rdma_devices_for(range(2, 3))
cls.process_prefill = popen_launch_server(
cls.model,
base_url=cls.prefill_url,
@@ -1243,7 +1243,7 @@ class TestEPDDisaggregationMultiEncoders(MMMUMixin, PDDisaggregationServerBase):
"--port",
cls.decode_port,
]
decode_args += cls.transfer_backend + cls.rdma_devices
decode_args += cls.transfer_backend + cls.rdma_devices_for(range(3, 4))
cls.process_decode = popen_launch_server(
cls.model,
base_url=cls.decode_url,
@@ -1352,7 +1352,7 @@ class TestEPDDisaggregationGrpcEncoderMMMU(MMMUMixin, PDDisaggregationServerBase
"--port",
cls.prefill_port,
]
prefill_args += cls.transfer_backend + cls.rdma_devices
prefill_args += cls.transfer_backend + cls.rdma_devices_for(range(1, 2))
prefill_env = os.environ.copy()
prefill_env["SGLANG_ENCODER_MM_RECEIVER_MODE"] = "grpc"
cls.process_prefill = popen_launch_server(
@@ -1378,7 +1378,7 @@ class TestEPDDisaggregationGrpcEncoderMMMU(MMMUMixin, PDDisaggregationServerBase
"--port",
cls.decode_port,
]
decode_args += cls.transfer_backend + cls.rdma_devices
decode_args += cls.transfer_backend + cls.rdma_devices_for(range(2, 3))
cls.process_decode = popen_launch_server(
cls.model,
base_url=cls.decode_url,
@@ -1642,7 +1642,7 @@ class TestEPDDisaggregationMooncake(MMMUMixin, PDDisaggregationServerBase):
"--port",
cls.prefill_port,
]
prefill_args += cls.transfer_backend + cls.rdma_devices
prefill_args += cls.transfer_backend + cls.rdma_devices_for(range(1, 2))
cls.process_prefill = popen_launch_server(
cls.model,
base_url=cls.prefill_url,
@@ -1666,7 +1666,7 @@ class TestEPDDisaggregationMooncake(MMMUMixin, PDDisaggregationServerBase):
"--port",
cls.decode_port,
]
decode_args += cls.transfer_backend + cls.rdma_devices
decode_args += cls.transfer_backend + cls.rdma_devices_for(range(2, 3))
cls.process_decode = popen_launch_server(
cls.model,
base_url=cls.decode_url,