@@ -1,378 +1,347 @@
|
||||
import math
|
||||
import unittest
|
||||
from unittest.mock import MagicMock, PropertyMock, patch
|
||||
#
|
||||
# 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.
|
||||
# This file is a part of the vllm-ascend project.
|
||||
#
|
||||
|
||||
|
||||
import inspect
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from transformers.configuration_utils import PretrainedConfig
|
||||
from vllm.config import ModelConfig, VllmConfig
|
||||
from vllm.model_executor.layers.rotary_embedding import (
|
||||
DeepseekScalingRotaryEmbedding, RotaryEmbedding)
|
||||
from vllm.model_executor.layers.rotary_embedding import RotaryEmbedding, YaRNScalingRotaryEmbedding
|
||||
|
||||
from tests.ut.base import TestBase
|
||||
from vllm_ascend.ascend_forward_context import set_ascend_forward_context
|
||||
from vllm_ascend.ops.rotary_embedding import _custom_rotary_embedding_enabled
|
||||
from vllm_ascend.ops.rotary_embedding import AscendRotaryEmbedding, AscendYaRNRotaryEmbedding
|
||||
|
||||
MODEL = "Qwen3-0.6B"
|
||||
MAX_NUM_BATCHED_TOKEND = 10000
|
||||
HEAD_SIZE = 64
|
||||
ROTARY_DIM = 64
|
||||
MAX_POS = 2048
|
||||
BASE = 10000.0
|
||||
DTYPE = torch.bfloat16
|
||||
SEQ_LEN = 4
|
||||
NUM_HEADS = 2
|
||||
|
||||
|
||||
class TestCustomRotaryEmbeddingEnabled(unittest.TestCase):
|
||||
|
||||
def setUp(self):
|
||||
# Common setup for tests
|
||||
self.positions = torch.tensor([1, 2, 3])
|
||||
self.query = torch.randn(3, 4, dtype=torch.float16)
|
||||
self.key = torch.randn(3, 4, dtype=torch.float16)
|
||||
self.head_size = 32
|
||||
self.cos_sin_cache = torch.randn(3, 4)
|
||||
|
||||
# Mock self object for rope_forward_oot
|
||||
self.mock_self = MagicMock()
|
||||
self.mock_self.head_size = self.head_size
|
||||
self.mock_self.cos_sin_cache = self.cos_sin_cache
|
||||
self.mock_self.is_neox_style = True
|
||||
self.mock_self.forward_native.return_value = (self.query, self.key)
|
||||
|
||||
def test_custom_rotary_embedding_enabled(self):
|
||||
# Test when all conditions are True
|
||||
with patch('vllm_ascend.ops.rotary_embedding.enable_custom_op',
|
||||
return_value=True):
|
||||
result = _custom_rotary_embedding_enabled(self.query, True,
|
||||
self.head_size)
|
||||
self.assertTrue(result)
|
||||
|
||||
# Test when dtype is not float16
|
||||
with patch('vllm_ascend.ops.rotary_embedding.enable_custom_op',
|
||||
return_value=True):
|
||||
query = self.query.to(torch.float32)
|
||||
result = _custom_rotary_embedding_enabled(query, True,
|
||||
self.head_size)
|
||||
self.assertFalse(result)
|
||||
|
||||
# Test when neox_style is False
|
||||
with patch('vllm_ascend.ops.rotary_embedding.enable_custom_op',
|
||||
return_value=True):
|
||||
result = _custom_rotary_embedding_enabled(self.query, False,
|
||||
self.head_size)
|
||||
self.assertFalse(result)
|
||||
|
||||
# Test when head_size is not divisible by 32
|
||||
with patch('vllm_ascend.ops.rotary_embedding.enable_custom_op',
|
||||
return_value=True):
|
||||
result = _custom_rotary_embedding_enabled(self.query, True,
|
||||
self.head_size + 1)
|
||||
self.assertFalse(result)
|
||||
|
||||
# Test when custom op is disabled
|
||||
with patch('vllm_ascend.ops.rotary_embedding.enable_custom_op',
|
||||
return_value=False):
|
||||
result = _custom_rotary_embedding_enabled(self.query, True,
|
||||
self.head_size)
|
||||
self.assertFalse(result)
|
||||
def _make_tensors(seq_len=SEQ_LEN, num_heads=NUM_HEADS, head_size=HEAD_SIZE):
|
||||
positions = torch.arange(seq_len, dtype=torch.long)
|
||||
query = torch.randn(seq_len, num_heads * head_size)
|
||||
key = torch.randn(seq_len, num_heads * head_size)
|
||||
return positions, query, key
|
||||
|
||||
|
||||
class TestAscendRotaryEmbedding(unittest.TestCase):
|
||||
def check_parent_init_signature_has_not_changed(parent_func, child_func):
|
||||
parent_sig = inspect.signature(parent_func)
|
||||
parent_params = set(parent_sig.parameters) - {"self"}
|
||||
|
||||
def setUp(self):
|
||||
# Common setup for tests
|
||||
self.positions = torch.tensor([1, 2, 3])
|
||||
self.query = torch.randn(3, 1, 32, dtype=torch.float16)
|
||||
self.key = torch.randn(3, 1, 32, dtype=torch.float16)
|
||||
self.head_size = 32
|
||||
self.rotary_dim = self.head_size
|
||||
self.max_position = 16
|
||||
self.rope_theta = 10000
|
||||
self.is_neox_style = True
|
||||
self.cos_sin_cache = torch.randn(3, 1, 32)
|
||||
self.layer = RotaryEmbedding(self.head_size, self.rotary_dim,
|
||||
self.max_position, self.rope_theta,
|
||||
self.is_neox_style, torch.float16)
|
||||
child_sig = inspect.signature(child_func)
|
||||
child_params = set(child_sig.parameters) - {"self"}
|
||||
|
||||
# Mock self object for rope_forward_oot
|
||||
self.mock_self = MagicMock()
|
||||
self.mock_self.head_size = self.head_size
|
||||
self.mock_self.cos_sin_cache = self.cos_sin_cache
|
||||
self.mock_self.is_neox_style = self.is_neox_style
|
||||
added = parent_params - child_params
|
||||
removed = child_params - parent_params
|
||||
|
||||
@patch('torch.ops._C_ascend')
|
||||
@patch('vllm_ascend.ops.rotary_embedding.is_310p', return_value=False)
|
||||
@patch('vllm_ascend.ops.rotary_embedding._custom_rotary_embedding_enabled',
|
||||
return_value=True)
|
||||
@patch('torch.ops._npu_rotary_embedding')
|
||||
@patch('vllm.config.ModelConfig.__post_init__', MagicMock())
|
||||
@patch('vllm.config.VllmConfig.__post_init__', MagicMock())
|
||||
@patch('vllm.distributed.parallel_state._DP', MagicMock(world_size=1))
|
||||
@patch('vllm.distributed.parallel_state._TP', MagicMock(world_size=1))
|
||||
def test_rope_forward_oot_custom_kernel(self, mock_rotary_embedding,
|
||||
mock_custom_enabled, mock_is_310p,
|
||||
mock__c):
|
||||
mock_config = MagicMock()
|
||||
mock_config.torchair_graph_config.enabled = False
|
||||
|
||||
# Setup mock for custom kernel path
|
||||
|
||||
mock__c.rotary_embedding.return_value = self.query, self.key
|
||||
vllm_config = VllmConfig()
|
||||
model_config = ModelConfig(MODEL,
|
||||
tokenizer=MODEL,
|
||||
max_model_len=MAX_NUM_BATCHED_TOKEND)
|
||||
model_config.hf_config = PretrainedConfig()
|
||||
vllm_config.model_config = model_config
|
||||
with set_ascend_forward_context(None, vllm_config):
|
||||
result_q, result_k = self.layer.forward(self.positions, self.query,
|
||||
self.key)
|
||||
|
||||
mock__c.rotary_embedding.assert_called_once()
|
||||
self.assertEqual(result_q.shape, self.query.shape)
|
||||
self.assertEqual(result_k.shape, self.key.shape)
|
||||
|
||||
@patch('vllm_ascend.ops.rotary_embedding._custom_rotary_embedding_enabled',
|
||||
return_value=False)
|
||||
@patch('torch_npu._npu_rotary_embedding')
|
||||
@patch('vllm.config.ModelConfig.__post_init__', MagicMock())
|
||||
@patch('vllm.config.VllmConfig.__post_init__', MagicMock())
|
||||
@patch('vllm.distributed.parallel_state._DP', MagicMock(world_size=1))
|
||||
@patch('vllm.distributed.parallel_state._TP', MagicMock(world_size=1))
|
||||
def test_rope_forward_oot_contiguous(self, mock_npu_rotary,
|
||||
mock_custom_enabled):
|
||||
mock_config = MagicMock()
|
||||
mock_config.torchair_graph_config.enabled = False
|
||||
|
||||
# Test contiguous path when custom is disabled
|
||||
non_contig_query = self.query.transpose(0, 1)
|
||||
non_contig_key = self.key.transpose(0, 1)
|
||||
vllm_config = VllmConfig()
|
||||
model_config = ModelConfig(MODEL,
|
||||
tokenizer=MODEL,
|
||||
max_model_len=MAX_NUM_BATCHED_TOKEND)
|
||||
model_config.hf_config = PretrainedConfig()
|
||||
vllm_config.model_config = model_config
|
||||
with set_ascend_forward_context(None, vllm_config):
|
||||
result_q, result_k = self.layer.forward(self.positions,
|
||||
non_contig_query,
|
||||
non_contig_key)
|
||||
|
||||
mock_npu_rotary.assert_called_once()
|
||||
self.assertEqual(result_q.shape, non_contig_query.shape)
|
||||
self.assertEqual(result_k.shape, non_contig_key.shape)
|
||||
|
||||
@patch('vllm.config.ModelConfig.__post_init__', MagicMock())
|
||||
@patch('vllm.config.VllmConfig.__post_init__', MagicMock())
|
||||
@patch('vllm.distributed.parallel_state._DP', MagicMock(world_size=1))
|
||||
@patch('vllm.distributed.parallel_state._TP', MagicMock(world_size=1))
|
||||
def test_rope_forward_oot_with_offsets(self):
|
||||
mock_config = MagicMock()
|
||||
mock_config.torchair_graph_config.enabled = False
|
||||
|
||||
# Test that NotImplementedError is raised when offsets is provided
|
||||
offsets = torch.tensor([1, 2, 3])
|
||||
with self.assertRaises(NotImplementedError):
|
||||
vllm_config = VllmConfig()
|
||||
model_config = ModelConfig(MODEL,
|
||||
tokenizer=MODEL,
|
||||
max_model_len=MAX_NUM_BATCHED_TOKEND)
|
||||
model_config.hf_config = PretrainedConfig()
|
||||
vllm_config.model_config = model_config
|
||||
with set_ascend_forward_context(None, vllm_config):
|
||||
self.layer.forward(self.positions, self.query, self.key,
|
||||
offsets)
|
||||
|
||||
@patch('vllm_ascend.ops.rotary_embedding._custom_rotary_embedding_enabled',
|
||||
return_value=False)
|
||||
@patch('torch_npu._npu_rotary_embedding')
|
||||
@patch('vllm.config.ModelConfig.__post_init__', MagicMock())
|
||||
@patch('vllm.config.VllmConfig.__post_init__', MagicMock())
|
||||
@patch('vllm.distributed.parallel_state._DP', MagicMock(world_size=1))
|
||||
@patch('vllm.distributed.parallel_state._TP', MagicMock(world_size=1))
|
||||
def test_rope_forward_oot_neox_style_override(self, mock_npu_rotary,
|
||||
mock_custom_enabled):
|
||||
mock_config = MagicMock()
|
||||
mock_config.torchair_graph_config.enabled = False
|
||||
|
||||
# Test neox_style override
|
||||
vllm_config = VllmConfig()
|
||||
model_config = ModelConfig(MODEL,
|
||||
tokenizer=MODEL,
|
||||
max_model_len=MAX_NUM_BATCHED_TOKEND)
|
||||
model_config.hf_config = PretrainedConfig()
|
||||
vllm_config.model_config = model_config
|
||||
with set_ascend_forward_context(None, vllm_config):
|
||||
result_q, result_k = self.layer.forward(
|
||||
self.positions,
|
||||
self.query,
|
||||
self.key,
|
||||
is_neox_style_override=False)
|
||||
# Check that neox_style=False was passed to the NPU function
|
||||
args, kwargs = mock_npu_rotary.call_args
|
||||
self.assertFalse(args[-1])
|
||||
|
||||
@patch('vllm_ascend.ops.rotary_embedding._custom_rotary_embedding_enabled',
|
||||
return_value=False)
|
||||
@patch('torch_npu._npu_rotary_embedding')
|
||||
@patch('vllm.config.ModelConfig.__post_init__', MagicMock())
|
||||
@patch('vllm.config.VllmConfig.__post_init__', MagicMock())
|
||||
@patch('vllm.distributed.parallel_state._DP', MagicMock(world_size=1))
|
||||
@patch('vllm.distributed.parallel_state._TP', MagicMock(world_size=1))
|
||||
def test_rope_forward_oot_rotary_dim_less_than_head_size(
|
||||
self, mock_npu_rotary, mock_custom_enabled):
|
||||
mock_config = MagicMock()
|
||||
mock_config.torchair_graph_config.enabled = False
|
||||
|
||||
# test case when rotary_dim < head_size
|
||||
org_rotary_dim = self.layer.rotary_dim
|
||||
self.layer.rotary_dim = self.layer.head_size // 2
|
||||
|
||||
vllm_config = VllmConfig()
|
||||
model_config = ModelConfig(MODEL,
|
||||
tokenizer=MODEL,
|
||||
max_model_len=MAX_NUM_BATCHED_TOKEND)
|
||||
model_config.hf_config = PretrainedConfig()
|
||||
vllm_config.model_config = model_config
|
||||
with set_ascend_forward_context(None, vllm_config):
|
||||
result_q, result_k = self.layer.forward(self.positions, self.query,
|
||||
self.key)
|
||||
|
||||
mock_npu_rotary.assert_called_once()
|
||||
self.assertEqual(result_q.shape, self.query.shape)
|
||||
self.assertEqual(result_k.shape, self.key.shape)
|
||||
|
||||
# restore rotary_dim
|
||||
self.layer.rotary_dim = org_rotary_dim
|
||||
assert not added, (
|
||||
f"{parent_func.__name__} added new parameter(s): {added}. "
|
||||
f"Check whether {child_func.__name__} needs to forward them."
|
||||
)
|
||||
assert not removed, (
|
||||
f"{parent_func.__name__} removed parameter(s): {removed}. "
|
||||
f"Check whether {child_func.__name__} needs to forward them."
|
||||
)
|
||||
|
||||
|
||||
class MockRopeModule:
|
||||
|
||||
def __init__(self, max_seq_len=2048, is_neox_style=True):
|
||||
self.max_seq_len = max_seq_len
|
||||
self.is_neox_style = is_neox_style
|
||||
self.cos_cached = None
|
||||
self.sin_cached = None
|
||||
self.rotary_dim = 1
|
||||
self.base = 1
|
||||
@pytest.fixture(autouse=True)
|
||||
def patch_init_side_effects():
|
||||
"""
|
||||
Suppress all side-effects that fire during __init__ so every test starts
|
||||
from a clean, predictable state without needing real NPU ops or vLLM
|
||||
global config.
|
||||
"""
|
||||
with (
|
||||
patch("vllm_ascend.ops.rotary_embedding._record_cos_sin_cache"),
|
||||
patch("vllm_ascend.ops.rotary_embedding._record_cos_and_sin_cache_interleaved"),
|
||||
patch("vllm_ascend.ops.rotary_embedding.get_current_vllm_config") as mock_cfg,
|
||||
):
|
||||
# Default: speculative_config is None → use_mtp = False
|
||||
mock_cfg.return_value.speculative_config = None
|
||||
yield mock_cfg
|
||||
|
||||
|
||||
class TestAscendDeepseekScalingRotaryEmbedding(TestBase):
|
||||
@pytest.fixture()
|
||||
def make_embedding(patch_init_side_effects):
|
||||
"""Factory that creates an AscendRotaryEmbedding with controllable use_mtp."""
|
||||
|
||||
def setUp(self):
|
||||
# Common setup for tests
|
||||
self.positions = torch.tensor([1, 2, 3])
|
||||
self.query = torch.randn(3, 1, 32, dtype=torch.float16)
|
||||
self.key = torch.randn(3, 1, 32, dtype=torch.float16)
|
||||
self.head_size = 32
|
||||
self.rotary_dim = self.head_size
|
||||
self.max_position = 16
|
||||
self.rope_theta = 10000
|
||||
self.is_neox_style = True
|
||||
self.scaling_factor = 1
|
||||
self.layer = None
|
||||
def _factory(use_mtp: bool = False, is_neox_style: bool = True):
|
||||
spec_cfg = MagicMock(method="mtp") if use_mtp else None
|
||||
patch_init_side_effects.return_value.speculative_config = spec_cfg
|
||||
|
||||
def _create_layer(self):
|
||||
self.layer = DeepseekScalingRotaryEmbedding(
|
||||
self.head_size, self.rotary_dim, self.max_position,
|
||||
self.rope_theta, self.is_neox_style, self.scaling_factor,
|
||||
torch.float16)
|
||||
return self.layer
|
||||
with patch("vllm_ascend.ops.rotary_embedding.RotaryEmbedding.__init__") as mock_parent_init:
|
||||
mock_parent_init.return_value = None
|
||||
from vllm_ascend.ops.rotary_embedding import AscendRotaryEmbedding
|
||||
|
||||
@patch("vllm.platforms.current_platform.device_type",
|
||||
new=torch.device("cpu"))
|
||||
@patch("vllm_ascend.ops.rotary_embedding.NPUPlatform",
|
||||
new_callable=PropertyMock)
|
||||
def test_native_rope_deepseek_forward_base(self, mock_npuplatform):
|
||||
mock_npuplatform.device_type = torch.device("cpu")
|
||||
self.layer = self._create_layer()
|
||||
with patch("vllm_ascend.ops.rotary_embedding._rope_forward_oot",
|
||||
return_value=(self.query,
|
||||
self.key)) as mock_rope_forward_oot:
|
||||
q_pe, k_pe = self.layer.forward(self.positions, self.query,
|
||||
self.key)
|
||||
mock_rope_forward_oot.assert_called_once()
|
||||
assert q_pe.shape == self.query.shape
|
||||
assert k_pe.shape == self.key.shape
|
||||
emb = AscendRotaryEmbedding.__new__(AscendRotaryEmbedding)
|
||||
# Manually set attrs that the real parent would set
|
||||
emb.head_size = HEAD_SIZE
|
||||
emb.rotary_dim = ROTARY_DIM
|
||||
emb.is_neox_style = is_neox_style
|
||||
emb.cos_sin_cache = torch.zeros(MAX_POS, ROTARY_DIM)
|
||||
# Call __init__ to exercise our code path
|
||||
AscendRotaryEmbedding.__init__(emb, HEAD_SIZE, ROTARY_DIM, MAX_POS, BASE, is_neox_style, DTYPE)
|
||||
return emb
|
||||
|
||||
@patch('vllm_ascend.ops.rotary_embedding._rope_forward_oot')
|
||||
@patch("vllm.platforms.current_platform.device_type",
|
||||
new=torch.device("cpu"))
|
||||
@patch("vllm_ascend.ops.rotary_embedding.NPUPlatform",
|
||||
new_callable=PropertyMock)
|
||||
def test_native_rope_deepseek_forward_key_reshaping(
|
||||
self, mock_npuplatform, mock_rope_forward_oot):
|
||||
mock_npuplatform.device_type = torch.device("cpu")
|
||||
self.layer = self._create_layer()
|
||||
return _factory
|
||||
|
||||
key = torch.randn(1, 32)
|
||||
|
||||
mock_rope_forward_oot.return_value = (self.query, key)
|
||||
@pytest.fixture()
|
||||
def make_yarn_embedding(patch_init_side_effects):
|
||||
"""
|
||||
Factory for AscendYaRNRotaryEmbedding with parent __init__ suppressed.
|
||||
patch_init_side_effects is the same autouse fixture as before.
|
||||
"""
|
||||
|
||||
q_pe, k_pe = self.layer.forward(self.positions, self.query, key)
|
||||
mock_rope_forward_oot.assert_called_once()
|
||||
assert q_pe.shape == self.query.shape
|
||||
assert k_pe.shape == key.shape
|
||||
def _factory(is_neox_style: bool = True):
|
||||
with patch("vllm_ascend.ops.rotary_embedding.YaRNScalingRotaryEmbedding.__init__") as mock_parent_init:
|
||||
mock_parent_init.return_value = None
|
||||
from vllm_ascend.ops.rotary_embedding import AscendYaRNRotaryEmbedding
|
||||
|
||||
@patch('vllm_ascend.ops.rotary_embedding._rope_forward_oot')
|
||||
@patch("vllm.platforms.current_platform.device_type",
|
||||
new=torch.device("cpu"))
|
||||
@patch("vllm_ascend.ops.rotary_embedding.NPUPlatform",
|
||||
new_callable=PropertyMock)
|
||||
def test_native_rope_deepseek_forward_non_neox_style(
|
||||
self, mock_npuplatform, mock_rope_forward_oot):
|
||||
mock_npuplatform.device_type = torch.device("cpu")
|
||||
self.layer = self._create_layer()
|
||||
emb = AscendYaRNRotaryEmbedding.__new__(AscendYaRNRotaryEmbedding)
|
||||
emb.head_size = HEAD_SIZE
|
||||
emb.rotary_dim = ROTARY_DIM
|
||||
emb.is_neox_style = is_neox_style
|
||||
emb.cos_sin_cache = torch.zeros(MAX_POS, ROTARY_DIM)
|
||||
AscendYaRNRotaryEmbedding.__init__(
|
||||
emb,
|
||||
head_size=HEAD_SIZE,
|
||||
rotary_dim=ROTARY_DIM,
|
||||
max_position_embeddings=MAX_POS,
|
||||
base=BASE,
|
||||
is_neox_style=is_neox_style,
|
||||
scaling_factor=1.0,
|
||||
dtype=DTYPE,
|
||||
)
|
||||
return emb
|
||||
|
||||
mock_rope_forward_oot.return_value = (self.query, self.key)
|
||||
return _factory
|
||||
|
||||
q_pe, k_pe = self.layer.forward(self.positions, self.query, self.key)
|
||||
|
||||
mock_rope_forward_oot.assert_called_once()
|
||||
assert q_pe.shape == self.query.shape
|
||||
assert k_pe.shape == self.key.shape
|
||||
class TestAscendEmbeddingForwardOOT:
|
||||
@patch("torch.ops.vllm.npu_rotary_embedding")
|
||||
@patch("vllm_ascend.ascend_forward_context.get_forward_context")
|
||||
def test_basic_call_delegates_to_npu_op(self, mock_get_forward_context, mock_npu_op, make_embedding):
|
||||
"""forward_oot always calls npu_rotary_embedding and returns its result."""
|
||||
mock_get_forward_context.return_value = MagicMock()
|
||||
mock_get_forward_context.return_value.is_draft_model = False
|
||||
mock_get_forward_context.return_value.flash_comm_v1_enabled = False
|
||||
expected_output = (torch.randn(SEQ_LEN, NUM_HEADS * HEAD_SIZE),) * 2
|
||||
mock_npu_op.return_value = expected_output
|
||||
|
||||
@patch("vllm.platforms.current_platform.device_type",
|
||||
new=torch.device("cpu"))
|
||||
@patch("vllm_ascend.ops.rotary_embedding.NPUPlatform",
|
||||
new_callable=PropertyMock)
|
||||
def test_basic_case(self, mock_npuplatform):
|
||||
# Test with standard values
|
||||
mock_npuplatform.device_type = torch.device("cpu")
|
||||
self.layer = self._create_layer()
|
||||
num_rotations = 100
|
||||
dim = 512
|
||||
base = 10000
|
||||
max_position_embeddings = 2048
|
||||
emb = make_embedding()
|
||||
positions, query, key = _make_tensors()
|
||||
|
||||
result = self.layer._yarn_find_correction_dim(num_rotations, dim, base,
|
||||
max_position_embeddings)
|
||||
result = emb.forward_oot(positions, query, key)
|
||||
|
||||
# Calculate expected value manually
|
||||
expected = (dim * torch.log(
|
||||
torch.tensor(max_position_embeddings) /
|
||||
(num_rotations * 2 * torch.pi))) / (2 *
|
||||
torch.log(torch.tensor(base)))
|
||||
mock_npu_op.assert_called_once_with(
|
||||
positions,
|
||||
query,
|
||||
key,
|
||||
emb.cos_sin_cache,
|
||||
HEAD_SIZE,
|
||||
ROTARY_DIM,
|
||||
emb.is_neox_style,
|
||||
)
|
||||
assert result is expected_output
|
||||
|
||||
self.assertTrue(torch.allclose(result, expected))
|
||||
@patch("torch.ops.vllm.npu_rotary_embedding")
|
||||
@patch("vllm_ascend.ascend_forward_context.get_forward_context")
|
||||
def test_neox_style_override_true(self, mock_get_forward_context, mock_npu_op, make_embedding):
|
||||
"""is_neox_style_override=True wins over self.is_neox_style=False."""
|
||||
mock_get_forward_context.return_value = MagicMock()
|
||||
mock_get_forward_context.return_value.is_draft_model = False
|
||||
mock_get_forward_context.return_value.flash_comm_v1_enabled = False
|
||||
mock_npu_op.return_value = MagicMock()
|
||||
|
||||
@patch("vllm.platforms.current_platform.device_type",
|
||||
new=torch.device("cpu"))
|
||||
@patch("vllm_ascend.ops.rotary_embedding.NPUPlatform",
|
||||
new_callable=PropertyMock)
|
||||
def test_yarn_get_mscale(self, mock_npuplatform):
|
||||
mock_npuplatform.device_type = torch.device("cpu")
|
||||
self.layer = self._create_layer()
|
||||
emb = make_embedding(is_neox_style=False)
|
||||
positions, query, key = _make_tensors()
|
||||
|
||||
# test_scale_less_than_or_equal_1
|
||||
self.assertEqual(self.layer._yarn_get_mscale(scale=0.5), 1.0)
|
||||
self.assertEqual(self.layer._yarn_get_mscale(scale=1.0), 1.0)
|
||||
self.assertEqual(self.layer._yarn_get_mscale(scale=0.999), 1.0)
|
||||
emb.forward_oot(positions, query, key, is_neox_style_override=True)
|
||||
|
||||
# test_scale_greater_than_1:
|
||||
test_cases = [(2.0, 1.0, 1.0 + 0.1 * math.log(2.0)),
|
||||
(10.0, 1.0, 1.0 + 0.1 * math.log(10.0)),
|
||||
(5.0, 2.0, 1.0 + 0.2 * math.log(5.0)),
|
||||
(math.e, 1.0, 1.0 + 0.1)]
|
||||
_, kwargs = mock_npu_op.call_args
|
||||
# Verify the override was forwarded correctly
|
||||
assert mock_npu_op.call_args[0][-1] is True # last positional arg = is_neox_style
|
||||
|
||||
for scale, mscale, expected in test_cases:
|
||||
result = self.layer._yarn_get_mscale(scale, mscale)
|
||||
self.assertAlmostEqual(
|
||||
result,
|
||||
expected,
|
||||
places=6,
|
||||
msg=f"Failed for scale={scale}, mscale={mscale}")
|
||||
@patch("torch.ops.vllm.npu_rotary_embedding")
|
||||
@patch("vllm_ascend.ascend_forward_context.get_forward_context")
|
||||
def test_neox_style_override_false(self, mock_get_forward_context, mock_npu_op, make_embedding):
|
||||
"""is_neox_style_override=False wins over self.is_neox_style=True."""
|
||||
mock_get_forward_context.return_value = MagicMock()
|
||||
mock_get_forward_context.return_value.is_draft_model = False
|
||||
mock_get_forward_context.return_value.flash_comm_v1_enabled = False
|
||||
mock_npu_op.return_value = MagicMock()
|
||||
|
||||
emb = make_embedding(is_neox_style=True)
|
||||
positions, query, key = _make_tensors()
|
||||
|
||||
emb.forward_oot(positions, query, key, is_neox_style_override=False)
|
||||
|
||||
assert mock_npu_op.call_args[0][-1] is False
|
||||
|
||||
@patch("torch.ops.vllm.npu_rotary_embedding")
|
||||
@patch("vllm_ascend.ascend_forward_context.get_forward_context")
|
||||
def test_neox_style_override_none_uses_self(self, mock_get_forward_context, mock_npu_op, make_embedding):
|
||||
"""When override is None, self.is_neox_style is used unchanged."""
|
||||
mock_get_forward_context.return_value = MagicMock()
|
||||
mock_get_forward_context.return_value.is_draft_model = False
|
||||
mock_get_forward_context.return_value.flash_comm_v1_enabled = False
|
||||
mock_npu_op.return_value = MagicMock()
|
||||
|
||||
emb = make_embedding(is_neox_style=True)
|
||||
positions, query, key = _make_tensors()
|
||||
|
||||
emb.forward_oot(positions, query, key, is_neox_style_override=None)
|
||||
|
||||
assert mock_npu_op.call_args[0][-1] is True
|
||||
|
||||
@patch("torch.ops.vllm.maybe_all_gather_and_maybe_unpad")
|
||||
@patch("torch.ops.vllm.npu_rotary_embedding")
|
||||
@patch("vllm_ascend.ascend_forward_context.get_forward_context")
|
||||
def test_gather_unpad_called_when_all_conditions_met(
|
||||
self, mock_get_forward_context, mock_npu_op, mock_gather, make_embedding
|
||||
):
|
||||
"""
|
||||
maybe_all_gather_and_maybe_unpad is called iff:
|
||||
is_draft_model=True AND use_mtp=True AND flash_comm_v1_enabled=True
|
||||
"""
|
||||
mock_get_forward_context.return_value = MagicMock()
|
||||
mock_get_forward_context.return_value.is_draft_model = True
|
||||
mock_get_forward_context.return_value.flash_comm_v1_enabled = True
|
||||
gathered_positions = torch.arange(SEQ_LEN, dtype=torch.long)
|
||||
mock_gather.return_value = gathered_positions
|
||||
mock_npu_op.return_value = MagicMock()
|
||||
|
||||
emb = make_embedding(use_mtp=True)
|
||||
positions, query, key = _make_tensors()
|
||||
|
||||
emb.forward_oot(positions, query, key)
|
||||
|
||||
mock_gather.assert_called_once()
|
||||
# npu op should receive the gathered positions, not the originals
|
||||
assert mock_npu_op.call_args[0][0] is gathered_positions
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"is_draft_model,flash_comm,use_mtp",
|
||||
[
|
||||
(False, True, True), # not draft
|
||||
(True, False, True), # flash_comm disabled
|
||||
(True, True, False), # use_mtp disabled
|
||||
],
|
||||
)
|
||||
@patch("torch.ops.vllm.maybe_all_gather_and_maybe_unpad")
|
||||
@patch("torch.ops.vllm.npu_rotary_embedding")
|
||||
@patch("vllm_ascend.ascend_forward_context.get_forward_context")
|
||||
def test_gather_unpad_skipped_unless_all_conditions_met(
|
||||
self,
|
||||
mock_get_forward_context,
|
||||
mock_npu_op,
|
||||
mock_gather,
|
||||
is_draft_model,
|
||||
flash_comm,
|
||||
use_mtp,
|
||||
make_embedding,
|
||||
):
|
||||
"""gather/unpad must NOT fire if any one of the three conditions is False."""
|
||||
mock_get_forward_context.return_value = MagicMock()
|
||||
mock_get_forward_context.return_value.is_draft_model = is_draft_model
|
||||
mock_get_forward_context.return_value.flash_comm_v1_enabled = flash_comm
|
||||
mock_npu_op.return_value = MagicMock()
|
||||
|
||||
emb = make_embedding(use_mtp=use_mtp)
|
||||
positions, query, key = _make_tensors()
|
||||
|
||||
emb.forward_oot(positions, query, key)
|
||||
|
||||
mock_gather.assert_not_called()
|
||||
# Original positions tensor is passed through untouched
|
||||
assert mock_npu_op.call_args[0][0] is positions
|
||||
|
||||
def test_parent_init_signature_has_not_changed(self):
|
||||
"""
|
||||
Fail loudly if RotaryEmbedding.__init__ adds, removes, or
|
||||
renames parameters, so a developer knows to update AscendRotaryEmbedding
|
||||
accordingly.
|
||||
"""
|
||||
check_parent_init_signature_has_not_changed(RotaryEmbedding.__init__, AscendRotaryEmbedding.__init__)
|
||||
|
||||
|
||||
class TestAscendYaRNRotaryEmbeddingForwardOOT:
|
||||
@patch("vllm_ascend.ops.rotary_embedding.AscendRotaryEmbedding.forward_oot")
|
||||
def test_delegates_to_ascend_rotary_forward_oot(self, mock_delegate, make_yarn_embedding):
|
||||
"""forward_oot must delegate to AscendRotaryEmbedding.forward_oot."""
|
||||
expected = MagicMock()
|
||||
mock_delegate.return_value = expected
|
||||
|
||||
emb = make_yarn_embedding()
|
||||
positions, query, key = _make_tensors()
|
||||
|
||||
result = emb.forward_oot(positions, query, key)
|
||||
|
||||
mock_delegate.assert_called_once_with(emb, positions, query, key, None, None)
|
||||
assert result is expected
|
||||
|
||||
@patch("vllm_ascend.ops.rotary_embedding.AscendRotaryEmbedding.forward_oot")
|
||||
def test_return_value_passed_through(self, mock_delegate, make_yarn_embedding):
|
||||
"""Return value from the delegate is returned unchanged."""
|
||||
sentinel = (torch.randn(SEQ_LEN, HEAD_SIZE), torch.randn(SEQ_LEN, HEAD_SIZE))
|
||||
mock_delegate.return_value = sentinel
|
||||
|
||||
emb = make_yarn_embedding()
|
||||
positions, query, key = _make_tensors()
|
||||
|
||||
result = emb.forward_oot(positions, query, key)
|
||||
|
||||
assert result is sentinel
|
||||
|
||||
@pytest.mark.parametrize("override", [True, False])
|
||||
@patch("vllm_ascend.ops.rotary_embedding.AscendRotaryEmbedding.forward_oot")
|
||||
def test_is_neox_style_override_forwarded(self, mock_delegate, override, make_yarn_embedding):
|
||||
"""is_neox_style_override must be forwarded verbatim, both True and False."""
|
||||
mock_delegate.return_value = MagicMock()
|
||||
|
||||
emb = make_yarn_embedding()
|
||||
positions, query, key = _make_tensors()
|
||||
|
||||
emb.forward_oot(positions, query, key, is_neox_style_override=override)
|
||||
|
||||
_, call_args, _ = mock_delegate.mock_calls[0]
|
||||
assert call_args[5] is override # 6th positional arg
|
||||
|
||||
@patch("vllm_ascend.ops.rotary_embedding.AscendRotaryEmbedding.forward_oot")
|
||||
def test_all_args_forwarded_together(self, mock_delegate, make_yarn_embedding):
|
||||
"""Smoke test: all args passed simultaneously are all forwarded correctly."""
|
||||
mock_delegate.return_value = MagicMock()
|
||||
|
||||
emb = make_yarn_embedding()
|
||||
positions, query, key = _make_tensors()
|
||||
offsets = torch.ones(SEQ_LEN, dtype=torch.long)
|
||||
|
||||
emb.forward_oot(positions, query, key, offsets=offsets, is_neox_style_override=False)
|
||||
|
||||
mock_delegate.assert_called_once_with(emb, positions, query, key, offsets, False)
|
||||
|
||||
def test_parent_init_signature_has_not_changed(self):
|
||||
"""
|
||||
Fail loudly if YaRNScalingRotaryEmbedding.__init__ adds, removes, or
|
||||
renames parameters, so a developer knows to update AscendYaRNRotaryEmbedding
|
||||
accordingly.
|
||||
"""
|
||||
check_parent_init_signature_has_not_changed(
|
||||
YaRNScalingRotaryEmbedding.__init__, AscendYaRNRotaryEmbedding.__init__
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user