Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 9 additions & 10 deletions tests/ut/ops/test_layernorm.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@
from vllm.config import set_current_vllm_config
from vllm.model_executor.layers.layernorm import RMSNorm

from vllm_ascend.utils import AscendDeviceType, enable_custom_op
from vllm_ascend.utils import enable_custom_op
from vllm_ascend.utils import is_310p as is_310p_hw

enable_custom_op()
Expand Down Expand Up @@ -39,8 +39,8 @@ def default_vllm_config():
with set_current_vllm_config(mock_config):
yield mock_config

@pytest.mark.skip(
"Skip as register_kernels has NPU SocName checking in CANN 8.5.0.")

@pytest.mark.skip("Skip as register_kernels has NPU SocName checking in CANN 8.5.0.")
Comment thread
Tflowers-0129 marked this conversation as resolved.
@pytest.mark.skipif(is_310p_hw(), reason="non_310P device unittest case.")
@pytest.mark.parametrize("residual", [None, torch.randn(4, 8, dtype=torch.float32)])
@patch("torch_npu.npu_rms_norm", side_effect=mock_rms_norm)
Expand Down Expand Up @@ -68,19 +68,18 @@ def test_RMSNorm_forward(
@pytest.mark.skipif(not is_310p_hw(), reason="310P device unittest case.")
@pytest.mark.parametrize("residual", [None, torch.randn(4, 8, dtype=torch.float16)])
@patch("torch_npu.npu_rms_norm", side_effect=mock_rms_norm)
def test_RMSNorm_forward_310p(
mock_rmsnorm, residual, dummy_tensor, default_vllm_config
):
@patch("torch_npu.npu_add_rms_norm", side_effect=mock_add_rms_norm)
def test_RMSNorm_forward_310p(mock_add_rmsnorm, mock_rmsnorm, residual, dummy_tensor, default_vllm_config):
Comment thread
Tflowers-0129 marked this conversation as resolved.
layer = RMSNorm(hidden_size=8, eps=1e-05)
if residual is not None:
out_x, out_residual = layer.forward_oot(dummy_tensor, residual)
expected_out_residual = dummy_tensor + residual
expected_out_x = expected_out_residual + 1
mock_rmsnorm.assert_called_once()
expected_out_x = 2 * dummy_tensor
expected_out_residual = 2 * residual
mock_add_rmsnorm.assert_called_once()
assert torch.allclose(out_x, expected_out_x)
assert torch.allclose(out_residual, expected_out_residual)
else:
out_x = layer.forward_oot(dummy_tensor, residual)
expected_out_x = dummy_tensor + 1
mock_rmsnorm.assert_called_once()
assert torch.allclose(out_x, expected_out_x)
assert torch.allclose(out_x, expected_out_x)
10 changes: 3 additions & 7 deletions vllm_ascend/_310p/ops/layernorm.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,13 +11,9 @@ def forward_oot(
residual: torch.Tensor | None = None,
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
if residual is not None:
if x is None or x.numel() == 0 or x.shape[-1] == 0:
x = residual
else:
x = x + residual

residual = x
x, _ = torch_npu.npu_rms_norm(x, self.weight, self.variance_epsilon)
x, _, residual = torch_npu.npu_add_rms_norm(x, residual, self.weight, self.variance_epsilon)
if self.bias is not None:
x.add_(self.bias)
return x, residual
Comment thread
Tflowers-0129 marked this conversation as resolved.

x, _ = torch_npu.npu_rms_norm(x, self.weight, self.variance_epsilon)
Expand Down
61 changes: 0 additions & 61 deletions vllm_ascend/_310p/ops/mm_encoder_attention.py

This file was deleted.

2 changes: 0 additions & 2 deletions vllm_ascend/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -628,13 +628,11 @@ def register_ascend_customop(vllm_config: VllmConfig | None = None):
from vllm_ascend._310p.fused_moe.fused_moe import AscendFusedMoE310, AscendSharedFusedMoE310
from vllm_ascend._310p.ops.activation import AscendSiluAndMul310
from vllm_ascend._310p.ops.layernorm import AscendGemmaRMSNorm310, AscendRMSNorm310
from vllm_ascend._310p.ops.mm_encoder_attention import AscendMMEncoderAttention310
from vllm_ascend._310p.ops.rotary_embedding import AscendRotaryEmbedding310

REGISTERED_ASCEND_OPS.update(
{
"SiluAndMul": AscendSiluAndMul310,
"MMEncoderAttention": AscendMMEncoderAttention310,
"RotaryEmbedding": AscendRotaryEmbedding310,
"RMSNorm": AscendRMSNorm310,
"GemmaRMSNorm": AscendGemmaRMSNorm310,
Expand Down