0
tests/ut/model_loader/__init__.py
Normal file
0
tests/ut/model_loader/__init__.py
Normal file
0
tests/ut/model_loader/netloader/__init__.py
Normal file
0
tests/ut/model_loader/netloader/__init__.py
Normal file
485
tests/ut/model_loader/netloader/test_netloader.py
Normal file
485
tests/ut/model_loader/netloader/test_netloader.py
Normal file
@@ -0,0 +1,485 @@
|
||||
#
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
|
||||
import json
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from vllm_ascend.model_loader.netloader.netloader import DRAFT_PORT_OFFSET, ModelNetLoaderElastic
|
||||
|
||||
|
||||
class DummyDeviceConfig:
|
||||
device = "cuda"
|
||||
device_type = "cuda"
|
||||
|
||||
|
||||
class DummyParallelConfig:
|
||||
tensor_parallel_size = 1
|
||||
pipeline_parallel_size = 1
|
||||
|
||||
|
||||
class DummyVllmConfig:
|
||||
device_config = DummyDeviceConfig()
|
||||
parallel_config = DummyParallelConfig()
|
||||
additional_config = None
|
||||
quant_config = None
|
||||
speculative_config: object | None = None
|
||||
|
||||
|
||||
class DummyModelConfig:
|
||||
model = "dummy-model"
|
||||
dtype = torch.float32
|
||||
runner_type: str | None = None
|
||||
|
||||
|
||||
class DummyDraftModelConfig(DummyModelConfig):
|
||||
model = "draft-model"
|
||||
runner_type = "draft"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def default_load_config():
|
||||
class DummyLoadConfig:
|
||||
model_loader_extra_config = None
|
||||
load_format = "default"
|
||||
|
||||
return DummyLoadConfig()
|
||||
|
||||
|
||||
def make_loader_with_config(extra):
|
||||
class DummyLoadConfig:
|
||||
model_loader_extra_config = extra
|
||||
load_format = "default"
|
||||
|
||||
return ModelNetLoaderElastic(DummyLoadConfig())
|
||||
|
||||
|
||||
def test_init_with_extra_config_file(tmp_path, monkeypatch):
|
||||
# Generate test JSON file
|
||||
config_content = {
|
||||
"SOURCE": [{"device_id": 0}],
|
||||
"MODEL": "foo-model",
|
||||
"LISTEN_PORT": 5001,
|
||||
"INT8_CACHE": "hbm",
|
||||
"OUTPUT_PREFIX": str(tmp_path),
|
||||
}
|
||||
config_file = tmp_path / "config.json"
|
||||
config_file.write_text(json.dumps(config_content))
|
||||
|
||||
dummy_logger = MagicMock()
|
||||
monkeypatch.setattr("vllm.logger.logger", dummy_logger)
|
||||
monkeypatch.setattr("vllm_ascend.model_loader.netloader.utils.is_valid_path_prefix", lambda x: True)
|
||||
|
||||
extra = {"CONFIG_FILE": str(config_file)}
|
||||
loader = make_loader_with_config(extra)
|
||||
assert loader.model_path == "foo-model"
|
||||
assert loader.source == [{"device_id": 0}]
|
||||
assert loader.listen_port == 5001
|
||||
assert loader.int8_cache == "hbm"
|
||||
assert loader.output_prefix == str(tmp_path)
|
||||
|
||||
|
||||
def test_init_with_extra_config(monkeypatch):
|
||||
dummy_logger = MagicMock()
|
||||
monkeypatch.setattr("vllm.logger.logger", dummy_logger)
|
||||
monkeypatch.setattr("vllm_ascend.model_loader.netloader.utils.is_valid_path_prefix", lambda x: True)
|
||||
|
||||
extra = {
|
||||
"SOURCE": [{"device_id": 0}],
|
||||
"MODEL": "foo",
|
||||
"LISTEN_PORT": "4000",
|
||||
"INT8_CACHE": "dram",
|
||||
"OUTPUT_PREFIX": "/tmp/",
|
||||
}
|
||||
loader = make_loader_with_config(extra)
|
||||
assert loader.model_path == "foo"
|
||||
assert loader.listen_port == 4000
|
||||
assert loader.int8_cache == "dram"
|
||||
assert loader.output_prefix == "/tmp/"
|
||||
assert loader.source == [{"device_id": 0}]
|
||||
|
||||
|
||||
def test_init_with_invalid_config(monkeypatch):
|
||||
dummy_logger = MagicMock()
|
||||
monkeypatch.setattr("vllm.logger.logger", dummy_logger)
|
||||
monkeypatch.setattr("vllm_ascend.model_loader.netloader.utils.is_valid_path_prefix", lambda x: False)
|
||||
# c
|
||||
extra = {
|
||||
"SOURCE": None,
|
||||
"MODEL": None,
|
||||
"LISTEN_PORT": None,
|
||||
"INT8_CACHE": "something",
|
||||
"OUTPUT_PREFIX": None,
|
||||
}
|
||||
loader = make_loader_with_config(extra)
|
||||
assert loader.model_path is None
|
||||
assert loader.listen_port is None
|
||||
assert loader.int8_cache == "no"
|
||||
assert loader.output_prefix is None
|
||||
|
||||
|
||||
def test_clear_static_forward_context_clears_current_vllm_config(monkeypatch):
|
||||
class DummyCompilationConfig:
|
||||
def __init__(self):
|
||||
self.static_forward_context = {}
|
||||
|
||||
class ConfigWithCompilation:
|
||||
def __init__(self):
|
||||
self.compilation_config = DummyCompilationConfig()
|
||||
|
||||
passed_config = ConfigWithCompilation()
|
||||
current_config = ConfigWithCompilation()
|
||||
passed_config.compilation_config.static_forward_context["passed.layer"] = object()
|
||||
current_config.compilation_config.static_forward_context["current.layer"] = object()
|
||||
monkeypatch.setattr(
|
||||
"vllm_ascend.model_loader.netloader.netloader.get_current_vllm_config",
|
||||
lambda: current_config,
|
||||
)
|
||||
|
||||
ModelNetLoaderElastic._clear_static_forward_context(passed_config)
|
||||
|
||||
assert passed_config.compilation_config.static_forward_context == {}
|
||||
assert current_config.compilation_config.static_forward_context == {}
|
||||
|
||||
|
||||
@patch("vllm_ascend.model_loader.netloader.netloader.logger")
|
||||
def test_load_model_elastic_success(mock_logger, monkeypatch, tmp_path):
|
||||
monkeypatch.setattr("torch.distributed.get_rank", lambda: 0)
|
||||
|
||||
class FakeContext:
|
||||
def __enter__(self):
|
||||
pass
|
||||
|
||||
def __exit__(self, a, b, c):
|
||||
pass
|
||||
|
||||
monkeypatch.setattr("torch.device", lambda d: FakeContext())
|
||||
# patch deep copy
|
||||
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.deepcopy", lambda x: x)
|
||||
# patch set_default_torch_dtype
|
||||
monkeypatch.setattr(
|
||||
"vllm_ascend.model_loader.netloader.netloader.set_default_torch_dtype", lambda dtype: FakeContext()
|
||||
)
|
||||
# patch initialize_model
|
||||
dummy_model = MagicMock(spec=nn.Module)
|
||||
dummy_model.eval.return_value = dummy_model
|
||||
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.initialize_model", lambda **kwargs: dummy_model)
|
||||
# patch elastic_load
|
||||
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.elastic_load", lambda **kwargs: dummy_model)
|
||||
# patch process_weights_after_loading
|
||||
monkeypatch.setattr(
|
||||
"vllm_ascend.model_loader.netloader.netloader.process_weights_after_loading", lambda *a, **k: None
|
||||
)
|
||||
# patch get_ip
|
||||
monkeypatch.setattr("vllm.utils.network_utils.get_ip", lambda: "127.0.0.1")
|
||||
# patch find_free_port
|
||||
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.find_free_port", lambda: 8888)
|
||||
|
||||
# patch ElasticServer
|
||||
class DummyElasticServer:
|
||||
def __init__(*a, **k):
|
||||
pass
|
||||
|
||||
def start(self):
|
||||
pass
|
||||
|
||||
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.ElasticServer", DummyElasticServer)
|
||||
# write output_prefix to the temporary directory
|
||||
extra = {
|
||||
"SOURCE": [{"device_id": 0}],
|
||||
"MODEL": "foo",
|
||||
"LISTEN_PORT": 5555,
|
||||
"OUTPUT_PREFIX": str(tmp_path) + "/output_",
|
||||
"INT8_CACHE": "no",
|
||||
}
|
||||
loader = make_loader_with_config(extra)
|
||||
vllm_config = DummyVllmConfig()
|
||||
model_config = DummyModelConfig()
|
||||
result = loader.load_model(vllm_config, model_config)
|
||||
assert isinstance(result, nn.Module)
|
||||
# Check file
|
||||
written_file = tmp_path / "output_0.txt"
|
||||
assert written_file.exists()
|
||||
|
||||
|
||||
def _patch_loader_common(monkeypatch):
|
||||
monkeypatch.setattr("torch.distributed.get_rank", lambda: 0)
|
||||
|
||||
class FakeContext:
|
||||
def __enter__(self):
|
||||
pass
|
||||
|
||||
def __exit__(self, a, b, c):
|
||||
pass
|
||||
|
||||
monkeypatch.setattr("torch.device", lambda d: FakeContext())
|
||||
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.deepcopy", lambda x: x)
|
||||
monkeypatch.setattr(
|
||||
"vllm_ascend.model_loader.netloader.netloader.set_default_torch_dtype", lambda dtype: FakeContext()
|
||||
)
|
||||
dummy_model = MagicMock(spec=nn.Module)
|
||||
dummy_model.eval.return_value = dummy_model
|
||||
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.initialize_model", lambda **kwargs: dummy_model)
|
||||
monkeypatch.setattr(
|
||||
"vllm_ascend.model_loader.netloader.netloader.process_weights_after_loading", lambda *a, **k: None
|
||||
)
|
||||
monkeypatch.setattr("vllm.utils.network_utils.get_ip", lambda: "127.0.0.1")
|
||||
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.find_free_port", lambda: 8888)
|
||||
return dummy_model
|
||||
|
||||
|
||||
@patch("vllm_ascend.model_loader.netloader.netloader.logger")
|
||||
def test_target_model_waits_for_all_netloader_ranks_before_draft(mock_logger, monkeypatch):
|
||||
dummy_model = _patch_loader_common(monkeypatch)
|
||||
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.elastic_load", lambda **kwargs: dummy_model)
|
||||
|
||||
class DummyElasticServer:
|
||||
def __init__(*a, **k):
|
||||
pass
|
||||
|
||||
def start(self):
|
||||
pass
|
||||
|
||||
barrier_calls = []
|
||||
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.ElasticServer", DummyElasticServer)
|
||||
monkeypatch.setattr("torch.distributed.is_available", lambda: True)
|
||||
monkeypatch.setattr("torch.distributed.is_initialized", lambda: True)
|
||||
monkeypatch.setattr("torch.distributed.barrier", lambda: barrier_calls.append("barrier"))
|
||||
|
||||
extra = {
|
||||
"SOURCE": [{"device_id": 0, "sources": ["127.0.0.1:5000"]}],
|
||||
"MODEL": "dummy-model",
|
||||
"LISTEN_PORT": 5555,
|
||||
"INT8_CACHE": "dram",
|
||||
}
|
||||
loader = make_loader_with_config(extra)
|
||||
vllm_config = DummyVllmConfig()
|
||||
vllm_config.speculative_config = object()
|
||||
loader.load_model(vllm_config, DummyModelConfig())
|
||||
|
||||
assert barrier_calls == ["barrier"]
|
||||
|
||||
|
||||
@patch("vllm_ascend.model_loader.netloader.netloader.logger")
|
||||
def test_failed_target_model_participates_in_barrier_before_error(mock_logger, monkeypatch):
|
||||
_patch_loader_common(monkeypatch)
|
||||
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.elastic_load", lambda **kwargs: None)
|
||||
monkeypatch.setattr(
|
||||
ModelNetLoaderElastic,
|
||||
"revert_to_default",
|
||||
lambda self, *args, **kwargs: (None, False),
|
||||
)
|
||||
|
||||
barrier_calls = []
|
||||
monkeypatch.setattr("torch.distributed.is_available", lambda: True)
|
||||
monkeypatch.setattr("torch.distributed.is_initialized", lambda: True)
|
||||
monkeypatch.setattr("torch.distributed.barrier", lambda: barrier_calls.append("barrier"))
|
||||
|
||||
extra = {
|
||||
"SOURCE": [{"device_id": 0, "sources": ["127.0.0.1:5000"]}],
|
||||
"MODEL": "dummy-model",
|
||||
"LISTEN_PORT": 5555,
|
||||
"INT8_CACHE": "dram",
|
||||
}
|
||||
loader = make_loader_with_config(extra)
|
||||
vllm_config = DummyVllmConfig()
|
||||
vllm_config.speculative_config = object()
|
||||
|
||||
with pytest.raises(RuntimeError, match="NetLoader elastic loads model fails"):
|
||||
loader.load_model(vllm_config, DummyModelConfig())
|
||||
|
||||
assert barrier_calls == ["barrier"]
|
||||
|
||||
|
||||
@patch("vllm_ascend.model_loader.netloader.netloader.logger")
|
||||
def test_draft_model_does_not_wait_for_target_netloader_barrier(mock_logger, monkeypatch):
|
||||
dummy_model = _patch_loader_common(monkeypatch)
|
||||
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.elastic_load", lambda **kwargs: dummy_model)
|
||||
|
||||
class DummyElasticServer:
|
||||
def __init__(*a, **k):
|
||||
pass
|
||||
|
||||
def start(self):
|
||||
pass
|
||||
|
||||
barrier_calls = []
|
||||
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.ElasticServer", DummyElasticServer)
|
||||
monkeypatch.setattr("torch.distributed.is_available", lambda: True)
|
||||
monkeypatch.setattr("torch.distributed.is_initialized", lambda: True)
|
||||
monkeypatch.setattr("torch.distributed.barrier", lambda: barrier_calls.append("barrier"))
|
||||
|
||||
extra = {
|
||||
"SOURCE": [{"device_id": 0, "sources": ["127.0.0.1:5000"]}],
|
||||
"MODEL": "draft-model",
|
||||
"LISTEN_PORT": 5555,
|
||||
"INT8_CACHE": "dram",
|
||||
}
|
||||
loader = make_loader_with_config(extra)
|
||||
vllm_config = DummyVllmConfig()
|
||||
vllm_config.speculative_config = object()
|
||||
loader.load_model(vllm_config, DummyDraftModelConfig())
|
||||
|
||||
assert barrier_calls == []
|
||||
|
||||
|
||||
@patch("vllm_ascend.model_loader.netloader.netloader.logger")
|
||||
def test_load_draft_model_elastic_success(mock_logger, monkeypatch, tmp_path):
|
||||
dummy_model = _patch_loader_common(monkeypatch)
|
||||
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.elastic_load", lambda **kwargs: dummy_model)
|
||||
|
||||
elastic_server_instances = []
|
||||
|
||||
class DummyElasticServer:
|
||||
def __init__(self, *args, **kwargs):
|
||||
elastic_server_instances.append(self)
|
||||
self.int8_cache = args[7]
|
||||
self.group_name = kwargs.get("group_name")
|
||||
|
||||
def start(self):
|
||||
pass
|
||||
|
||||
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.ElasticServer", DummyElasticServer)
|
||||
|
||||
extra = {
|
||||
"SOURCE": [{"device_id": 0, "sources": ["127.0.0.1:5000"]}],
|
||||
"MODEL": "draft-model",
|
||||
"LISTEN_PORT": 5555,
|
||||
"OUTPUT_PREFIX": str(tmp_path) + "/output_",
|
||||
"INT8_CACHE": "no",
|
||||
}
|
||||
loader = make_loader_with_config(extra)
|
||||
result = loader.load_model(DummyVllmConfig(), DummyDraftModelConfig())
|
||||
|
||||
assert isinstance(result, nn.Module)
|
||||
assert loader._draft_elastic_server is elastic_server_instances[0]
|
||||
assert loader._draft_elastic_server.int8_cache == "no"
|
||||
assert loader._draft_elastic_server.group_name == "netloader_draft"
|
||||
assert not (tmp_path / "output_0.txt").exists()
|
||||
|
||||
|
||||
@patch("vllm_ascend.model_loader.netloader.netloader.logger")
|
||||
def test_load_draft_model_uses_hbm_when_int8_cache_is_dram(mock_logger, monkeypatch):
|
||||
dummy_model = _patch_loader_common(monkeypatch)
|
||||
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.elastic_load", lambda **kwargs: dummy_model)
|
||||
|
||||
captured = {}
|
||||
|
||||
class DummyElasticServer:
|
||||
def __init__(self, *args, **kwargs):
|
||||
captured["int8_cache"] = args[7]
|
||||
captured["group_name"] = kwargs.get("group_name")
|
||||
|
||||
def start(self):
|
||||
pass
|
||||
|
||||
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.ElasticServer", DummyElasticServer)
|
||||
|
||||
extra = {
|
||||
"SOURCE": [{"device_id": 0, "sources": ["127.0.0.1:5000"]}],
|
||||
"MODEL": "draft-model",
|
||||
"LISTEN_PORT": 5555,
|
||||
"INT8_CACHE": "dram",
|
||||
}
|
||||
loader = make_loader_with_config(extra)
|
||||
loader.load_model(DummyVllmConfig(), DummyDraftModelConfig())
|
||||
|
||||
assert captured["int8_cache"] == "hbm"
|
||||
assert captured["group_name"] == "netloader_draft"
|
||||
|
||||
|
||||
@patch("vllm_ascend.model_loader.netloader.netloader.logger")
|
||||
def test_load_draft_model_port_offset_and_group_name(mock_logger, monkeypatch, tmp_path):
|
||||
dummy_model = _patch_loader_common(monkeypatch)
|
||||
captured = {}
|
||||
|
||||
def capture_elastic_load(**kwargs):
|
||||
captured.update(kwargs)
|
||||
return dummy_model
|
||||
|
||||
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.elastic_load", capture_elastic_load)
|
||||
|
||||
class DummyElasticServer:
|
||||
def __init__(*a, **k):
|
||||
pass
|
||||
|
||||
def start(self):
|
||||
pass
|
||||
|
||||
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.ElasticServer", DummyElasticServer)
|
||||
|
||||
extra = {
|
||||
"SOURCE": [{"device_id": 0, "sources": ["127.0.0.1:5000", "10.0.0.1:6000"]}],
|
||||
"MODEL": "draft-model",
|
||||
"LISTEN_PORT": 5555,
|
||||
"INT8_CACHE": "no",
|
||||
}
|
||||
loader = make_loader_with_config(extra)
|
||||
loader.load_model(DummyVllmConfig(), DummyDraftModelConfig())
|
||||
|
||||
assert captured["group_name"] == "netloader_draft"
|
||||
assert captured["model_path"] == "draft-model"
|
||||
assert captured["sources"] == [
|
||||
{
|
||||
"device_id": 0,
|
||||
"sources": [
|
||||
f"127.0.0.1:{5000 + DRAFT_PORT_OFFSET}",
|
||||
f"10.0.0.1:{6000 + DRAFT_PORT_OFFSET}",
|
||||
],
|
||||
}
|
||||
]
|
||||
assert loader.listen_port == 5555 + DRAFT_PORT_OFFSET
|
||||
|
||||
|
||||
@patch("vllm_ascend.model_loader.netloader.netloader.logger")
|
||||
def test_load_draft_model_skips_invalid_source_addresses(mock_logger, monkeypatch):
|
||||
dummy_model = _patch_loader_common(monkeypatch)
|
||||
captured = {}
|
||||
|
||||
def capture_elastic_load(**kwargs):
|
||||
captured.update(kwargs)
|
||||
return dummy_model
|
||||
|
||||
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.elastic_load", capture_elastic_load)
|
||||
|
||||
class DummyElasticServer:
|
||||
def __init__(*a, **k):
|
||||
pass
|
||||
|
||||
def start(self):
|
||||
pass
|
||||
|
||||
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.ElasticServer", DummyElasticServer)
|
||||
|
||||
extra = {
|
||||
"SOURCE": [{"device_id": 0, "sources": ["127.0.0.1:5000", "invalid", "10.0.0.1:not_port"]}],
|
||||
"MODEL": "draft-model",
|
||||
"LISTEN_PORT": 5555,
|
||||
"INT8_CACHE": "no",
|
||||
}
|
||||
loader = make_loader_with_config(extra)
|
||||
loader.load_model(DummyVllmConfig(), DummyDraftModelConfig())
|
||||
|
||||
assert captured["sources"] == [
|
||||
{"device_id": 0, "sources": [f"127.0.0.1:{5000 + DRAFT_PORT_OFFSET}"]},
|
||||
]
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main()
|
||||
437
tests/ut/model_loader/netloader/test_netloader_elastic.py
Normal file
437
tests/ut/model_loader/netloader/test_netloader_elastic.py
Normal file
@@ -0,0 +1,437 @@
|
||||
#
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
|
||||
import io
|
||||
import json
|
||||
import logging
|
||||
import socket
|
||||
from contextlib import contextmanager
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from vllm_ascend.model_loader.netloader.interaction import elastic
|
||||
from vllm_ascend.model_loader.netloader.interaction.elastic import ElasticClient, ElasticServer
|
||||
|
||||
|
||||
# Simulate server's normal response
|
||||
def mock_server_response(data):
|
||||
return json.dumps({"label": "JOIN_ACK", "content": {"name": "mocked_name"}}).encode("utf-8")
|
||||
|
||||
|
||||
# Simulate server's error response
|
||||
def mock_server_error_response(data):
|
||||
return json.dumps({"label": "JOIN_ACK", "content": None}).encode("utf-8")
|
||||
|
||||
|
||||
# Simulated server's abnormal response
|
||||
def mock_server_exception_response(data):
|
||||
raise Exception("Mocked server exception")
|
||||
|
||||
|
||||
@contextmanager
|
||||
def capture_elastic_logs(level=logging.DEBUG):
|
||||
log_capture_string = io.StringIO()
|
||||
handler = logging.StreamHandler(log_capture_string)
|
||||
handler.setLevel(level)
|
||||
original_level = elastic.logger.level
|
||||
elastic.logger.setLevel(level)
|
||||
elastic.logger.addHandler(handler)
|
||||
try:
|
||||
yield log_capture_string
|
||||
finally:
|
||||
elastic.logger.removeHandler(handler)
|
||||
elastic.logger.setLevel(original_level)
|
||||
log_capture_string.close()
|
||||
|
||||
|
||||
# Test the initialization of ElasticClient
|
||||
def test_elastic_client_init():
|
||||
sources = ["127.0.0.1:12345"]
|
||||
device_id = 0
|
||||
model_path = "mocked_model_path"
|
||||
tp = 1
|
||||
pp = 1
|
||||
|
||||
with patch("socket.socket") as mock_socket:
|
||||
mock_socket_instance = MagicMock()
|
||||
mock_socket.return_value = mock_socket_instance
|
||||
mock_socket_instance.recv.return_value = mock_server_response(None)
|
||||
|
||||
mock_socket_instance.getsockname.return_value = ("127.0.0.1", 12346)
|
||||
mock_socket_instance.__enter__.return_value = mock_socket_instance
|
||||
|
||||
with ElasticClient(sources, device_id, model_path, tp, pp) as client:
|
||||
assert client.server_addr == "127.0.0.1"
|
||||
assert client.server_port == 12345
|
||||
assert client.ack == ("mocked_name", 12346)
|
||||
mock_socket_instance.close.assert_called_once()
|
||||
|
||||
|
||||
# Test the register method of ElasticClient
|
||||
def test_elastic_client_register():
|
||||
sources = ["127.0.0.1:12345"]
|
||||
device_id = 0
|
||||
model_path = "mocked_model_path"
|
||||
tp = 1
|
||||
pp = 1
|
||||
|
||||
with patch("socket.socket") as mock_socket:
|
||||
mock_socket_instance = MagicMock()
|
||||
mock_socket.return_value = mock_socket_instance
|
||||
mock_socket_instance.connect.return_value = None
|
||||
mock_socket_instance.recv.return_value = mock_server_response(None)
|
||||
|
||||
mock_socket_instance.getsockname.return_value = ("127.0.0.1", 12346)
|
||||
mock_socket_instance.__enter__.return_value = mock_socket_instance
|
||||
|
||||
client = ElasticClient(sources, device_id, model_path, tp, pp)
|
||||
assert client.register(device_id, model_path, tp, pp) == ("mocked_name", 12346)
|
||||
|
||||
|
||||
# Test the behavior of the `register` method of ElasticClient when the server returns an error response.
|
||||
def test_elastic_client_register_error_response():
|
||||
sources = ["127.0.0.1:12345"]
|
||||
device_id = 0
|
||||
model_path = "mocked_model_path"
|
||||
tp = 1
|
||||
pp = 1
|
||||
|
||||
with patch("socket.socket") as mock_socket:
|
||||
mock_socket_instance = MagicMock()
|
||||
mock_socket.return_value = mock_socket_instance
|
||||
mock_socket_instance.connect.return_value = None
|
||||
mock_socket_instance.recv.return_value = mock_server_error_response(None)
|
||||
|
||||
with ElasticClient(sources, device_id, model_path, tp, pp) as client, pytest.raises(RuntimeError):
|
||||
client.register(device_id, model_path, tp, pp)
|
||||
mock_socket_instance.close.assert_called_once()
|
||||
|
||||
|
||||
# Test the behavior of the `register` method of ElasticClient when an exception is thrown on the server.
|
||||
def test_elastic_client_register_exception():
|
||||
sources = ["127.0.0.1:12345"]
|
||||
device_id = 0
|
||||
model_path = "mocked_model_path"
|
||||
tp = 1
|
||||
pp = 1
|
||||
|
||||
with patch("socket.socket") as mock_socket:
|
||||
mock_socket_instance = MagicMock()
|
||||
mock_socket.return_value = mock_socket_instance
|
||||
mock_socket_instance.connect.return_value = None
|
||||
mock_socket_instance.recv.side_effect = mock_server_exception_response
|
||||
mock_socket_instance.__enter__.return_value = mock_socket_instance
|
||||
mock_socket_instance.__exit__.return_value = None
|
||||
|
||||
with ElasticClient(sources, device_id, model_path, tp, pp) as client, pytest.raises(RuntimeError):
|
||||
client.register(device_id, model_path, tp, pp)
|
||||
mock_socket_instance.close.assert_called_once()
|
||||
|
||||
|
||||
class FakeInt8Param:
|
||||
def __init__(self, name="param", device="npu", dtype=torch.int8):
|
||||
self.dtype = dtype
|
||||
self.device = torch.device(device)
|
||||
|
||||
@property
|
||||
def data(self):
|
||||
return self # Simulate .data returning self so .cpu() etc. can be chained
|
||||
|
||||
def clone(self):
|
||||
return self
|
||||
|
||||
def detach(self):
|
||||
return self
|
||||
|
||||
def cpu(self):
|
||||
self.device = torch.device("cpu")
|
||||
return self
|
||||
|
||||
|
||||
class FakeModel:
|
||||
def __init__(self):
|
||||
self.params = {
|
||||
"param1": MagicMock(dtype=torch.float32), # This will be ignored
|
||||
"param2": FakeInt8Param(), # This simulates a real int8 param
|
||||
}
|
||||
|
||||
def named_parameters(self):
|
||||
return self.params.items()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_model():
|
||||
return FakeModel()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def server_config():
|
||||
return {
|
||||
"addr": "127.0.0.1",
|
||||
"port": 8080,
|
||||
"model": MagicMock(),
|
||||
"device_id": 0,
|
||||
"model_path": "/test/model",
|
||||
"tp": 1,
|
||||
"pp": 1,
|
||||
"int8_cache": "dram",
|
||||
"int8_cache_name": None,
|
||||
}
|
||||
|
||||
|
||||
# Test server initialization
|
||||
def test_server_initialization(server_config, mock_model):
|
||||
server_config["model"] = mock_model
|
||||
with patch("socket.socket") as mock_socket, capture_elastic_logs() as log_capture_string:
|
||||
server = ElasticServer(**server_config)
|
||||
|
||||
# Check the socket configuration
|
||||
mock_socket.assert_called_with(socket.AF_INET, socket.SOCK_STREAM)
|
||||
mock_socket.return_value.setsockopt.assert_called_with(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
|
||||
mock_socket.return_value.bind.assert_called_with(("127.0.0.1", 8080))
|
||||
mock_socket.return_value.listen.assert_called_with(256)
|
||||
|
||||
# Check int8 cache
|
||||
assert "param2" in server.original_int8
|
||||
assert server.original_int8["param2"].device.type == "cpu" # Verifying DRAM Cache
|
||||
|
||||
assert server.addr == server_config["addr"]
|
||||
assert server.port == server_config["port"]
|
||||
assert server.device_id == server_config["device_id"]
|
||||
assert server.model_path == server_config["model_path"]
|
||||
assert server.tp == server_config["tp"]
|
||||
assert server.pp == server_config["pp"]
|
||||
|
||||
# Get captured logs
|
||||
log_output = log_capture_string.getvalue()
|
||||
|
||||
# Check output
|
||||
assert "Server 127.0.0.1:8080 starts" in log_output
|
||||
|
||||
|
||||
# Test the int8 cache option
|
||||
@pytest.mark.parametrize("cache_option,expected_device", [("dram", "cpu"), ("no", None), ("invalid", None)])
|
||||
def test_int8_cache_handling(server_config, mock_model, cache_option, expected_device):
|
||||
server_config["int8_cache"] = cache_option
|
||||
server_config["model"] = mock_model
|
||||
|
||||
with patch("socket.socket"), capture_elastic_logs() as log_capture_string:
|
||||
server = ElasticServer(**server_config)
|
||||
|
||||
log_output = log_capture_string.getvalue()
|
||||
|
||||
if cache_option == "invalid":
|
||||
assert "int8_cache should be selected in [HBM, DRAM]" in log_output
|
||||
|
||||
if expected_device is None:
|
||||
assert len(server.original_int8) == 0
|
||||
else:
|
||||
assert server.original_int8["param2"].device.type == expected_device
|
||||
|
||||
|
||||
# Test client processing
|
||||
def test_client_handler_valid_join(server_config, mock_model):
|
||||
server_config["model"] = mock_model
|
||||
with patch("vllm_ascend.model_loader.netloader.interaction.elastic.P2PSend") as mock_p2p_send:
|
||||
# Create a simulated connection
|
||||
mock_conn = MagicMock()
|
||||
mock_addr = ("192.168.1.1", 12345)
|
||||
|
||||
# Configuring Client Data
|
||||
valid_data = {
|
||||
"label": "JOIN",
|
||||
"content": {"device_id": 0, "model_path": "/test/model", "tp": 1, "pp": 1, "port": 9090},
|
||||
}
|
||||
mock_conn.recv.return_value = json.dumps(valid_data).encode("utf-8")
|
||||
|
||||
# Start the server
|
||||
server = ElasticServer(**server_config)
|
||||
server.register_handler(mock_conn, mock_addr)
|
||||
|
||||
# Verify response
|
||||
expected_ack = {"label": "JOIN_ACK", "content": {"name": "192.168.1.1:12345"}}
|
||||
mock_conn.send.assert_called_once_with(json.dumps(expected_ack).encode("utf-8"))
|
||||
mock_p2p_send.assert_called_once_with("127.0.0.1", 9090, "192.168.1.1:12345", "netloader")
|
||||
mock_conn.close.assert_called_once()
|
||||
|
||||
|
||||
# Test mismatched JOIN requests
|
||||
def test_client_handler_mismatch(server_config):
|
||||
with patch("socket.socket"):
|
||||
server = ElasticServer(**server_config)
|
||||
mock_conn = MagicMock()
|
||||
mock_addr = ("192.168.1.1", 12345)
|
||||
|
||||
# Send mismatched data
|
||||
mismatch_data = {
|
||||
"label": "JOIN",
|
||||
"content": {
|
||||
"device_id": 1, # 不匹配的ID
|
||||
"model_path": "/wrong/model",
|
||||
"tp": 2,
|
||||
"pp": 2,
|
||||
"port": 9090,
|
||||
},
|
||||
}
|
||||
mock_conn.recv.return_value = json.dumps(mismatch_data).encode("utf-8")
|
||||
|
||||
server.register_handler(mock_conn, mock_addr)
|
||||
|
||||
assert isinstance(mismatch_data["content"], dict)
|
||||
|
||||
# Verify response
|
||||
mismatch_tuple = (
|
||||
mismatch_data["content"]["device_id"],
|
||||
mismatch_data["content"]["model_path"],
|
||||
mismatch_data["content"]["tp"],
|
||||
mismatch_data["content"]["pp"],
|
||||
)
|
||||
|
||||
server_tuple = (
|
||||
server_config["device_id"],
|
||||
server_config["model_path"],
|
||||
server_config["tp"],
|
||||
server_config["pp"],
|
||||
)
|
||||
|
||||
expected_ack = {
|
||||
"label": "JOIN_NACK",
|
||||
"content": (f"Received data {mismatch_tuple} does not consist with this server {server_tuple}"),
|
||||
}
|
||||
mock_conn.send.assert_called_once_with(json.dumps(expected_ack).encode("utf-8"))
|
||||
mock_conn.close.assert_called_once()
|
||||
|
||||
|
||||
# Test Invalid Request
|
||||
@pytest.mark.parametrize(
|
||||
"invalid_data,should_send",
|
||||
[
|
||||
({"label": "WRONG_LABEL"}, True), # Incorrect label, can be decoded as JSON, but the content is invalid.
|
||||
(
|
||||
{"content": {"missing_fields": True}},
|
||||
True,
|
||||
), # Missing field, can be decoded as JSON, but the content is invalid.
|
||||
("plain text", False), # Non-JSON data, json.loads failed
|
||||
(b"invalid_bytes", False), # Invalid byte, decode or json.loads failed
|
||||
],
|
||||
)
|
||||
def test_client_handler_invalid_requests(server_config, invalid_data, should_send):
|
||||
with patch("socket.socket"), capture_elastic_logs() as log_capture_string:
|
||||
server = ElasticServer(**server_config)
|
||||
mock_conn = MagicMock()
|
||||
mock_addr = ("192.168.1.1", 12345)
|
||||
|
||||
if isinstance(invalid_data, (str, bytes)):
|
||||
mock_conn.recv.return_value = invalid_data if isinstance(invalid_data, bytes) else invalid_data.encode()
|
||||
else:
|
||||
mock_conn.recv.return_value = json.dumps(invalid_data).encode("utf-8")
|
||||
|
||||
server.register_handler(mock_conn, mock_addr)
|
||||
|
||||
if should_send:
|
||||
expected_ack = {
|
||||
"label": "JOIN_NACK",
|
||||
"content": f"Received data does not contain required fields: {invalid_data}",
|
||||
}
|
||||
mock_conn.send.assert_called_once_with(json.dumps(expected_ack).encode("utf-8"))
|
||||
else:
|
||||
mock_conn.send.assert_not_called()
|
||||
|
||||
log_output = log_capture_string.getvalue()
|
||||
|
||||
# Any warning in the log is acceptable
|
||||
assert "Failed to load" in log_output or "does not contain" in log_output
|
||||
mock_conn.close.assert_called_once()
|
||||
|
||||
|
||||
# Test the thread startup.
|
||||
def test_server_start(server_config):
|
||||
with patch("socket.socket"), patch("threading.Thread") as mock_thread:
|
||||
handler_thread_instance = mock_thread.return_value
|
||||
|
||||
server = ElasticServer(**server_config)
|
||||
server.start()
|
||||
|
||||
# Assert that the correct target parameter was passed when instantiating the Thread instance.
|
||||
mock_thread.assert_called_once()
|
||||
args, kwargs = mock_thread.call_args
|
||||
assert kwargs["target"] == server.elastic_client_handler
|
||||
|
||||
# Verify the daemon attribute is set to True (the attribute value will be recorded after MagicMock assignment).
|
||||
assert handler_thread_instance.daemon is True
|
||||
|
||||
# Check if the start() method is called.
|
||||
handler_thread_instance.start.assert_called_once()
|
||||
|
||||
|
||||
# Test resource clearing
|
||||
def test_server_cleanup(server_config):
|
||||
with patch("socket.socket") as mock_socket:
|
||||
server = ElasticServer(**server_config)
|
||||
del server
|
||||
mock_socket.return_value.close.assert_called_once()
|
||||
|
||||
|
||||
def test_draft_group_name_in_client_register():
|
||||
sent_payloads = []
|
||||
|
||||
with (
|
||||
patch("socket.socket") as mock_socket,
|
||||
patch("vllm_ascend.model_loader.netloader.interaction.elastic.find_free_port", return_value=12346),
|
||||
):
|
||||
mock_socket_instance = MagicMock()
|
||||
mock_socket.return_value = mock_socket_instance
|
||||
mock_socket_instance.recv.return_value = mock_server_response(None)
|
||||
mock_socket_instance.send.side_effect = lambda data: sent_payloads.append(json.loads(data.decode()))
|
||||
|
||||
client = ElasticClient(["127.0.0.1:12345"], 0, "draft-model", 1, 1, "netloader_draft")
|
||||
client.s = mock_socket_instance
|
||||
client.register(0, "draft-model", 1, 1)
|
||||
|
||||
assert client.group_name == "netloader_draft"
|
||||
assert sent_payloads[0]["content"]["group_name"] == "netloader_draft"
|
||||
|
||||
|
||||
def test_draft_group_name_in_server_p2p_send(server_config, mock_model):
|
||||
server_config["model"] = mock_model
|
||||
join_data = {
|
||||
"label": "JOIN",
|
||||
"content": {
|
||||
"device_id": 0,
|
||||
"model_path": "/test/model",
|
||||
"tp": 1,
|
||||
"pp": 1,
|
||||
"port": 9090,
|
||||
"group_name": "netloader_draft",
|
||||
},
|
||||
}
|
||||
|
||||
with (
|
||||
patch("socket.socket"),
|
||||
patch("vllm_ascend.model_loader.netloader.interaction.elastic.P2PSend") as mock_p2p_send,
|
||||
):
|
||||
mock_conn = MagicMock()
|
||||
mock_conn.recv.return_value = json.dumps(join_data).encode("utf-8")
|
||||
|
||||
ElasticServer(**server_config, group_name="netloader_draft").register_handler(mock_conn, ("192.168.1.1", 12345))
|
||||
|
||||
mock_p2p_send.assert_called_once_with("127.0.0.1", 9090, "192.168.1.1:12345", "netloader_draft")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main()
|
||||
129
tests/ut/model_loader/netloader/test_netloader_load.py
Normal file
129
tests/ut/model_loader/netloader/test_netloader_load.py
Normal file
@@ -0,0 +1,129 @@
|
||||
#
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from vllm_ascend.model_loader.netloader.load import elastic_load
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_sources():
|
||||
return [
|
||||
{"device_id": 0, "sources": ["a", "b"]},
|
||||
{"device_id": 1, "sources": ["c"]},
|
||||
]
|
||||
|
||||
|
||||
@patch("vllm_ascend.model_loader.netloader.interaction.elastic.ElasticClient")
|
||||
@patch("vllm_ascend.model_loader.netloader.executor.elastic_load.P2PLoad")
|
||||
def test_sources_this_device_empty(mock_p2p, mock_client):
|
||||
sources = [{"device_id": 1, "sources": ["c"]}]
|
||||
result = elastic_load("model", 0, "model_path", sources, 1, 1)
|
||||
assert result is None
|
||||
mock_client.assert_not_called()
|
||||
mock_p2p.assert_not_called()
|
||||
|
||||
|
||||
@patch("vllm_ascend.model_loader.netloader.interaction.elastic.ElasticClient")
|
||||
@patch("vllm_ascend.model_loader.netloader.executor.elastic_load.P2PLoad")
|
||||
def test_client_s_none(mock_p2p, mock_client, mock_sources):
|
||||
# Simulate ElasticClient.s as None
|
||||
mock_instance = MagicMock()
|
||||
mock_instance.s = None
|
||||
mock_client.return_value = mock_instance
|
||||
result = elastic_load("model", 0, "model_path", mock_sources, 1, 1)
|
||||
assert result is None
|
||||
|
||||
|
||||
@patch("vllm_ascend.model_loader.netloader.interaction.elastic.ElasticClient")
|
||||
@patch("vllm_ascend.model_loader.netloader.executor.elastic_load.P2PLoad")
|
||||
def test_client_ack_none(mock_p2p, mock_client, mock_sources):
|
||||
# Simulate ElasticClient.ack as None
|
||||
mock_instance = MagicMock()
|
||||
mock_instance.s = True
|
||||
mock_instance.ack = None
|
||||
mock_client.return_value = mock_instance
|
||||
result = elastic_load("model", 0, "model_path", mock_sources, 1, 1)
|
||||
assert result is None
|
||||
|
||||
|
||||
@patch("vllm_ascend.model_loader.netloader.load.P2PLoad")
|
||||
@patch("vllm_ascend.model_loader.netloader.load.logger")
|
||||
def test_model_load_fail(mock_logger, mock_p2p):
|
||||
mock_client = MagicMock()
|
||||
mock_client.s = True
|
||||
mock_client.ack = ["foo", "bar"]
|
||||
mock_client.server_addr = "addr"
|
||||
|
||||
with patch("vllm_ascend.model_loader.netloader.load.ElasticClient", return_value=mock_client):
|
||||
# P2PLoad.load returns None
|
||||
mock_p2p_instance = MagicMock()
|
||||
mock_p2p_instance.load.return_value = None
|
||||
mock_p2p.return_value = mock_p2p_instance
|
||||
|
||||
sources = [{"device_id": 0, "sources": ["whatever"]}]
|
||||
result = elastic_load("model", 0, "model_path", sources, 1, 1)
|
||||
assert result is None
|
||||
mock_logger.error.assert_called_once()
|
||||
|
||||
|
||||
@patch("vllm_ascend.model_loader.netloader.load.P2PLoad")
|
||||
@patch("vllm_ascend.model_loader.netloader.load.logger")
|
||||
def test_model_load_success(mock_logger, mock_p2p):
|
||||
mock_client = MagicMock()
|
||||
mock_client.s = True
|
||||
mock_client.ack = ["foo", "bar"]
|
||||
mock_client.server_addr = "addr"
|
||||
|
||||
with patch("vllm_ascend.model_loader.netloader.load.ElasticClient", return_value=mock_client):
|
||||
expected_model = object()
|
||||
mock_p2p_instance = MagicMock()
|
||||
mock_p2p_instance.load.return_value = expected_model
|
||||
mock_p2p.return_value = mock_p2p_instance
|
||||
|
||||
sources = [{"device_id": 0, "sources": ["whatever"]}]
|
||||
result = elastic_load("model", 0, "model_path", sources, 1, 1)
|
||||
assert result is expected_model
|
||||
mock_logger.info.assert_called_once()
|
||||
|
||||
|
||||
@patch("vllm_ascend.model_loader.netloader.load.P2PLoad")
|
||||
@patch("vllm_ascend.model_loader.netloader.load.ElasticClient")
|
||||
def test_elastic_load_passes_draft_group_name(mock_client, mock_p2p):
|
||||
mock_client_instance = MagicMock()
|
||||
mock_client_instance.s = True
|
||||
mock_client_instance.ack = ["foo", "bar"]
|
||||
mock_client_instance.server_addr = "addr"
|
||||
mock_client_instance.__enter__.return_value = mock_client_instance
|
||||
mock_client.return_value = mock_client_instance
|
||||
|
||||
expected_model = object()
|
||||
mock_p2p_instance = MagicMock()
|
||||
mock_p2p_instance.load.return_value = expected_model
|
||||
mock_p2p.return_value = mock_p2p_instance
|
||||
|
||||
sources = [{"device_id": 0, "sources": ["127.0.0.1:15000"]}]
|
||||
result = elastic_load("model", 0, "draft-model", sources, 1, 1, group_name="netloader_draft")
|
||||
|
||||
assert result is expected_model
|
||||
mock_client.assert_called_once_with(["127.0.0.1:15000"], 0, "draft-model", 1, 1, "netloader_draft")
|
||||
mock_p2p.assert_called_once_with("foo", "addr", "bar", "netloader_draft")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main()
|
||||
60
tests/ut/model_loader/netloader/test_netloader_utils.py
Normal file
60
tests/ut/model_loader/netloader/test_netloader_utils.py
Normal file
@@ -0,0 +1,60 @@
|
||||
#
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
|
||||
import os
|
||||
import tempfile
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
from vllm_ascend.model_loader.netloader.utils import find_free_port, is_valid_path_prefix
|
||||
|
||||
|
||||
def test_find_free_port():
|
||||
port = find_free_port()
|
||||
assert isinstance(port, int)
|
||||
assert port > 0
|
||||
|
||||
|
||||
def test_is_valid_path_prefix_empty():
|
||||
assert not is_valid_path_prefix("")
|
||||
|
||||
|
||||
def test_is_valid_path_prefixIllegal_characters():
|
||||
assert not is_valid_path_prefix('test<>:"|?*')
|
||||
|
||||
|
||||
def test_is_valid_path_prefixRelative_path():
|
||||
assert is_valid_path_prefix("test")
|
||||
|
||||
|
||||
def test_is_valid_path_prefixAbsolute_path():
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
assert is_valid_path_prefix(os.path.join(tmpdir, "test"))
|
||||
|
||||
|
||||
@patch("os.path.exists", return_value=False)
|
||||
def test_is_valid_path_prefix_no_directory(mock_exists):
|
||||
assert not is_valid_path_prefix("/nonexistent_dir/test")
|
||||
|
||||
|
||||
@patch("os.path.exists", return_value=True)
|
||||
def test_is_valid_path_prefix_directory_exists(mock_exists):
|
||||
assert is_valid_path_prefix("/existing_dir/test")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main()
|
||||
531
tests/ut/model_loader/rfork/test_rfork_loader.py
Normal file
531
tests/ut/model_loader/rfork/test_rfork_loader.py
Normal file
@@ -0,0 +1,531 @@
|
||||
#
|
||||
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from vllm_ascend.model_loader.rfork.rfork_loader import (
|
||||
RForkModelLoader,
|
||||
_get_ep_rank,
|
||||
_get_pp_rank,
|
||||
_get_rfork_worker_attr,
|
||||
_is_draft_model,
|
||||
_is_dynamic_eplb_enabled,
|
||||
_is_layer_sharding_enabled,
|
||||
_make_fallback_load_config,
|
||||
)
|
||||
from vllm_ascend.model_loader.rfork.seed_protocol import get_local_seed_key
|
||||
|
||||
|
||||
class DummyLoadConfig:
|
||||
device = None
|
||||
load_format = "rfork"
|
||||
|
||||
def __init__(self, model_loader_extra_config):
|
||||
self.model_loader_extra_config = model_loader_extra_config
|
||||
|
||||
|
||||
@pytest.mark.parametrize("config_value", [True, False])
|
||||
def test_rfork_seed_timeout_bool_falls_back_to_env(monkeypatch, config_value):
|
||||
monkeypatch.setenv("RFORK_SEED_TIMEOUT_SEC", "7.5")
|
||||
|
||||
loader = RForkModelLoader(
|
||||
DummyLoadConfig(
|
||||
{
|
||||
"rfork_seed_timeout_sec": config_value,
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
assert loader.seed_timeout_sec == 7.5
|
||||
|
||||
|
||||
@pytest.mark.parametrize("config_value", [True, False])
|
||||
def test_rfork_seed_timeout_bool_falls_back_to_default(monkeypatch, config_value):
|
||||
monkeypatch.delenv("RFORK_SEED_TIMEOUT_SEC", raising=False)
|
||||
|
||||
loader = RForkModelLoader(
|
||||
DummyLoadConfig(
|
||||
{
|
||||
"rfork_seed_timeout_sec": config_value,
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
assert loader.seed_timeout_sec == 5.0
|
||||
|
||||
|
||||
def _parallel_config(
|
||||
*,
|
||||
enable_eplb=False,
|
||||
enable_expert_parallel=False,
|
||||
pipeline_parallel_size=1,
|
||||
is_moe_model=True,
|
||||
):
|
||||
return SimpleNamespace(
|
||||
enable_eplb=enable_eplb,
|
||||
enable_expert_parallel=enable_expert_parallel,
|
||||
pipeline_parallel_size=pipeline_parallel_size,
|
||||
is_moe_model=is_moe_model,
|
||||
)
|
||||
|
||||
|
||||
def _vllm_config(model_config=None, scheduler_config=None, parallel_config=None):
|
||||
return SimpleNamespace(
|
||||
additional_config=None,
|
||||
device_config=SimpleNamespace(device="cpu"),
|
||||
model_config=model_config or SimpleNamespace(),
|
||||
parallel_config=parallel_config or _parallel_config(),
|
||||
scheduler_config=scheduler_config or SimpleNamespace(),
|
||||
)
|
||||
|
||||
|
||||
def _parallel_vllm_config(
|
||||
*,
|
||||
enable_expert_parallel=False,
|
||||
pipeline_parallel_size=1,
|
||||
is_moe_model=True,
|
||||
):
|
||||
return SimpleNamespace(
|
||||
parallel_config=_parallel_config(
|
||||
enable_expert_parallel=enable_expert_parallel,
|
||||
pipeline_parallel_size=pipeline_parallel_size,
|
||||
is_moe_model=is_moe_model,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def test_rfork_ep_rank_is_not_added_when_expert_parallel_is_disabled(monkeypatch):
|
||||
def fail_if_ep_group_is_accessed():
|
||||
pytest.fail("EP group should not be accessed when expert parallelism is disabled.")
|
||||
|
||||
monkeypatch.setattr(
|
||||
"vllm_ascend.model_loader.rfork.rfork_loader.get_ep_group",
|
||||
fail_if_ep_group_is_accessed,
|
||||
)
|
||||
|
||||
assert _get_ep_rank(_parallel_vllm_config()) is None
|
||||
|
||||
|
||||
def test_rfork_ep_rank_comes_from_ep_group(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"vllm_ascend.model_loader.rfork.rfork_loader.get_ep_group",
|
||||
lambda: SimpleNamespace(rank_in_group=7),
|
||||
)
|
||||
|
||||
assert _get_ep_rank(_parallel_vllm_config(enable_expert_parallel=True)) == 7
|
||||
|
||||
|
||||
def test_rfork_ep_rank_is_not_added_for_dense_model(monkeypatch):
|
||||
def fail_if_ep_group_is_accessed():
|
||||
pytest.fail("EP group should not be accessed for a dense model.")
|
||||
|
||||
monkeypatch.setattr(
|
||||
"vllm_ascend.model_loader.rfork.rfork_loader.get_ep_group",
|
||||
fail_if_ep_group_is_accessed,
|
||||
)
|
||||
|
||||
assert _get_ep_rank(_parallel_vllm_config(enable_expert_parallel=True, is_moe_model=False)) is None
|
||||
|
||||
|
||||
def test_rfork_requires_initialized_ep_group(monkeypatch):
|
||||
def raise_uninitialized_ep_group():
|
||||
raise AssertionError("expert parallel group is not initialized")
|
||||
|
||||
monkeypatch.setattr(
|
||||
"vllm_ascend.model_loader.rfork.rfork_loader.get_ep_group",
|
||||
raise_uninitialized_ep_group,
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError, match="EP group is not initialized"):
|
||||
_get_ep_rank(_parallel_vllm_config(enable_expert_parallel=True))
|
||||
|
||||
|
||||
def test_rfork_pp_rank_is_not_added_when_pipeline_parallelism_is_disabled(monkeypatch):
|
||||
def fail_if_pp_group_is_accessed():
|
||||
pytest.fail("PP group should not be accessed when pipeline parallelism is disabled.")
|
||||
|
||||
monkeypatch.setattr(
|
||||
"vllm_ascend.model_loader.rfork.rfork_loader.get_pp_group",
|
||||
fail_if_pp_group_is_accessed,
|
||||
)
|
||||
|
||||
assert _get_pp_rank(_parallel_vllm_config()) is None
|
||||
|
||||
|
||||
def test_rfork_pp_rank_comes_from_pp_group(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"vllm_ascend.model_loader.rfork.rfork_loader.get_pp_group",
|
||||
lambda: SimpleNamespace(rank_in_group=3),
|
||||
)
|
||||
|
||||
assert _get_pp_rank(_parallel_vllm_config(pipeline_parallel_size=2)) == 3
|
||||
|
||||
|
||||
def test_rfork_requires_initialized_pp_group(monkeypatch):
|
||||
def raise_uninitialized_pp_group():
|
||||
raise AssertionError("pipeline parallel group is not initialized")
|
||||
|
||||
monkeypatch.setattr(
|
||||
"vllm_ascend.model_loader.rfork.rfork_loader.get_pp_group",
|
||||
raise_uninitialized_pp_group,
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError, match="PP group is not initialized"):
|
||||
_get_pp_rank(_parallel_vllm_config(pipeline_parallel_size=2))
|
||||
|
||||
|
||||
def test_rfork_seed_key_preserves_non_ep_format():
|
||||
assert (
|
||||
get_local_seed_key(
|
||||
disaggregation_mode="kv_consumer",
|
||||
node_rank=0,
|
||||
tp_rank=3,
|
||||
model_url="/models/dsv4",
|
||||
model_deploy_strategy_name="decode",
|
||||
)
|
||||
== "/models/dsv4$decode$kv_consumer$0$3"
|
||||
)
|
||||
|
||||
|
||||
def test_rfork_seed_key_isolated_by_ep_rank():
|
||||
common_config = {
|
||||
"disaggregation_mode": "kv_consumer",
|
||||
"node_rank": 0,
|
||||
"tp_rank": 0,
|
||||
"model_url": "/models/dsv4",
|
||||
"model_deploy_strategy_name": "decode",
|
||||
}
|
||||
|
||||
assert get_local_seed_key(**common_config, ep_rank=0) == "/models/dsv4$decode$kv_consumer$0$0$ep0"
|
||||
assert get_local_seed_key(**common_config, ep_rank=1) == "/models/dsv4$decode$kv_consumer$0$0$ep1"
|
||||
|
||||
|
||||
def test_rfork_seed_key_isolated_by_pp_rank():
|
||||
common_config = {
|
||||
"disaggregation_mode": "kv_consumer",
|
||||
"node_rank": 0,
|
||||
"tp_rank": 0,
|
||||
"ep_rank": 0,
|
||||
"model_url": "/models/dsv4",
|
||||
"model_deploy_strategy_name": "decode",
|
||||
}
|
||||
|
||||
assert get_local_seed_key(**common_config, pp_rank=0) == "/models/dsv4$decode$kv_consumer$0$pp0$0$ep0"
|
||||
assert get_local_seed_key(**common_config, pp_rank=1) == "/models/dsv4$decode$kv_consumer$0$pp1$0$ep0"
|
||||
|
||||
|
||||
def test_rfork_seed_key_distinguishes_parallel_rank_types():
|
||||
common_config = {
|
||||
"disaggregation_mode": "kv_consumer",
|
||||
"node_rank": 0,
|
||||
"model_url": "/models/dsv4",
|
||||
"model_deploy_strategy_name": "decode",
|
||||
}
|
||||
|
||||
pp_key = get_local_seed_key(**common_config, pp_rank=3, tp_rank=1)
|
||||
ep_key = get_local_seed_key(**common_config, tp_rank=3, ep_rank=1)
|
||||
|
||||
assert pp_key == "/models/dsv4$decode$kv_consumer$0$pp3$1"
|
||||
assert ep_key == "/models/dsv4$decode$kv_consumer$0$3$ep1"
|
||||
assert pp_key != ep_key
|
||||
|
||||
|
||||
def test_rfork_draft_seed_key_isolated_by_ep_rank():
|
||||
assert (
|
||||
get_local_seed_key(
|
||||
disaggregation_mode="kv_consumer",
|
||||
node_rank=0,
|
||||
tp_rank=0,
|
||||
model_url="/models/dsv4",
|
||||
model_deploy_strategy_name="decode",
|
||||
is_draft_worker=True,
|
||||
ep_rank=5,
|
||||
)
|
||||
== "/models/dsv4$decode$kv_consumer$0$0$ep5$draft"
|
||||
)
|
||||
|
||||
|
||||
def test_rfork_worker_receives_parallel_ranks(monkeypatch):
|
||||
load_config = DummyLoadConfig({"model_url": "model", "model_deploy_strategy_name": "strategy"})
|
||||
loader = RForkModelLoader(load_config)
|
||||
model_config = SimpleNamespace()
|
||||
vllm_config = SimpleNamespace(
|
||||
kv_transfer_config=None,
|
||||
model_config=model_config,
|
||||
scheduler_config=SimpleNamespace(),
|
||||
parallel_config=SimpleNamespace(node_rank=2),
|
||||
)
|
||||
captured = {}
|
||||
expected_worker = SimpleNamespace()
|
||||
|
||||
def fake_rfork_worker(**kwargs):
|
||||
captured.update(kwargs)
|
||||
return expected_worker
|
||||
|
||||
monkeypatch.setattr("vllm_ascend.model_loader.rfork.rfork_loader.RForkWorker", fake_rfork_worker)
|
||||
monkeypatch.setattr("vllm_ascend.model_loader.rfork.rfork_loader._get_pp_rank", lambda config: 3)
|
||||
monkeypatch.setattr("vllm_ascend.model_loader.rfork.rfork_loader._get_ep_rank", lambda config: 7)
|
||||
monkeypatch.setattr("vllm_ascend.model_loader.rfork.rfork_loader.get_tensor_model_parallel_rank", lambda: 5)
|
||||
monkeypatch.setattr(torch.distributed, "get_rank", lambda: 11)
|
||||
|
||||
worker = loader._ensure_rfork_worker(vllm_config, model_config)
|
||||
|
||||
assert worker is expected_worker
|
||||
assert captured["node_rank"] == 2
|
||||
assert captured["tp_rank"] == 5
|
||||
assert captured["pp_rank"] == 3
|
||||
assert captured["ep_rank"] == 7
|
||||
assert captured["device_id"] == 11
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_config",
|
||||
[
|
||||
SimpleNamespace(runner_type="draft"),
|
||||
SimpleNamespace(hf_config=SimpleNamespace(model_type="deepseek_mtp")),
|
||||
SimpleNamespace(hf_config=SimpleNamespace(architectures=["DeepSeekV4MTPModel"])),
|
||||
SimpleNamespace(hf_text_config=SimpleNamespace(architectures=["OpenPanguMTPModel"])),
|
||||
],
|
||||
)
|
||||
def test_rfork_detects_draft_model(model_config):
|
||||
assert _is_draft_model(_vllm_config(model_config=model_config))
|
||||
|
||||
|
||||
def test_rfork_detects_draft_model_from_scheduler_config():
|
||||
scheduler_config = SimpleNamespace(runner_type="draft")
|
||||
|
||||
assert _is_draft_model(_vllm_config(scheduler_config=scheduler_config))
|
||||
|
||||
|
||||
def test_rfork_does_not_treat_target_model_as_draft():
|
||||
target_model_config = SimpleNamespace(
|
||||
hf_config=SimpleNamespace(
|
||||
model_type="deepseek_v4",
|
||||
architectures=["DeepSeekV4ForCausalLM"],
|
||||
)
|
||||
)
|
||||
|
||||
assert not _is_draft_model(_vllm_config(model_config=target_model_config))
|
||||
|
||||
|
||||
def test_rfork_detects_explicit_draft_model_config():
|
||||
target_vllm_config = _vllm_config(
|
||||
model_config=SimpleNamespace(
|
||||
hf_config=SimpleNamespace(
|
||||
model_type="deepseek_v4",
|
||||
architectures=["DeepSeekV4ForCausalLM"],
|
||||
)
|
||||
)
|
||||
)
|
||||
draft_model_config = SimpleNamespace(
|
||||
hf_config=SimpleNamespace(
|
||||
model_type="deepseek_mtp",
|
||||
architectures=["DeepSeekV4MTPModel"],
|
||||
)
|
||||
)
|
||||
|
||||
assert _is_draft_model(target_vllm_config, draft_model_config)
|
||||
|
||||
|
||||
def test_rfork_uses_separate_worker_attr_for_explicit_draft_model_config():
|
||||
target_vllm_config = _vllm_config(
|
||||
model_config=SimpleNamespace(
|
||||
hf_config=SimpleNamespace(
|
||||
model_type="deepseek_v4",
|
||||
architectures=["DeepSeekV4ForCausalLM"],
|
||||
)
|
||||
)
|
||||
)
|
||||
draft_model_config = SimpleNamespace(
|
||||
hf_config=SimpleNamespace(
|
||||
model_type="deepseek_mtp",
|
||||
architectures=["DeepSeekV4MTPModel"],
|
||||
)
|
||||
)
|
||||
|
||||
assert _get_rfork_worker_attr(target_vllm_config, target_vllm_config.model_config) == "rfork_worker"
|
||||
assert _get_rfork_worker_attr(target_vllm_config, draft_model_config) == "rfork_draft_worker"
|
||||
|
||||
|
||||
def test_rfork_fallback_load_config_copy_does_not_mutate_original():
|
||||
original_extra_config = {"model_url": "model", "model_deploy_strategy_name": "tp8"}
|
||||
load_config = DummyLoadConfig(original_extra_config)
|
||||
|
||||
fallback_load_config = _make_fallback_load_config(load_config)
|
||||
|
||||
assert fallback_load_config is not load_config
|
||||
assert fallback_load_config.load_format == "auto"
|
||||
assert fallback_load_config.model_loader_extra_config == {}
|
||||
assert load_config.load_format == "rfork"
|
||||
assert load_config.model_loader_extra_config == original_extra_config
|
||||
|
||||
|
||||
def test_rfork_detects_layer_sharding_config():
|
||||
assert _is_layer_sharding_enabled(
|
||||
SimpleNamespace(
|
||||
additional_config={
|
||||
"layer_sharding": ["o_proj"],
|
||||
}
|
||||
)
|
||||
)
|
||||
assert not _is_layer_sharding_enabled(SimpleNamespace(additional_config={}))
|
||||
assert not _is_layer_sharding_enabled(SimpleNamespace(additional_config=None))
|
||||
|
||||
|
||||
def test_rfork_detects_dynamic_eplb_config():
|
||||
assert _is_dynamic_eplb_enabled(
|
||||
SimpleNamespace(
|
||||
parallel_config=SimpleNamespace(enable_eplb=True),
|
||||
additional_config=None,
|
||||
)
|
||||
)
|
||||
assert _is_dynamic_eplb_enabled(
|
||||
SimpleNamespace(
|
||||
parallel_config=SimpleNamespace(enable_eplb=False),
|
||||
additional_config={
|
||||
"eplb_config": {
|
||||
"dynamic_eplb": True,
|
||||
}
|
||||
},
|
||||
)
|
||||
)
|
||||
assert _is_dynamic_eplb_enabled(
|
||||
SimpleNamespace(
|
||||
parallel_config=SimpleNamespace(enable_eplb=False),
|
||||
additional_config={
|
||||
"eplb_config": {
|
||||
"expert_map_record_path": "/tmp/expert-map.json",
|
||||
}
|
||||
},
|
||||
)
|
||||
)
|
||||
assert not _is_dynamic_eplb_enabled(
|
||||
SimpleNamespace(
|
||||
parallel_config=SimpleNamespace(enable_eplb=False),
|
||||
additional_config={"eplb_config": {}},
|
||||
)
|
||||
)
|
||||
assert not _is_dynamic_eplb_enabled(
|
||||
SimpleNamespace(
|
||||
parallel_config=SimpleNamespace(enable_eplb=False),
|
||||
additional_config=None,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def test_rfork_layer_sharding_uses_default_loader(monkeypatch):
|
||||
import vllm.model_executor.model_loader as model_loader
|
||||
|
||||
load_config = DummyLoadConfig({"model_url": "model", "model_deploy_strategy_name": "tp8"})
|
||||
loader = RForkModelLoader(load_config)
|
||||
model_config = SimpleNamespace(dtype=torch.float32, model="/models/test")
|
||||
vllm_config = _vllm_config(model_config=model_config)
|
||||
vllm_config.additional_config = {"layer_sharding": ["o_proj"]}
|
||||
|
||||
def fail_if_rfork_worker_is_created(*args, **kwargs):
|
||||
raise AssertionError("RFork worker should not be initialized when layer_sharding is enabled.")
|
||||
|
||||
expected_model = SimpleNamespace()
|
||||
captured = {}
|
||||
|
||||
def fake_get_model(**kwargs):
|
||||
captured.update(kwargs)
|
||||
return expected_model
|
||||
|
||||
monkeypatch.setattr(loader, "_ensure_rfork_worker", fail_if_rfork_worker_is_created)
|
||||
monkeypatch.setattr(model_loader, "get_model", fake_get_model)
|
||||
|
||||
model = loader.load_model(vllm_config=vllm_config, model_config=model_config)
|
||||
|
||||
assert model is expected_model
|
||||
assert captured["vllm_config"] is vllm_config
|
||||
assert captured["model_config"] is model_config
|
||||
assert captured["prefix"] == ""
|
||||
assert captured["load_config"] is not load_config
|
||||
assert captured["load_config"].load_format == "auto"
|
||||
assert captured["load_config"].model_loader_extra_config == {}
|
||||
|
||||
|
||||
def test_rfork_dynamic_eplb_uses_default_loader(monkeypatch):
|
||||
import vllm.model_executor.model_loader as model_loader
|
||||
|
||||
load_config = DummyLoadConfig({"model_url": "model", "model_deploy_strategy_name": "tp8"})
|
||||
loader = RForkModelLoader(load_config)
|
||||
model_config = SimpleNamespace(dtype=torch.float32, model="/models/test")
|
||||
vllm_config = _vllm_config(model_config=model_config)
|
||||
vllm_config.additional_config = {"eplb_config": {"dynamic_eplb": True}}
|
||||
|
||||
def fail_if_rfork_worker_is_created(*args, **kwargs):
|
||||
raise AssertionError("RFork worker should not be initialized when dynamic EPLB is enabled.")
|
||||
|
||||
expected_model = SimpleNamespace()
|
||||
captured = {}
|
||||
|
||||
def fake_get_model(**kwargs):
|
||||
captured.update(kwargs)
|
||||
return expected_model
|
||||
|
||||
monkeypatch.setattr(loader, "_ensure_rfork_worker", fail_if_rfork_worker_is_created)
|
||||
monkeypatch.setattr(model_loader, "get_model", fake_get_model)
|
||||
|
||||
model = loader.load_model(vllm_config=vllm_config, model_config=model_config)
|
||||
|
||||
assert model is expected_model
|
||||
assert captured["vllm_config"] is vllm_config
|
||||
assert captured["model_config"] is model_config
|
||||
assert captured["prefix"] == ""
|
||||
assert captured["load_config"] is not load_config
|
||||
assert captured["load_config"].load_format == "auto"
|
||||
assert captured["load_config"].model_loader_extra_config == {}
|
||||
|
||||
|
||||
def test_rfork_native_eplb_uses_default_loader(monkeypatch):
|
||||
import vllm.model_executor.model_loader as model_loader
|
||||
|
||||
load_config = DummyLoadConfig({"model_url": "model", "model_deploy_strategy_name": "tp8"})
|
||||
loader = RForkModelLoader(load_config)
|
||||
model_config = SimpleNamespace(dtype=torch.float32, model="/models/test")
|
||||
vllm_config = _vllm_config(
|
||||
model_config=model_config,
|
||||
parallel_config=_parallel_config(enable_eplb=True),
|
||||
)
|
||||
vllm_config.additional_config = None
|
||||
|
||||
def fail_if_rfork_worker_is_created(*args, **kwargs):
|
||||
raise AssertionError("RFork worker should not be initialized when native EPLB is enabled.")
|
||||
|
||||
expected_model = SimpleNamespace()
|
||||
captured = {}
|
||||
|
||||
def fake_get_model(**kwargs):
|
||||
captured.update(kwargs)
|
||||
return expected_model
|
||||
|
||||
monkeypatch.setattr(loader, "_ensure_rfork_worker", fail_if_rfork_worker_is_created)
|
||||
monkeypatch.setattr(model_loader, "get_model", fake_get_model)
|
||||
|
||||
model = loader.load_model(vllm_config=vllm_config, model_config=model_config)
|
||||
|
||||
assert model is expected_model
|
||||
assert captured["vllm_config"] is vllm_config
|
||||
assert captured["model_config"] is model_config
|
||||
assert captured["prefix"] == ""
|
||||
assert captured["load_config"] is not load_config
|
||||
assert captured["load_config"].load_format == "auto"
|
||||
assert captured["load_config"].model_loader_extra_config == {}
|
||||
118
tests/ut/model_loader/rfork/test_transfer_backend.py
Normal file
118
tests/ut/model_loader/rfork/test_transfer_backend.py
Normal file
@@ -0,0 +1,118 @@
|
||||
#
|
||||
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
|
||||
from types import SimpleNamespace
|
||||
|
||||
import torch
|
||||
|
||||
import vllm_ascend.model_loader.rfork.transfer_backend as transfer_backend
|
||||
from vllm_ascend.model_loader.rfork.transfer_backend import (
|
||||
RForkTransferBackend,
|
||||
_parse_weight_info,
|
||||
_reshape_tensor_to_seed_shape,
|
||||
get_remote_instance_transfer_engine_info,
|
||||
)
|
||||
|
||||
|
||||
def test_parse_weight_info_keeps_backward_compatibility():
|
||||
assert _parse_weight_info([1, 2, 4]) == (1, 2, 4, None)
|
||||
|
||||
|
||||
def test_parse_weight_info_accepts_shape_metadata_from_json():
|
||||
assert _parse_weight_info([1, 6, 2, [2, 3]]) == (1, 6, 2, (2, 3))
|
||||
|
||||
|
||||
def test_parse_weight_info_rejects_invalid_shape_metadata():
|
||||
assert _parse_weight_info([1, 6, 2, ["2", 3]]) is None
|
||||
assert _parse_weight_info([1, 6, 2, -1]) is None
|
||||
|
||||
|
||||
def test_reshape_tensor_to_seed_shape_updates_tensor_metadata_only():
|
||||
tensor = torch.arange(6).reshape(2, 3)
|
||||
original_ptr = tensor.data_ptr()
|
||||
|
||||
assert _reshape_tensor_to_seed_shape("weight", tensor, (1, 2, 3))
|
||||
|
||||
assert tuple(tensor.shape) == (1, 2, 3)
|
||||
assert tensor.data_ptr() == original_ptr
|
||||
|
||||
|
||||
def test_reshape_tensor_to_seed_shape_rejects_numel_mismatch():
|
||||
tensor = torch.arange(6).reshape(2, 3)
|
||||
|
||||
assert not _reshape_tensor_to_seed_shape("weight", tensor, (2, 2))
|
||||
assert tuple(tensor.shape) == (2, 3)
|
||||
|
||||
|
||||
def test_recv_from_source_refreshes_registered_shape_after_reshape(monkeypatch):
|
||||
tensor = torch.arange(6).reshape(2, 3)
|
||||
backend = RForkTransferBackend.__new__(RForkTransferBackend)
|
||||
backend.rfork_transfer_engine = SimpleNamespace(
|
||||
batch_transfer_sync_read=lambda *args: SimpleNamespace(is_error=lambda: False)
|
||||
)
|
||||
backend.rfork_transfer_engine_weights_shape_dict = {"weight": (2, 3)}
|
||||
|
||||
monkeypatch.setattr(transfer_backend, "_iter_transferable_tensors", lambda model: iter([("weight", tensor)]))
|
||||
monkeypatch.setattr(
|
||||
transfer_backend,
|
||||
"get_remote_instance_transfer_engine_info",
|
||||
lambda *args: (
|
||||
"seed-session",
|
||||
{"weight": [1, tensor.numel(), tensor.element_size()]},
|
||||
{"weight": [1, 2, 3]},
|
||||
),
|
||||
)
|
||||
|
||||
assert backend.recv_from_source(object(), "127.0.0.1", 8000, "seed-key")
|
||||
assert tuple(tensor.shape) == (1, 2, 3)
|
||||
assert backend.rfork_transfer_engine_weights_shape_dict["weight"] == (1, 2, 3)
|
||||
|
||||
|
||||
def test_recv_from_source_reuses_registered_transferable_tensors(monkeypatch):
|
||||
tensor = torch.arange(6).reshape(2, 3)
|
||||
backend = RForkTransferBackend.__new__(RForkTransferBackend)
|
||||
backend.rfork_transfer_engine = SimpleNamespace(
|
||||
batch_transfer_sync_read=lambda *args: SimpleNamespace(is_error=lambda: False)
|
||||
)
|
||||
backend.rfork_transfer_engine_weights_shape_dict = {"weight": (2, 3)}
|
||||
backend._registered_transferable_tensors = [("weight", tensor)]
|
||||
|
||||
def fail_if_rescanned(model):
|
||||
raise AssertionError("recv_from_source should reuse the registered tensor cache")
|
||||
|
||||
monkeypatch.setattr(transfer_backend, "_iter_transferable_tensors", fail_if_rescanned)
|
||||
monkeypatch.setattr(
|
||||
transfer_backend,
|
||||
"get_remote_instance_transfer_engine_info",
|
||||
lambda *args: (
|
||||
"seed-session",
|
||||
{"weight": [1, tensor.numel(), tensor.element_size(), [2, 3]]},
|
||||
None,
|
||||
),
|
||||
)
|
||||
|
||||
assert backend.recv_from_source(object(), "127.0.0.1", 8000, "seed-key")
|
||||
assert backend._registered_transferable_tensors is None
|
||||
|
||||
|
||||
def test_get_remote_instance_transfer_engine_info_non_200_returns_three_values(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
transfer_backend.requests,
|
||||
"get",
|
||||
lambda *args, **kwargs: SimpleNamespace(status_code=503),
|
||||
)
|
||||
|
||||
assert get_remote_instance_transfer_engine_info("http://seed", "seed-key") == (None, None, None)
|
||||
Reference in New Issue
Block a user