@@ -13,139 +13,369 @@
|
||||
# This file is a part of the vllm-ascend project.
|
||||
#
|
||||
|
||||
import json
|
||||
import os
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
from transformers import PretrainedConfig
|
||||
from vllm.config import ModelConfig, ParallelConfig, VllmConfig
|
||||
from vllm.config import KVTransferConfig, VllmConfig
|
||||
|
||||
from tests.ut.base import TestBase
|
||||
from vllm_ascend.ascend_config import (_check_torchair_supported,
|
||||
check_ascend_config,
|
||||
clear_ascend_config, get_ascend_config,
|
||||
init_ascend_config)
|
||||
from vllm_ascend.ascend_config import clear_ascend_config, get_ascend_config, init_ascend_config
|
||||
from vllm_ascend.utils import clear_enable_sp, enable_sp, get_flashcomm2_config_and_validate
|
||||
|
||||
|
||||
class TestAscendConfig(TestBase):
|
||||
|
||||
@staticmethod
|
||||
def _clean_up_ascend_config(func):
|
||||
|
||||
def wrapper(*args, **kwargs):
|
||||
clear_ascend_config()
|
||||
func(*args, **kwargs)
|
||||
clear_ascend_config()
|
||||
clear_enable_sp()
|
||||
try:
|
||||
func(*args, **kwargs)
|
||||
finally:
|
||||
clear_ascend_config()
|
||||
clear_enable_sp()
|
||||
|
||||
return wrapper
|
||||
|
||||
@staticmethod
|
||||
def _make_model_config(
|
||||
total_num_attention_heads: int = 32,
|
||||
total_num_kv_heads: int = 8,
|
||||
is_deepseek_mla: bool = False,
|
||||
):
|
||||
return SimpleNamespace(
|
||||
is_deepseek_mla=is_deepseek_mla,
|
||||
use_mla=is_deepseek_mla,
|
||||
enforce_eager=True,
|
||||
model_arch_config=SimpleNamespace(total_num_attention_heads=total_num_attention_heads),
|
||||
get_total_num_kv_heads=lambda: total_num_kv_heads,
|
||||
)
|
||||
|
||||
@_clean_up_ascend_config
|
||||
def test_init_ascend_config_without_additional_config(self):
|
||||
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")
|
||||
def test_init_ascend_config_without_additional_config(self, mock_fix_incompatible_config):
|
||||
test_vllm_config = VllmConfig()
|
||||
# No additional config given, check the default value here.
|
||||
ascend_config = init_ascend_config(test_vllm_config)
|
||||
self.assertIsNone(ascend_config.expert_map_path)
|
||||
self.assertFalse(ascend_config.multistream_overlap_shared_expert)
|
||||
self.assertFalse(ascend_config.enable_kv_nz)
|
||||
|
||||
torchair_graph_config = ascend_config.torchair_graph_config
|
||||
self.assertFalse(torchair_graph_config.enabled)
|
||||
self.assertEqual(torchair_graph_config.mode, '')
|
||||
self.assertFalse(torchair_graph_config.use_cached_graph)
|
||||
self.assertEqual(torchair_graph_config.graph_batch_sizes, [])
|
||||
self.assertFalse(torchair_graph_config.graph_batch_sizes_init)
|
||||
self.assertFalse(torchair_graph_config.enable_multistream_mla)
|
||||
self.assertTrue(torchair_graph_config.enable_view_optimize)
|
||||
self.assertTrue(torchair_graph_config.enable_frozen_parameter)
|
||||
self.assertFalse(torchair_graph_config.enable_kv_nz)
|
||||
ascend_compilation_config = ascend_config.ascend_compilation_config
|
||||
self.assertTrue(ascend_compilation_config.fuse_norm_quant)
|
||||
|
||||
ascend_scheduler_config = ascend_config.ascend_scheduler_config
|
||||
self.assertFalse(ascend_scheduler_config.enabled)
|
||||
ascend_fusion_config = ascend_config.ascend_fusion_config
|
||||
self.assertTrue(ascend_fusion_config.fusion_ops_gmmswigluquant)
|
||||
|
||||
@_clean_up_ascend_config
|
||||
def test_init_ascend_config_with_additional_config(self):
|
||||
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")
|
||||
def test_init_ascend_config_with_additional_config(self, mock_fix_incompatible_config):
|
||||
test_vllm_config = VllmConfig()
|
||||
test_vllm_config.additional_config = {
|
||||
"torchair_graph_config": {
|
||||
"enabled": True,
|
||||
"use_cached_graph": True,
|
||||
"graph_batch_sizes": [1, 2, 4],
|
||||
"graph_batch_sizes_init": False,
|
||||
"enable_multistream_mla": True,
|
||||
"enable_view_optimize": True,
|
||||
"enable_frozen_parameter": True,
|
||||
"enable_kv_nz": True
|
||||
"ascend_compilation_config": {
|
||||
"fuse_norm_quant": False,
|
||||
},
|
||||
"ascend_fusion_config": {
|
||||
"fusion_ops_gmmswigluquant": False,
|
||||
},
|
||||
"multistream_overlap_shared_expert": True,
|
||||
"ascend_scheduler_config": {
|
||||
"enabled": True
|
||||
},
|
||||
"expert_map_path": "test_expert_map_path",
|
||||
"eplb_config": {"num_redundant_experts": 2},
|
||||
"refresh": True,
|
||||
"enable_kv_nz": False,
|
||||
}
|
||||
ascend_config = init_ascend_config(test_vllm_config)
|
||||
self.assertEqual(ascend_config.expert_map_path, "test_expert_map_path")
|
||||
self.assertEqual(ascend_config.eplb_config.num_redundant_experts, 2)
|
||||
self.assertTrue(ascend_config.multistream_overlap_shared_expert)
|
||||
|
||||
torchair_graph_config = ascend_config.torchair_graph_config
|
||||
self.assertTrue(torchair_graph_config.enabled)
|
||||
self.assertTrue(torchair_graph_config.use_cached_graph)
|
||||
self.assertEqual(torchair_graph_config.graph_batch_sizes, [1, 2, 4])
|
||||
self.assertFalse(torchair_graph_config.graph_batch_sizes_init)
|
||||
self.assertTrue(torchair_graph_config.enable_multistream_mla)
|
||||
self.assertTrue(torchair_graph_config.enable_view_optimize)
|
||||
self.assertTrue(torchair_graph_config.enable_frozen_parameter)
|
||||
self.assertTrue(torchair_graph_config.enable_kv_nz)
|
||||
ascend_compilation_config = ascend_config.ascend_compilation_config
|
||||
self.assertFalse(ascend_compilation_config.fuse_norm_quant)
|
||||
self.assertFalse(ascend_config.enable_kv_nz)
|
||||
self.assertTrue(ascend_compilation_config.enable_npugraph_ex)
|
||||
self.assertFalse(ascend_compilation_config.enable_static_kernel)
|
||||
|
||||
ascend_scheduler_config = ascend_config.ascend_scheduler_config
|
||||
self.assertTrue(ascend_scheduler_config.enabled)
|
||||
ascend_fusion_config = ascend_config.ascend_fusion_config
|
||||
self.assertFalse(ascend_fusion_config.fusion_ops_gmmswigluquant)
|
||||
|
||||
@_clean_up_ascend_config
|
||||
def test_init_ascend_config_with_refresh(self):
|
||||
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")
|
||||
def test_init_ascend_config_enable_npugraph_ex(self, mock_fix_incompatible_config):
|
||||
test_vllm_config = VllmConfig()
|
||||
ascend_config = init_ascend_config(test_vllm_config)
|
||||
self.assertFalse(ascend_config.torchair_graph_config.enabled)
|
||||
|
||||
test_vllm_config.additional_config = {
|
||||
"torchair_graph_config": {
|
||||
"enabled": True,
|
||||
},
|
||||
}
|
||||
ascend_config = init_ascend_config(test_vllm_config)
|
||||
self.assertFalse(ascend_config.torchair_graph_config.enabled)
|
||||
|
||||
test_vllm_config.additional_config = {
|
||||
"torchair_graph_config": {
|
||||
"enabled": True,
|
||||
},
|
||||
"ascend_compilation_config": {"enable_npugraph_ex": True, "enable_static_kernel": True},
|
||||
"refresh": True,
|
||||
}
|
||||
ascend_config = init_ascend_config(test_vllm_config)
|
||||
self.assertTrue(ascend_config.torchair_graph_config.enabled)
|
||||
ascend_compilation_config = init_ascend_config(test_vllm_config).ascend_compilation_config
|
||||
self.assertTrue(ascend_compilation_config.enable_npugraph_ex)
|
||||
self.assertTrue(ascend_compilation_config.enable_static_kernel)
|
||||
|
||||
@_clean_up_ascend_config
|
||||
def test_init_ascend_config_with_wrong_input(self):
|
||||
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")
|
||||
def test_init_ascend_config_rejects_mooncake_c8_kv_cache_consumer(self, mock_fix_incompatible_config):
|
||||
test_vllm_config = VllmConfig()
|
||||
test_vllm_config.additional_config = {
|
||||
"torchair_graph_config": {
|
||||
"enabled": True,
|
||||
"graph_batch_sizes": "fake_size",
|
||||
},
|
||||
"refresh": True,
|
||||
}
|
||||
with self.assertRaises(TypeError):
|
||||
init_ascend_config(test_vllm_config)
|
||||
test_vllm_config.kv_transfer_config = KVTransferConfig(
|
||||
kv_connector="MooncakeConnectorV1",
|
||||
kv_role="kv_consumer",
|
||||
)
|
||||
test_vllm_config.quant_config = SimpleNamespace(enable_c8_quant=True)
|
||||
test_vllm_config.model_config = self._make_model_config()
|
||||
|
||||
test_vllm_config.additional_config = {
|
||||
"torchair_graph_config": {
|
||||
"enabled": False,
|
||||
"graph_batch_sizes": [1, 2, 4, 8],
|
||||
"graph_batch_sizes_init": True,
|
||||
},
|
||||
"refresh": True,
|
||||
}
|
||||
with self.assertRaises(ValueError):
|
||||
with self.assertRaisesRegex(ValueError, "does not support C8 KV cache quantization"):
|
||||
init_ascend_config(test_vllm_config)
|
||||
|
||||
@_clean_up_ascend_config
|
||||
def test_get_ascend_config(self):
|
||||
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")
|
||||
def test_init_ascend_config_rejects_multi_connector_mooncake_c8_consumer(self, mock_fix_incompatible_config):
|
||||
test_vllm_config = VllmConfig()
|
||||
test_vllm_config.kv_transfer_config = KVTransferConfig(
|
||||
kv_connector="MultiConnector",
|
||||
kv_role="kv_consumer",
|
||||
kv_connector_extra_config={
|
||||
"connectors": [
|
||||
{
|
||||
"kv_connector": "MooncakeConnectorV1",
|
||||
"kv_role": "kv_consumer",
|
||||
}
|
||||
]
|
||||
},
|
||||
)
|
||||
test_vllm_config.quant_config = SimpleNamespace(enable_c8_quant=True)
|
||||
test_vllm_config.model_config = self._make_model_config()
|
||||
|
||||
with self.assertRaisesRegex(ValueError, "does not support C8 KV cache quantization"):
|
||||
init_ascend_config(test_vllm_config)
|
||||
|
||||
@_clean_up_ascend_config
|
||||
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")
|
||||
def test_init_ascend_config_allows_layerwise_c8_kv_cache_consumer(self, mock_fix_incompatible_config):
|
||||
test_vllm_config = VllmConfig()
|
||||
test_vllm_config.kv_transfer_config = KVTransferConfig(
|
||||
kv_connector="MooncakeLayerwiseConnector",
|
||||
kv_role="kv_consumer",
|
||||
)
|
||||
test_vllm_config.quant_config = SimpleNamespace(enable_c8_quant=True)
|
||||
test_vllm_config.model_config = self._make_model_config()
|
||||
|
||||
ascend_config = init_ascend_config(test_vllm_config)
|
||||
|
||||
self.assertIsNotNone(ascend_config)
|
||||
|
||||
@_clean_up_ascend_config
|
||||
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")
|
||||
def test_init_ascend_config_allows_mha_mooncake_c8_kv_cache_consumer(self, mock_fix_incompatible_config):
|
||||
test_vllm_config = VllmConfig()
|
||||
test_vllm_config.kv_transfer_config = KVTransferConfig(
|
||||
kv_connector="MooncakeConnectorV1",
|
||||
kv_role="kv_consumer",
|
||||
)
|
||||
test_vllm_config.quant_config = SimpleNamespace(enable_c8_quant=True)
|
||||
test_vllm_config.model_config = self._make_model_config(
|
||||
total_num_attention_heads=8,
|
||||
total_num_kv_heads=8,
|
||||
)
|
||||
|
||||
ascend_config = init_ascend_config(test_vllm_config)
|
||||
|
||||
self.assertIsNotNone(ascend_config)
|
||||
|
||||
@_clean_up_ascend_config
|
||||
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")
|
||||
def test_init_ascend_config_rejects_mooncake_c8_kv_cache_producer(self, mock_fix_incompatible_config):
|
||||
test_vllm_config = VllmConfig()
|
||||
test_vllm_config.kv_transfer_config = KVTransferConfig(
|
||||
kv_connector="MooncakeConnectorV1",
|
||||
kv_role="kv_producer",
|
||||
)
|
||||
test_vllm_config.quant_config = SimpleNamespace(enable_c8_quant=True)
|
||||
test_vllm_config.model_config = self._make_model_config()
|
||||
|
||||
with self.assertRaisesRegex(ValueError, "does not support C8 KV cache quantization"):
|
||||
init_ascend_config(test_vllm_config)
|
||||
|
||||
@_clean_up_ascend_config
|
||||
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")
|
||||
def test_init_ascend_config_rejects_mooncake_c8_kv_cache_both_role(self, mock_fix_incompatible_config):
|
||||
test_vllm_config = VllmConfig()
|
||||
test_vllm_config.kv_transfer_config = KVTransferConfig(
|
||||
kv_connector="MooncakeConnectorV1",
|
||||
kv_role="kv_both",
|
||||
)
|
||||
test_vllm_config.quant_config = SimpleNamespace(enable_c8_quant=True)
|
||||
test_vllm_config.model_config = self._make_model_config()
|
||||
|
||||
with self.assertRaisesRegex(ValueError, "does not support C8 KV cache quantization"):
|
||||
init_ascend_config(test_vllm_config)
|
||||
|
||||
@_clean_up_ascend_config
|
||||
@patch("vllm_ascend.ascend_config.logger.warning")
|
||||
@patch("vllm_ascend.utils.is_310p", return_value=True)
|
||||
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")
|
||||
def test_init_ascend_config_disable_npugraph_ex_on_310p(
|
||||
self, mock_fix_incompatible_config, mock_is_310p, mock_warning
|
||||
):
|
||||
test_vllm_config = VllmConfig()
|
||||
test_vllm_config.additional_config = {
|
||||
"ascend_compilation_config": {"enable_npugraph_ex": True, "enable_static_kernel": True},
|
||||
"refresh": True,
|
||||
}
|
||||
|
||||
ascend_compilation_config = init_ascend_config(test_vllm_config).ascend_compilation_config
|
||||
|
||||
self.assertFalse(ascend_compilation_config.enable_npugraph_ex)
|
||||
self.assertFalse(ascend_compilation_config.enable_static_kernel)
|
||||
warning_messages = [call.args[0] for call in mock_warning.call_args_list]
|
||||
self.assertIn("npugraph_ex is not supported on Ascend 310P. Disabling it.", warning_messages)
|
||||
self.assertIn(
|
||||
"static kernel requires npugraph_ex, which is not supported on Ascend 310P. Disabling it.",
|
||||
warning_messages,
|
||||
)
|
||||
|
||||
@_clean_up_ascend_config
|
||||
@patch("vllm_ascend.ascend_config.logger.info_once")
|
||||
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")
|
||||
def test_migrated_config_falls_back_to_envs(self, mock_fix_incompatible_config, mock_info_once):
|
||||
test_vllm_config = VllmConfig()
|
||||
test_vllm_config.parallel_config.tensor_parallel_size = 4
|
||||
with patch.dict(
|
||||
os.environ,
|
||||
{
|
||||
"VLLM_ASCEND_ENABLE_MATMUL_ALLREDUCE": "1",
|
||||
"VLLM_ASCEND_ENABLE_FUSED_MC2": "2",
|
||||
"VLLM_ASCEND_ENABLE_MLAPO": "0",
|
||||
"VLLM_ASCEND_ENABLE_FLASHCOMM1": "1",
|
||||
"VLLM_ASCEND_FLASHCOMM2_PARALLEL_SIZE": "2",
|
||||
"MSMONITOR_USE_DAEMON": "1",
|
||||
"VLLM_ASCEND_FUSION_OP_TRANSPOSE_KV_CACHE_BY_BLOCK": "0",
|
||||
"VLLM_ASCEND_ENABLE_NZ": "2",
|
||||
},
|
||||
):
|
||||
ascend_config = init_ascend_config(test_vllm_config)
|
||||
|
||||
self.assertTrue(ascend_config.enable_matmul_allreduce)
|
||||
self.assertEqual(ascend_config.enable_fused_mc2, 2)
|
||||
self.assertFalse(ascend_config.enable_mlapo)
|
||||
self.assertTrue(ascend_config.enable_flashcomm1)
|
||||
self.assertEqual(ascend_config.enable_flashcomm2_parallel_size, 2)
|
||||
self.assertTrue(ascend_config.msmonitor_use_daemon)
|
||||
self.assertFalse(ascend_config.enable_transpose_kv_cache_by_block)
|
||||
self.assertEqual(ascend_config.weight_nz_mode, 2)
|
||||
mock_info_once.assert_any_call(
|
||||
"AscendConfig.enable_mlapo falls back to environment variable VLLM_ASCEND_ENABLE_MLAPO with value False. "
|
||||
"Please use additional_config.enable_mlapo instead, because VLLM_ASCEND_ENABLE_MLAPO will be "
|
||||
"removed in the next release."
|
||||
)
|
||||
mock_info_once.assert_any_call(
|
||||
"AscendConfig.weight_nz_mode falls back to environment variable VLLM_ASCEND_ENABLE_NZ with value 2. "
|
||||
"Please use additional_config.weight_nz_mode instead, because VLLM_ASCEND_ENABLE_NZ will be removed "
|
||||
"in the next release."
|
||||
)
|
||||
|
||||
@_clean_up_ascend_config
|
||||
@patch("vllm_ascend.ascend_config.logger.info_once")
|
||||
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")
|
||||
def test_migrated_config_skips_default_env_fallback_logs(self, mock_fix_incompatible_config, mock_info_once):
|
||||
test_vllm_config = VllmConfig()
|
||||
with patch.dict(os.environ, {}, clear=True):
|
||||
init_ascend_config(test_vllm_config)
|
||||
|
||||
fallback_logs = [
|
||||
call.args[0]
|
||||
for call in mock_info_once.call_args_list
|
||||
if "falls back to environment variable" in call.args[0]
|
||||
]
|
||||
self.assertEqual(fallback_logs, [])
|
||||
|
||||
@_clean_up_ascend_config
|
||||
@patch("vllm_ascend.ascend_config.logger.info_once")
|
||||
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")
|
||||
def test_migrated_config_overrides_envs(self, mock_fix_incompatible_config, mock_info_once):
|
||||
test_vllm_config = VllmConfig()
|
||||
test_vllm_config.additional_config = {
|
||||
"enable_matmul_allreduce": False,
|
||||
"enable_fused_mc2": 0,
|
||||
"enable_mlapo": True,
|
||||
"enable_flashcomm1": False,
|
||||
"enable_flashcomm2_parallel_size": 0,
|
||||
"msmonitor_use_daemon": False,
|
||||
"enable_transpose_kv_cache_by_block": True,
|
||||
"weight_nz_mode": 1,
|
||||
}
|
||||
with patch.dict(
|
||||
os.environ,
|
||||
{
|
||||
"VLLM_ASCEND_ENABLE_MATMUL_ALLREDUCE": "1",
|
||||
"VLLM_ASCEND_ENABLE_FUSED_MC2": "2",
|
||||
"VLLM_ASCEND_ENABLE_MLAPO": "0",
|
||||
"VLLM_ASCEND_ENABLE_FLASHCOMM1": "1",
|
||||
"VLLM_ASCEND_FLASHCOMM2_PARALLEL_SIZE": "2",
|
||||
"MSMONITOR_USE_DAEMON": "1",
|
||||
"VLLM_ASCEND_FUSION_OP_TRANSPOSE_KV_CACHE_BY_BLOCK": "0",
|
||||
"VLLM_ASCEND_ENABLE_NZ": "2",
|
||||
},
|
||||
):
|
||||
ascend_config = init_ascend_config(test_vllm_config)
|
||||
|
||||
self.assertFalse(ascend_config.enable_matmul_allreduce)
|
||||
self.assertEqual(ascend_config.enable_fused_mc2, 0)
|
||||
self.assertTrue(ascend_config.enable_mlapo)
|
||||
self.assertFalse(ascend_config.enable_flashcomm1)
|
||||
self.assertEqual(ascend_config.enable_flashcomm2_parallel_size, 0)
|
||||
self.assertFalse(ascend_config.msmonitor_use_daemon)
|
||||
self.assertTrue(ascend_config.enable_transpose_kv_cache_by_block)
|
||||
self.assertEqual(ascend_config.weight_nz_mode, 1)
|
||||
mock_info_once.assert_any_call("AscendConfig.enable_mlapo is set from additional_config with value True.")
|
||||
mock_info_once.assert_any_call("AscendConfig.weight_nz_mode is set from additional_config with value 1.")
|
||||
|
||||
@_clean_up_ascend_config
|
||||
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")
|
||||
@patch.dict(os.environ, {"VLLM_ASCEND_ENABLE_FLASHCOMM1": "1"}, clear=True)
|
||||
def test_enable_flashcomm1_config_overrides_disabled_env(self, mock_fix_incompatible_config):
|
||||
test_vllm_config = VllmConfig()
|
||||
test_vllm_config.additional_config = {"enable_flashcomm1": True}
|
||||
with patch.dict(os.environ, {"VLLM_ASCEND_ENABLE_FLASHCOMM1": "0"}, clear=True):
|
||||
ascend_config = init_ascend_config(test_vllm_config)
|
||||
self.assertTrue(ascend_config.enable_flashcomm1)
|
||||
self.assertTrue(enable_sp(test_vllm_config))
|
||||
|
||||
@_clean_up_ascend_config
|
||||
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")
|
||||
def test_enable_sp_falls_back_to_env_without_current_config(self, mock_check_and_update_config):
|
||||
clear_enable_sp()
|
||||
with (
|
||||
patch.dict(os.environ, {"VLLM_ASCEND_ENABLE_FLASHCOMM1": "1"}),
|
||||
patch("vllm.config.get_current_vllm_config", side_effect=AssertionError),
|
||||
):
|
||||
self.assertTrue(enable_sp())
|
||||
|
||||
@_clean_up_ascend_config
|
||||
@patch("vllm_ascend.utils.logger.warning_once")
|
||||
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")
|
||||
def test_flashcomm2_warning_uses_enable_flashcomm1_config(self, mock_check_and_update_config, mock_warning_once):
|
||||
test_vllm_config = VllmConfig()
|
||||
test_vllm_config.parallel_config.tensor_parallel_size = 4
|
||||
test_vllm_config.kv_transfer_config = None
|
||||
ascend_config = type(
|
||||
"MockAscendConfig",
|
||||
(),
|
||||
{
|
||||
"enable_flashcomm2_parallel_size": 2,
|
||||
"layer_sharding": None,
|
||||
"enable_flashcomm1": True,
|
||||
"finegrained_tp_config": type("MockFinegrainedTPConfig", (), {"oproj_tensor_parallel_size": 0})(),
|
||||
},
|
||||
)()
|
||||
|
||||
with patch.dict(os.environ, {"VLLM_ASCEND_ENABLE_FLASHCOMM1": "0"}):
|
||||
self.assertEqual(get_flashcomm2_config_and_validate(ascend_config, test_vllm_config), 2)
|
||||
|
||||
flashcomm1_warning = (
|
||||
"It is recommended to enable FLASHCOMM1 simultaneously when starting FLASHCOMM2 for optimal performance."
|
||||
)
|
||||
self.assertNotIn(flashcomm1_warning, [call.args[0] for call in mock_warning_once.call_args_list])
|
||||
|
||||
@_clean_up_ascend_config
|
||||
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")
|
||||
def test_get_ascend_config(self, mock_fix_incompatible_config):
|
||||
test_vllm_config = VllmConfig()
|
||||
ascend_config = init_ascend_config(test_vllm_config)
|
||||
self.assertEqual(get_ascend_config(), ascend_config)
|
||||
@@ -156,7 +386,8 @@ class TestAscendConfig(TestBase):
|
||||
get_ascend_config()
|
||||
|
||||
@_clean_up_ascend_config
|
||||
def test_clear_ascend_config(self):
|
||||
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")
|
||||
def test_clear_ascend_config(self, mock_fix_incompatible_config):
|
||||
test_vllm_config = VllmConfig()
|
||||
ascend_config = init_ascend_config(test_vllm_config)
|
||||
self.assertEqual(get_ascend_config(), ascend_config)
|
||||
@@ -165,198 +396,51 @@ class TestAscendConfig(TestBase):
|
||||
get_ascend_config()
|
||||
|
||||
@_clean_up_ascend_config
|
||||
def test_check_ascend_config_pass(self):
|
||||
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")
|
||||
def test_init_ascend_config_with_dump_config_materializes_fixed_file(self, mock_fix_incompatible_config):
|
||||
test_vllm_config = VllmConfig()
|
||||
init_ascend_config(test_vllm_config)
|
||||
check_ascend_config(test_vllm_config, False)
|
||||
dump_config = {"task": "tensor", "level": "L1", "dump_path": "/tmp/msprobe_dump"}
|
||||
test_vllm_config.additional_config = {"dump_config": dump_config}
|
||||
|
||||
test_vllm_config.additional_config = {
|
||||
"torchair_graph_config": {
|
||||
"enabled": True,
|
||||
},
|
||||
"refresh": True
|
||||
}
|
||||
init_ascend_config(test_vllm_config)
|
||||
check_ascend_config(test_vllm_config, False)
|
||||
|
||||
test_vllm_config.additional_config = {
|
||||
"torchair_graph_config": {
|
||||
"enabled": False,
|
||||
},
|
||||
"refresh": True
|
||||
}
|
||||
init_ascend_config(test_vllm_config)
|
||||
check_ascend_config(test_vllm_config, False)
|
||||
ascend_config = init_ascend_config(test_vllm_config)
|
||||
self.assertIsNotNone(ascend_config.dump_config_path)
|
||||
assert ascend_config.dump_config_path is not None
|
||||
expected_path = os.path.join(os.getcwd(), ".vllm_ascend", "msprobe", "msprobe_dump_config.json")
|
||||
self.assertEqual(ascend_config.dump_config_path, expected_path)
|
||||
self.assertTrue(os.path.exists(ascend_config.dump_config_path))
|
||||
with open(ascend_config.dump_config_path, encoding="utf-8") as file:
|
||||
persisted = json.load(file)
|
||||
self.assertEqual(persisted, dump_config)
|
||||
|
||||
@_clean_up_ascend_config
|
||||
def test_check_ascend_config_wrong_case(self):
|
||||
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")
|
||||
def test_init_ascend_config_dump_config_and_path_conflict(self, mock_fix_incompatible_config):
|
||||
test_vllm_config = VllmConfig()
|
||||
|
||||
# torchair + eager mode
|
||||
with self.assertRaises(RuntimeError):
|
||||
test_vllm_config.additional_config = {
|
||||
"torchair_graph_config": {
|
||||
"enabled": True,
|
||||
},
|
||||
"refresh": True
|
||||
}
|
||||
test_vllm_config.additional_config = {"dump_config_path": "/tmp/config.json", "dump_config": {"task": "tensor"}}
|
||||
with self.assertRaises(ValueError):
|
||||
init_ascend_config(test_vllm_config)
|
||||
enforce_eager = True
|
||||
check_ascend_config(test_vllm_config, enforce_eager)
|
||||
# torchair + non deepseek model
|
||||
with self.assertRaises(NotImplementedError):
|
||||
test_vllm_config.additional_config = {
|
||||
"torchair_graph_config": {
|
||||
"enabled": True,
|
||||
},
|
||||
"refresh": True
|
||||
}
|
||||
model_path = os.path.join(os.path.dirname(__file__), "fake_weight")
|
||||
fake_model_config = ModelConfig(model=model_path)
|
||||
fake_model_config.hf_config = PretrainedConfig()
|
||||
fake_model_config.hf_config.model_type = "llama"
|
||||
test_vllm_config.model_config = fake_model_config
|
||||
init_ascend_config(test_vllm_config)
|
||||
check_ascend_config(test_vllm_config, False)
|
||||
|
||||
def test_check_torchair_supported(self):
|
||||
test_cases = [('deepseek_v3', True), ('PanguProMoE', True),
|
||||
('qwen', True), ('llama', False)]
|
||||
for model_type, expected_output in test_cases:
|
||||
self.assertEqual(_check_torchair_supported(model_type),
|
||||
expected_output)
|
||||
|
||||
@_clean_up_ascend_config
|
||||
def test_ascend_config_load_error(self):
|
||||
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")
|
||||
def test_init_ascend_config_dump_config_type_validation(self, mock_fix_incompatible_config):
|
||||
test_vllm_config = VllmConfig()
|
||||
# graph_batch_sizes should be list.
|
||||
with self.assertRaises(TypeError):
|
||||
test_vllm_config.additional_config = {
|
||||
"torchair_graph_config": {
|
||||
"graph_batch_sizes": "fake_size",
|
||||
},
|
||||
"refresh": True
|
||||
}
|
||||
test_vllm_config.additional_config = {"dump_config": "/tmp/config.json"}
|
||||
with self.assertRaises(ValueError):
|
||||
init_ascend_config(test_vllm_config)
|
||||
|
||||
# use_cached_graph should not be enabled without torchair graph mode
|
||||
with self.assertRaises(RuntimeError):
|
||||
test_vllm_config.additional_config = {
|
||||
"torchair_graph_config": {
|
||||
"enabled": False,
|
||||
"use_cached_graph": True,
|
||||
},
|
||||
"refresh": True
|
||||
@_clean_up_ascend_config
|
||||
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")
|
||||
def test_init_ascend_config_recreates_for_new_vllm_config(self, mock_fix_incompatible_config):
|
||||
first_vllm_config = VllmConfig()
|
||||
first_vllm_config.additional_config = {
|
||||
"ascend_compilation_config": {
|
||||
"enable_npugraph_ex": False,
|
||||
}
|
||||
init_ascend_config(test_vllm_config)
|
||||
}
|
||||
first_ascend_config = init_ascend_config(first_vllm_config)
|
||||
self.assertFalse(first_ascend_config.ascend_compilation_config.enable_npugraph_ex)
|
||||
|
||||
# use_cached_kv_cache_bytes should not be enabled without torchair graph mode
|
||||
with self.assertRaises(RuntimeError):
|
||||
test_vllm_config.additional_config = {
|
||||
"torchair_graph_config": {
|
||||
"enabled": False,
|
||||
"use_cached_kv_cache_bytes": True,
|
||||
},
|
||||
"refresh": True
|
||||
}
|
||||
init_ascend_config(test_vllm_config)
|
||||
|
||||
# graph_batch_sizes should not be set without torchair graph mode
|
||||
with self.assertRaises(RuntimeError):
|
||||
test_vllm_config.additional_config = {
|
||||
"torchair_graph_config": {
|
||||
"enabled": False,
|
||||
"graph_batch_sizes": [1, 2, 4],
|
||||
},
|
||||
"refresh": True
|
||||
}
|
||||
init_ascend_config(test_vllm_config)
|
||||
|
||||
# use_cached_kv_cache_bytes is valid only when torchair graph mode and use_cached_graph are enabled
|
||||
with self.assertRaises(RuntimeError):
|
||||
test_vllm_config.additional_config = {
|
||||
"torchair_graph_config": {
|
||||
"enabled": True,
|
||||
"use_cached_graph": False,
|
||||
"use_cached_kv_cache_bytes": True,
|
||||
},
|
||||
"refresh": True
|
||||
}
|
||||
init_ascend_config(test_vllm_config)
|
||||
|
||||
# graph_batch_sizes_init should not be enabled without torchair graph mode
|
||||
with self.assertRaises(RuntimeError):
|
||||
test_vllm_config.additional_config = {
|
||||
"torchair_graph_config": {
|
||||
"enabled": False,
|
||||
"graph_batch_sizes_init": True,
|
||||
},
|
||||
"refresh": True
|
||||
}
|
||||
init_ascend_config(test_vllm_config)
|
||||
|
||||
# enable_multistream_mla should not be enabled without torchair graph mode
|
||||
with self.assertRaises(RuntimeError):
|
||||
test_vllm_config.additional_config = {
|
||||
"torchair_graph_config": {
|
||||
"enabled": False,
|
||||
"enable_multistream_mla": True,
|
||||
},
|
||||
"refresh": True
|
||||
}
|
||||
init_ascend_config(test_vllm_config)
|
||||
|
||||
# mode should not be configured without torchair graph mode
|
||||
with self.assertRaises(RuntimeError):
|
||||
test_vllm_config.additional_config = {
|
||||
"torchair_graph_config": {
|
||||
"enabled": False,
|
||||
"mode": 'max-autotune',
|
||||
},
|
||||
"refresh": True
|
||||
}
|
||||
init_ascend_config(test_vllm_config)
|
||||
|
||||
# enable_kv_nz should not be enabled without torchair graph mode
|
||||
with self.assertRaises(RuntimeError):
|
||||
test_vllm_config.additional_config = {
|
||||
"torchair_graph_config": {
|
||||
"enabled": False,
|
||||
"enable_kv_nz": True,
|
||||
},
|
||||
"refresh": True
|
||||
}
|
||||
init_ascend_config(test_vllm_config)
|
||||
|
||||
with self.assertRaises(AssertionError):
|
||||
test_vllm_config.additional_config = {
|
||||
"lmhead_tensor_parallel_size": 2,
|
||||
"refresh": True
|
||||
}
|
||||
test_vllm_config.parallel_config = ParallelConfig(
|
||||
data_parallel_size=4, tensor_parallel_size=2)
|
||||
init_ascend_config(test_vllm_config)
|
||||
|
||||
with self.assertRaises(AssertionError):
|
||||
test_vllm_config.additional_config = {
|
||||
"torchair_graph_config": {
|
||||
"enabled": True,
|
||||
},
|
||||
"oproj_tensor_parallel_size": 2,
|
||||
"refresh": True
|
||||
}
|
||||
test_vllm_config.parallel_config = ParallelConfig(
|
||||
data_parallel_size=4, tensor_parallel_size=2)
|
||||
init_ascend_config(test_vllm_config)
|
||||
|
||||
with self.assertRaises(AssertionError):
|
||||
test_vllm_config.additional_config = {
|
||||
"torchair_graph_config": {
|
||||
"enabled": False,
|
||||
},
|
||||
"oproj_tensor_parallel_size": 2,
|
||||
"refresh": True
|
||||
}
|
||||
test_vllm_config.parallel_config = ParallelConfig(
|
||||
data_parallel_size=4, tensor_parallel_size=1)
|
||||
init_ascend_config(test_vllm_config)
|
||||
second_vllm_config = VllmConfig()
|
||||
second_ascend_config = init_ascend_config(second_vllm_config)
|
||||
self.assertIsNot(first_ascend_config, second_ascend_config)
|
||||
self.assertTrue(second_ascend_config.ascend_compilation_config.enable_npugraph_ex)
|
||||
|
||||
Reference in New Issue
Block a user