@@ -14,39 +14,68 @@
|
||||
# Adapted from vllm/tests/lora/test_layers.py
|
||||
|
||||
import unittest
|
||||
from unittest import mock
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import torch
|
||||
|
||||
from vllm_ascend.ascend_config import init_ascend_config
|
||||
from vllm_ascend.distributed import parallel_state
|
||||
from vllm_ascend.ops.vocab_parallel_embedding import (
|
||||
AscendLogitsProcessor, AscendParallelLMHead, AscendVocabParallelEmbedding)
|
||||
AscendLogitsProcessor,
|
||||
AscendParallelLMHead,
|
||||
AscendVocabParallelEmbedding,
|
||||
)
|
||||
|
||||
VOCAB_PARALLEL_EMBEDDING_TEST_NUM_RANDOM_SEEDS = 128
|
||||
|
||||
|
||||
class TestCustomVocabParallelEmbedding(unittest.TestCase):
|
||||
|
||||
def setUp(self):
|
||||
self.num_embeddings = 50
|
||||
self.embedding_dim = 10
|
||||
self.org_num_embeddings = 40
|
||||
self.padding_size = 8
|
||||
|
||||
self.mock_group = mock.MagicMock()
|
||||
self.mock_group.world_size = 2
|
||||
self.mock_group.rank_in_group = 0
|
||||
|
||||
parallel_state._MLP_TP = self.mock_group
|
||||
parallel_state._OTP = self.mock_group
|
||||
|
||||
mock_vllm_config = MagicMock()
|
||||
mock_vllm_config.additional_config = {}
|
||||
init_ascend_config(mock_vllm_config)
|
||||
self.mock_ascend_config = MagicMock()
|
||||
self.mock_ascend_config.finegrained_tp_config.lmhead_tensor_parallel_size = 2
|
||||
self.mock_ascend_config.finegrained_tp_config.embedding_tensor_parallel_size = 2
|
||||
|
||||
self.patches = [
|
||||
patch("vllm_ascend.utils.get_ascend_config", return_value=self.mock_ascend_config),
|
||||
patch("vllm_ascend.distributed.parallel_state.get_lmhead_tp_group", return_value=self.mock_group),
|
||||
patch(
|
||||
"vllm.distributed.parallel_state.get_tp_group",
|
||||
return_value=self.mock_group,
|
||||
),
|
||||
]
|
||||
|
||||
for p in self.patches:
|
||||
p.start()
|
||||
|
||||
def _create_layer(self):
|
||||
# Patch methods and dependencies for VocabParallelEmbedding
|
||||
mock_group = MagicMock()
|
||||
mock_group.world_size = 2
|
||||
mock_group.rank_in_group = 0
|
||||
with patch("vllm_ascend.ops.vocab_parallel_embedding.get_tp_group", return_value=mock_group), \
|
||||
patch("vllm.model_executor.layers.vocab_parallel_embedding.get_tensor_model_parallel_rank", return_value=0), \
|
||||
patch("vllm.model_executor.layers.vocab_parallel_embedding.get_tensor_model_parallel_world_size", return_value=2), \
|
||||
patch("vllm.model_executor.layers.vocab_parallel_embedding.pad_vocab_size", side_effect=lambda x, y: x + y), \
|
||||
patch("vllm.model_executor.layers.vocab_parallel_embedding.divide", side_effect=lambda x, y: x // y):
|
||||
|
||||
with (
|
||||
patch("vllm_ascend.ops.vocab_parallel_embedding.get_tp_group", return_value=mock_group),
|
||||
patch("vllm.model_executor.layers.vocab_parallel_embedding.get_tensor_model_parallel_rank", return_value=0),
|
||||
patch(
|
||||
"vllm.model_executor.layers.vocab_parallel_embedding.get_tensor_model_parallel_world_size",
|
||||
return_value=2,
|
||||
),
|
||||
patch("vllm.model_executor.layers.vocab_parallel_embedding.pad_vocab_size", side_effect=lambda x, y: x + y),
|
||||
patch("vllm.model_executor.layers.vocab_parallel_embedding.divide", side_effect=lambda x, y: x // y),
|
||||
):
|
||||
# Create an instance of VocabParallelEmbedding
|
||||
layer = AscendVocabParallelEmbedding(
|
||||
num_embeddings=self.num_embeddings,
|
||||
@@ -54,7 +83,8 @@ class TestCustomVocabParallelEmbedding(unittest.TestCase):
|
||||
org_num_embeddings=self.org_num_embeddings,
|
||||
padding_size=self.padding_size,
|
||||
quant_config=None, # Mock quantization config
|
||||
prefix="")
|
||||
prefix="",
|
||||
)
|
||||
|
||||
layer.shard_indices = MagicMock()
|
||||
layer.shard_indices.org_vocab_start_index = 10
|
||||
@@ -65,8 +95,8 @@ class TestCustomVocabParallelEmbedding(unittest.TestCase):
|
||||
|
||||
# Mock the quantization method
|
||||
layer.quant_method.embedding = MagicMock(
|
||||
side_effect=lambda _, x: torch.randn(x.shape[0], self.
|
||||
embedding_dim))
|
||||
side_effect=lambda _, x: torch.randn(x.shape[0], self.embedding_dim)
|
||||
)
|
||||
return layer
|
||||
|
||||
def test_get_masked_input_and_mask(self):
|
||||
@@ -81,17 +111,16 @@ class TestCustomVocabParallelEmbedding(unittest.TestCase):
|
||||
org_vocab_end_index=20,
|
||||
num_org_vocab_padding=5,
|
||||
added_vocab_start_index=30,
|
||||
added_vocab_end_index=40)
|
||||
added_vocab_end_index=40,
|
||||
)
|
||||
|
||||
expected_mask = torch.tensor([True, False, True, False, True])
|
||||
self.assertTrue(
|
||||
torch.equal(mask, expected_mask),
|
||||
f"Mask mismatch. Expected {expected_mask}, got {mask}")
|
||||
self.assertTrue(torch.equal(mask, expected_mask), f"Mask mismatch. Expected {expected_mask}, got {mask}")
|
||||
|
||||
expected_masked = torch.tensor([0, 5, 0, 20, 0])
|
||||
self.assertTrue(
|
||||
torch.equal(masked_input, expected_masked),
|
||||
f"Masked input mismatch. Expected {expected_masked}, got {masked_input}"
|
||||
f"Masked input mismatch. Expected {expected_masked}, got {masked_input}",
|
||||
)
|
||||
|
||||
def test_forward_with_tp_size_1(self):
|
||||
@@ -99,19 +128,15 @@ class TestCustomVocabParallelEmbedding(unittest.TestCase):
|
||||
# Create a fresh mock embedding with tp_size=1
|
||||
layer = self._create_layer()
|
||||
layer.tp_size = 1
|
||||
layer.quant_method.embedding = MagicMock(
|
||||
return_value=torch.randn(3, layer.embedding_dim))
|
||||
layer.quant_method.embedding = MagicMock(return_value=torch.randn(3, layer.embedding_dim))
|
||||
|
||||
input_ = torch.tensor([1, 2, 3])
|
||||
|
||||
with patch(
|
||||
"vllm_ascend.ops.vocab_parallel_embedding.tensor_model_parallel_all_reduce",
|
||||
side_effect=lambda x: x) as mock_reduce_tp1:
|
||||
with patch("torch.ops.vllm.maybe_pad_and_reduce", side_effect=lambda x: x) as mock_reduce_tp1:
|
||||
output = layer.forward(input_)
|
||||
|
||||
# Should just pass through without masking
|
||||
layer.quant_method.embedding.assert_called_once_with(
|
||||
layer, input_.long())
|
||||
layer.quant_method.embedding.assert_called_once_with(layer, input_.long())
|
||||
self.assertEqual(output.shape, (3, layer.embedding_dim))
|
||||
|
||||
# Verify all_reduce was called once
|
||||
@@ -123,9 +148,7 @@ class TestCustomVocabParallelEmbedding(unittest.TestCase):
|
||||
|
||||
input_ = torch.tensor([15, 35]) # one org vocab, one added vocab
|
||||
|
||||
with patch(
|
||||
"vllm_ascend.ops.vocab_parallel_embedding.tensor_model_parallel_all_reduce",
|
||||
side_effect=lambda x: x) as mock_reduce_tp:
|
||||
with patch("torch.ops.vllm.maybe_pad_and_reduce", side_effect=lambda x: x) as mock_reduce_tp:
|
||||
# Call the forward method
|
||||
output = layer.forward(input_)
|
||||
|
||||
@@ -146,13 +169,10 @@ class TestCustomVocabParallelEmbedding(unittest.TestCase):
|
||||
input_ = torch.tensor([5, 15, 25, 35, 45]) # includes invalid cases
|
||||
# Create predictable mock output
|
||||
mock_output = torch.randn(5, self.embedding_dim)
|
||||
layer.quant_method.embedding = MagicMock(
|
||||
return_value=mock_output.clone())
|
||||
layer.quant_method.embedding = MagicMock(return_value=mock_output.clone())
|
||||
|
||||
# Patch tensor_model_parallel_all_reduce to mock its behavior
|
||||
with patch(
|
||||
"vllm_ascend.ops.vocab_parallel_embedding.tensor_model_parallel_all_reduce",
|
||||
side_effect=lambda x: x):
|
||||
with patch("torch.ops.vllm.maybe_pad_and_reduce", side_effect=lambda x: x):
|
||||
# Call the forward method
|
||||
output = layer.forward(input_)
|
||||
# Check that invalid positions (0, 2, 4) were zeroed out
|
||||
@@ -176,17 +196,23 @@ class TestCustomVocabParallelEmbedding(unittest.TestCase):
|
||||
|
||||
for input_, expected_shape in test_cases:
|
||||
with self.subTest(input=input_):
|
||||
with patch(
|
||||
"vllm_ascend.ops.vocab_parallel_embedding.tensor_model_parallel_all_reduce",
|
||||
side_effect=lambda x: x):
|
||||
with patch("torch.ops.vllm.maybe_pad_and_reduce", side_effect=lambda x: x):
|
||||
# Call the forward method
|
||||
output = layer.forward(input_)
|
||||
self.assertEqual(output.shape, expected_shape)
|
||||
|
||||
|
||||
class TestAscendLogitsProcessor(unittest.TestCase):
|
||||
|
||||
def setUp(self):
|
||||
self.mock_vllm_config = MagicMock()
|
||||
self.mock_vllm_config.compilation_config.custom_ops = ["all"]
|
||||
|
||||
from vllm.config.vllm import set_current_vllm_config
|
||||
|
||||
set_current_vllm_config(self.mock_vllm_config)
|
||||
|
||||
self.config_patch = patch("vllm.config.vllm.get_current_vllm_config", return_value=self.mock_vllm_config)
|
||||
self.config_patch.start()
|
||||
self.vocab_size = 50
|
||||
self.num_embeddings = 50
|
||||
self.embedding_dim = 10
|
||||
@@ -198,27 +224,19 @@ class TestAscendLogitsProcessor(unittest.TestCase):
|
||||
self.mock_group.rank_in_group = 0
|
||||
self.mock_ascend_config = MagicMock()
|
||||
self.mock_quant_method = MagicMock()
|
||||
self.mock_quant_method.apply = MagicMock(
|
||||
return_value=torch.randn(1, self.vocab_size))
|
||||
self.mock_quant_method.apply = MagicMock(return_value=torch.randn(1, self.vocab_size))
|
||||
self.patches = [
|
||||
patch("vllm_ascend.ascend_config.get_ascend_config",
|
||||
return_value=self.mock_ascend_config),
|
||||
patch(
|
||||
"vllm_ascend.ops.vocab_parallel_embedding.get_lmhead_tp_group",
|
||||
return_value=self.mock_group),
|
||||
patch("vllm_ascend.ops.vocab_parallel_embedding.lmhead_tp_enable",
|
||||
return_value=True),
|
||||
patch("vllm_ascend.ops.vocab_parallel_embedding.get_ascend_config", return_value=self.mock_ascend_config),
|
||||
patch("vllm_ascend.ops.vocab_parallel_embedding.get_lmhead_tp_group", return_value=self.mock_group),
|
||||
patch("vllm_ascend.ops.vocab_parallel_embedding.lmhead_tp_enable", return_value=True),
|
||||
patch(
|
||||
"vllm_ascend.ops.vocab_parallel_embedding.get_lmhead_tp_group.all_to_all",
|
||||
return_value=torch.randn(1, self.vocab_size)),
|
||||
return_value=torch.randn(1, self.vocab_size),
|
||||
),
|
||||
patch(
|
||||
"vllm_ascend.ops.vocab_parallel_embedding.get_lmhead_tp_group.all_gather",
|
||||
return_value=torch.randn(1, self.vocab_size)),
|
||||
patch(
|
||||
"vllm_ascend.core.schedule_config.AscendSchedulerConfig.initialize_from_config",
|
||||
return_value=MagicMock(max_num_batched_tokens=1000,
|
||||
max_model_len=512,
|
||||
enable_chunked_prefill=False))
|
||||
return_value=torch.randn(1, self.vocab_size),
|
||||
),
|
||||
]
|
||||
|
||||
for p in self.patches:
|
||||
@@ -234,9 +252,9 @@ class TestAscendLogitsProcessor(unittest.TestCase):
|
||||
|
||||
def test_get_logits(self):
|
||||
processor = AscendLogitsProcessor(vocab_size=self.vocab_size)
|
||||
lmhead = AscendParallelLMHead(num_embeddings=self.num_embeddings,
|
||||
embedding_dim=self.embedding_dim,
|
||||
prefix="lm_head")
|
||||
lmhead = AscendParallelLMHead(
|
||||
num_embeddings=self.num_embeddings, embedding_dim=self.embedding_dim, prefix="lm_head"
|
||||
)
|
||||
lmhead.quant_method = self.mock_quant_method
|
||||
lmhead.quant_method.apply = self.mock_quant_method.apply
|
||||
hidden_state = torch.randn(1, self.org_num_embeddings)
|
||||
|
||||
Reference in New Issue
Block a user