[PD] Fix IB device validation for JSON mappings (#26114)

This commit is contained in:
Teng Ma
2026-05-29 16:44:49 +08:00
committed by GitHub
parent 9062f583db
commit 544f3039d5
4 changed files with 182 additions and 81 deletions
@@ -1,7 +1,7 @@
import json import json
import logging import logging
import os import os
from typing import List, Optional from typing import Dict, List, Optional, Union
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.utils.network import NetworkAddress, get_free_port from sglang.srt.utils.network import NetworkAddress, get_free_port
@@ -12,6 +12,58 @@ logger = logging.getLogger(__name__)
_mooncake_transfer_engine: Optional["MooncakeTransferEngine"] = None _mooncake_transfer_engine: Optional["MooncakeTransferEngine"] = None
def parse_ib_device_config(
ib_device_str: Optional[str],
) -> Optional[Union[str, Dict[int, str]]]:
"""Parse IB device config from a shared string, JSON mapping, or JSON file."""
if ib_device_str is None or not ib_device_str.strip():
return None
normalized_input = ib_device_str.strip()
if not normalized_input.endswith(".json") and not normalized_input.startswith("{"):
return normalized_input
if normalized_input.endswith(".json"):
if not os.path.isfile(normalized_input):
raise RuntimeError(f"File {normalized_input} does not exist.")
try:
with open(normalized_input, "r", encoding="utf-8") as file:
mapping = json.load(file)
except json.JSONDecodeError as exc:
raise RuntimeError(
f"Failed to parse JSON content from file {normalized_input}"
) from exc
except (IOError, OSError) as exc:
raise RuntimeError(
f"Failed to read JSON file {normalized_input}: {exc}"
) from exc
else:
try:
mapping = json.loads(normalized_input)
except json.JSONDecodeError as exc:
raise ValueError(f"Invalid JSON mapping: {normalized_input}") from exc
if not isinstance(mapping, dict):
raise ValueError(
"Invalid format: expected a mapping from GPU id to IB device string"
)
normalized_mapping: Dict[int, str] = {}
for gpu_key, ib_devices in mapping.items():
normalized_key = int(gpu_key) if str(gpu_key).isdigit() else None
if normalized_key is None or not isinstance(ib_devices, str):
raise ValueError(
"Invalid format: keys must be integers (or string "
"representations of integers) and values must be strings"
)
normalized_mapping[normalized_key] = ib_devices.strip()
if not normalized_mapping:
raise ValueError("No valid GPU mappings found in JSON")
return normalized_mapping
def get_ib_devices_for_gpu(ib_device_str: Optional[str], gpu_id: int) -> Optional[str]: def get_ib_devices_for_gpu(ib_device_str: Optional[str], gpu_id: int) -> Optional[str]:
""" """
Parse IB device string and get IB devices for a specific GPU ID. Parse IB device string and get IB devices for a specific GPU ID.
@@ -28,66 +80,20 @@ def get_ib_devices_for_gpu(ib_device_str: Optional[str], gpu_id: int) -> Optiona
Returns: Returns:
IB devices string for the GPU, or None if not available IB devices string for the GPU, or None if not available
""" """
if ib_device_str is None or not ib_device_str.strip(): parsed_config = parse_ib_device_config(ib_device_str)
if parsed_config is None:
return None return None
ib_device_str = ib_device_str.strip() if isinstance(parsed_config, str):
return parsed_config
# Check if it's a JSON file first and load its content if gpu_id in parsed_config:
is_json_file = ib_device_str.endswith(".json") return parsed_config[gpu_id]
if is_json_file:
try:
if os.path.isfile(ib_device_str):
with open(ib_device_str, "r") as f:
ib_device_str = f.read()
else:
# File doesn't exist, treat as old format
raise RuntimeError(f"File {ib_device_str} does not exist.")
except (IOError, OSError) as e:
# File reading failed, raise exception
raise RuntimeError(f"Failed to read JSON file {ib_device_str}: {e}") from e
# Check if it's JSON format (new format) raise ValueError(
try: f"No IB devices configured for GPU {gpu_id}. "
parsed_json = json.loads(ib_device_str) f"Available GPUs: {list(parsed_config.keys())}"
if isinstance(parsed_json, dict): )
# Validate format - keys should be integers (or string rep), values should be strings
gpu_mapping = {}
for gpu_key, ib_devices in parsed_json.items():
if (
isinstance(gpu_key, str)
and gpu_key.isdigit()
and isinstance(ib_devices, str)
):
gpu_mapping[int(gpu_key)] = ib_devices.strip()
elif isinstance(gpu_key, int) and isinstance(ib_devices, str):
gpu_mapping[gpu_key] = ib_devices.strip()
else:
raise ValueError(
"Invalid format: keys must be integers (or string "
"representations of integers) and values must be strings"
)
if not gpu_mapping:
raise ValueError("No valid GPU mappings found in JSON")
# Return devices for specific GPU
if gpu_id in gpu_mapping:
return gpu_mapping[gpu_id]
else:
raise ValueError(
f"No IB devices configured for GPU {gpu_id}. "
f"Available GPUs: {list(gpu_mapping.keys())}"
)
except json.JSONDecodeError:
if is_json_file:
# It was supposed to be a JSON file but failed to parse
raise RuntimeError(
f"Failed to parse JSON content from file {ib_device_str}"
)
# Not JSON format, treat as old format - return same devices for all GPUs
return ib_device_str
class MooncakeTransferEngine: class MooncakeTransferEngine:
@@ -1055,7 +1055,16 @@ class ModelRunner(ModelRunnerKVCacheMixin):
if self.device == "cuda" and self.server_args.elastic_ep_backend == "mooncake": if self.device == "cuda" and self.server_args.elastic_ep_backend == "mooncake":
backend = "mooncake" backend = "mooncake"
if self.server_args.mooncake_ib_device: if self.server_args.mooncake_ib_device:
mooncake_ib_device = self.server_args.mooncake_ib_device.split(",") from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import (
get_ib_devices_for_gpu,
)
ib_device_for_gpu = get_ib_devices_for_gpu(
self.server_args.mooncake_ib_device, self.gpu_id
)
mooncake_ib_device = (
ib_device_for_gpu.split(",") if ib_device_for_gpu else []
)
try: try:
from mooncake import ep as mooncake_ep from mooncake import ep as mooncake_ep
+48 -23
View File
@@ -36,6 +36,9 @@ from sglang.srt.arg_groups.argparse_actions import (
) )
from sglang.srt.configs.linear_attn_model_registry import get_linear_attn_spec_by_arch from sglang.srt.configs.linear_attn_model_registry import get_linear_attn_spec_by_arch
from sglang.srt.connector import ConnectorType from sglang.srt.connector import ConnectorType
from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import (
parse_ib_device_config,
)
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.function_call.function_call_parser import FunctionCallParser from sglang.srt.function_call.function_call_parser import FunctionCallParser
from sglang.srt.layers.attention.fla.chunk_delta_h import CHUNK_SIZE as FLA_CHUNK_SIZE from sglang.srt.layers.attention.fla.chunk_delta_h import CHUNK_SIZE as FLA_CHUNK_SIZE
@@ -3880,15 +3883,16 @@ class ServerArgs:
f"Supported architectures: Qwen2VL, Qwen3VL, Qwen3.5, InternS2, Qwen2Audio, Qwen2.5Omni, Kimi, MiMoV2." f"Supported architectures: Qwen2VL, Qwen3VL, Qwen3.5, InternS2, Qwen2Audio, Qwen2.5Omni, Kimi, MiMoV2."
) )
def _validate_ib_devices(self, device_str: str) -> Optional[str]: def _validate_ib_devices(self, device_str: Optional[str]) -> Optional[str]:
""" """
Validate IB devices before passing to mooncake. Validate IB devices before passing to mooncake.
Args: Args:
device_str: Comma-separated IB device names (e.g., "mlx5_0,mlx5_1") device_str: Comma-separated IB device names, a per-GPU JSON mapping,
or a path to a JSON file containing that mapping.
Returns: Returns:
Normalized comma-separated string of validated device names, or None if input is None. A normalized comma-separated string or per-GPU JSON mapping string, or None if input is None.
""" """
if device_str is None: if device_str is None:
logger.warning( logger.warning(
@@ -3896,20 +3900,34 @@ class ServerArgs:
) )
return None return None
# Strip whitespace from device names def _normalize_device_group(raw_value: str, context: str) -> str:
devices = [d.strip() for d in device_str.split(",") if d.strip()] if not isinstance(raw_value, str):
if len(devices) == 0: raise ValueError(
raise ValueError("No valid IB devices specified") f"Invalid IB device format for {context}: expected a string. "
f"Got {type(raw_value)}"
)
devices = [d.strip() for d in raw_value.split(",") if d.strip()]
if not devices:
raise ValueError(f"No valid IB devices specified for {context}")
unique_devices = list(dict.fromkeys(devices))
if len(unique_devices) != len(devices):
logger.warning(
"Duplicate IB devices specified for %s: %s. Deduplicating to: %s",
context,
raw_value,
",".join(unique_devices),
)
invalid_devices = [d for d in unique_devices if d not in available_devices]
if len(invalid_devices) != 0:
raise ValueError(
f"Invalid IB devices specified for {context}: {invalid_devices}. "
f"Available devices: {sorted(available_devices)}"
)
return ",".join(unique_devices)
# Deduplicate while preserving order normalized_input = device_str.strip()
unique_devices = list(dict.fromkeys(devices)) if not normalized_input:
if len(unique_devices) != len(devices): raise ValueError("No valid IB devices specified")
logger.warning(
"Duplicate IB devices specified: %s. Deduplicating to: %s",
device_str,
",".join(unique_devices),
)
devices = unique_devices
# Get available IB devices from sysfs # Get available IB devices from sysfs
ib_sysfs_path = "/sys/class/infiniband" ib_sysfs_path = "/sys/class/infiniband"
@@ -3923,15 +3941,22 @@ class ServerArgs:
if len(available_devices) == 0: if len(available_devices) == 0:
raise RuntimeError(f"No IB devices found in {ib_sysfs_path}") raise RuntimeError(f"No IB devices found in {ib_sysfs_path}")
# Check for invalid devices parsed_config = parse_ib_device_config(normalized_input)
invalid_devices = [d for d in devices if d not in available_devices] if isinstance(parsed_config, str):
if len(invalid_devices) != 0: return _normalize_device_group(normalized_input, "all GPUs")
raise ValueError( assert parsed_config is not None
f"Invalid IB devices specified: {invalid_devices}. "
f"Available devices: {sorted(available_devices)}" normalized_mapping: Dict[str, str] = {}
for gpu_key, gpu_devices in parsed_config.items():
normalized_key = str(gpu_key)
normalized_mapping[normalized_key] = _normalize_device_group(
gpu_devices, f"GPU {normalized_key}"
) )
return ",".join(devices) if not normalized_mapping:
raise ValueError("No valid GPU mappings found in IB device JSON")
return json.dumps(normalized_mapping, separators=(",", ":"))
def _handle_tokenizer_batching(self): def _handle_tokenizer_batching(self):
if self.enable_tokenizer_batch_encode and self.enable_dynamic_batch_tokenizer: if self.enable_tokenizer_batch_encode and self.enable_dynamic_batch_tokenizer:
@@ -1,3 +1,6 @@
import json
import os
import tempfile
import unittest import unittest
from unittest.mock import patch from unittest.mock import patch
@@ -34,5 +37,63 @@ class TestServerArgsCPUBackend(unittest.TestCase):
self.assertEqual(server_args.sampling_backend, "pytorch") self.assertEqual(server_args.sampling_backend, "pytorch")
class TestServerArgsIBDeviceValidation(unittest.TestCase):
def _validate_ib_devices(self, device_str, available_devices=None):
server_args = ServerArgs.__new__(ServerArgs)
available_devices = available_devices or [
"mlx5_0",
"mlx5_1",
"mlx5_2",
"mlx5_3",
]
real_isdir = os.path.isdir
real_listdir = os.listdir
with patch(
"sglang.srt.server_args.os.path.isdir",
side_effect=lambda path: (
True if path == "/sys/class/infiniband" else real_isdir(path)
),
), patch(
"sglang.srt.server_args.os.listdir",
side_effect=lambda path: (
available_devices
if path == "/sys/class/infiniband"
else real_listdir(path)
),
):
return ServerArgs._validate_ib_devices(server_args, device_str)
def test_validate_ib_devices_accepts_comma_separated(self):
self.assertEqual(
self._validate_ib_devices("mlx5_0, mlx5_1"),
"mlx5_0,mlx5_1",
)
def test_validate_ib_devices_accepts_json_object(self):
result = self._validate_ib_devices(
'{"0": "mlx5_0, mlx5_1", "1": "mlx5_2, mlx5_3"}'
)
self.assertEqual(
json.loads(result),
{"0": "mlx5_0,mlx5_1", "1": "mlx5_2,mlx5_3"},
)
def test_validate_ib_devices_accepts_json_file(self):
with tempfile.NamedTemporaryFile("w", suffix=".json", delete=False) as file:
json.dump({"0": "mlx5_0, mlx5_1", "1": "mlx5_2"}, file)
json_file = file.name
try:
result = self._validate_ib_devices(json_file)
finally:
os.unlink(json_file)
self.assertEqual(
json.loads(result),
{"0": "mlx5_0,mlx5_1", "1": "mlx5_2"},
)
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()