init v0.23.0

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

View File

View File

View File

@@ -0,0 +1,194 @@
import unittest
from unittest.mock import MagicMock, patch
import torch
from transformers import DeepseekV2Config
from vllm_ascend.eplb.adaptor.vllm_adaptor import EPLB_EXPERT_WEIGHT_NAMES, VllmEplbAdaptor
from vllm_ascend.quantization.quant_type import QuantType
class TestVllmAdaptor(unittest.TestCase):
def setUp(self):
VllmEplbAdaptor._registered_moe_layers = []
n_routed_experts = 256
self.mock_layer = MagicMock()
self.mock_layer.local_num_experts = n_routed_experts
self.mock_layer.ep_rank = 0
self.mock_layer.quant_type = QuantType.W8A8
self.mock_layer.w13_weight_list = [torch.randn(256, 128) for _ in range(n_routed_experts)]
self.mock_layer.w2_weight_list = [torch.randn(128, 256) for _ in range(n_routed_experts)]
self.mock_layer.w13_weight_scale_fp32_list = [torch.tensor([1.0]) for _ in range(n_routed_experts)]
self.mock_layer.w2_weight_scale_list = [torch.tensor([1.0]) for _ in range(n_routed_experts)]
self.mock_layer.w13_weight = torch.randn(n_routed_experts, 256, 128)
self.mock_layer.w2_weight = torch.randn(n_routed_experts, 128, 256)
self.mock_layer.moe_load = torch.randn(n_routed_experts)
self.mock_layer.global_expert_map = torch.arange(n_routed_experts * 4).reshape(n_routed_experts, 4)
self.mock_layer.get_log2phy_map.return_value = torch.arange(4)
self.mock_layer.clear_moe_load = MagicMock()
VllmEplbAdaptor.register_layer(self.mock_layer)
mock_model = MagicMock()
mock_model.model.named_parameters.return_value = dict()
config = DeepseekV2Config(n_routed_experts=n_routed_experts)
mock_model.config = config
del mock_model.language_model
self.model = mock_model
num_dense_layers = getattr(config, "first_k_dense_replace", 0)
self.model.model.layers[num_dense_layers].mlp.experts.quant_type = QuantType.W8A8
self.mock_rank = patch("vllm_ascend.eplb.adaptor.vllm_adaptor.dist.get_rank", return_value=0).start()
self.mock_size = patch("vllm_ascend.eplb.adaptor.vllm_adaptor.dist.get_world_size", return_value=4).start()
@patch("torch.empty_like", return_value=torch.zeros(16, 32))
@patch("vllm_ascend.eplb.adaptor.vllm_adaptor.get_ascend_config")
def test_init_fp16(self, mock_get_config, mock_func):
mock_config = MagicMock()
mock_config.enable_fused_mc2 = 1
mock_get_config.return_value = mock_config
self.model.quant_config = None
adaptor = VllmEplbAdaptor(self.model)
self.assertEqual(adaptor.expert_weight_key_per_layer[0], (QuantType.NONE, True))
self.assertIs(adaptor.expert_param_per_layer[0][0][0], self.mock_layer.w13_weight_list[0])
self.assertIs(adaptor.expert_param_per_layer[0][0][1], self.mock_layer.w2_weight_list[0])
@patch("torch.empty_like", return_value=torch.zeros(16, 32))
@patch("vllm_ascend.eplb.adaptor.vllm_adaptor.get_ascend_config")
def test_init_w8a8(self, mock_get_config, mock_func):
mock_config = MagicMock()
mock_config.enable_fused_mc2 = 0
mock_get_config.return_value = mock_config
VllmEplbAdaptor(self.model)
@patch("torch.empty_like", return_value=torch.zeros(16, 32))
@patch("vllm_ascend.eplb.adaptor.vllm_adaptor.get_ascend_config")
def test_language_model_w8a8(self, mock_get_config, mock_func):
mock_config = MagicMock()
mock_config.enable_fused_mc2 = 0
mock_get_config.return_value = mock_config
model = MagicMock()
model.language_model = self.model
model.config.text_config = self.model.config
VllmEplbAdaptor(model)
def test_pp_eplb_adaptor_init_with_registered_layer(self):
"""PP+EPLB: adaptor picks up MoE layers registered via register_layer."""
VllmEplbAdaptor._registered_moe_layers = []
layer = MagicMock()
layer.local_num_experts = 4
layer.ep_rank = 0
layer.quant_type = QuantType.W8A8
layer.w13_weight_list = [torch.randn(256, 128) for _ in range(4)]
layer.w2_weight_list = [torch.randn(128, 256) for _ in range(4)]
layer.w13_weight_scale_fp32_list = [torch.tensor([1.0]) for _ in range(4)]
layer.w2_weight_scale_list = [torch.tensor([1.0]) for _ in range(4)]
layer.moe_load = torch.randn(4)
layer.global_expert_map = torch.arange(16).reshape(4, 4)
layer.get_log2phy_map.return_value = torch.arange(4)
VllmEplbAdaptor.register_layer(layer)
with patch("vllm_ascend.eplb.adaptor.vllm_adaptor.get_ascend_config") as mock_get_config:
mock_config = MagicMock()
mock_config.enable_fused_mc2 = 0
mock_get_config.return_value = mock_config
model = MagicMock()
model.quant_config = MagicMock()
model.config.first_k_dense_replace = 0
del model.language_model
adaptor = VllmEplbAdaptor(model)
self.assertEqual(adaptor.num_moe_layers, 1)
self.assertEqual(adaptor.num_local_experts, 4)
self.assertEqual(adaptor.ep_rank, 0)
@patch("vllm_ascend.eplb.adaptor.vllm_adaptor.get_ascend_config")
def test_init_mixed_quant_type_per_layer(self, mock_get_config):
mock_config = MagicMock()
mock_config.enable_fused_mc2 = 1
mock_get_config.return_value = mock_config
VllmEplbAdaptor._registered_moe_layers = []
num_local_experts = 2
w8a8_layer = MagicMock()
w8a8_layer.local_num_experts = num_local_experts
w8a8_layer.ep_rank = 0
w8a8_layer.quant_type = QuantType.W8A8
w8a8_layer.w13_weight_list = [torch.randn(2, 2) for _ in range(num_local_experts)]
w8a8_layer.w2_weight_list = [torch.randn(2, 2) for _ in range(num_local_experts)]
w8a8_layer.w13_weight_scale_fp32_list = [torch.randn(1) for _ in range(num_local_experts)]
w8a8_layer.w2_weight_scale_list = [torch.randn(1) for _ in range(num_local_experts)]
w8a8_layer.fused_w1_scale_list = [torch.randn(1) for _ in range(num_local_experts)]
w8a8_layer.fused_w2_scale_list = [torch.randn(1) for _ in range(num_local_experts)]
w8a8_layer.moe_load = torch.zeros(num_local_experts)
w8a8_layer.global_expert_map = torch.arange(num_local_experts * 4).reshape(num_local_experts, 4)
w8a8_layer.get_log2phy_map.return_value = torch.arange(4)
mxfp8_layer = MagicMock()
mxfp8_layer.local_num_experts = num_local_experts
mxfp8_layer.ep_rank = 0
mxfp8_layer.quant_type = QuantType.MXFP8
mxfp8_layer.w13_weight = torch.randn(num_local_experts, 2, 2)
mxfp8_layer.w2_weight = torch.randn(num_local_experts, 2, 2)
mxfp8_layer.w13_weight_scale = torch.randn(num_local_experts, 1)
mxfp8_layer.w2_weight_scale = torch.randn(num_local_experts, 1)
mxfp8_layer.moe_load = torch.zeros(num_local_experts)
mxfp8_layer.global_expert_map = torch.arange(num_local_experts * 4).reshape(num_local_experts, 4)
mxfp8_layer.get_log2phy_map.return_value = torch.arange(4)
VllmEplbAdaptor.register_layer(w8a8_layer)
VllmEplbAdaptor.register_layer(mxfp8_layer)
model = MagicMock()
model.quant_config = MagicMock()
model.config.first_k_dense_replace = 0
del model.language_model
adaptor = VllmEplbAdaptor(model)
w8a8_key = (QuantType.W8A8, True)
mxfp8_key = (QuantType.MXFP8, True)
self.assertEqual(adaptor.expert_weight_key_per_layer[0], w8a8_key)
self.assertEqual(adaptor.expert_weight_key_per_layer[1], mxfp8_key)
self.assertEqual(len(adaptor.buffer_tensor_list[w8a8_key][0]), len(EPLB_EXPERT_WEIGHT_NAMES[w8a8_key]))
self.assertEqual(len(adaptor.buffer_tensor_list[mxfp8_key][0]), len(EPLB_EXPERT_WEIGHT_NAMES[mxfp8_key]))
self.assertEqual(len(adaptor.expert_param_per_layer[0][0]), len(EPLB_EXPERT_WEIGHT_NAMES[w8a8_key]))
self.assertEqual(len(adaptor.expert_param_per_layer[1][0]), len(EPLB_EXPERT_WEIGHT_NAMES[mxfp8_key]))
@patch("vllm_ascend.eplb.adaptor.vllm_adaptor.get_ascend_config")
def test_reused_buffer_requires_same_expert_weight_shape(self, mock_get_config):
mock_config = MagicMock()
mock_config.enable_fused_mc2 = 0
mock_get_config.return_value = mock_config
VllmEplbAdaptor._registered_moe_layers = []
num_local_experts = 2
for weight_shape in [(2, 2), (3, 2)]:
layer = MagicMock()
layer.local_num_experts = num_local_experts
layer.ep_rank = 0
layer.quant_type = QuantType.W8A8
layer.w13_weight_list = [torch.randn(*weight_shape) for _ in range(num_local_experts)]
layer.w2_weight_list = [torch.randn(2, 2) for _ in range(num_local_experts)]
layer.w13_weight_scale_fp32_list = [torch.randn(1) for _ in range(num_local_experts)]
layer.w2_weight_scale_list = [torch.randn(1) for _ in range(num_local_experts)]
layer.moe_load = torch.zeros(num_local_experts)
layer.global_expert_map = torch.arange(num_local_experts * 4).reshape(num_local_experts, 4)
layer.get_log2phy_map.return_value = torch.arange(4)
VllmEplbAdaptor.register_layer(layer)
model = MagicMock()
model.quant_config = MagicMock()
model.config.first_k_dense_replace = 0
del model.language_model
with self.assertRaisesRegex(AssertionError, "EPLB expert weight shapes mismatch"):
VllmEplbAdaptor(model)
def tearDown(self):
self.mock_rank.stop()
self.mock_size.stop()
VllmEplbAdaptor._registered_moe_layers = []
if __name__ == "__main__":
unittest.main()

View File

View File

View File

@@ -0,0 +1,17 @@
{
"moe_layer_count":
1,
"layer_list": [{
"layer_id":
0,
"device_count":
2,
"device_list": [{
"device_id": 0,
"device_expert": [7, 2, 0, 3, 5]
}, {
"device_id": 1,
"device_expert": [6, 1, 4, 7, 2]
}]
}]
}

View File

@@ -0,0 +1,117 @@
import os
import unittest
from unittest.mock import MagicMock, patch
# isort: off
import torch
from vllm.config import VllmConfig
from vllm.model_executor.layers.fused_moe.config import FusedMoEConfig, FusedMoEParallelConfig
from vllm_ascend.ascend_config import init_ascend_config
from vllm_ascend.eplb.core.eplb_utils import generate_log2phy_map, init_eplb_config
from vllm_ascend.utils import vllm_version_is
# isort: on
class TestAscendConfig(unittest.TestCase):
@patch("vllm.config.VllmConfig.__post_init__", MagicMock())
@patch("vllm_ascend.platform.NPUPlatform._fix_incompatible_config")
def setUp(self, mock_fix_incompatible_config):
vllm_config = VllmConfig()
vllm_config.model_config = MagicMock()
vllm_config.additional_config = {
"refresh": True,
"eplb_config": {"dynamic_eplb": True, "num_redundant_experts": 2},
}
from vllm.model_executor.layers.fused_moe.config import RoutingMethodType
moe_parallel_config = FusedMoEParallelConfig(2, 0, 1, 2, 1, 1, 1, 1, 1, True, "hccl", enable_eplb=True)
if vllm_version_is("0.23.0"):
moe_config = FusedMoEConfig(
num_experts=8,
experts_per_token=8,
hidden_dim=8192,
intermediate_size_per_partition=5,
num_local_experts=8,
num_logical_experts=8,
activation="silu",
device="npu",
routing_method=RoutingMethodType.Simulated,
moe_parallel_config=moe_parallel_config,
in_dtype=torch.float16,
)
else:
from vllm.model_executor.layers.fused_moe.activation import MoEActivation
moe_config = FusedMoEConfig(
num_experts=8,
experts_per_token=8,
hidden_dim=8192,
intermediate_size=10,
num_local_experts=8,
num_logical_experts=8,
activation=MoEActivation.SILU,
device="npu",
routing_method=RoutingMethodType.Simulated,
moe_parallel_config=moe_parallel_config,
in_dtype=torch.float16,
)
moe_config.supports_eplb = True
self.vllm_config = vllm_config
self.moe_config = moe_config
self.mock_npu_patcher = patch("torch.Tensor.npu", new=lambda self: self)
self.mock_npu_patcher.start()
os.environ["DYNAMIC_EPLB"] = "true"
def tearDown(self):
self.mock_npu_patcher.stop()
os.environ.pop("DYNAMIC_EPLB", None)
def test_init_eplb_config_with_eplb(self):
eplb_config = init_ascend_config(self.vllm_config).eplb_config
_, expert_map, log2phy, redundant_experts = init_eplb_config(eplb_config, 0, self.moe_config)
gt_expert_map = torch.tensor([4, -1, -1, -1, 0, 1, 2, 3])
gt_log2phy = torch.tensor([9, 1, 2, 3, 5, 6, 7, 8])
self.assertTrue(torch.equal(expert_map, gt_expert_map))
self.assertTrue(torch.equal(log2phy, gt_log2phy))
self.assertEqual(redundant_experts, 2)
def test_init_eplb_config_with_eplb_withmap(self):
_TEST_DIR = os.path.dirname(__file__)
self.vllm_config.additional_config["eplb_config"]["expert_map_path"] = _TEST_DIR + "/expert_map.json"
eplb_config = init_ascend_config(self.vllm_config).eplb_config
_, expert_map, log2phy, redundant_experts = init_eplb_config(eplb_config, 0, self.moe_config)
gt_expert_map = torch.tensor([-1, 1, 4, -1, 2, -1, 0, 3])
gt_log2phy = torch.tensor([2, 6, 9, 3, 7, 4, 5, 8])
self.assertTrue(torch.equal(expert_map, gt_expert_map))
self.assertTrue(torch.equal(log2phy, gt_log2phy))
self.assertEqual(redundant_experts, 2)
def test_generate_log2phy_map_rotates_tail_tp_rank_with_tp_size(self):
global_expert_map = [
torch.tensor([0, -1], dtype=torch.int32),
torch.tensor([0, -1], dtype=torch.int32),
torch.tensor([0, -1], dtype=torch.int32),
torch.tensor([0, -1], dtype=torch.int32),
torch.tensor([-1, 0], dtype=torch.int32),
torch.tensor([-1, 0], dtype=torch.int32),
torch.tensor([-1, 0], dtype=torch.int32),
torch.tensor([-1, 0], dtype=torch.int32),
]
fallback_tail_dp1 = generate_log2phy_map(global_expert_map, ep_rank=7)
rotated_tail_dp0 = generate_log2phy_map(global_expert_map, ep_rank=3, tp_size=4)
rotated_tail_dp1 = generate_log2phy_map(global_expert_map, ep_rank=7, tp_size=4)
self.assertTrue(torch.equal(fallback_tail_dp1, torch.tensor([3, 7], dtype=torch.int32)))
self.assertTrue(torch.equal(rotated_tail_dp0, torch.tensor([3, 4], dtype=torch.int32)))
self.assertTrue(torch.equal(rotated_tail_dp1, torch.tensor([0, 5], dtype=torch.int32)))
def test_init_eplb_config_without_eplb(self):
self.vllm_config.additional_config = {"refresh": True}
eplb_config = init_ascend_config(self.vllm_config).eplb_config
_, expert_map, log2phy, redundant_experts = init_eplb_config(eplb_config, 0, self.moe_config)
gt_expert_map = torch.tensor([-1, -1, -1, -1, 0, 1, 2, 3])
self.assertIsNone(log2phy)
self.assertTrue(torch.equal(expert_map, gt_expert_map))
self.assertEqual(redundant_experts, 0)

View File

View File

@@ -0,0 +1,36 @@
import unittest
import torch
from vllm_ascend.eplb.core.eplb_worker import EplbWorker
from vllm_ascend.eplb.core.policy.policy_factory import PolicyFactory
from vllm_ascend.eplb.core.policy.policy_flashlb import generate_layered_experts
class TestEplbRebalancePolicies(unittest.TestCase):
def setUp(self):
torch.manual_seed(42)
self.current_expert_table = generate_layered_experts()
x = torch.rand(100, 58, 32, 9)
x = x**10
self.expert_workload = (x * 999 + 1).long()
self.hotness = EplbWorker._calculate_hotness(self.current_expert_table, self.expert_workload.sum(0))
@unittest.mock.patch("torch.npu.device_count", return_value=16)
def test_swift_balance_rebalance_experts(self, mock_count):
swift_policy = PolicyFactory.generate_policy(2)
_, _, new_placement = swift_policy.rebalance_experts(self.current_expert_table, self.expert_workload.sum(0))
update_mean, _ = EplbWorker._compute_imbalance(new_placement, self.hotness)
self.assertLessEqual(update_mean, 1.08)
def test_flashlb_rebalance_experts(self):
flashlb_policy = PolicyFactory.generate_policy(3)
_, _, new_placement = flashlb_policy.rebalance_experts(self.current_expert_table, self.expert_workload)
update_mean, _ = EplbWorker._compute_imbalance(new_placement, self.hotness)
self.assertLessEqual(update_mean, 1.1)
if __name__ == "__main__":
unittest.main(verbosity=2)

View File

@@ -11,54 +11,61 @@ import vllm_ascend.eplb.core.eplb_device_transfer_loader as loader
def mock_adaptor():
adaptor = MagicMock()
adaptor.expert_map_per_layer_cpu = {
0: {
10: torch.tensor(1),
20: torch.tensor(0)
}
}
adaptor.expert_map_per_layer_cpu = {0: {10: torch.tensor(1), 20: torch.tensor(0)}}
adaptor.expert_param_per_layer = {
0: {
0: [[torch.tensor([1.0])]],
1: [[torch.tensor([2.0])]]
}
}
adaptor.expert_param_per_layer = {0: {0: [[torch.tensor([1.0])]], 1: [[torch.tensor([2.0])]]}}
adaptor.buffer_tensor_list = [[[torch.tensor([3.0])],
[torch.tensor([4.0])]]]
adaptor.expert_weight_key_per_layer = {0: "weight_key"}
adaptor.buffer_tensor_list = {
"weight_key": [[torch.tensor([3.0]), torch.tensor([4.0])], [torch.tensor([5.0]), torch.tensor([6.0])]]
}
return adaptor
def test_generate_task_and_state_flow(mock_adaptor):
loader_obj = loader.D2DExpertWeightLoader()
with patch("vllm_ascend.eplb.core.eplb_device_transfer_loader.get_dynamic_eplb_group", return_value=None):
loader_obj = loader.D2DExpertWeightLoader()
loader_obj.set_adator(mock_adaptor)
with patch("torch.distributed.P2POp") as mock_p2p, \
patch("torch.distributed.isend", return_value="isend_op"), \
patch("torch.distributed.irecv", return_value="irecv_op"):
with (
patch("torch.distributed.P2POp") as mock_p2p,
patch("torch.distributed.isend", return_value="isend_op"),
patch("torch.distributed.irecv", return_value="irecv_op"),
):
mock_p2p.side_effect = lambda op, tensor, rank: (op, tensor, rank)
loader_obj.state = loader.ExpertWeightUpdateState.READY
loader_obj.generate_expert_d2d_transfer_task([(1, 10)], [(2, 20)],
{20: torch.tensor(0)}, 0)
loader_obj.generate_expert_d2d_transfer_task([(1, 10)], [(2, 20)], {20: torch.tensor(0)}, 0)
assert loader_obj.comm_op_list is None
loader_obj.state = loader.ExpertWeightUpdateState.WAITING
loader_obj.generate_expert_d2d_transfer_task([], [], {}, 0)
assert loader_obj.comm_op_list is None
updated_map = {20: torch.tensor(0)}
loader_obj.generate_expert_d2d_transfer_task([(1, 10)], [(2, 20)],
updated_map, 0)
assert not loader_obj.comm_op_list
assert loader_obj.state == loader.ExpertWeightUpdateState.READY
assert loader_obj.comm_op_list
assert loader_obj.recv_expert_list
def test_generate_task_uses_layer_weight_key_buffer(mock_adaptor):
comm_group = MagicMock()
comm_group.ranks = {2: 20}
comm_group.device_group = object()
with patch("vllm_ascend.eplb.core.eplb_device_transfer_loader.get_dynamic_eplb_group", return_value=comm_group):
loader_obj = loader.D2DExpertWeightLoader()
loader_obj.set_adator(mock_adaptor)
with (
patch("torch.distributed.P2POp") as mock_p2p,
patch("torch.distributed.irecv", return_value="irecv_op"),
):
mock_p2p.side_effect = lambda op, tensor, rank, group=None: (op, tensor, rank, group)
loader_obj.generate_expert_d2d_transfer_task([], [(2, 20)], {20: torch.tensor(0)}, 0)
assert mock_p2p.call_args_list[0].args[1] is mock_adaptor.buffer_tensor_list["weight_key"][0][0]
assert mock_p2p.call_args_list[1].args[1] is mock_adaptor.buffer_tensor_list["weight_key"][0][1]
def test_asyn_transfer_and_update(mock_adaptor):
loader_obj = loader.D2DExpertWeightLoader()
with patch("vllm_ascend.eplb.core.eplb_device_transfer_loader.get_dynamic_eplb_group", return_value=None):
loader_obj = loader.D2DExpertWeightLoader()
loader_obj.set_adator(mock_adaptor)
loader_obj.comm_op_list = ["fake_op"]
@@ -66,8 +73,7 @@ def test_asyn_transfer_and_update(mock_adaptor):
reqs: list[MagicMock] = []
with patch("torch.distributed.batch_isend_irecv",
return_value=[MagicMock(), MagicMock()]):
with patch("torch.distributed.batch_isend_irecv", return_value=[MagicMock(), MagicMock()]):
loader_obj.asyn_expert_weight_transfer(reqs)
assert loader_obj.state == loader.ExpertWeightUpdateState.TRANSFERRING
@@ -94,14 +100,16 @@ def test_asyn_transfer_and_update(mock_adaptor):
def test_set_log2phy_map(mock_adaptor):
loader_obj = loader.D2DExpertWeightLoader()
with patch("vllm_ascend.eplb.core.eplb_device_transfer_loader.get_dynamic_eplb_group", return_value=None):
loader_obj = loader.D2DExpertWeightLoader()
loader_obj.set_adator(mock_adaptor)
loader_obj.set_log2phy_map({"a": 1})
assert loader_obj.updated_log2phy_map == {"a": 1}
def test_invalid_state_asyn_update(mock_adaptor):
loader_obj = loader.D2DExpertWeightLoader()
with patch("vllm_ascend.eplb.core.eplb_device_transfer_loader.get_dynamic_eplb_group", return_value=None):
loader_obj = loader.D2DExpertWeightLoader()
loader_obj.set_adator(mock_adaptor)
loader_obj.state = loader.ExpertWeightUpdateState.WAITING
@@ -113,10 +121,3 @@ def test_invalid_state_asyn_update(mock_adaptor):
loader_obj.update_expert_map_and_weight([])
assert not mock_adaptor.do_update_expert_map.called
def test_load_impl_not_implemented(mock_adaptor):
loader_obj = loader.D2DExpertWeightLoader()
loader_obj.set_adator(mock_adaptor)
with pytest.raises(NotImplementedError):
loader_obj.load_impl({}, {})

View File

@@ -0,0 +1,87 @@
import unittest
from unittest.mock import MagicMock, patch
import torch
from vllm_ascend.eplb.eplb_updator import EplbUpdator
class TestEplbUpdatorComputeAndSetMoeLoad(unittest.TestCase):
def setUp(self):
# ====================== 1. Mock environment ======================
self.rank = 0
self.world_size = 4
self.device = torch.device("cpu")
# mock dist
p1 = patch("torch.distributed.get_rank", return_value=self.rank)
p2 = patch("torch.distributed.get_world_size", return_value=self.world_size)
self.addCleanup(p1.stop)
self.addCleanup(p2.stop)
p1.start()
p2.start()
# ====================== 2. Mock comm group ======================
self.mock_comm_group = MagicMock()
def mock_all_gather(tensor, dim):
gathered = torch.cat([tensor for _ in range(self.world_size)], dim=dim)
return gathered
self.mock_comm_group.all_gather = mock_all_gather
p3 = patch("vllm_ascend.eplb.eplb_updator.get_dynamic_eplb_group", return_value=self.mock_comm_group)
self.addCleanup(p3.stop)
p3.start()
# mock _PP in vllm.distributed.parallel_state (PP+EPLB support)
# Patching the variable directly so that even the real get_pp_group()
# (already imported into eplb_updator's namespace) reads a non-None _PP.
self.mock_pp = MagicMock()
self.mock_pp.rank_in_group = 0
p4 = patch("vllm.distributed.parallel_state._PP", self.mock_pp)
self.addCleanup(p4.stop)
p4.start()
# ====================== 3. Mock EplbUpdator ======================
self.eplb_config = MagicMock()
self.loader = MagicMock()
self.eplb_process = MagicMock()
self.process = MagicMock()
self.eplb_process.shared_dict = {}
self.updator = EplbUpdator(
eplb_config=self.eplb_config, loader=self.loader, eplb_process=self.eplb_process, process=self.process
)
# ====================== 4. Mock adaptor ======================
self.adaptor = MagicMock()
self.adaptor.num_moe_layers = 4
self.adaptor.num_dense_layers = 2
self.mock_local_load = torch.randn(58, 100, 8, device=self.device)
self.adaptor.get_rank_expert_workload.return_value = self.mock_local_load
self.updator.set_adaptor(self.adaptor)
def test_compute_and_set_moe_load_normal(self):
self.updator.multi_stage = False
moe_load = self.updator.compute_and_set_moe_load()
self.assertEqual(moe_load.shape, (58, self.world_size, 100, 8))
self.assertTrue("moe_load" in self.updator.shared_dict)
self.assertEqual(moe_load.device.type, "cpu")
self.assertEqual(moe_load.shape[1], self.world_size)
def test_compute_and_set_moe_load_multi_stage(self):
self.updator.multi_stage = True
moe_load = self.updator.compute_and_set_moe_load()
self.assertEqual(moe_load.shape, (100, 58, self.world_size, 8))
self.assertTrue("moe_load" in self.updator.shared_dict)
self.assertEqual(moe_load.device.type, "cpu")
if __name__ == "__main__":
unittest.main()