Files
enginex-ascend-910-vllm/tests/ut/quantization/methods/test_w8a8_mxfp8.py
Sun Ruoxi 7f8a1b1f7a init v0.23.0
Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
2026-08-27 15:11:51 +08:00

170 lines
8.3 KiB
Python

from unittest.mock import Mock, patch
import torch
import torch.nn as nn
from tests.ut.base import TestBase
from tests.ut.quantization.conftest_quantization import (
create_mock_ascend_config,
create_mock_vllm_config,
create_mxfp_moe_layer,
)
from vllm_ascend.quantization.methods.w8a8_mxfp8 import (
AscendW8A8MXFP8DynamicFusedMoEMethod,
AscendW8A8MXFP8DynamicLinearMethod,
)
class TestAscendW8A8MXFP8LinearMethod(TestBase):
@patch("vllm_ascend.quantization.methods.w8a8_mxfp8.ensure_mxfp8_linear_available")
@patch("vllm_ascend.quantization.methods.w8a8_mxfp8.get_current_vllm_config")
def setUp(self, mock_vllm, mock_ensure):
mock_vllm.return_value = create_mock_vllm_config()
mock_ensure.return_value = None
self.scheme = AscendW8A8MXFP8DynamicLinearMethod()
def test_get_weight_various_input_sizes(self):
sizes = [(128, 64), (512, 256), (1024, 512)]
for input_size, output_size in sizes:
result = self.scheme.get_weight(input_size, output_size, torch.bfloat16)
self.assertEqual(result["weight"].shape, (output_size, input_size))
self.assertEqual(result["weight"].dtype, torch.float8_e4m3fn)
def test_get_pergroup_param_group_size_variations(self):
group_sizes = [16, 32, 64, 128]
for gs in group_sizes:
self.scheme.group_size = gs
result = self.scheme.get_pergroup_param(256, 128, torch.bfloat16)
self.assertEqual(result["weight_scale"].shape, (128, 256 // gs))
self.assertEqual(result["weight_scale"].dtype, torch.uint8)
def test_process_weights_stores_original_shapes(self):
layer = nn.Module()
layer.weight = nn.Parameter(torch.randn(128, 256).to(torch.float8_e4m3fn), requires_grad=False)
layer.weight_scale = nn.Parameter(torch.randint(0, 255, (128, 8), dtype=torch.uint8), requires_grad=False)
self.scheme.process_weights_after_loading(layer)
self.assertTrue(hasattr(layer, "_mxfp8_original_shapes"))
self.assertEqual(layer._mxfp8_original_shapes["weight"], (128, 256))
self.assertTrue(layer._mxfp8_transformed)
self.assertEqual(layer.weight_scale.shape, (4, 128, 2))
self.assertTrue(layer.weight.data.is_contiguous())
self.assertTrue(layer.weight_scale.data.is_contiguous())
def test_restore_after_process_returns_original_shape(self):
layer = nn.Module()
layer.weight = nn.Parameter(torch.randn(128, 256).to(torch.float8_e4m3fn), requires_grad=False)
layer.weight_scale = nn.Parameter(torch.randint(0, 255, (128, 8), dtype=torch.uint8), requires_grad=False)
original_weight_shape = layer.weight.shape
original_scale_shape = layer.weight_scale.shape
self.scheme.process_weights_after_loading(layer)
self.scheme.restore_weights_for_rl_loading(layer)
self.assertEqual(layer.weight.shape, original_weight_shape)
self.assertEqual(layer.weight_scale.shape, original_scale_shape)
self.assertFalse(layer._mxfp8_transformed)
@patch("vllm_ascend.quantization.methods.w8a8_mxfp8.torch_npu")
def test_apply(self, mock_torch_npu):
from vllm_ascend.device.mxfp_compat import FLOAT8_E8M0FNU_DTYPE
dynamic_scale = torch.randint(0, 255, (32, 8), dtype=torch.uint8)
mock_torch_npu.npu_dynamic_mx_quant.return_value = (
torch.randint(0, 255, (32, 256), dtype=torch.uint8),
dynamic_scale,
)
mock_torch_npu.npu_quant_matmul.return_value = torch.randn(32, 128, dtype=torch.float16)
layer = nn.Module()
layer.weight = nn.Parameter(torch.randn(256, 128).to(torch.float8_e4m3fn), requires_grad=False)
layer.weight_scale = nn.Parameter(torch.randint(0, 255, (4, 128, 2), dtype=torch.uint8), requires_grad=False)
x = torch.randn(32, 1, 256, dtype=torch.float16)
bias = torch.randn(128, dtype=torch.float16)
output = self.scheme.apply(layer, x, bias)
self.assertEqual(output.shape, (32, 1, 128))
call_kwargs = mock_torch_npu.npu_quant_matmul.call_args.kwargs
self.assertEqual(call_kwargs["bias"].dtype, torch.float32)
self.assertEqual(call_kwargs["group_sizes"], [1, 1, self.scheme.group_size])
self.assertEqual(call_kwargs["scale_dtype"], FLOAT8_E8M0FNU_DTYPE)
self.assertEqual(call_kwargs["output_dtype"], torch.float16)
class TestAscendW8A8MXFP8MoEMethod(TestBase):
num_experts = 8
hidden_size = 128
intermediate_size = 256
@patch("vllm_ascend.quantization.methods.w8a8_mxfp8.ensure_mxfp8_moe_available")
@patch("vllm_ascend.quantization.methods.w8a8_mxfp8.get_current_vllm_config")
@patch("vllm_ascend.quantization.methods.w8a8_mxfp8.get_ascend_config")
def setUp(self, mock_ascend, mock_vllm, mock_ensure):
mock_vllm.return_value = create_mock_vllm_config()
mock_ascend.return_value = create_mock_ascend_config()
mock_ensure.return_value = None
self.scheme = AscendW8A8MXFP8DynamicFusedMoEMethod()
def test_get_weight_various_expert_counts(self):
for num_experts in [4, 8, 16]:
result = self.scheme.get_weight(num_experts, self.intermediate_size, self.hidden_size, torch.bfloat16)
self.assertEqual(result["w13_weight"].shape[0], num_experts)
self.assertEqual(result["w2_weight"].dtype, torch.float8_e4m3fn)
def test_get_dynamic_quant_param_dtype_uint8(self):
result = self.scheme.get_dynamic_quant_param(
self.num_experts, self.intermediate_size, self.hidden_size, torch.bfloat16
)
self.assertEqual(result["w13_weight_scale"].shape, (8, 512, 4))
self.assertEqual(result["w2_weight_scale"].dtype, torch.uint8)
def test_process_weights_stores_original_shapes(self):
layer = create_mxfp_moe_layer(
num_experts=self.num_experts, hidden_size=self.hidden_size, intermediate_size=self.intermediate_size
)
original_shape = layer.w13_weight.shape
self.scheme.process_weights_after_loading(layer)
self.assertTrue(hasattr(layer, "_mxfp8_original_shapes"))
self.assertIn("w13_weight", layer._mxfp8_original_shapes)
self.assertEqual(layer.w13_weight.shape, (original_shape[0], original_shape[2], original_shape[1]))
self.assertFalse(layer.w13_weight.data.is_contiguous())
self.assertFalse(layer.w2_weight.data.is_contiguous())
self.assertFalse(layer.w13_weight_scale.data.is_contiguous())
self.assertFalse(layer.w2_weight_scale.data.is_contiguous())
def test_restore_weights_for_rl_loading(self):
layer = create_mxfp_moe_layer(
num_experts=self.num_experts, hidden_size=self.hidden_size, intermediate_size=self.intermediate_size
)
original_w13_shape = layer.w13_weight.shape
self.scheme.process_weights_after_loading(layer)
self.assertNotEqual(layer.w13_weight.shape, original_w13_shape)
self.scheme.restore_weights_for_rl_loading(layer)
self.assertEqual(layer.w13_weight.shape, original_w13_shape)
@patch("vllm_ascend.quantization.methods.w8a8_mxfp8._EXTRA_CTX")
@patch("vllm_ascend.quantization.methods.w8a8_mxfp8.select_experts")
def test_apply_full_params(self, mock_select, mock_ctx):
tokens = 4
layer = create_mxfp_moe_layer(
num_experts=self.num_experts, hidden_size=self.hidden_size, intermediate_size=self.intermediate_size
)
self.scheme.process_weights_after_loading(layer)
layer.swiglu_limit = 1000000
x = torch.randn(tokens, self.hidden_size, dtype=torch.bfloat16)
router_logits = torch.randn(tokens, self.num_experts, dtype=torch.float32)
topk_weights = torch.randn(tokens, 2)
topk_ids = torch.randint(0, self.num_experts, (tokens, 2))
mock_select.return_value = (topk_weights, topk_ids)
mock_comm = Mock()
mock_comm.fused_experts.return_value = torch.randn(tokens, self.hidden_size)
mock_ctx.moe_comm_method = mock_comm
mock_ctx.moe_comm_type = Mock()
self.scheme.apply(
layer,
x,
router_logits,
top_k=2,
renormalize=True,
num_experts=self.num_experts,
activation="silu",
pertoken_scale=torch.randn(tokens),
)
mock_select.assert_called_once()
mock_comm.fused_experts.assert_called_once()