0
tests/ut/_tools/__init__.py
Normal file
0
tests/ut/_tools/__init__.py
Normal file
300
tests/ut/_tools/test_ai_qos_tool.py
Normal file
300
tests/ut/_tools/test_ai_qos_tool.py
Normal file
@@ -0,0 +1,300 @@
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# This file is a part of the vllm-ascend project.
|
||||
#
|
||||
|
||||
import importlib.util
|
||||
import json
|
||||
import sys
|
||||
from io import StringIO
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
TOOL_PATH = REPO_ROOT / "tools" / "ai_qos.py"
|
||||
MODULE_NAME = "vllm_ascend_tools_ai_qos"
|
||||
|
||||
MASTER_IDS = (11, 12, 13, 7)
|
||||
|
||||
|
||||
def _load_ai_qos_tool(mock_ai: MagicMock | None = None):
|
||||
if mock_ai is None:
|
||||
mock_ai = MagicMock()
|
||||
|
||||
def get_qos_fn(device_id, master_id):
|
||||
return (0, master_id, 42, 0, 0, 0)
|
||||
|
||||
mock_ai.get_qos.side_effect = get_qos_fn
|
||||
mock_ai.get_bw.return_value = (0, 1, 2, 0)
|
||||
mock_ai.get_fuse_mode.return_value = (0, 1, 1, 0)
|
||||
mock_ai.set_bw.return_value = 0
|
||||
mock_ai.set_qos.return_value = 0
|
||||
mock_ai.set_fuse_gbl_config.return_value = 0
|
||||
|
||||
ascend = MagicMock()
|
||||
ascend.ai_qos = mock_ai
|
||||
|
||||
sys.modules.pop(MODULE_NAME, None)
|
||||
with patch.dict(
|
||||
sys.modules,
|
||||
{
|
||||
"vllm_ascend": ascend,
|
||||
},
|
||||
):
|
||||
spec = importlib.util.spec_from_file_location(MODULE_NAME, TOOL_PATH)
|
||||
assert spec and spec.loader
|
||||
mod = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(mod)
|
||||
return mod, mock_ai
|
||||
|
||||
|
||||
def test_device_list_uses_all_visible_devices_when_env_unset(monkeypatch):
|
||||
monkeypatch.delenv("ASCEND_RT_VISIBLE_DEVICES", raising=False)
|
||||
mod, _ = _load_ai_qos_tool()
|
||||
mock_torch = MagicMock()
|
||||
mock_torch.npu.device_count.return_value = 4
|
||||
with patch.dict(sys.modules, {"torch": mock_torch}):
|
||||
assert mod._device_list() == [0, 1, 2, 3]
|
||||
|
||||
|
||||
def test_device_list_exits_when_env_unset_and_torch_query_fails(monkeypatch):
|
||||
monkeypatch.delenv("ASCEND_RT_VISIBLE_DEVICES", raising=False)
|
||||
mod, _ = _load_ai_qos_tool()
|
||||
mock_torch = MagicMock()
|
||||
mock_torch.npu.device_count.side_effect = RuntimeError("query failed")
|
||||
with patch.dict(sys.modules, {"torch": mock_torch}), pytest.raises(SystemExit) as e:
|
||||
mod._device_list()
|
||||
assert e.value.code == 1
|
||||
|
||||
|
||||
def test_device_list_parses_visible_devices(monkeypatch):
|
||||
mod, _ = _load_ai_qos_tool()
|
||||
monkeypatch.setenv("ASCEND_RT_VISIBLE_DEVICES", "0,2")
|
||||
assert mod._device_list() == [0, 2]
|
||||
|
||||
|
||||
def test_device_list_parses_single_id(monkeypatch):
|
||||
mod, _ = _load_ai_qos_tool()
|
||||
monkeypatch.setenv("ASCEND_RT_VISIBLE_DEVICES", "3")
|
||||
assert mod._device_list() == [3]
|
||||
|
||||
|
||||
def test_print_config_block(capsys):
|
||||
mod, _ = _load_ai_qos_tool()
|
||||
mod._print_config_block(["line a", "line b"])
|
||||
out = capsys.readouterr().out
|
||||
assert "system-view" in out and "line a" in out and "line b" in out
|
||||
assert out.strip().endswith("commit")
|
||||
|
||||
|
||||
def test_load_first_apply_baseline_no_file(tmp_path):
|
||||
mod, _ = _load_ai_qos_tool()
|
||||
p = tmp_path / "missing.json"
|
||||
assert mod._load_first_apply_baseline(p) is None
|
||||
|
||||
|
||||
def test_load_first_apply_baseline_malformed_json(tmp_path):
|
||||
mod, _ = _load_ai_qos_tool()
|
||||
p = tmp_path / "x.json"
|
||||
p.write_text("{", encoding="utf-8")
|
||||
assert mod._load_first_apply_baseline(p) is None
|
||||
|
||||
|
||||
def test_load_first_apply_baseline_invalid_original_qos(tmp_path):
|
||||
mod, _ = _load_ai_qos_tool()
|
||||
p = tmp_path / "x.json"
|
||||
p.write_text(json.dumps({"original_qos": "bad"}), encoding="utf-8")
|
||||
assert mod._load_first_apply_baseline(p) is None
|
||||
|
||||
|
||||
def test_load_first_apply_baseline_success(tmp_path):
|
||||
mod, _ = _load_ai_qos_tool()
|
||||
p = tmp_path / "x.json"
|
||||
body = {
|
||||
"original_qos": {"0": {"7": [7, 0, 0, 0, 0]}},
|
||||
"original_sdma_mata": {"0": [0, 1, 2, 0]},
|
||||
"original_fuse": {"0": [1, 1, 0]},
|
||||
}
|
||||
p.write_text(json.dumps(body), encoding="utf-8")
|
||||
b = mod._load_first_apply_baseline(p)
|
||||
assert b is not None
|
||||
oq, osm, ofu = b
|
||||
assert oq == body["original_qos"]
|
||||
assert osm == body["original_sdma_mata"]
|
||||
assert ofu == body["original_fuse"]
|
||||
|
||||
|
||||
def test_run_unset_exits_without_state_file(capsys):
|
||||
mod, _ = _load_ai_qos_tool()
|
||||
with pytest.raises(SystemExit) as e:
|
||||
mod.run_unset(Path("/nonexistent/ai_qos_state.json"))
|
||||
assert e.value.code == 1
|
||||
err = capsys.readouterr().err
|
||||
assert "No state file" in err
|
||||
|
||||
|
||||
def test_run_unset_parse_failed_bad_json_deletes_file(tmp_path, capsys):
|
||||
mod, _ = _load_ai_qos_tool()
|
||||
state = tmp_path / "ai_qos_state.json"
|
||||
state.write_text("{", encoding="utf-8")
|
||||
with pytest.raises(SystemExit) as e:
|
||||
mod.run_unset(state)
|
||||
assert e.value.code == 1
|
||||
assert not state.is_file()
|
||||
err = capsys.readouterr().err
|
||||
assert "Failed to parse the state file." in err
|
||||
|
||||
|
||||
def test_run_unset_parse_failed_invalid_structure_deletes_file(tmp_path, capsys):
|
||||
mod, _ = _load_ai_qos_tool()
|
||||
state = tmp_path / "ai_qos_state.json"
|
||||
state.write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"original_qos": {},
|
||||
"printed_commands": [123],
|
||||
"original_sdma_mata": {},
|
||||
"original_fuse": {},
|
||||
}
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
with pytest.raises(SystemExit) as e:
|
||||
mod.run_unset(state)
|
||||
assert e.value.code == 1
|
||||
assert not state.is_file()
|
||||
assert "Failed to parse the state file." in capsys.readouterr().err
|
||||
|
||||
|
||||
def test_run_unset_restores_and_deletes_file(tmp_path):
|
||||
mod, mock_ai = _load_ai_qos_tool()
|
||||
state = tmp_path / "ai_qos_state.json"
|
||||
data = {
|
||||
"original_qos": {
|
||||
"0": {
|
||||
str(MASTER_IDS[0]): [11, 1, 0, 0, 0],
|
||||
}
|
||||
},
|
||||
"original_sdma_mata": {"0": [0, 1, 2, 0]},
|
||||
"original_fuse": {"0": [1, 0, 0]},
|
||||
"printed_commands": ["hccs qos remap 1 0 0"],
|
||||
}
|
||||
state.write_text(json.dumps(data), encoding="utf-8")
|
||||
out_buf = StringIO()
|
||||
with patch("sys.stdout", out_buf):
|
||||
mod.run_unset(state)
|
||||
assert not state.is_file()
|
||||
assert mock_ai.set_bw.called
|
||||
assert mock_ai.set_qos.called
|
||||
assert mock_ai.set_fuse_gbl_config.called
|
||||
uo = out_buf.getvalue()
|
||||
assert "system-view" in uo
|
||||
assert "undo hccs qos remap 1 0 0" in uo
|
||||
assert "commit" in uo
|
||||
|
||||
|
||||
def test_AiqosConfig_set_qos_captures_and_writes_state(tmp_path, monkeypatch):
|
||||
monkeypatch.setenv("ASCEND_RT_VISIBLE_DEVICES", "0")
|
||||
mod, mock_ai = _load_ai_qos_tool()
|
||||
cfg = {
|
||||
"mode": "auto",
|
||||
"aiqos_priority": {
|
||||
"AIV_D2D": "high",
|
||||
"AIV_H2D": "high",
|
||||
"SDMA_D2D": "high",
|
||||
"SDMA_H2D": "low",
|
||||
"PCIEDMA_H2D": "high",
|
||||
},
|
||||
}
|
||||
out_buf = StringIO()
|
||||
with patch("sys.stdout", out_buf):
|
||||
mod.AiqosConfig(cfg).set_qos(tmp_path / "state.json")
|
||||
state = tmp_path / "state.json"
|
||||
assert state.is_file()
|
||||
j = json.loads(state.read_text(encoding="utf-8"))
|
||||
assert "original_qos" in j and "printed_commands" in j
|
||||
assert mock_ai.get_qos.call_count == 9
|
||||
assert mock_ai.set_fuse_gbl_config.called
|
||||
assert mock_ai.get_fuse_mode.called
|
||||
assert mock_ai.get_bw.called
|
||||
assert mock_ai.set_bw.called
|
||||
|
||||
|
||||
def test_AiqosConfig_second_apply_reuses_baseline_fewer_capture_qos(tmp_path, monkeypatch):
|
||||
monkeypatch.setenv("ASCEND_RT_VISIBLE_DEVICES", "0")
|
||||
mod, mock_ai = _load_ai_qos_tool()
|
||||
body = {
|
||||
"original_qos": {
|
||||
"0": {str(m): [m, 1, 0, 0, 0] for m in MASTER_IDS},
|
||||
},
|
||||
"original_sdma_mata": {"0": [0, 1, 1, 0]},
|
||||
"original_fuse": {"0": [1, 1, 0]},
|
||||
"printed_commands": [],
|
||||
}
|
||||
p = tmp_path / "s.json"
|
||||
p.write_text(json.dumps(body), encoding="utf-8")
|
||||
mock_ai.reset_mock()
|
||||
cfg = {
|
||||
"mode": "auto",
|
||||
"aiqos_priority": {
|
||||
"AIV_D2D": "high",
|
||||
"AIV_H2D": "high",
|
||||
"SDMA_D2D": "high",
|
||||
"SDMA_H2D": "low",
|
||||
"PCIEDMA_H2D": "high",
|
||||
},
|
||||
}
|
||||
with patch("sys.stdout", StringIO()):
|
||||
mod.AiqosConfig(cfg).set_qos(p)
|
||||
n_with_baseline = mock_ai.get_qos.call_count
|
||||
assert n_with_baseline == 4, "apply loop only: 4 masters, no capture get_qos"
|
||||
p.unlink()
|
||||
mock_ai.reset_mock()
|
||||
with patch("sys.stdout", StringIO()):
|
||||
mod.AiqosConfig(cfg).set_qos(p)
|
||||
n_cold = mock_ai.get_qos.call_count
|
||||
assert n_cold == 9, "1 dev cold: capture 4 + SDMA 1 + apply 4 = 9"
|
||||
|
||||
|
||||
def test_AiqosConfig_merges_baseline_when_device_list_grows(tmp_path, monkeypatch):
|
||||
"""Second apply with more NPU ids than the first-apply state must save baseline for new ids."""
|
||||
monkeypatch.setenv("ASCEND_RT_VISIBLE_DEVICES", "0,1")
|
||||
mod, _ = _load_ai_qos_tool()
|
||||
body = {
|
||||
"original_qos": {
|
||||
"0": {str(m): [m, 1, 0, 0, 0] for m in MASTER_IDS},
|
||||
},
|
||||
"original_sdma_mata": {"0": [0, 1, 1, 0]},
|
||||
"original_fuse": {"0": [1, 1, 0]},
|
||||
"printed_commands": [],
|
||||
}
|
||||
p = tmp_path / "state.json"
|
||||
p.write_text(json.dumps(body), encoding="utf-8")
|
||||
cfg = {
|
||||
"mode": "auto",
|
||||
"aiqos_priority": {
|
||||
"AIV_D2D": "high",
|
||||
"AIV_H2D": "high",
|
||||
"SDMA_D2D": "high",
|
||||
"SDMA_H2D": "low",
|
||||
"PCIEDMA_H2D": "high",
|
||||
},
|
||||
}
|
||||
with patch("sys.stdout", StringIO()):
|
||||
mod.AiqosConfig(cfg).set_qos(p)
|
||||
j = json.loads(p.read_text(encoding="utf-8"))
|
||||
assert "0" in j["original_qos"] and "1" in j["original_qos"]
|
||||
assert "0" in j["original_sdma_mata"] and "1" in j["original_sdma_mata"]
|
||||
assert "0" in j["original_fuse"] and "1" in j["original_fuse"]
|
||||
584
tests/ut/_tools/test_docs_codegen.py
Normal file
584
tests/ut/_tools/test_docs_codegen.py
Normal file
@@ -0,0 +1,584 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from io import StringIO
|
||||
from pathlib import Path
|
||||
from textwrap import dedent
|
||||
|
||||
import pytest
|
||||
|
||||
from tools.docs_codegen.cli import main
|
||||
from tools.docs_codegen.converters import RUN_DP_TEMPLATE_POSITIONALS
|
||||
from tools.docs_codegen.errors import DocsCodegenError
|
||||
from tools.docs_codegen.generator import GeneratorService
|
||||
from tools.docs_codegen.scanner import BlockScanner
|
||||
from tools.docs_codegen.utils import substitute_template_positionals
|
||||
|
||||
|
||||
def _write_text(path: Path, content: str) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
path.write_text(dedent(content).lstrip(), encoding="utf-8")
|
||||
|
||||
|
||||
def _write_single_node_yaml(tmp_path: Path) -> Path:
|
||||
case_path = tmp_path / "cases" / "single_node.yaml"
|
||||
_write_text(
|
||||
case_path,
|
||||
"""
|
||||
test_cases:
|
||||
- name: default-case
|
||||
model: default/model
|
||||
envs:
|
||||
SERVER_PORT: 8123
|
||||
server_cmd: []
|
||||
- name: selected-case
|
||||
model: "Qwen/Test Model"
|
||||
envs:
|
||||
HCCL_BUFFSIZE: 1024
|
||||
PROMPT: "hello world"
|
||||
SERVER_PORT: DEFAULT_PORT
|
||||
server_cmd:
|
||||
- "--tensor-parallel-size"
|
||||
- 2
|
||||
- "--kv-transfer-config"
|
||||
- '{"foo": "bar", "enabled": true}'
|
||||
server_cmd_extra: "--trust-remote-code --enable-expert-parallel"
|
||||
""",
|
||||
)
|
||||
return case_path.relative_to(tmp_path)
|
||||
|
||||
|
||||
def _write_multi_node_yaml(tmp_path: Path, *, invalid_command: bool = False) -> Path:
|
||||
case_path = tmp_path / "cases" / "multi_node.yaml"
|
||||
second_command = (
|
||||
'"vllm serve multi-node/model extra-positional"'
|
||||
if invalid_command
|
||||
else """
|
||||
- vllm
|
||||
- serve
|
||||
- multi-node/model
|
||||
- "--headless"
|
||||
- "--port"
|
||||
- "$SERVER_PORT"
|
||||
"""
|
||||
)
|
||||
_write_text(
|
||||
case_path,
|
||||
f"""
|
||||
deployment:
|
||||
- envs:
|
||||
LOCAL_IP: 127.0.0.1
|
||||
server_cmd: "vllm serve first-host --port 8000"
|
||||
- envs:
|
||||
MASTER_IP: 10.0.0.1
|
||||
SERVER_PORT: 9000
|
||||
server_cmd: {second_command.rstrip()}
|
||||
""",
|
||||
)
|
||||
return case_path.relative_to(tmp_path)
|
||||
|
||||
|
||||
def _write_model_code_doc(tmp_path: Path, content: str, *, name: str = "Demo.md") -> Path:
|
||||
doc_path = tmp_path / "docs" / "models" / name
|
||||
_write_text(doc_path, content)
|
||||
return doc_path.relative_to(tmp_path)
|
||||
|
||||
|
||||
def _generate_block(tmp_path: Path, doc_path: Path, block_name: str) -> str:
|
||||
service = GeneratorService(artifact_root="artifacts")
|
||||
return service.generate_block(doc_path, block_name, dry_run=True)[1].content
|
||||
|
||||
|
||||
# Explicit CPU smart-UT routing keeps this guard out of the --run-all-cpu bucket.
|
||||
def test_block_scanner_parses_metadata_and_trims_raw_block(tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
|
||||
single_yaml = _write_single_node_yaml(tmp_path)
|
||||
doc_path = _write_model_code_doc(
|
||||
tmp_path,
|
||||
f"""
|
||||
# Demo
|
||||
|
||||
```{{model-code}}
|
||||
:block_name: single
|
||||
:converter_tag: single_node
|
||||
:test_case_path: {single_yaml}
|
||||
:case_index: 1
|
||||
|
||||
set -eux
|
||||
{{{{ generated }}}}
|
||||
|
||||
```
|
||||
""",
|
||||
)
|
||||
monkeypatch.chdir(tmp_path)
|
||||
|
||||
blocks = BlockScanner().scan_document_blocks(doc_path)
|
||||
|
||||
assert len(blocks) == 1
|
||||
block = blocks[0]
|
||||
assert block.doc_path == doc_path
|
||||
assert block.block_name == "single"
|
||||
assert block.converter_tag == "single_node"
|
||||
assert block.test_case_path == single_yaml.as_posix()
|
||||
assert block.extra_options == (("case_index", "1"),)
|
||||
assert block.directive_line == 3
|
||||
assert block.raw_block_lines == ("set -eux", "{{ generated }}")
|
||||
|
||||
|
||||
def test_block_scanner_rejects_duplicate_block_names(tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
|
||||
single_yaml = _write_single_node_yaml(tmp_path)
|
||||
doc_path = _write_model_code_doc(
|
||||
tmp_path,
|
||||
f"""
|
||||
```{{model-code}}
|
||||
:block_name: serve
|
||||
:converter_tag: single_node
|
||||
:test_case_path: {single_yaml}
|
||||
```
|
||||
|
||||
```{{model-code}}
|
||||
:block_name: serve
|
||||
:converter_tag: single_node
|
||||
:test_case_path: {single_yaml}
|
||||
```
|
||||
""",
|
||||
name="Duplicate.md",
|
||||
)
|
||||
monkeypatch.chdir(tmp_path)
|
||||
|
||||
with pytest.raises(DocsCodegenError) as exc_info:
|
||||
BlockScanner().scan_document_blocks(doc_path)
|
||||
|
||||
error_message = str(exc_info.value)
|
||||
assert "docs/models/Duplicate.md:7: model-code generation error" in error_message
|
||||
assert "block_name: serve" in error_message
|
||||
assert "duplicated block_name 'serve'" in error_message
|
||||
assert "previous declaration is on line 1" in error_message
|
||||
|
||||
|
||||
def test_block_scanner_rejects_unsupported_metadata(tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
|
||||
single_yaml = _write_single_node_yaml(tmp_path)
|
||||
doc_path = _write_model_code_doc(
|
||||
tmp_path,
|
||||
f"""
|
||||
```{{model-code}}
|
||||
:block_name: serve
|
||||
:converter_tag: single_node
|
||||
:test_case_path: {single_yaml}
|
||||
:unknown_option: value
|
||||
```
|
||||
""",
|
||||
)
|
||||
monkeypatch.chdir(tmp_path)
|
||||
|
||||
with pytest.raises(DocsCodegenError) as exc_info:
|
||||
BlockScanner().scan_document_blocks(doc_path)
|
||||
|
||||
assert "unsupported metadata: unknown_option" in str(exc_info.value)
|
||||
|
||||
|
||||
def test_single_node_converter_uses_case_index_defaults_and_extra_args(
|
||||
tmp_path: Path,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
):
|
||||
single_yaml = _write_single_node_yaml(tmp_path)
|
||||
doc_path = _write_model_code_doc(
|
||||
tmp_path,
|
||||
f"""
|
||||
```{{model-code}}
|
||||
:block_name: single
|
||||
:converter_tag: single_node
|
||||
:test_case_path: {single_yaml}
|
||||
:case_index: 1
|
||||
|
||||
set -eux
|
||||
{{{{ generated }}}}
|
||||
```
|
||||
""",
|
||||
)
|
||||
monkeypatch.chdir(tmp_path)
|
||||
|
||||
script = _generate_block(tmp_path, doc_path, "single")
|
||||
|
||||
assert script.startswith("set -eux\nexport HCCL_BUFFSIZE=1024")
|
||||
assert 'export PROMPT="hello world"' in script
|
||||
assert "export SERVER_PORT=8000" in script
|
||||
assert "ignored/model" not in script
|
||||
assert "vllm serve 'Qwen/Test Model' \\" in script
|
||||
assert "--tensor-parallel-size 2 \\" in script
|
||||
assert '"foo": "bar"' in script
|
||||
assert '"enabled": true' in script
|
||||
assert "--trust-remote-code \\" in script
|
||||
assert "--enable-expert-parallel" in script
|
||||
|
||||
|
||||
def test_single_node_converter_defaults_to_first_case_and_preserves_port(
|
||||
tmp_path: Path,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
):
|
||||
single_yaml = _write_single_node_yaml(tmp_path)
|
||||
doc_path = _write_model_code_doc(
|
||||
tmp_path,
|
||||
f"""
|
||||
```{{model-code}}
|
||||
:block_name: default_case
|
||||
:converter_tag: single_node
|
||||
:test_case_path: {single_yaml}
|
||||
```
|
||||
""",
|
||||
)
|
||||
monkeypatch.chdir(tmp_path)
|
||||
|
||||
script = _generate_block(tmp_path, doc_path, "default_case")
|
||||
|
||||
assert script == "export SERVER_PORT=8123\n\nvllm serve default/model\n"
|
||||
|
||||
|
||||
def test_single_node_converter_reports_invalid_case_index(tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
|
||||
single_yaml = _write_single_node_yaml(tmp_path)
|
||||
doc_path = _write_model_code_doc(
|
||||
tmp_path,
|
||||
f"""
|
||||
```{{model-code}}
|
||||
:block_name: missing_case
|
||||
:converter_tag: single_node
|
||||
:test_case_path: {single_yaml}
|
||||
:case_index: 3
|
||||
```
|
||||
""",
|
||||
)
|
||||
monkeypatch.chdir(tmp_path)
|
||||
|
||||
with pytest.raises(DocsCodegenError) as exc_info:
|
||||
_generate_block(tmp_path, doc_path, "missing_case")
|
||||
|
||||
assert "case_index 3 is out of range for 'test_cases' with 2 items" in str(exc_info.value)
|
||||
|
||||
|
||||
def test_multi_node_converter_uses_host_index_and_token_list(tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
|
||||
multi_yaml = _write_multi_node_yaml(tmp_path)
|
||||
doc_path = _write_model_code_doc(
|
||||
tmp_path,
|
||||
f"""
|
||||
```{{model-code}}
|
||||
:block_name: worker
|
||||
:converter_tag: multi_node
|
||||
:test_case_path: {multi_yaml}
|
||||
:host_index: 1
|
||||
```
|
||||
""",
|
||||
)
|
||||
monkeypatch.chdir(tmp_path)
|
||||
|
||||
script = _generate_block(tmp_path, doc_path, "worker")
|
||||
|
||||
assert script.startswith("export MASTER_IP=10.0.0.1\nexport SERVER_PORT=9000")
|
||||
assert "vllm serve multi-node/model \\" in script
|
||||
assert "--headless \\" in script
|
||||
assert "--port $SERVER_PORT" in script
|
||||
assert "first-host" not in script
|
||||
|
||||
|
||||
def test_multi_node_converter_requires_host_index(tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
|
||||
multi_yaml = _write_multi_node_yaml(tmp_path)
|
||||
doc_path = _write_model_code_doc(
|
||||
tmp_path,
|
||||
f"""
|
||||
```{{model-code}}
|
||||
:block_name: worker
|
||||
:converter_tag: multi_node
|
||||
:test_case_path: {multi_yaml}
|
||||
```
|
||||
""",
|
||||
)
|
||||
monkeypatch.chdir(tmp_path)
|
||||
|
||||
with pytest.raises(DocsCodegenError) as exc_info:
|
||||
_generate_block(tmp_path, doc_path, "worker")
|
||||
|
||||
assert "converter_tag 'multi_node' requires host_index" in str(exc_info.value)
|
||||
|
||||
|
||||
def test_multi_node_converter_rejects_extra_positional_args(tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
|
||||
multi_yaml = _write_multi_node_yaml(tmp_path, invalid_command=True)
|
||||
doc_path = _write_model_code_doc(
|
||||
tmp_path,
|
||||
f"""
|
||||
```{{model-code}}
|
||||
:block_name: worker
|
||||
:converter_tag: multi_node
|
||||
:test_case_path: {multi_yaml}
|
||||
:host_index: 1
|
||||
```
|
||||
""",
|
||||
)
|
||||
monkeypatch.chdir(tmp_path)
|
||||
|
||||
with pytest.raises(DocsCodegenError) as exc_info:
|
||||
_generate_block(tmp_path, doc_path, "worker")
|
||||
|
||||
assert "unsupported positional argument 'extra-positional'" in str(exc_info.value)
|
||||
|
||||
|
||||
def test_generator_service_writes_selected_artifact(tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
|
||||
single_yaml = _write_single_node_yaml(tmp_path)
|
||||
doc_path = _write_model_code_doc(
|
||||
tmp_path,
|
||||
f"""
|
||||
```{{model-code}}
|
||||
:block_name: default_case
|
||||
:converter_tag: single_node
|
||||
:test_case_path: {single_yaml}
|
||||
```
|
||||
""",
|
||||
)
|
||||
monkeypatch.chdir(tmp_path)
|
||||
|
||||
service = GeneratorService(artifact_root="artifacts")
|
||||
output_path, generated_script = service.generate_block(doc_path, "default_case", dry_run=False)
|
||||
|
||||
assert output_path == Path("artifacts/Demo/default_case.sh")
|
||||
assert output_path.read_text(encoding="utf-8") == generated_script.content
|
||||
assert generated_script.content == "export SERVER_PORT=8123\n\nvllm serve default/model\n"
|
||||
|
||||
|
||||
def test_cli_generates_block_to_stdout(tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
|
||||
single_yaml = _write_single_node_yaml(tmp_path)
|
||||
doc_path = _write_model_code_doc(
|
||||
tmp_path,
|
||||
f"""
|
||||
```{{model-code}}
|
||||
:block_name: default_case
|
||||
:converter_tag: single_node
|
||||
:test_case_path: {single_yaml}
|
||||
```
|
||||
""",
|
||||
)
|
||||
monkeypatch.chdir(tmp_path)
|
||||
stdout = StringIO()
|
||||
stderr = StringIO()
|
||||
|
||||
exit_code = main(["--block", f"{doc_path}::default_case", "--dry-run", "--stdout"], stdout=stdout, stderr=stderr)
|
||||
|
||||
assert exit_code == 0
|
||||
assert stderr.getvalue() == ""
|
||||
assert stdout.getvalue() == (
|
||||
"docs/_build/doc_codegen/Demo/default_case.sh\nexport SERVER_PORT=8123\n\nvllm serve default/model\n"
|
||||
)
|
||||
|
||||
|
||||
def test_cli_rejects_invalid_block_reference():
|
||||
stdout = StringIO()
|
||||
stderr = StringIO()
|
||||
|
||||
exit_code = main(["--block", "docs/model.md"], stdout=stdout, stderr=stderr)
|
||||
|
||||
assert exit_code == 1
|
||||
assert stdout.getvalue() == ""
|
||||
assert "block reference must use '<doc_path>::<block_name>'" in stderr.getvalue()
|
||||
|
||||
|
||||
def _write_external_dp_yaml(tmp_path: Path, *, routing_type: str = "disaggregated_prefill") -> Path:
|
||||
case_path = tmp_path / "cases" / "external_dp.yaml"
|
||||
_write_text(
|
||||
case_path,
|
||||
f"""
|
||||
model: "Eco-Tech/GLM-Test"
|
||||
num_nodes: 2
|
||||
|
||||
routing:
|
||||
type: "{routing_type}"
|
||||
groups:
|
||||
prefiller: [0]
|
||||
decoder: [1]
|
||||
|
||||
config:
|
||||
- node_index: 0
|
||||
port_start: 7100
|
||||
dp_rpc_port: 12321
|
||||
dp_size: 4
|
||||
dp_size_local: 2
|
||||
dp_rank_start: 0
|
||||
tp_size: 8
|
||||
dp_address: "${{NODE_0_IP}}"
|
||||
- node_index: 1
|
||||
port_start: 7200
|
||||
dp_rpc_port: 12321
|
||||
dp_size: 8
|
||||
dp_size_local: 4
|
||||
dp_rank_start: 0
|
||||
tp_size: 4
|
||||
dp_address: "${{NODE_1_IP}}"
|
||||
|
||||
env_common: &env_common
|
||||
HCCL_BUFFSIZE: "1024"
|
||||
OMP_PROC_BIND: "false"
|
||||
|
||||
templates:
|
||||
- node_index: 0
|
||||
envs:
|
||||
<<: *env_common
|
||||
ASCEND_RT_VISIBLE_DEVICES: "${{VISIBLE_DEVICES}}"
|
||||
server_cmd_template:
|
||||
- --host
|
||||
- "0.0.0.0"
|
||||
- --port
|
||||
- ${{PORT}}
|
||||
- --data-parallel-size
|
||||
- ${{DP_SIZE}}
|
||||
- --data-parallel-rank
|
||||
- ${{DP_RANK}}
|
||||
- --tensor-parallel-size
|
||||
- ${{TP_SIZE}}
|
||||
- --profiler-config
|
||||
- '{{"profiler":"torch","with_stack":false}}'
|
||||
- --kv-transfer-config
|
||||
- '{{"kv_connector": "MooncakeConnectorV1", "kv_role": "kv_producer", "kv_port": "30000"}}'
|
||||
- node_index: 1
|
||||
envs:
|
||||
<<: *env_common
|
||||
ASCEND_RT_VISIBLE_DEVICES: "${{VISIBLE_DEVICES}}"
|
||||
server_cmd_template:
|
||||
- --host
|
||||
- "0.0.0.0"
|
||||
- --port
|
||||
- ${{PORT}}
|
||||
- --data-parallel-size
|
||||
- ${{DP_SIZE}}
|
||||
- --tensor-parallel-size
|
||||
- ${{TP_SIZE}}
|
||||
""",
|
||||
)
|
||||
return case_path.relative_to(tmp_path)
|
||||
|
||||
|
||||
def test_substitute_template_positionals():
|
||||
positionals = RUN_DP_TEMPLATE_POSITIONALS
|
||||
assert substitute_template_positionals("${DP_SIZE}", positionals=positionals) == "$3"
|
||||
assert substitute_template_positionals("--port ${PORT}", positionals=positionals) == "--port $2"
|
||||
# Unknown braced variables and unbraced refs are left untouched.
|
||||
assert substitute_template_positionals("${UNKNOWN}", positionals=positionals) == "${UNKNOWN}"
|
||||
assert substitute_template_positionals("$SERVER_PORT", positionals=positionals) == "$SERVER_PORT"
|
||||
|
||||
|
||||
def test_external_dp_template_converter_maps_positionals(tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
|
||||
external_yaml = _write_external_dp_yaml(tmp_path)
|
||||
doc_path = _write_model_code_doc(
|
||||
tmp_path,
|
||||
f"""
|
||||
```{{model-code}}
|
||||
:block_name: prefill_n0
|
||||
:converter_tag: external_dp_template
|
||||
:test_case_path: {external_yaml}
|
||||
:host_index: 0
|
||||
```
|
||||
""",
|
||||
)
|
||||
monkeypatch.chdir(tmp_path)
|
||||
|
||||
script = _generate_block(tmp_path, doc_path, "prefill_n0")
|
||||
|
||||
assert "export HCCL_BUFFSIZE=1024" in script
|
||||
assert "export ASCEND_RT_VISIBLE_DEVICES=$1" in script
|
||||
assert "SERVER_PORT" not in script
|
||||
assert "vllm serve Eco-Tech/GLM-Test \\" in script
|
||||
assert "--port $2 \\" in script
|
||||
assert "--data-parallel-size $3 \\" in script
|
||||
assert "--data-parallel-rank $4 \\" in script
|
||||
assert "--tensor-parallel-size $7 \\" in script
|
||||
# Space-free JSON values are quoted (not just whitespace-containing ones).
|
||||
assert '--profiler-config \'{"profiler":"torch","with_stack":false}\' \\' in script
|
||||
assert '"kv_role": "kv_producer"' in script
|
||||
assert "${DP_SIZE}" not in script
|
||||
|
||||
|
||||
def test_external_dp_template_converter_requires_host_index(tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
|
||||
external_yaml = _write_external_dp_yaml(tmp_path)
|
||||
doc_path = _write_model_code_doc(
|
||||
tmp_path,
|
||||
f"""
|
||||
```{{model-code}}
|
||||
:block_name: prefill_n0
|
||||
:converter_tag: external_dp_template
|
||||
:test_case_path: {external_yaml}
|
||||
```
|
||||
""",
|
||||
)
|
||||
monkeypatch.chdir(tmp_path)
|
||||
|
||||
with pytest.raises(DocsCodegenError) as exc_info:
|
||||
_generate_block(tmp_path, doc_path, "prefill_n0")
|
||||
|
||||
assert "converter_tag 'external_dp_template' requires host_index" in str(exc_info.value)
|
||||
|
||||
|
||||
def test_external_dp_launch_converter_combines_all_nodes(tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
|
||||
external_yaml = _write_external_dp_yaml(tmp_path)
|
||||
doc_path = _write_model_code_doc(
|
||||
tmp_path,
|
||||
f"""
|
||||
```{{model-code}}
|
||||
:block_name: launch
|
||||
:converter_tag: external_dp_launch
|
||||
:test_case_path: {external_yaml}
|
||||
```
|
||||
""",
|
||||
)
|
||||
monkeypatch.chdir(tmp_path)
|
||||
|
||||
script = _generate_block(tmp_path, doc_path, "launch")
|
||||
|
||||
assert script == (
|
||||
"python launch_online_dp.py --dp-size 4 --tp-size 8 --dp-size-local 2 --dp-rank-start 0 "
|
||||
"--dp-address ${NODE_0_IP} --dp-rpc-port 12321 --vllm-start-port 7100\n\n"
|
||||
"python launch_online_dp.py --dp-size 8 --tp-size 4 --dp-size-local 4 --dp-rank-start 0 "
|
||||
"--dp-address ${NODE_1_IP} --dp-rpc-port 12321 --vllm-start-port 7200\n"
|
||||
)
|
||||
|
||||
|
||||
def test_external_dp_proxy_converter_expands_groups(tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
|
||||
external_yaml = _write_external_dp_yaml(tmp_path)
|
||||
doc_path = _write_model_code_doc(
|
||||
tmp_path,
|
||||
f"""
|
||||
```{{model-code}}
|
||||
:block_name: proxy
|
||||
:converter_tag: external_dp_proxy
|
||||
:test_case_path: {external_yaml}
|
||||
```
|
||||
""",
|
||||
)
|
||||
monkeypatch.chdir(tmp_path)
|
||||
|
||||
script = _generate_block(tmp_path, doc_path, "proxy")
|
||||
|
||||
assert script.startswith("python load_balance_proxy_server_example.py \\")
|
||||
# Single-value flags stay inline; multi-value flags expand one value per line.
|
||||
assert " --host ${NODE_0_IP} \\" in script
|
||||
assert " --port 1999 \\" in script
|
||||
assert " --prefiller-hosts \\\n ${NODE_0_IP} \\\n ${NODE_0_IP} \\" in script
|
||||
assert " --prefiller-ports \\\n 7100 \\\n 7101 \\" in script
|
||||
assert (
|
||||
" --decoder-hosts \\\n ${NODE_1_IP} \\\n ${NODE_1_IP} \\\n ${NODE_1_IP} \\\n ${NODE_1_IP} \\"
|
||||
in script
|
||||
)
|
||||
assert " --decoder-ports \\\n 7200 \\\n 7201 \\\n 7202 \\\n 7203" in script
|
||||
assert script.rstrip().endswith(" 7203")
|
||||
|
||||
|
||||
def test_external_dp_proxy_converter_rejects_unsupported_routing(tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
|
||||
external_yaml = _write_external_dp_yaml(tmp_path, routing_type="generic_dp")
|
||||
doc_path = _write_model_code_doc(
|
||||
tmp_path,
|
||||
f"""
|
||||
```{{model-code}}
|
||||
:block_name: proxy
|
||||
:converter_tag: external_dp_proxy
|
||||
:test_case_path: {external_yaml}
|
||||
```
|
||||
""",
|
||||
)
|
||||
monkeypatch.chdir(tmp_path)
|
||||
|
||||
with pytest.raises(DocsCodegenError) as exc_info:
|
||||
_generate_block(tmp_path, doc_path, "proxy")
|
||||
|
||||
assert "only supports routing.type 'disaggregated_prefill'" in str(exc_info.value)
|
||||
Reference in New Issue
Block a user