@@ -0,0 +1,331 @@
|
||||
import logging
|
||||
import os
|
||||
import subprocess
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import regex as re
|
||||
|
||||
from tests.e2e.nightly.multi_node.scripts.utils import (
|
||||
get_available_port,
|
||||
get_net_interface,
|
||||
load_yaml_mapping,
|
||||
resolve_cluster_ips,
|
||||
resolve_current_node_index,
|
||||
setup_logger,
|
||||
)
|
||||
|
||||
setup_logger()
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
DEFAULT_CONFIG_BASE_PATH = "tests/e2e/nightly/multi_node/internal_dp/config/"
|
||||
DEFAULT_SERVER_PORT = 8080
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class NodeInfo:
|
||||
index: int
|
||||
ip: str
|
||||
server_cmd: str
|
||||
envs: dict[str, Any] | None = None
|
||||
headless: bool = False
|
||||
|
||||
def __post_init__(self):
|
||||
if not self.ip:
|
||||
raise ValueError("NodeInfo.ip must not be empty")
|
||||
|
||||
def __str__(self) -> str:
|
||||
return f"NodeInfo(\n index={self.index},\n ip={self.ip},\n headless={self.headless},\n)"
|
||||
|
||||
|
||||
class DisaggregatedPrefillCfg:
|
||||
def __init__(self, raw_cfg: dict, num_nodes: int):
|
||||
self.prefiller_indices: list[int] = raw_cfg.get("prefiller_host_index", [])
|
||||
self.decoder_indices: list[int] = raw_cfg.get("decoder_host_index", [])
|
||||
|
||||
if not self.decoder_indices:
|
||||
raise RuntimeError("decoder_host_index must be provided")
|
||||
|
||||
self._validate(num_nodes)
|
||||
|
||||
self.decode_start_index = self.decoder_indices[0]
|
||||
self.num_prefillers = len(self.prefiller_indices)
|
||||
self.num_decoders = len(self.decoder_indices)
|
||||
|
||||
def _validate(self, num_nodes: int):
|
||||
overlap = set(self.prefiller_indices) & set(self.decoder_indices)
|
||||
if overlap:
|
||||
raise AssertionError(f"Prefiller and decoder overlap: {overlap}")
|
||||
|
||||
all_indices = self.prefiller_indices + self.decoder_indices
|
||||
if any(i >= num_nodes for i in all_indices):
|
||||
raise ValueError("Disaggregated prefill index out of range")
|
||||
|
||||
def is_prefiller(self, index: int) -> bool:
|
||||
return index in self.prefiller_indices
|
||||
|
||||
def is_decoder(self, index: int) -> bool:
|
||||
return index in self.decoder_indices
|
||||
|
||||
def master_ip_for_node(self, index: int, nodes: list[NodeInfo]) -> str:
|
||||
if self.is_prefiller(index):
|
||||
return nodes[0].ip
|
||||
return nodes[self.decode_start_index].ip
|
||||
|
||||
|
||||
class DistEnvBuilder:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
cur_node: NodeInfo,
|
||||
master_ip: str,
|
||||
):
|
||||
self.cur_ip = cur_node.ip
|
||||
self.nic_name = get_net_interface(self.cur_ip)
|
||||
self.master_ip = master_ip
|
||||
|
||||
self.base_envs = dict(cur_node.envs or {})
|
||||
|
||||
def build(self) -> dict:
|
||||
envs = dict(self.base_envs)
|
||||
|
||||
envs.update(
|
||||
{
|
||||
"HCCL_IF_IP": self.cur_ip,
|
||||
"HCCL_SOCKET_IFNAME": self.nic_name,
|
||||
"GLOO_SOCKET_IFNAME": self.nic_name,
|
||||
"TP_SOCKET_IFNAME": self.nic_name,
|
||||
"LOCAL_IP": self.cur_ip,
|
||||
"NIC_NAME": self.nic_name,
|
||||
"MASTER_IP": self.master_ip,
|
||||
}
|
||||
)
|
||||
|
||||
return {k: str(v) for k, v in envs.items()}
|
||||
|
||||
|
||||
class ProxyLauncher:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
nodes: list[NodeInfo],
|
||||
envs: dict,
|
||||
proxy_port: int,
|
||||
cur_index: int,
|
||||
disagg_cfg: DisaggregatedPrefillCfg | None = None,
|
||||
):
|
||||
self.nodes = nodes
|
||||
self.cfg = disagg_cfg
|
||||
self.server_port = envs.get("SERVER_PORT", DEFAULT_SERVER_PORT)
|
||||
self.proxy_port = proxy_port
|
||||
self.proxy_script = envs.get(
|
||||
"DISAGGREGATED_PREFILL_PROXY_SCRIPT",
|
||||
"examples/disaggregated_prefill_v1/load_balance_proxy_server_example.py",
|
||||
)
|
||||
self.envs = envs
|
||||
self.is_master = cur_index == 0
|
||||
self.cur_ip = nodes[cur_index].ip
|
||||
self.process: subprocess.Popen[bytes] | None = None
|
||||
|
||||
def __enter__(self):
|
||||
if not self.is_master or self.cfg is None:
|
||||
logger.info("Not launching proxy on non-master node")
|
||||
return self
|
||||
prefiller_ips = [self.nodes[i].ip for i in self.cfg.prefiller_indices if not self.nodes[i].headless]
|
||||
decoder_ips = [self.nodes[i].ip for i in self.cfg.decoder_indices if not self.nodes[i].headless]
|
||||
|
||||
cmd = [
|
||||
"python",
|
||||
self.proxy_script,
|
||||
"--host",
|
||||
self.cur_ip,
|
||||
"--port",
|
||||
str(self.proxy_port),
|
||||
"--prefiller-hosts",
|
||||
*prefiller_ips,
|
||||
"--prefiller-ports",
|
||||
*[str(self.server_port)] * len(prefiller_ips),
|
||||
"--decoder-hosts",
|
||||
*decoder_ips,
|
||||
"--decoder-ports",
|
||||
*[str(self.server_port)] * len(decoder_ips),
|
||||
]
|
||||
|
||||
logger.info("Launching proxy: %s", " ".join(cmd))
|
||||
self.process = subprocess.Popen(cmd, env={**os.environ, **self.envs})
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc, tb):
|
||||
if not self.process:
|
||||
return
|
||||
logger.info("Stopping proxy server...")
|
||||
self.process.terminate()
|
||||
try:
|
||||
self.process.wait(timeout=5)
|
||||
except subprocess.TimeoutExpired:
|
||||
self.process.kill()
|
||||
|
||||
|
||||
class MultiNodeConfig:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
model: str,
|
||||
test_name: str,
|
||||
nodes: list[NodeInfo],
|
||||
npu_per_node: int,
|
||||
disaggregated_prefill: dict | None,
|
||||
benchmark_cases: list[dict],
|
||||
special_dependencies: dict,
|
||||
):
|
||||
self.model = model
|
||||
self.test_name = test_name
|
||||
self.nodes = nodes
|
||||
self.npu_per_node = npu_per_node
|
||||
self.benchmark_cases = benchmark_cases
|
||||
|
||||
self.cur_index = self._resolve_cur_index()
|
||||
self.cur_node = self.nodes[self.cur_index]
|
||||
self.special_dependencies = special_dependencies
|
||||
|
||||
self.disagg_cfg = DisaggregatedPrefillCfg(disaggregated_prefill, len(nodes)) if disaggregated_prefill else None
|
||||
|
||||
master_ip = (
|
||||
self.disagg_cfg.master_ip_for_node(self.cur_index, self.nodes) if self.disagg_cfg else self.nodes[0].ip
|
||||
)
|
||||
self.proxy_port = get_available_port()
|
||||
|
||||
self.envs = DistEnvBuilder(
|
||||
cur_node=self.cur_node,
|
||||
master_ip=master_ip,
|
||||
).build()
|
||||
logger.info("Node %d envs: %s", self.cur_index, self.envs)
|
||||
|
||||
self.server_cmd = self._expand_env(self.cur_node.server_cmd)
|
||||
|
||||
def _resolve_cur_index(self) -> int:
|
||||
return resolve_current_node_index([node.ip for node in self.nodes])
|
||||
|
||||
def _expand_env(self, cmd: str) -> str:
|
||||
pattern = re.compile(r"\$(\w+)|\$\{(\w+)\}")
|
||||
|
||||
def repl(m):
|
||||
key = m.group(1) or m.group(2)
|
||||
return self.envs.get(key, m.group(0))
|
||||
|
||||
return pattern.sub(repl, cmd)
|
||||
|
||||
@property
|
||||
def world_size(self) -> int:
|
||||
return len(self.nodes) * self.npu_per_node
|
||||
|
||||
@property
|
||||
def is_master(self) -> bool:
|
||||
return self.cur_index == 0
|
||||
|
||||
@property
|
||||
def server_port(self) -> int:
|
||||
return self.envs.get("SERVER_PORT", DEFAULT_SERVER_PORT)
|
||||
|
||||
@property
|
||||
def master_ip(self) -> str:
|
||||
return self.nodes[0].ip
|
||||
|
||||
@property
|
||||
def benchmark_endpoint(self) -> tuple[str, int]:
|
||||
"""
|
||||
Endpoint used by benchmark clients.
|
||||
"""
|
||||
master_ip = self.nodes[0].ip
|
||||
server_port = self.envs.get("SERVER_PORT", DEFAULT_SERVER_PORT)
|
||||
if self.disagg_cfg:
|
||||
return master_ip, self.proxy_port
|
||||
return master_ip, server_port
|
||||
|
||||
|
||||
class MultiNodeConfigLoader:
|
||||
"""Load MultiNodeConfig from yaml file."""
|
||||
|
||||
DEFAULT_CONFIG_NAME = "DeepSeek-V3.yaml"
|
||||
|
||||
@classmethod
|
||||
def from_yaml(cls, yaml_path: str | None = None) -> MultiNodeConfig:
|
||||
config = cls._load_yaml(yaml_path)
|
||||
cls._validate_root(config)
|
||||
|
||||
nodes = cls._parse_nodes(config)
|
||||
benchmarks = cls._parse_benchmarks(config)
|
||||
|
||||
return MultiNodeConfig(
|
||||
model=config["model"],
|
||||
test_name=config.get("test_name", "untitled_test"),
|
||||
nodes=nodes,
|
||||
npu_per_node=config.get("npu_per_node", 16),
|
||||
disaggregated_prefill=config.get("disaggregated_prefill"),
|
||||
special_dependencies=config.get("special_dependencies", {}),
|
||||
benchmark_cases=list(benchmarks.values()),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _load_yaml(cls, yaml_path: str | None) -> dict:
|
||||
return load_yaml_mapping(
|
||||
yaml_path,
|
||||
default_name=cls.DEFAULT_CONFIG_NAME,
|
||||
default_base_path=DEFAULT_CONFIG_BASE_PATH,
|
||||
description="config",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _validate_root(cfg: dict):
|
||||
required = ["model", "deployment", "num_nodes", "npu_per_node", "benchmarks"]
|
||||
missing = [k for k in required if k not in cfg]
|
||||
if missing:
|
||||
raise KeyError(f"Missing required config fields: {missing}")
|
||||
|
||||
@classmethod
|
||||
def _parse_nodes(cls, cfg: dict) -> list[NodeInfo]:
|
||||
num_nodes = cfg["num_nodes"]
|
||||
deployments = cfg["deployment"]
|
||||
|
||||
if len(deployments) != num_nodes:
|
||||
raise AssertionError(f"deployment size ({len(deployments)}) != num_nodes ({num_nodes})")
|
||||
|
||||
for idx, deploy in enumerate(deployments):
|
||||
if deploy.get("envs") is None:
|
||||
raise KeyError(f"deployment[{idx}].envs is required for multi-node configs")
|
||||
|
||||
cluster_ips = cls._resolve_cluster_ips(cfg, num_nodes)
|
||||
|
||||
nodes: list[NodeInfo] = []
|
||||
for idx, deploy in enumerate(deployments):
|
||||
cmd = deploy.get("server_cmd", "")
|
||||
envs = deploy["envs"]
|
||||
nodes.append(
|
||||
NodeInfo(
|
||||
index=idx,
|
||||
ip=cluster_ips[idx],
|
||||
server_cmd=cmd,
|
||||
envs=envs,
|
||||
headless="--headless" in cmd,
|
||||
)
|
||||
)
|
||||
return nodes
|
||||
|
||||
@staticmethod
|
||||
def _parse_benchmarks(cfg: dict) -> dict:
|
||||
benchmarks = cfg.get("benchmarks") or {}
|
||||
for name, case in benchmarks.items():
|
||||
case["case_name"] = name
|
||||
return benchmarks
|
||||
|
||||
@staticmethod
|
||||
def _resolve_cluster_ips(cfg: dict, num_nodes: int) -> list[str]:
|
||||
return resolve_cluster_ips(
|
||||
cfg,
|
||||
num_nodes,
|
||||
cluster_hosts_log_message=(
|
||||
"Using cluster_hosts from config. This typically indicates that your current environment is a "
|
||||
"non-Kubernetes environment."
|
||||
),
|
||||
dns_log_message="Resolving cluster IPs via DNS...",
|
||||
)
|
||||
@@ -0,0 +1,207 @@
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import shlex
|
||||
import subprocess
|
||||
import sys
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
import vllm
|
||||
|
||||
from tests.e2e.conftest import RemoteOpenAIServer
|
||||
from tests.e2e.nightly.multi_node.internal_dp.scripts.multi_node_config import (
|
||||
MultiNodeConfig,
|
||||
MultiNodeConfigLoader,
|
||||
ProxyLauncher,
|
||||
)
|
||||
from tests.e2e.nightly.multi_node.scripts.benchmark_results import (
|
||||
build_task_entry,
|
||||
extract_hardware,
|
||||
filter_environment,
|
||||
write_results_json,
|
||||
)
|
||||
from tools.aisbench import run_aisbench_cases
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_FEATURE_ENVS: dict[str, str] = {
|
||||
"VLLM_ASCEND_ENABLE_FLASHCOMM": "flashcomm",
|
||||
"VLLM_ASCEND_ENABLE_FLASHCOMM1": "flashcomm1",
|
||||
"VLLM_ASCEND_ENABLE_TOPK_OPTIMIZE": "topk_optimize",
|
||||
"VLLM_ASCEND_ENABLE_MATMUL_ALLREDUCE": "matmul_allreduce",
|
||||
"VLLM_ASCEND_ENABLE_MLAPO": "mlapo",
|
||||
"VLLM_ASCEND_ENABLE_FUSED_MC2": "fused_mc2",
|
||||
}
|
||||
|
||||
|
||||
def _extract_dtype(config: MultiNodeConfig) -> str:
|
||||
"""Determine weight dtype: w8a8 if model name contains 'w8a8' and any node uses --quantization ascend."""
|
||||
has_w8a8 = "w8a8" in config.model.lower()
|
||||
has_quant_ascend = any("--quantization ascend" in node.server_cmd for node in config.nodes)
|
||||
return "w8a8" if (has_w8a8 and has_quant_ascend) else "bf16"
|
||||
|
||||
|
||||
def _cmd_to_list(server_cmd: list[str] | str) -> list[str]:
|
||||
"""Normalize server_cmd to a list of argument strings."""
|
||||
if isinstance(server_cmd, str):
|
||||
try:
|
||||
return shlex.split(server_cmd)
|
||||
except ValueError:
|
||||
return server_cmd.split()
|
||||
return list(server_cmd)
|
||||
|
||||
|
||||
def _extract_server_cmd_value(cmd_list: list[str], flag: str) -> str | None:
|
||||
"""Return the value following `flag` in a command list, or None."""
|
||||
try:
|
||||
idx = cmd_list.index(flag)
|
||||
return cmd_list[idx + 1]
|
||||
except (ValueError, IndexError):
|
||||
return None
|
||||
|
||||
|
||||
def _parse_json_flag(cmd_list: list[str], flag: str) -> dict[str, Any]:
|
||||
"""Extract and JSON-parse the value following `flag` in a command list."""
|
||||
val = _extract_server_cmd_value(cmd_list, flag)
|
||||
if not val:
|
||||
return {}
|
||||
try:
|
||||
return json.loads(val)
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
return {}
|
||||
|
||||
|
||||
def _extract_features(server_cmd: list[str] | str, envs: dict[str, Any]) -> list[str]:
|
||||
"""Extract enabled feature names from server_cmd and environment variables."""
|
||||
cmd_list = _cmd_to_list(server_cmd)
|
||||
features: list[str] = []
|
||||
|
||||
# Features from --additional-config JSON
|
||||
additional = _parse_json_flag(cmd_list, "--additional-config")
|
||||
if additional.get("enable_weight_nz_layout"):
|
||||
features.append("weight_nz_layout")
|
||||
wp = additional.get("weight_prefetch_config") or {}
|
||||
if isinstance(wp, dict) and wp.get("enabled"):
|
||||
features.append("weight_prefetch")
|
||||
tc = additional.get("torchair_graph_config") or {}
|
||||
if isinstance(tc, dict) and tc.get("enabled"):
|
||||
features.append("torchair_graph")
|
||||
asc = additional.get("ascend_scheduler_config") or {}
|
||||
if isinstance(asc, dict) and asc.get("enabled"):
|
||||
features.append("ascend_scheduler")
|
||||
|
||||
# Features from --compilation-config JSON
|
||||
compilation = _parse_json_flag(cmd_list, "--compilation-config")
|
||||
if compilation.get("cudagraph_mode"):
|
||||
features.append("aclgraph")
|
||||
|
||||
# Features from --speculative-config JSON
|
||||
speculative = _parse_json_flag(cmd_list, "--speculative-config")
|
||||
if speculative:
|
||||
features.append(speculative.get("method", "speculative"))
|
||||
|
||||
# Features from direct flags
|
||||
if "--enable-expert-parallel" in cmd_list:
|
||||
features.append("expert_parallel")
|
||||
|
||||
# Features from environment variables
|
||||
for env_key, feature_name in _FEATURE_ENVS.items():
|
||||
val = str(envs.get(env_key, "0"))
|
||||
if val not in ("0", "", "false", "False"):
|
||||
features.append(feature_name)
|
||||
if int(envs.get("VLLM_ASCEND_FLASHCOMM2_PARALLEL_SIZE", 0)) > 0:
|
||||
features.append("flashcomm2")
|
||||
|
||||
return features
|
||||
|
||||
|
||||
def _build_serve_cmd(config: MultiNodeConfig) -> dict[str, Any]:
|
||||
"""Build serve_cmd dict: pd format for disaggregated, dp format for multi-node."""
|
||||
if config.disagg_cfg:
|
||||
pd: dict[str, str] = {}
|
||||
for node in config.nodes:
|
||||
idx = node.index
|
||||
if config.disagg_cfg.is_prefiller(idx):
|
||||
n = config.disagg_cfg.prefiller_indices.index(idx)
|
||||
pd[f"prefill-{n}"] = node.server_cmd
|
||||
elif config.disagg_cfg.is_decoder(idx):
|
||||
n = config.disagg_cfg.decoder_indices.index(idx)
|
||||
pd[f"decode-{n}"] = node.server_cmd
|
||||
return {"pd": pd}
|
||||
return {"dp": {f"node{node.index}": node.server_cmd for node in config.nodes}}
|
||||
|
||||
|
||||
def _save_benchmark_results_json(config: MultiNodeConfig, results: list[Any]) -> None:
|
||||
"""Serialize acc & perf benchmark results to a JSON file under benchmark_results/."""
|
||||
runner = os.environ.get("VLLM_CI_RUNNER", "")
|
||||
|
||||
# Filter out None benchmark cases; results align with the non-None ones in order
|
||||
valid_items = [(case["case_name"], case) for case in config.benchmark_cases]
|
||||
|
||||
tasks = [build_task_entry(key, case_cfg, result) for (key, case_cfg), result in zip(valid_items, results)]
|
||||
|
||||
output: dict[str, Any] = {
|
||||
"model_name": config.model,
|
||||
"hardware": extract_hardware(runner),
|
||||
"dtype": _extract_dtype(config),
|
||||
"feature": _extract_features(config.nodes[0].server_cmd, config.envs),
|
||||
"vllm_version": vllm.__version__,
|
||||
"vllm_ascend_version": os.environ.get("VLLM_ASCEND_REF", ""),
|
||||
"tasks": tasks,
|
||||
"serve_cmd": _build_serve_cmd(config),
|
||||
"environment": filter_environment(config.envs),
|
||||
}
|
||||
|
||||
job_name = os.environ.get("BENCHMARK_JOB_NAME", "")
|
||||
write_results_json(output, job_name=job_name)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_multi_node() -> None:
|
||||
config = MultiNodeConfigLoader.from_yaml()
|
||||
if config.special_dependencies:
|
||||
for k, v in config.special_dependencies.items():
|
||||
command = [
|
||||
sys.executable,
|
||||
"-m",
|
||||
"pip",
|
||||
"install",
|
||||
f"{k}=={v}",
|
||||
]
|
||||
subprocess.call(command)
|
||||
|
||||
with (
|
||||
ProxyLauncher(
|
||||
nodes=config.nodes,
|
||||
disagg_cfg=config.disagg_cfg,
|
||||
envs=config.envs,
|
||||
proxy_port=config.proxy_port,
|
||||
cur_index=config.cur_index,
|
||||
) as proxy,
|
||||
RemoteOpenAIServer(
|
||||
model=config.model,
|
||||
vllm_serve_args=config.server_cmd,
|
||||
server_port=config.server_port,
|
||||
server_host=config.master_ip,
|
||||
env_dict=config.envs,
|
||||
auto_port=False,
|
||||
proxy_port=proxy.proxy_port,
|
||||
disaggregated_prefill=config.disagg_cfg,
|
||||
nodes_info=config.nodes,
|
||||
max_wait_seconds=2800,
|
||||
) as server,
|
||||
):
|
||||
host, port = config.benchmark_endpoint
|
||||
|
||||
if config.is_master:
|
||||
results = run_aisbench_cases(
|
||||
model=config.model,
|
||||
port=port,
|
||||
aisbench_cases=config.benchmark_cases,
|
||||
host_ip=host,
|
||||
)
|
||||
_save_benchmark_results_json(config, results)
|
||||
else:
|
||||
# We should keep listening on the master node's server url determining when to exit.
|
||||
server.hang_until_terminated(f"http://{host}:{config.server_port}/health")
|
||||
28
tests/e2e/nightly/multi_node/internal_dp/scripts/utils.py
Normal file
28
tests/e2e/nightly/multi_node/internal_dp/scripts/utils.py
Normal file
@@ -0,0 +1,28 @@
|
||||
import os
|
||||
|
||||
from tests.e2e.nightly.multi_node.scripts.utils import (
|
||||
get_all_ipv4,
|
||||
get_available_port,
|
||||
get_cluster_ips,
|
||||
get_net_interface,
|
||||
setup_logger,
|
||||
temp_env,
|
||||
)
|
||||
|
||||
DISAGGEGATED_PREFILL_PORT = 5333
|
||||
DEFAULT_CONFIG_BASE_PATH = "tests/e2e/nightly/multi_node/internal_dp/config/"
|
||||
CONFIG_BASE_PATH = os.getenv("CONFIG_BASE_PATH") or DEFAULT_CONFIG_BASE_PATH
|
||||
DEFAULT_SERVER_PORT = 8080
|
||||
|
||||
__all__ = [
|
||||
"CONFIG_BASE_PATH",
|
||||
"DEFAULT_CONFIG_BASE_PATH",
|
||||
"DEFAULT_SERVER_PORT",
|
||||
"DISAGGEGATED_PREFILL_PORT",
|
||||
"get_all_ipv4",
|
||||
"get_available_port",
|
||||
"get_cluster_ips",
|
||||
"get_net_interface",
|
||||
"setup_logger",
|
||||
"temp_env",
|
||||
]
|
||||
Reference in New Issue
Block a user