From 87c9f78e994f6e956b1e25291766288c7928d319 Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Tue, 15 Sep 2026 15:07:42 -0700 Subject: [PATCH] [Fix][PD] Give prefill and decode their own RDMA NICs in disaggregation tests (#39584) --- .../server_fixtures/disaggregation_fixture.py | 102 +++++++----------- .../test_disaggregation_aarch64.py | 4 +- .../test_disaggregation_decode_offload.py | 4 +- .../test_disaggregation_different_tp.py | 32 +++--- .../test_disaggregation_dp_attention.py | 8 +- .../test_disaggregation_dsv4.py | 4 +- .../test_disaggregation_hisparse.py | 4 +- .../test_disaggregation_hybrid_attention.py | 20 ++-- .../test_disaggregation_inkling_mxfp8.py | 5 +- .../test_disaggregation_nixl.py | 8 +- .../disaggregation/test_disaggregation_pp.py | 8 +- .../disaggregation/test_epd_disaggregation.py | 24 ++--- 12 files changed, 103 insertions(+), 120 deletions(-) diff --git a/python/sglang/test/server_fixtures/disaggregation_fixture.py b/python/sglang/test/server_fixtures/disaggregation_fixture.py index 3f7067e31..9956e4728 100644 --- a/python/sglang/test/server_fixtures/disaggregation_fixture.py +++ b/python/sglang/test/server_fixtures/disaggregation_fixture.py @@ -114,6 +114,13 @@ class PDDisaggregationServerBase(CustomTestCase): msg = "No RDMA devices specified for disaggregation test, using default settings." warnings.warn(msg) + @classmethod + def rdma_devices_for(cls, gpu_indices) -> list: + """`--disaggregation-ib-device` args for a server pinned to these GPUs.""" + if not is_in_ci(): + return cls.rdma_devices + return ["--disaggregation-ib-device", get_rdma_devices_args(gpu_indices)] + # Subclasses can set these to customize server args extra_prefill_args = [] extra_decode_args = [] @@ -131,7 +138,9 @@ class PDDisaggregationServerBase(CustomTestCase): "--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, @@ -160,7 +169,9 @@ class PDDisaggregationServerBase(CustomTestCase): "--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, @@ -338,93 +349,56 @@ def _get_available_ib_devices(): return devices if devices else None -def get_rdma_devices_args(): +def get_rdma_devices_args(gpu_indices=None) -> str: + """RDMA devices for a server pinned to `gpu_indices`, as absolute node ids. + + Ids relative to a group base would map a decode server pinned with + `--base-gpu-id 4` onto the first NICs, across the socket boundary. + """ + def _parse_list_env(var_name: str): val = os.getenv(var_name) - if not val: - return None - items = [x.strip() for x in val.split(",") if x.strip()] + items = [x.strip() for x in (val or "").split(",") if x.strip()] return items or None - def _pick_default_pair(rdma_all_devices): - return [rdma_all_devices[0], rdma_all_devices[len(rdma_all_devices) // 2]] - - # Priority: env var > auto-detect > hardcoded fallback rdma_all_devices = ( _parse_list_env("SGLANG_CI_RDMA_ALL_DEVICES") or _get_available_ib_devices() or [f"mlx5_roce{i}" for i in range(8)] ) - logger.warning("Resolved rdma_all_devices=%s", rdma_all_devices) - n_rdma = len(rdma_all_devices) - # 1. Get visible GPU indices - cuda_visible_devices = os.getenv("CUDA_VISIBLE_DEVICES") - if not cuda_visible_devices: - warnings.warn("CUDA_VISIBLE_DEVICES is not set. Using default RDMA devices.") - return ",".join(_pick_default_pair(rdma_all_devices)) - + if gpu_indices is None: + gpu_indices = _parse_list_env("CUDA_VISIBLE_DEVICES") or [] try: - # Convert to list of integers (handling possible spaces and empty strings) - gpu_indices = [ - int(idx.strip()) for idx in cuda_visible_devices.split(",") if idx.strip() - ] - if not gpu_indices or len(gpu_indices) > 4: - return ",".join(_pick_default_pair(rdma_all_devices)) + gpu_indices = sorted({int(g) for g in gpu_indices}) except ValueError: - warnings.warn(f"Invalid CUDA_VISIBLE_DEVICES format: {cuda_visible_devices}") - return ",".join(_pick_default_pair(rdma_all_devices)) + warnings.warn(f"Invalid GPU indices: {gpu_indices}") + gpu_indices = [] + if not gpu_indices: + return ",".join([rdma_all_devices[0], rdma_all_devices[n_rdma // 2]]) - # 2. Calculate base RDMA index group (each group of 4 GPUs uses consecutive devices) - base_rdma_group = (min(gpu_indices) // 4) * 4 - for gpu_idx in gpu_indices: - if not (base_rdma_group <= gpu_idx < base_rdma_group + 4): - warnings.warn( - f"GPU index {gpu_idx} is outside expected group " - f"{base_rdma_group}-{base_rdma_group + 3}" - ) - - # 3. Generate RDMA device names - # Detect total GPUs on the node (not just visible ones) try: import torch total_gpus = torch.cuda.device_count() except Exception: - total_gpus = 8 # Fallback to common 8-GPU setup + total_gpus = 0 + total_gpus = total_gpus or max(8, gpu_indices[-1] + 1) - # Handle edge cases - if total_gpus == 0: - total_gpus = 8 - if n_rdma > total_gpus: - logger.warning( - "More RDMA devices (%d) than GPUs (%d), using first and middle device", - n_rdma, - total_gpus, - ) - return ",".join(_pick_default_pair(rdma_all_devices)) - - # Calculate how many GPUs share each RDMA device gpus_per_rdma = max(1, total_gpus // n_rdma) + devices = [ + rdma_all_devices[min(gpu // gpus_per_rdma, n_rdma - 1)] for gpu in gpu_indices + ] + resolved = ",".join(dict.fromkeys(devices)) logger.warning( - "GPU-to-RDMA mapping: total_gpus=%d, n_rdma=%d, gpus_per_rdma=%d", + "RDMA for gpus=%s: total_gpus=%d n_rdma=%d -> %s", + gpu_indices, total_gpus, n_rdma, - gpus_per_rdma, + resolved, ) - - rdma_devices = [] - base_gpu = min(gpu_indices) - for gpu_idx in gpu_indices: - nic_index = min((gpu_idx - base_gpu) // gpus_per_rdma, n_rdma - 1) - rdma_devices.append(rdma_all_devices[nic_index]) - - if not rdma_devices: - return ",".join(_pick_default_pair(rdma_all_devices)) - - # Deduplicate while preserving order - return ",".join(dict.fromkeys(rdma_devices)) + return resolved _IB_SYSFS = "/sys/class/infiniband" diff --git a/test/registered/disaggregation/test_disaggregation_aarch64.py b/test/registered/disaggregation/test_disaggregation_aarch64.py index 18ad5b344..9c6a2aa79 100644 --- a/test/registered/disaggregation/test_disaggregation_aarch64.py +++ b/test/registered/disaggregation/test_disaggregation_aarch64.py @@ -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, diff --git a/test/registered/disaggregation/test_disaggregation_decode_offload.py b/test/registered/disaggregation/test_disaggregation_decode_offload.py index 05f4eb1f7..f4dfd0219 100644 --- a/test/registered/disaggregation/test_disaggregation_decode_offload.py +++ b/test/registered/disaggregation/test_disaggregation_decode_offload.py @@ -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, diff --git a/test/registered/disaggregation/test_disaggregation_different_tp.py b/test/registered/disaggregation/test_disaggregation_different_tp.py index add3c6a1d..67145d595 100644 --- a/test/registered/disaggregation/test_disaggregation_different_tp.py +++ b/test/registered/disaggregation/test_disaggregation_different_tp.py @@ -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, diff --git a/test/registered/disaggregation/test_disaggregation_dp_attention.py b/test/registered/disaggregation/test_disaggregation_dp_attention.py index e22b15e67..50a8a4f1c 100644 --- a/test/registered/disaggregation/test_disaggregation_dp_attention.py +++ b/test/registered/disaggregation/test_disaggregation_dp_attention.py @@ -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, diff --git a/test/registered/disaggregation/test_disaggregation_dsv4.py b/test/registered/disaggregation/test_disaggregation_dsv4.py index c4a3b6bff..4d5722f54 100644 --- a/test/registered/disaggregation/test_disaggregation_dsv4.py +++ b/test/registered/disaggregation/test_disaggregation_dsv4.py @@ -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, diff --git a/test/registered/disaggregation/test_disaggregation_hisparse.py b/test/registered/disaggregation/test_disaggregation_hisparse.py index 602619a1e..1b212012b 100644 --- a/test/registered/disaggregation/test_disaggregation_hisparse.py +++ b/test/registered/disaggregation/test_disaggregation_hisparse.py @@ -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, diff --git a/test/registered/disaggregation/test_disaggregation_hybrid_attention.py b/test/registered/disaggregation/test_disaggregation_hybrid_attention.py index b97c51c2c..66f19afee 100644 --- a/test/registered/disaggregation/test_disaggregation_hybrid_attention.py +++ b/test/registered/disaggregation/test_disaggregation_hybrid_attention.py @@ -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, diff --git a/test/registered/disaggregation/test_disaggregation_inkling_mxfp8.py b/test/registered/disaggregation/test_disaggregation_inkling_mxfp8.py index 9a2f4c555..6821e91b8 100644 --- a/test/registered/disaggregation/test_disaggregation_inkling_mxfp8.py +++ b/test/registered/disaggregation/test_disaggregation_inkling_mxfp8.py @@ -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, diff --git a/test/registered/disaggregation/test_disaggregation_nixl.py b/test/registered/disaggregation/test_disaggregation_nixl.py index 8754d5145..04850bc32 100644 --- a/test/registered/disaggregation/test_disaggregation_nixl.py +++ b/test/registered/disaggregation/test_disaggregation_nixl.py @@ -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, diff --git a/test/registered/disaggregation/test_disaggregation_pp.py b/test/registered/disaggregation/test_disaggregation_pp.py index a08c678ff..47be3e223 100644 --- a/test/registered/disaggregation/test_disaggregation_pp.py +++ b/test/registered/disaggregation/test_disaggregation_pp.py @@ -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, diff --git a/test/registered/disaggregation/test_epd_disaggregation.py b/test/registered/disaggregation/test_epd_disaggregation.py index 3b0364f46..ecfb7a075 100644 --- a/test/registered/disaggregation/test_epd_disaggregation.py +++ b/test/registered/disaggregation/test_epd_disaggregation.py @@ -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,