init v0.23.0

Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
This commit is contained in:
2026-08-27 15:11:51 +08:00
parent b582a8e7d1
commit 7f8a1b1f7a
2849 changed files with 712887 additions and 22001 deletions

View File

@@ -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...",
)

View File

@@ -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")

View 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",
]