[UT][NPU] add NPU attention unit tests for ascend_backend and ascend_dsv4_backend (#32294)
This commit is contained in:
@@ -0,0 +1,767 @@
|
||||
"""
|
||||
Unit tests for sglang.srt.hardware_backend.npu.attention.ascend_backend.
|
||||
"""
|
||||
|
||||
import sys
|
||||
import unittest
|
||||
from dataclasses import fields, is_dataclass
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.test.ci.ci_register import register_npu_ci
|
||||
|
||||
register_npu_ci(est_time=5, suite="stage-a-unit-test-npu")
|
||||
|
||||
# Mock NPU-only modules before importing the source module.
|
||||
for _ in (
|
||||
"torch_npu",
|
||||
"torch_npu.contrib",
|
||||
"sgl_kernel_npu",
|
||||
"sgl_kernel_npu.attention",
|
||||
"sgl_kernel_npu.attention.sinks_attention",
|
||||
"sglang.srt.speculative",
|
||||
"sglang.srt.speculative.decoupled_spec_io",
|
||||
"sglang.srt.speculative.spec_info",
|
||||
"sglang.srt.speculative.eagle_info",
|
||||
):
|
||||
sys.modules.setdefault(_, MagicMock())
|
||||
|
||||
from sglang.srt.hardware_backend.npu.attention.ascend_backend import (
|
||||
AscendAttnBackend,
|
||||
AscendAttnMaskBuilder,
|
||||
AscendAttnMultiStepDraftBackend,
|
||||
ForwardMetadata,
|
||||
_expand_dsa_sparse_indices,
|
||||
_reshape_kv_for_fia_nz,
|
||||
)
|
||||
|
||||
|
||||
class TestExpandDsaSparseIndices(unittest.TestCase):
|
||||
def test_2d_input_adds_unsqueeze(self):
|
||||
"""A [T, K] tensor becomes [T, 1, K]."""
|
||||
topk = torch.tensor([[1, 2, 3], [4, 5, 6]])
|
||||
result = _expand_dsa_sparse_indices(topk)
|
||||
self.assertEqual(result.shape, (2, 1, 3))
|
||||
self.assertTrue(torch.equal(result.squeeze(1), topk))
|
||||
|
||||
def test_3d_input_passthrough(self):
|
||||
topk = torch.tensor([[[1, 2], [3, 4]], [[5, 6], [7, 8]]])
|
||||
result = _expand_dsa_sparse_indices(topk)
|
||||
self.assertEqual(result.shape, (2, 2, 2))
|
||||
self.assertTrue(torch.equal(result, topk))
|
||||
|
||||
def test_2d_shape_correctness(self):
|
||||
topk = torch.zeros(5, 8)
|
||||
result = _expand_dsa_sparse_indices(topk)
|
||||
self.assertEqual(result.dim(), 3)
|
||||
self.assertEqual(result.shape[0], 5)
|
||||
self.assertEqual(result.shape[1], 1)
|
||||
self.assertEqual(result.shape[2], 8)
|
||||
|
||||
def test_2d_single_row(self):
|
||||
topk = torch.tensor([[1, 2, 3, 4]])
|
||||
result = _expand_dsa_sparse_indices(topk)
|
||||
self.assertEqual(result.shape, (1, 1, 4))
|
||||
|
||||
|
||||
class TestReshapeKvForFiaNz(unittest.TestCase):
|
||||
def test_output_shape(self):
|
||||
"""Output shape is (-1, 1, num_heads*head_dim//16, page_size, 16)."""
|
||||
num_heads = 2
|
||||
head_dim = 64
|
||||
page_size = 16
|
||||
total = 1 * 1 * (num_heads * head_dim // 16) * page_size * 16
|
||||
tensor = torch.arange(total, dtype=torch.float32)
|
||||
result = _reshape_kv_for_fia_nz(tensor, num_heads, head_dim, page_size)
|
||||
self.assertEqual(
|
||||
result.shape, (1, 1, num_heads * head_dim // 16, page_size, 16)
|
||||
)
|
||||
|
||||
def test_element_preservation(self):
|
||||
num_heads = 2
|
||||
head_dim = 64
|
||||
page_size = 16
|
||||
total = 2 * 1 * (num_heads * head_dim // 16) * page_size * 16
|
||||
tensor = torch.arange(total, dtype=torch.float32)
|
||||
result = _reshape_kv_for_fia_nz(tensor, num_heads, head_dim, page_size)
|
||||
self.assertEqual(result.numel(), tensor.numel())
|
||||
self.assertTrue(torch.equal(result.flatten(), tensor))
|
||||
|
||||
def test_different_parameters(self):
|
||||
num_heads = 4
|
||||
head_dim = 128
|
||||
page_size = 32
|
||||
total = 3 * 1 * (num_heads * head_dim // 16) * page_size * 16
|
||||
tensor = torch.randn(total)
|
||||
result = _reshape_kv_for_fia_nz(tensor, num_heads, head_dim, page_size)
|
||||
self.assertEqual(
|
||||
result.shape, (3, 1, num_heads * head_dim // 16, page_size, 16)
|
||||
)
|
||||
|
||||
def test_view_relationship(self):
|
||||
num_heads = 2
|
||||
head_dim = 64
|
||||
page_size = 16
|
||||
total = 1 * 1 * (num_heads * head_dim // 16) * page_size * 16
|
||||
tensor = torch.arange(total, dtype=torch.float32)
|
||||
result = _reshape_kv_for_fia_nz(tensor, num_heads, head_dim, page_size)
|
||||
self.assertEqual(result.data_ptr(), tensor.data_ptr())
|
||||
|
||||
|
||||
class TestForwardMetadata(unittest.TestCase):
|
||||
def test_is_dataclass(self):
|
||||
self.assertTrue(is_dataclass(ForwardMetadata))
|
||||
|
||||
def test_all_fields_default_none(self):
|
||||
metadata = ForwardMetadata()
|
||||
for f in fields(ForwardMetadata):
|
||||
self.assertIsNone(
|
||||
getattr(metadata, f.name),
|
||||
f"Field '{f.name}' should default to None",
|
||||
)
|
||||
|
||||
def test_create_with_values(self):
|
||||
block_tables = torch.tensor([[1, 2], [3, 4]])
|
||||
seq_lens = torch.tensor([10, 20])
|
||||
metadata = ForwardMetadata(
|
||||
block_tables=block_tables,
|
||||
seq_lens=seq_lens,
|
||||
seq_lens_cpu_list=[10, 20],
|
||||
)
|
||||
self.assertTrue(torch.equal(metadata.block_tables, block_tables))
|
||||
self.assertTrue(torch.equal(metadata.seq_lens, seq_lens))
|
||||
self.assertEqual(metadata.seq_lens_cpu_list, [10, 20])
|
||||
|
||||
def test_partial_assignment(self):
|
||||
metadata = ForwardMetadata(swa_mask=torch.ones(3, 3))
|
||||
self.assertIsNotNone(metadata.swa_mask)
|
||||
self.assertIsNone(metadata.block_tables)
|
||||
self.assertIsNone(metadata.seq_lens)
|
||||
|
||||
def test_field_names(self):
|
||||
names = {f.name for f in fields(ForwardMetadata)}
|
||||
expected = {
|
||||
"block_tables",
|
||||
"block_tables_swa",
|
||||
"swa_out_cache_loc",
|
||||
"extend_seq_lens_cpu_int",
|
||||
"seq_lens_cpu_int",
|
||||
"seq_lens_cpu_list",
|
||||
"seq_lens_list_cumsum",
|
||||
"seq_lens",
|
||||
"actual_seq_lengths_q",
|
||||
"actual_seq_lengths_q_pa",
|
||||
"actual_seq_lengths_kv",
|
||||
"swa_mask",
|
||||
"prefix_lens",
|
||||
"flatten_prefix_block_tables",
|
||||
}
|
||||
self.assertEqual(names, expected)
|
||||
|
||||
|
||||
class TestGenerateMaskFlag(unittest.TestCase):
|
||||
def test_shape(self):
|
||||
mask = AscendAttnMaskBuilder.generate_mask_flag(8)
|
||||
self.assertEqual(mask.shape, (8, 8))
|
||||
|
||||
def test_dtype_bool(self):
|
||||
mask = AscendAttnMaskBuilder.generate_mask_flag(4)
|
||||
self.assertEqual(mask.dtype, torch.bool)
|
||||
|
||||
def test_upper_triangular_pattern(self):
|
||||
"""generate_mask_flag returns ~tril, i.e. True above the diagonal."""
|
||||
mask = AscendAttnMaskBuilder.generate_mask_flag(4)
|
||||
self.assertTrue(mask[0, 1].item())
|
||||
self.assertTrue(mask[0, 3].item())
|
||||
self.assertTrue(mask[2, 3].item())
|
||||
self.assertFalse(mask[0, 0].item())
|
||||
self.assertFalse(mask[1, 0].item())
|
||||
self.assertFalse(mask[3, 3].item())
|
||||
|
||||
def test_diagonal_is_false(self):
|
||||
mask = AscendAttnMaskBuilder.generate_mask_flag(5)
|
||||
for i in range(5):
|
||||
self.assertFalse(mask[i, i].item())
|
||||
|
||||
def test_1x1(self):
|
||||
mask = AscendAttnMaskBuilder.generate_mask_flag(1)
|
||||
self.assertEqual(mask.shape, (1, 1))
|
||||
self.assertFalse(mask[0, 0].item())
|
||||
|
||||
def test_symmetric_upper(self):
|
||||
n = 6
|
||||
mask = AscendAttnMaskBuilder.generate_mask_flag(n)
|
||||
for i in range(n):
|
||||
for j in range(i + 1, n):
|
||||
self.assertTrue(mask[i, j].item())
|
||||
self.assertFalse(mask[j, i].item())
|
||||
|
||||
|
||||
class TestGenerateAttnMask(unittest.TestCase):
|
||||
def test_shape(self):
|
||||
mask = AscendAttnMaskBuilder.generate_attn_mask(8, "norm", torch.float16)
|
||||
self.assertEqual(mask.shape, (8, 8))
|
||||
|
||||
def test_dtype(self):
|
||||
mask = AscendAttnMaskBuilder.generate_attn_mask(4, "norm", torch.bfloat16)
|
||||
self.assertEqual(mask.dtype, torch.bfloat16)
|
||||
|
||||
def test_default_dtype_float16(self):
|
||||
mask = AscendAttnMaskBuilder.generate_attn_mask(4, "norm")
|
||||
self.assertEqual(mask.dtype, torch.float16)
|
||||
|
||||
def test_mix_mode_float16(self):
|
||||
"""mix + float16 -> upper triangle is -inf, lower is 0."""
|
||||
mask = AscendAttnMaskBuilder.generate_attn_mask(4, "mix", torch.float16)
|
||||
self.assertEqual(mask.dtype, torch.float16)
|
||||
self.assertTrue(torch.isinf(mask[0, 1]))
|
||||
self.assertTrue(mask[0, 1] < 0)
|
||||
self.assertEqual(mask[0, 0].item(), 0.0)
|
||||
self.assertEqual(mask[1, 0].item(), 0.0)
|
||||
|
||||
def test_mix_mode_bfloat16(self):
|
||||
"""mix + bfloat16 -> upper triangle is -inf, lower is 0."""
|
||||
mask = AscendAttnMaskBuilder.generate_attn_mask(4, "mix", torch.bfloat16)
|
||||
self.assertEqual(mask.dtype, torch.bfloat16)
|
||||
self.assertTrue(torch.isinf(mask[0, 1]))
|
||||
self.assertTrue(mask[0, 1] < 0)
|
||||
self.assertEqual(mask[1, 1].item(), 0.0)
|
||||
|
||||
def test_norm_mode_float16(self):
|
||||
"""norm + float16 -> upper triangle is -inf (overflow), lower is 0."""
|
||||
mask = AscendAttnMaskBuilder.generate_attn_mask(4, "norm", torch.float16)
|
||||
self.assertEqual(mask.dtype, torch.float16)
|
||||
self.assertTrue(torch.isinf(mask[0, 1]))
|
||||
self.assertTrue(mask[0, 1] < 0)
|
||||
self.assertEqual(mask[0, 0].item(), 0.0)
|
||||
|
||||
def test_norm_mode_bfloat16(self):
|
||||
"""norm + bfloat16 -> upper triangle is 1, lower is 0."""
|
||||
mask = AscendAttnMaskBuilder.generate_attn_mask(4, "norm", torch.bfloat16)
|
||||
self.assertEqual(mask.dtype, torch.bfloat16)
|
||||
self.assertEqual(mask[0, 1].item(), 1.0)
|
||||
self.assertEqual(mask[0, 0].item(), 0.0)
|
||||
|
||||
def test_norm_mode_float32(self):
|
||||
"""norm + float32 -> upper triangle is 1, lower is 0."""
|
||||
mask = AscendAttnMaskBuilder.generate_attn_mask(4, "norm", torch.float32)
|
||||
self.assertEqual(mask.dtype, torch.float32)
|
||||
self.assertEqual(mask[0, 1].item(), 1.0)
|
||||
self.assertEqual(mask[0, 0].item(), 0.0)
|
||||
|
||||
def test_mix_mode_float32(self):
|
||||
"""mix + float32 -> upper triangle is 1, lower is 0."""
|
||||
mask = AscendAttnMaskBuilder.generate_attn_mask(4, "mix", torch.float32)
|
||||
self.assertEqual(mask.dtype, torch.float32)
|
||||
self.assertEqual(mask[0, 1].item(), 1.0)
|
||||
self.assertEqual(mask[0, 0].item(), 0.0)
|
||||
|
||||
def test_lower_triangle_all_zero(self):
|
||||
"""The lower triangle (including diagonal) is always zero."""
|
||||
n = 6
|
||||
mask = AscendAttnMaskBuilder.generate_attn_mask(n, "norm", torch.float32)
|
||||
for i in range(n):
|
||||
for j in range(i + 1):
|
||||
self.assertEqual(mask[i, j].item(), 0.0)
|
||||
|
||||
def test_upper_triangle_all_masked(self):
|
||||
n = 6
|
||||
mask = AscendAttnMaskBuilder.generate_attn_mask(n, "norm", torch.float32)
|
||||
for i in range(n):
|
||||
for j in range(i + 1, n):
|
||||
self.assertEqual(mask[i, j].item(), 1.0)
|
||||
|
||||
def test_diagonal_is_zero(self):
|
||||
"""Diagonal is always zero regardless of mode/dtype."""
|
||||
for mode in ("mix", "norm"):
|
||||
for dtype in (torch.float16, torch.bfloat16, torch.float32):
|
||||
mask = AscendAttnMaskBuilder.generate_attn_mask(4, mode, dtype)
|
||||
for i in range(4):
|
||||
self.assertEqual(mask[i, i].item(), 0.0)
|
||||
|
||||
|
||||
class TestGetAttentionMaskId(unittest.TestCase):
|
||||
def test_flat_tensor(self):
|
||||
"""Produces a flat tensor of arange ranges concatenated."""
|
||||
seq_lens = torch.tensor([10, 20])
|
||||
extend_lens = torch.tensor([3, 5])
|
||||
result = AscendAttnMaskBuilder.get_attention_mask_id(seq_lens, extend_lens)
|
||||
expected = torch.tensor([7, 8, 9, 15, 16, 17, 18, 19])
|
||||
self.assertEqual(result.dim(), 1)
|
||||
self.assertTrue(torch.equal(result, expected))
|
||||
|
||||
def test_single_sequence(self):
|
||||
seq_lens = torch.tensor([5])
|
||||
extend_lens = torch.tensor([2])
|
||||
result = AscendAttnMaskBuilder.get_attention_mask_id(seq_lens, extend_lens)
|
||||
expected = torch.tensor([3, 4])
|
||||
self.assertTrue(torch.equal(result, expected))
|
||||
|
||||
def test_multiple_sequences(self):
|
||||
seq_lens = torch.tensor([3, 6, 10])
|
||||
extend_lens = torch.tensor([1, 2, 4])
|
||||
result = AscendAttnMaskBuilder.get_attention_mask_id(seq_lens, extend_lens)
|
||||
expected = torch.tensor([2, 4, 5, 6, 7, 8, 9])
|
||||
self.assertTrue(torch.equal(result, expected))
|
||||
|
||||
def test_total_length(self):
|
||||
"""Result length equals sum of extend_lens per row."""
|
||||
seq_lens = torch.tensor([10, 20])
|
||||
extend_lens = torch.tensor([3, 5])
|
||||
result = AscendAttnMaskBuilder.get_attention_mask_id(seq_lens, extend_lens)
|
||||
expected_len = 3 + 5
|
||||
self.assertEqual(result.numel(), expected_len)
|
||||
|
||||
def test_values_correct(self):
|
||||
seq_lens = torch.tensor([4, 8])
|
||||
extend_lens = torch.tensor([2, 3])
|
||||
result = AscendAttnMaskBuilder.get_attention_mask_id(seq_lens, extend_lens)
|
||||
self.assertEqual(result[0].item(), 2)
|
||||
self.assertEqual(result[1].item(), 3)
|
||||
self.assertEqual(result[2].item(), 5)
|
||||
self.assertEqual(result[4].item(), 7)
|
||||
|
||||
|
||||
class TestUpdateAttnCache(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.builder = object.__new__(AscendAttnMaskBuilder)
|
||||
self.builder.device = "cpu"
|
||||
|
||||
def test_seqlen_greater_than_cached(self):
|
||||
"""When seqlen > cached, the mask is regenerated and cache updated."""
|
||||
old_mask = torch.zeros(4, 4)
|
||||
result_mask, result_len = self.builder.update_attn_cache(
|
||||
seqlen=8,
|
||||
mask_cache=old_mask,
|
||||
seq_len_cached=4,
|
||||
dtype=torch.float16,
|
||||
mode="norm",
|
||||
)
|
||||
self.assertEqual(result_len, 8)
|
||||
self.assertEqual(result_mask.shape, (8, 8))
|
||||
self.assertEqual(result_mask.dtype, torch.float16)
|
||||
|
||||
def test_seqlen_less_equal_cached(self):
|
||||
"""When seqlen <= cached, the existing mask is kept."""
|
||||
old_mask = torch.ones(8, 8, dtype=torch.float16)
|
||||
result_mask, result_len = self.builder.update_attn_cache(
|
||||
seqlen=4,
|
||||
mask_cache=old_mask,
|
||||
seq_len_cached=8,
|
||||
dtype=torch.float16,
|
||||
mode="norm",
|
||||
)
|
||||
self.assertEqual(result_len, 8)
|
||||
self.assertIs(result_mask, old_mask)
|
||||
|
||||
def test_dtype_change(self):
|
||||
"""When dtype differs, the mask is converted but not regenerated."""
|
||||
old_mask = torch.ones(8, 8, dtype=torch.float16)
|
||||
result_mask, result_len = self.builder.update_attn_cache(
|
||||
seqlen=4,
|
||||
mask_cache=old_mask,
|
||||
seq_len_cached=8,
|
||||
dtype=torch.float32,
|
||||
mode="norm",
|
||||
)
|
||||
self.assertEqual(result_len, 8)
|
||||
self.assertEqual(result_mask.dtype, torch.float32)
|
||||
self.assertIsNot(result_mask, old_mask)
|
||||
|
||||
def test_no_change(self):
|
||||
"""When seqlen <= cached and dtype matches, nothing changes."""
|
||||
old_mask = torch.ones(8, 8, dtype=torch.float16)
|
||||
result_mask, result_len = self.builder.update_attn_cache(
|
||||
seqlen=8,
|
||||
mask_cache=old_mask,
|
||||
seq_len_cached=8,
|
||||
dtype=torch.float16,
|
||||
mode="norm",
|
||||
)
|
||||
self.assertEqual(result_len, 8)
|
||||
self.assertIs(result_mask, old_mask)
|
||||
|
||||
def test_seqlen_greater_and_dtype_change(self):
|
||||
"""Both regeneration and dtype conversion happen."""
|
||||
old_mask = torch.zeros(4, 4, dtype=torch.float16)
|
||||
result_mask, result_len = self.builder.update_attn_cache(
|
||||
seqlen=16,
|
||||
mask_cache=old_mask,
|
||||
seq_len_cached=4,
|
||||
dtype=torch.float32,
|
||||
mode="norm",
|
||||
)
|
||||
self.assertEqual(result_len, 16)
|
||||
self.assertEqual(result_mask.shape, (16, 16))
|
||||
self.assertEqual(result_mask.dtype, torch.float32)
|
||||
|
||||
def test_seqlen_equal_cached_no_regen(self):
|
||||
"""seqlen == cached should not trigger regeneration."""
|
||||
old_mask = torch.ones(8, 8, dtype=torch.float32)
|
||||
result_mask, result_len = self.builder.update_attn_cache(
|
||||
seqlen=8,
|
||||
mask_cache=old_mask,
|
||||
seq_len_cached=8,
|
||||
dtype=torch.float32,
|
||||
mode="mix",
|
||||
)
|
||||
self.assertEqual(result_len, 8)
|
||||
self.assertIs(result_mask, old_mask)
|
||||
|
||||
|
||||
class TestGetSplitfuseAttnMask(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.builder = object.__new__(AscendAttnMaskBuilder)
|
||||
self.builder.device = "cpu"
|
||||
|
||||
def test_output_shape(self):
|
||||
mask = self.builder.get_splitfuse_attn_mask(8)
|
||||
self.assertEqual(mask.shape, (8, 8))
|
||||
|
||||
def test_dtype_int8(self):
|
||||
mask = self.builder.get_splitfuse_attn_mask(4)
|
||||
self.assertEqual(mask.dtype, torch.int8)
|
||||
|
||||
def test_upper_triangular(self):
|
||||
"""Upper triangle (excluding diagonal) is 1, rest is 0."""
|
||||
mask = self.builder.get_splitfuse_attn_mask(4)
|
||||
self.assertEqual(mask[0, 1].item(), 1)
|
||||
self.assertEqual(mask[0, 3].item(), 1)
|
||||
self.assertEqual(mask[0, 0].item(), 0)
|
||||
self.assertEqual(mask[1, 0].item(), 0)
|
||||
self.assertEqual(mask[1, 1].item(), 0)
|
||||
|
||||
def test_lower_triangle_zero(self):
|
||||
n = 5
|
||||
mask = self.builder.get_splitfuse_attn_mask(n)
|
||||
for i in range(n):
|
||||
for j in range(i + 1):
|
||||
self.assertEqual(mask[i, j].item(), 0)
|
||||
|
||||
def test_upper_triangle_one(self):
|
||||
n = 5
|
||||
mask = self.builder.get_splitfuse_attn_mask(n)
|
||||
for i in range(n):
|
||||
for j in range(i + 1, n):
|
||||
self.assertEqual(mask[i, j].item(), 1)
|
||||
|
||||
|
||||
class TestGetSwaMask(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.builder = object.__new__(AscendAttnMaskBuilder)
|
||||
self.builder.device = "cpu"
|
||||
|
||||
def test_output_shape(self):
|
||||
"""Output shape is (batch, 1, s2)."""
|
||||
seq_lens = torch.tensor([5, 10])
|
||||
mask = self.builder.get_swa_mask(seq_lens, s2=15, left_context=512)
|
||||
self.assertEqual(mask.shape, (2, 1, 15))
|
||||
|
||||
def test_1d_input_unsqueezed(self):
|
||||
"""1-D input of shape (B,) is handled; output still (B, 1, s2)."""
|
||||
seq_lens = torch.tensor([3, 6, 9])
|
||||
mask = self.builder.get_swa_mask(seq_lens, s2=12, left_context=512)
|
||||
self.assertEqual(mask.shape, (3, 1, 12))
|
||||
|
||||
def test_2d_input(self):
|
||||
"""2-D input of shape (B, 1) is also accepted."""
|
||||
seq_lens = torch.tensor([[5], [10]])
|
||||
mask = self.builder.get_swa_mask(seq_lens, s2=15, left_context=512)
|
||||
self.assertEqual(mask.shape, (2, 1, 15))
|
||||
|
||||
def test_mask_values_large_left_context(self):
|
||||
"""With left_context >= max seq_len, only indices >= seq_lens are True."""
|
||||
seq_lens = torch.tensor([3, 5])
|
||||
mask = self.builder.get_swa_mask(seq_lens, s2=8, left_context=512)
|
||||
row0 = mask[0, 0]
|
||||
self.assertFalse(row0[0].item())
|
||||
self.assertFalse(row0[2].item())
|
||||
self.assertTrue(row0[3].item())
|
||||
self.assertTrue(row0[7].item())
|
||||
row1 = mask[1, 0]
|
||||
self.assertFalse(row1[4].item())
|
||||
self.assertTrue(row1[5].item())
|
||||
self.assertTrue(row1[7].item())
|
||||
|
||||
def test_mask_values_small_left_context(self):
|
||||
"""With a small left_context, earlier positions are also masked."""
|
||||
seq_lens = torch.tensor([10, 20])
|
||||
mask = self.builder.get_swa_mask(seq_lens, s2=30, left_context=5)
|
||||
row0 = mask[0, 0]
|
||||
self.assertTrue(row0[0].item())
|
||||
self.assertTrue(row0[4].item())
|
||||
self.assertFalse(row0[5].item())
|
||||
self.assertFalse(row0[9].item())
|
||||
self.assertTrue(row0[10].item())
|
||||
self.assertTrue(row0[29].item())
|
||||
row1 = mask[1, 0]
|
||||
self.assertTrue(row1[0].item())
|
||||
self.assertTrue(row1[14].item())
|
||||
self.assertFalse(row1[15].item())
|
||||
self.assertFalse(row1[19].item())
|
||||
self.assertTrue(row1[20].item())
|
||||
|
||||
def test_clamp_to_zero(self):
|
||||
"""When seq_len < left_context, start is clamped to 0."""
|
||||
seq_lens = torch.tensor([2])
|
||||
mask = self.builder.get_swa_mask(seq_lens, s2=5, left_context=10)
|
||||
row = mask[0, 0]
|
||||
self.assertFalse(row[0].item())
|
||||
self.assertFalse(row[1].item())
|
||||
self.assertTrue(row[2].item())
|
||||
self.assertTrue(row[4].item())
|
||||
|
||||
def test_default_left_context(self):
|
||||
"""Default left_context is 512."""
|
||||
seq_lens = torch.tensor([10])
|
||||
mask = self.builder.get_swa_mask(seq_lens, s2=15)
|
||||
self.assertEqual(mask.shape, (1, 1, 15))
|
||||
row = mask[0, 0]
|
||||
self.assertFalse(row[9].item())
|
||||
self.assertTrue(row[10].item())
|
||||
|
||||
def test_dtype_bool(self):
|
||||
seq_lens = torch.tensor([5, 10])
|
||||
mask = self.builder.get_swa_mask(seq_lens, s2=15, left_context=512)
|
||||
self.assertEqual(mask.dtype, torch.bool)
|
||||
|
||||
|
||||
class TestCanUseTnd(unittest.TestCase):
|
||||
def test_128_128(self):
|
||||
self.assertTrue(
|
||||
AscendAttnBackend._can_use_tnd(
|
||||
SimpleNamespace(qk_head_dim=128, v_head_dim=128)
|
||||
)
|
||||
)
|
||||
|
||||
def test_192_192(self):
|
||||
self.assertTrue(
|
||||
AscendAttnBackend._can_use_tnd(
|
||||
SimpleNamespace(qk_head_dim=192, v_head_dim=192)
|
||||
)
|
||||
)
|
||||
|
||||
def test_256_256(self):
|
||||
self.assertTrue(
|
||||
AscendAttnBackend._can_use_tnd(
|
||||
SimpleNamespace(qk_head_dim=256, v_head_dim=256)
|
||||
)
|
||||
)
|
||||
|
||||
def test_192_128(self):
|
||||
self.assertTrue(
|
||||
AscendAttnBackend._can_use_tnd(
|
||||
SimpleNamespace(qk_head_dim=192, v_head_dim=128)
|
||||
)
|
||||
)
|
||||
|
||||
def test_64_64(self):
|
||||
self.assertFalse(
|
||||
AscendAttnBackend._can_use_tnd(
|
||||
SimpleNamespace(qk_head_dim=64, v_head_dim=64)
|
||||
)
|
||||
)
|
||||
|
||||
def test_128_256(self):
|
||||
self.assertFalse(
|
||||
AscendAttnBackend._can_use_tnd(
|
||||
SimpleNamespace(qk_head_dim=128, v_head_dim=256)
|
||||
)
|
||||
)
|
||||
|
||||
def test_256_128(self):
|
||||
self.assertFalse(
|
||||
AscendAttnBackend._can_use_tnd(
|
||||
SimpleNamespace(qk_head_dim=256, v_head_dim=128)
|
||||
)
|
||||
)
|
||||
|
||||
def test_128_192(self):
|
||||
self.assertFalse(
|
||||
AscendAttnBackend._can_use_tnd(
|
||||
SimpleNamespace(qk_head_dim=128, v_head_dim=192)
|
||||
)
|
||||
)
|
||||
|
||||
def test_192_256(self):
|
||||
self.assertFalse(
|
||||
AscendAttnBackend._can_use_tnd(
|
||||
SimpleNamespace(qk_head_dim=192, v_head_dim=256)
|
||||
)
|
||||
)
|
||||
|
||||
def test_96_96(self):
|
||||
self.assertFalse(
|
||||
AscendAttnBackend._can_use_tnd(
|
||||
SimpleNamespace(qk_head_dim=96, v_head_dim=96)
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
class TestGenerateAlibiBias(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.backend = object.__new__(AscendAttnBackend)
|
||||
|
||||
def test_shape(self):
|
||||
"""Output shape is (num_heads, 1, seq_len)."""
|
||||
slopes = torch.tensor([0.1, 0.2, 0.3, 0.4])
|
||||
result = self.backend._generate_alibi_bias(
|
||||
seq_len=8,
|
||||
slopes=slopes,
|
||||
num_heads=4,
|
||||
device=torch.device("cpu"),
|
||||
dtype=torch.float32,
|
||||
)
|
||||
self.assertEqual(result.shape, (4, 1, 8))
|
||||
|
||||
def test_values(self):
|
||||
"""Each element is slopes[h] * position."""
|
||||
slopes = torch.tensor([1.0, 2.0, 3.0, 4.0])
|
||||
seq_len = 5
|
||||
result = self.backend._generate_alibi_bias(
|
||||
seq_len=seq_len,
|
||||
slopes=slopes,
|
||||
num_heads=4,
|
||||
device=torch.device("cpu"),
|
||||
dtype=torch.float32,
|
||||
)
|
||||
for h in range(4):
|
||||
for p in range(seq_len):
|
||||
expected = slopes[h].item() * p
|
||||
self.assertAlmostEqual(result[h, 0, p].item(), expected, places=5)
|
||||
|
||||
def test_dtype(self):
|
||||
slopes = torch.tensor([0.5, 1.0])
|
||||
result = self.backend._generate_alibi_bias(
|
||||
seq_len=4,
|
||||
slopes=slopes,
|
||||
num_heads=2,
|
||||
device=torch.device("cpu"),
|
||||
dtype=torch.bfloat16,
|
||||
)
|
||||
self.assertEqual(result.dtype, torch.bfloat16)
|
||||
|
||||
def test_single_head(self):
|
||||
slopes = torch.tensor([1.5])
|
||||
result = self.backend._generate_alibi_bias(
|
||||
seq_len=3,
|
||||
slopes=slopes,
|
||||
num_heads=1,
|
||||
device=torch.device("cpu"),
|
||||
dtype=torch.float32,
|
||||
)
|
||||
self.assertEqual(result.shape, (1, 1, 3))
|
||||
self.assertAlmostEqual(result[0, 0, 0].item(), 0.0)
|
||||
self.assertAlmostEqual(result[0, 0, 1].item(), 1.5)
|
||||
self.assertAlmostEqual(result[0, 0, 2].item(), 3.0)
|
||||
|
||||
def test_zero_position_is_zero(self):
|
||||
"""Position 0 always yields 0 regardless of slopes."""
|
||||
slopes = torch.tensor([1.0, 2.0, 3.0])
|
||||
result = self.backend._generate_alibi_bias(
|
||||
seq_len=4,
|
||||
slopes=slopes,
|
||||
num_heads=3,
|
||||
device=torch.device("cpu"),
|
||||
dtype=torch.float32,
|
||||
)
|
||||
for h in range(3):
|
||||
self.assertAlmostEqual(result[h, 0, 0].item(), 0.0)
|
||||
|
||||
def test_default_dtype_bfloat16(self):
|
||||
slopes = torch.tensor([0.5, 1.0])
|
||||
result = self.backend._generate_alibi_bias(
|
||||
seq_len=4,
|
||||
slopes=slopes,
|
||||
num_heads=2,
|
||||
device=torch.device("cpu"),
|
||||
)
|
||||
self.assertEqual(result.dtype, torch.bfloat16)
|
||||
|
||||
|
||||
class TestGetCudaGraphSeqLenFillValue(unittest.TestCase):
|
||||
def test_returns_zero(self):
|
||||
backend = object.__new__(AscendAttnBackend)
|
||||
self.assertEqual(backend.get_cuda_graph_seq_len_fill_value(), 0)
|
||||
|
||||
|
||||
class TestGetVerifyBuffers(unittest.TestCase):
|
||||
def test_returns_none_none(self):
|
||||
backend = object.__new__(AscendAttnBackend)
|
||||
result = backend.get_verify_buffers_to_fill_after_draft()
|
||||
self.assertEqual(result, [None, None])
|
||||
self.assertEqual(len(result), 2)
|
||||
|
||||
def test_update_is_noop(self):
|
||||
backend = object.__new__(AscendAttnBackend)
|
||||
backend.update_verify_buffers_to_fill_after_draft(None, None)
|
||||
backend.update_verify_buffers_to_fill_after_draft(MagicMock(), 4)
|
||||
backend.update_verify_buffers_to_fill_after_draft(None, 16)
|
||||
|
||||
|
||||
class TestCommonTemplate(unittest.TestCase):
|
||||
@staticmethod
|
||||
def _make_draft_backend(speculative_num_steps):
|
||||
backend = object.__new__(AscendAttnMultiStepDraftBackend)
|
||||
backend.speculative_num_steps = speculative_num_steps
|
||||
return backend
|
||||
|
||||
def test_calls_fn_for_each_step(self):
|
||||
"""call_fn is invoked for steps 0..speculative_num_steps-2."""
|
||||
backend = self._make_draft_backend(speculative_num_steps=4)
|
||||
forward_batch = MagicMock()
|
||||
forward_batch.spec_info = MagicMock()
|
||||
call_fn = MagicMock()
|
||||
backend.common_template(forward_batch, call_fn)
|
||||
self.assertEqual(call_fn.call_count, 3)
|
||||
for i in range(3):
|
||||
call_fn.assert_any_call(i, forward_batch)
|
||||
|
||||
def test_zero_steps(self):
|
||||
"""speculative_num_steps=1 -> no calls (range(0))."""
|
||||
backend = self._make_draft_backend(speculative_num_steps=1)
|
||||
forward_batch = MagicMock()
|
||||
forward_batch.spec_info = MagicMock()
|
||||
call_fn = MagicMock()
|
||||
backend.common_template(forward_batch, call_fn)
|
||||
call_fn.assert_not_called()
|
||||
|
||||
def test_two_steps(self):
|
||||
"""speculative_num_steps=2 -> exactly one call with index 0."""
|
||||
backend = self._make_draft_backend(speculative_num_steps=2)
|
||||
forward_batch = MagicMock()
|
||||
forward_batch.spec_info = MagicMock()
|
||||
call_fn = MagicMock()
|
||||
backend.common_template(forward_batch, call_fn)
|
||||
call_fn.assert_called_once_with(0, forward_batch)
|
||||
|
||||
def test_call_indices(self):
|
||||
backend = self._make_draft_backend(speculative_num_steps=5)
|
||||
forward_batch = MagicMock()
|
||||
forward_batch.spec_info = MagicMock()
|
||||
indices = []
|
||||
backend.common_template(forward_batch, lambda i, fb: indices.append(i))
|
||||
self.assertEqual(indices, [0, 1, 2, 3])
|
||||
|
||||
def test_assert_spec_info_not_none(self):
|
||||
"""Raises AssertionError when forward_batch.spec_info is None."""
|
||||
backend = self._make_draft_backend(speculative_num_steps=4)
|
||||
forward_batch = MagicMock()
|
||||
forward_batch.spec_info = None
|
||||
with self.assertRaises(AssertionError):
|
||||
backend.common_template(forward_batch, MagicMock())
|
||||
|
||||
def test_passes_same_forward_batch(self):
|
||||
backend = self._make_draft_backend(speculative_num_steps=3)
|
||||
forward_batch = MagicMock()
|
||||
forward_batch.spec_info = MagicMock()
|
||||
call_fn = MagicMock()
|
||||
backend.common_template(forward_batch, call_fn)
|
||||
for call in call_fn.call_args_list:
|
||||
self.assertIs(call.args[1], forward_batch)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,460 @@
|
||||
"""
|
||||
Unit tests for sglang.srt.hardware_backend.npu.attention.ascend_dsv4_backend.
|
||||
"""
|
||||
|
||||
import math
|
||||
import sys
|
||||
import unittest
|
||||
from types import ModuleType, SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.test.ci.ci_register import register_npu_ci
|
||||
|
||||
register_npu_ci(est_time=4, suite="stage-a-unit-test-npu")
|
||||
|
||||
for mod in (
|
||||
"torch_npu",
|
||||
"torch_npu.contrib",
|
||||
"sgl_kernel_npu",
|
||||
"sgl_kernel_npu.attention",
|
||||
"sgl_kernel_npu.attention.sinks_attention",
|
||||
"sgl_kernel_npu.norm",
|
||||
"sgl_kernel_npu.norm.add_rmsnorm_bias",
|
||||
"sglang.srt.speculative",
|
||||
"sglang.srt.speculative.decoupled_spec_io",
|
||||
"sglang.srt.speculative.spec_info",
|
||||
"sglang.srt.speculative.eagle_info",
|
||||
):
|
||||
sys.modules.setdefault(mod, MagicMock())
|
||||
|
||||
# Stub deepseek_v2._is_hip to avoid importing the heavy model.
|
||||
_ds2_stub = ModuleType("sglang.srt.models.deepseek_v2")
|
||||
_ds2_stub._is_hip = False
|
||||
sys.modules.setdefault("sglang.srt.models.deepseek_v2", _ds2_stub)
|
||||
|
||||
# Stub eagle_utils with a faithful per_step_draft_out_cache_loc.
|
||||
_eagle_stub = ModuleType("sglang.srt.speculative.eagle_utils")
|
||||
|
||||
|
||||
def _per_step_draft_out_cache_loc(out_cache_loc, batch_size, topk, num_steps):
|
||||
expected = batch_size * topk * num_steps
|
||||
assert out_cache_loc.shape[0] == expected, (
|
||||
f"out_cache_loc.shape[0]={out_cache_loc.shape[0]} != "
|
||||
f"batch_size * topk * num_steps = {batch_size}*{topk}*{num_steps}={expected}"
|
||||
)
|
||||
return (
|
||||
out_cache_loc.view(batch_size, topk, num_steps)
|
||||
.permute(2, 0, 1)
|
||||
.reshape(num_steps, -1)
|
||||
)
|
||||
|
||||
|
||||
_eagle_stub.per_step_draft_out_cache_loc = _per_step_draft_out_cache_loc
|
||||
sys.modules.setdefault("sglang.srt.speculative", ModuleType("sglang.srt.speculative"))
|
||||
sys.modules.setdefault("sglang.srt.speculative.eagle_utils", _eagle_stub)
|
||||
|
||||
from sglang.srt.hardware_backend.npu.attention.ascend_dsv4_backend import (
|
||||
DeepseekV4AscendMultiStepDraftBackend,
|
||||
_apply_hadamard,
|
||||
_get_kv_indices,
|
||||
_overlap_transform,
|
||||
_walsh_hadamard_matrix,
|
||||
)
|
||||
|
||||
|
||||
class TestWalshHadamardMatrix(unittest.TestCase):
|
||||
def test_shape_n1(self):
|
||||
had = _walsh_hadamard_matrix(1, torch.float32, "cpu")
|
||||
self.assertEqual(had.shape, (1, 1))
|
||||
|
||||
def test_shape_n2(self):
|
||||
had = _walsh_hadamard_matrix(2, torch.float32, "cpu")
|
||||
self.assertEqual(had.shape, (2, 2))
|
||||
|
||||
def test_shape_n4(self):
|
||||
had = _walsh_hadamard_matrix(4, torch.float32, "cpu")
|
||||
self.assertEqual(had.shape, (4, 4))
|
||||
|
||||
def test_value_error_n3_not_power_of_two(self):
|
||||
with self.assertRaises(ValueError):
|
||||
_walsh_hadamard_matrix(3, torch.float32, "cpu")
|
||||
|
||||
def test_value_error_n0(self):
|
||||
with self.assertRaises(ValueError):
|
||||
_walsh_hadamard_matrix(0, torch.float32, "cpu")
|
||||
|
||||
def test_value_error_negative(self):
|
||||
with self.assertRaises(ValueError):
|
||||
_walsh_hadamard_matrix(-2, torch.float32, "cpu")
|
||||
|
||||
def test_orthonormality_n2(self):
|
||||
had = _walsh_hadamard_matrix(2, torch.float32, "cpu").float()
|
||||
# bfloat16 truncates 1/sqrt(2), so use a looser tolerance
|
||||
self.assertTrue(torch.allclose(had @ had.T, torch.eye(2), atol=1e-2))
|
||||
|
||||
def test_orthonormality_n4(self):
|
||||
had = _walsh_hadamard_matrix(4, torch.float32, "cpu").float()
|
||||
self.assertTrue(torch.allclose(had @ had.T, torch.eye(4), atol=1e-2))
|
||||
|
||||
def test_orthonormality_n1(self):
|
||||
had = _walsh_hadamard_matrix(1, torch.float32, "cpu").float()
|
||||
self.assertTrue(torch.allclose(had @ had.T, torch.eye(1)))
|
||||
|
||||
def test_caching_returns_same_object(self):
|
||||
h1 = _walsh_hadamard_matrix(4, torch.float32, "cpu")
|
||||
h2 = _walsh_hadamard_matrix(4, torch.float32, "cpu")
|
||||
self.assertIs(h1, h2)
|
||||
|
||||
def test_caching_different_n_returns_different_object(self):
|
||||
h1 = _walsh_hadamard_matrix(2, torch.float32, "cpu")
|
||||
h2 = _walsh_hadamard_matrix(4, torch.float32, "cpu")
|
||||
self.assertIsNot(h1, h2)
|
||||
|
||||
def test_dtype_always_bfloat16(self):
|
||||
had = _walsh_hadamard_matrix(4, torch.float32, "cpu")
|
||||
self.assertEqual(had.dtype, torch.bfloat16)
|
||||
|
||||
def test_dtype_argument_ignored_for_cache_key(self):
|
||||
h1 = _walsh_hadamard_matrix(4, torch.float32, "cpu")
|
||||
h2 = _walsh_hadamard_matrix(4, torch.bfloat16, "cpu")
|
||||
self.assertIs(h1, h2)
|
||||
|
||||
def test_entries_are_plus_minus_norm(self):
|
||||
n = 4
|
||||
had = _walsh_hadamard_matrix(n, torch.float32, "cpu").float()
|
||||
expected_abs = 1.0 / math.sqrt(n)
|
||||
self.assertTrue(
|
||||
torch.allclose(had.abs(), torch.full_like(had, expected_abs), atol=1e-2)
|
||||
)
|
||||
|
||||
|
||||
class TestApplyHadamard(unittest.TestCase):
|
||||
def test_shape_preserved_2d(self):
|
||||
n = 4
|
||||
H = _walsh_hadamard_matrix(n, torch.float32, "cpu")
|
||||
inp = torch.randn(3, n, dtype=H.dtype)
|
||||
out = _apply_hadamard(inp, H)
|
||||
self.assertEqual(out.shape, inp.shape)
|
||||
|
||||
def test_shape_preserved_3d(self):
|
||||
n = 4
|
||||
H = _walsh_hadamard_matrix(n, torch.float32, "cpu")
|
||||
inp = torch.randn(2, 5, n, dtype=H.dtype)
|
||||
out = _apply_hadamard(inp, H)
|
||||
self.assertEqual(out.shape, inp.shape)
|
||||
|
||||
def test_identity_times_hadamard_equals_hadamard(self):
|
||||
n = 4
|
||||
H = _walsh_hadamard_matrix(n, torch.float32, "cpu")
|
||||
eye = torch.eye(n, dtype=H.dtype)
|
||||
out = _apply_hadamard(eye, H)
|
||||
self.assertTrue(torch.equal(out, H))
|
||||
|
||||
def test_identity_times_hadamard_n2(self):
|
||||
n = 2
|
||||
H = _walsh_hadamard_matrix(n, torch.float32, "cpu")
|
||||
eye = torch.eye(n, dtype=H.dtype)
|
||||
out = _apply_hadamard(eye, H)
|
||||
self.assertTrue(torch.equal(out, H))
|
||||
|
||||
def test_output_dtype_is_bfloat16(self):
|
||||
n = 4
|
||||
H = _walsh_hadamard_matrix(n, torch.float32, "cpu")
|
||||
inp = torch.randn(3, n, dtype=H.dtype)
|
||||
out = _apply_hadamard(inp, H)
|
||||
self.assertEqual(out.dtype, torch.bfloat16)
|
||||
|
||||
def test_output_dtype_bfloat16_from_float32_input(self):
|
||||
n = 4
|
||||
H = _walsh_hadamard_matrix(n, torch.bfloat16, "cpu")
|
||||
inp = torch.randn(3, n, dtype=torch.bfloat16)
|
||||
out = _apply_hadamard(inp, H)
|
||||
self.assertEqual(out.dtype, torch.bfloat16)
|
||||
|
||||
def test_3d_values(self):
|
||||
n = 2
|
||||
H = _walsh_hadamard_matrix(n, torch.float32, "cpu").float()
|
||||
inp = torch.randn(2, 3, n, dtype=torch.float32)
|
||||
expected = inp.matmul(H).to(torch.bfloat16)
|
||||
out = _apply_hadamard(inp, H)
|
||||
self.assertTrue(torch.equal(out, expected))
|
||||
|
||||
|
||||
class TestOverlapTransform(unittest.TestCase):
|
||||
def test_shape(self):
|
||||
# (n_chunks, ratio, 2*d) -> (n_chunks, 2*ratio, d)
|
||||
n_chunks, r, d = 3, 2, 4
|
||||
tensor = torch.randn(n_chunks, r, 2 * d)
|
||||
out = _overlap_transform(tensor, value=0.0, head_dim=d)
|
||||
self.assertEqual(out.shape, (n_chunks, 2 * r, d))
|
||||
|
||||
def test_first_chunk_left_half_filled_with_value(self):
|
||||
n_chunks, r, d = 3, 2, 4
|
||||
tensor = torch.randn(n_chunks, r, 2 * d)
|
||||
fill = float("-inf")
|
||||
out = _overlap_transform(tensor, value=fill, head_dim=d)
|
||||
self.assertTrue(torch.equal(out[0, :r], torch.full((r, d), fill)))
|
||||
|
||||
def test_first_chunk_left_half_filled_with_zero(self):
|
||||
n_chunks, r, d = 2, 2, 4
|
||||
tensor = torch.randn(n_chunks, r, 2 * d)
|
||||
out = _overlap_transform(tensor, value=0.0, head_dim=d)
|
||||
self.assertTrue(torch.equal(out[0, :r], torch.zeros(r, d)))
|
||||
|
||||
def test_right_half_mirrors_tensor_second_half(self):
|
||||
n_chunks, r, d = 3, 2, 4
|
||||
tensor = torch.randn(n_chunks, r, 2 * d)
|
||||
out = _overlap_transform(tensor, value=0.0, head_dim=d)
|
||||
self.assertTrue(torch.equal(out[:, r:], tensor[..., d:]))
|
||||
|
||||
def test_previous_chunk_left_half(self):
|
||||
n_chunks, r, d = 3, 2, 4
|
||||
tensor = torch.randn(n_chunks, r, 2 * d)
|
||||
out = _overlap_transform(tensor, value=0.0, head_dim=d)
|
||||
self.assertTrue(torch.equal(out[1:, :r], tensor[:-1, :, :d]))
|
||||
|
||||
def test_single_chunk(self):
|
||||
n_chunks, r, d = 1, 2, 4
|
||||
tensor = torch.randn(n_chunks, r, 2 * d)
|
||||
fill = 7.0
|
||||
out = _overlap_transform(tensor, value=fill, head_dim=d)
|
||||
self.assertEqual(out.shape, (1, 2 * r, d))
|
||||
self.assertTrue(torch.equal(out[0, :r], torch.full((r, d), fill)))
|
||||
self.assertTrue(torch.equal(out[0, r:], tensor[0, :, d:]))
|
||||
|
||||
def test_full_element_mapping(self):
|
||||
n_chunks, r, d = 2, 2, 3
|
||||
tensor = torch.arange(n_chunks * r * 2 * d, dtype=torch.float32).reshape(
|
||||
n_chunks, r, 2 * d
|
||||
)
|
||||
fill = -1.0
|
||||
out = _overlap_transform(tensor, value=fill, head_dim=d)
|
||||
|
||||
for c in range(n_chunks):
|
||||
for row in range(2 * r):
|
||||
for col in range(d):
|
||||
if c == 0 and row < r:
|
||||
expected = fill
|
||||
elif row >= r:
|
||||
expected = tensor[c, row - r, d + col].item()
|
||||
else:
|
||||
expected = tensor[c - 1, row, col].item()
|
||||
self.assertEqual(
|
||||
out[c, row, col].item(),
|
||||
expected,
|
||||
f"mismatch at (c={c}, row={row}, col={col})",
|
||||
)
|
||||
|
||||
def test_preserves_input_dtype(self):
|
||||
n_chunks, r, d = 2, 2, 4
|
||||
tensor = torch.randn(n_chunks, r, 2 * d, dtype=torch.bfloat16)
|
||||
out = _overlap_transform(tensor, value=0.0, head_dim=d)
|
||||
self.assertEqual(out.dtype, torch.bfloat16)
|
||||
|
||||
|
||||
class TestGetKvIndices(unittest.TestCase):
|
||||
_PATCH_TARGET = (
|
||||
"sglang.srt.hardware_backend.npu.attention.ascend_dsv4_backend.get_attn_backend"
|
||||
)
|
||||
|
||||
@patch(_PATCH_TARGET)
|
||||
def test_page_size_one_plain_slice(self, mock_get_attn_backend):
|
||||
mock_get_attn_backend.return_value = SimpleNamespace(page_size=1)
|
||||
page_table = torch.arange(16, dtype=torch.int32).reshape(2, 8)
|
||||
# req_idx=0, seqlen=5, kv_len=5 -> logic_start=0, logic_end=5
|
||||
result = _get_kv_indices(MagicMock(), 5, page_table, 0, 5)
|
||||
expected = page_table[0, 0:5]
|
||||
self.assertEqual(result.tolist(), expected.tolist())
|
||||
|
||||
@patch(_PATCH_TARGET)
|
||||
def test_page_size_one_partial_window(self, mock_get_attn_backend):
|
||||
mock_get_attn_backend.return_value = SimpleNamespace(page_size=1)
|
||||
page_table = torch.arange(16, dtype=torch.int32).reshape(2, 8)
|
||||
# req_idx=0, seqlen=10, kv_len=4 -> logic_start=6, logic_end=10
|
||||
result = _get_kv_indices(MagicMock(), 4, page_table, 0, 10)
|
||||
expected = page_table[0, 6:10]
|
||||
self.assertEqual(result.tolist(), expected.tolist())
|
||||
|
||||
@patch(_PATCH_TARGET)
|
||||
def test_page_size_gt_one_paged(self, mock_get_attn_backend):
|
||||
page_size = 4
|
||||
mock_get_attn_backend.return_value = SimpleNamespace(page_size=page_size)
|
||||
page_table = torch.tensor(
|
||||
[[10, 20, 30, 40], [50, 60, 70, 80]], dtype=torch.int32
|
||||
)
|
||||
# req_idx=0, seqlen=6, kv_len=6 -> logic_pos=[0..5]; block_id=[0,0,0,0,1,1]
|
||||
# page_table[0, block_id]=[10,10,10,10,20,20]; physical=[40,41,42,43,80,81]
|
||||
result = _get_kv_indices(MagicMock(), 6, page_table, 0, 6)
|
||||
expected = [40, 41, 42, 43, 80, 81]
|
||||
self.assertEqual(result.tolist(), expected)
|
||||
|
||||
@patch(_PATCH_TARGET)
|
||||
def test_page_size_gt_one_partial_window(self, mock_get_attn_backend):
|
||||
page_size = 4
|
||||
mock_get_attn_backend.return_value = SimpleNamespace(page_size=page_size)
|
||||
page_table = torch.tensor(
|
||||
[[10, 20, 30, 40], [50, 60, 70, 80]], dtype=torch.int32
|
||||
)
|
||||
# req_idx=0, seqlen=10, kv_len=4 -> logic_pos=[6,7,8,9]; block_id=[1,1,2,2]
|
||||
# physical=[82,83,120,121]
|
||||
result = _get_kv_indices(MagicMock(), 4, page_table, 0, 10)
|
||||
expected = [82, 83, 120, 121]
|
||||
self.assertEqual(result.tolist(), expected)
|
||||
|
||||
@patch(_PATCH_TARGET)
|
||||
def test_page_size_gt_one_second_request(self, mock_get_attn_backend):
|
||||
page_size = 4
|
||||
mock_get_attn_backend.return_value = SimpleNamespace(page_size=page_size)
|
||||
page_table = torch.tensor(
|
||||
[[10, 20, 30, 40], [50, 60, 70, 80]], dtype=torch.int32
|
||||
)
|
||||
# req_idx=1, seqlen=6 -> page_table[1, block_id]=[50,50,50,50,60,60]
|
||||
# physical=[200,201,202,203,240,241]
|
||||
result = _get_kv_indices(MagicMock(), 6, page_table, 1, 6)
|
||||
expected = [200, 201, 202, 203, 240, 241]
|
||||
self.assertEqual(result.tolist(), expected)
|
||||
|
||||
@patch(_PATCH_TARGET)
|
||||
def test_kv_len_clamped_to_zero(self, mock_get_attn_backend):
|
||||
# kv_len > seqlen -> logic_start = max(0, seqlen - kv_len) = 0
|
||||
mock_get_attn_backend.return_value = SimpleNamespace(page_size=1)
|
||||
page_table = torch.arange(16, dtype=torch.int32).reshape(2, 8)
|
||||
result = _get_kv_indices(MagicMock(), 100, page_table, 0, 3)
|
||||
expected = page_table[0, 0:3]
|
||||
self.assertEqual(result.tolist(), expected.tolist())
|
||||
|
||||
|
||||
class TestStepOutCacheLoc(unittest.TestCase):
|
||||
def _make_backend(self, topk, speculative_num_steps):
|
||||
backend = object.__new__(DeepseekV4AscendMultiStepDraftBackend)
|
||||
backend.topk = topk
|
||||
backend.speculative_num_steps = speculative_num_steps
|
||||
return backend
|
||||
|
||||
def test_none_out_cache_loc_returns_none(self):
|
||||
backend = self._make_backend(topk=2, speculative_num_steps=3)
|
||||
forward_batch = SimpleNamespace(out_cache_loc=None, batch_size=4)
|
||||
self.assertIsNone(backend._step_out_cache_loc(forward_batch, 0))
|
||||
|
||||
def test_short_out_cache_loc_returns_as_is(self):
|
||||
# numel <= single_step_width (batch_size * topk) -> returned unchanged
|
||||
backend = self._make_backend(topk=2, speculative_num_steps=3)
|
||||
loc = torch.tensor([10, 20, 30], dtype=torch.int32)
|
||||
forward_batch = SimpleNamespace(out_cache_loc=loc, batch_size=4)
|
||||
result = backend._step_out_cache_loc(forward_batch, 0)
|
||||
self.assertIs(result, loc)
|
||||
|
||||
def test_short_out_cache_loc_boundary_equal(self):
|
||||
backend = self._make_backend(topk=2, speculative_num_steps=3)
|
||||
loc = torch.arange(8, dtype=torch.int32)
|
||||
forward_batch = SimpleNamespace(out_cache_loc=loc, batch_size=4)
|
||||
# single_step_width = 4*2 = 8; numel=8 <= 8
|
||||
result = backend._step_out_cache_loc(forward_batch, 0)
|
||||
self.assertIs(result, loc)
|
||||
|
||||
def test_indivisible_returns_as_is(self):
|
||||
backend = self._make_backend(topk=2, speculative_num_steps=3)
|
||||
# step_layout_width = 2*3 = 6; numel=10, 10 % 6 != 0
|
||||
loc = torch.arange(10, dtype=torch.int32)
|
||||
forward_batch = SimpleNamespace(out_cache_loc=loc, batch_size=2)
|
||||
result = backend._step_out_cache_loc(forward_batch, 0)
|
||||
self.assertIs(result, loc)
|
||||
|
||||
def test_step_layout_width_zero_returns_as_is(self):
|
||||
backend = self._make_backend(topk=0, speculative_num_steps=3)
|
||||
loc = torch.arange(5, dtype=torch.int32)
|
||||
forward_batch = SimpleNamespace(out_cache_loc=loc, batch_size=2)
|
||||
result = backend._step_out_cache_loc(forward_batch, 0)
|
||||
self.assertIs(result, loc)
|
||||
|
||||
def test_normal_case_step0(self):
|
||||
backend = self._make_backend(topk=2, speculative_num_steps=3)
|
||||
loc = torch.arange(12, dtype=torch.int32)
|
||||
forward_batch = SimpleNamespace(out_cache_loc=loc, batch_size=2)
|
||||
# view(2,2,3).permute(2,0,1).reshape(3,-1); step 0: [0,3,6,9]
|
||||
result = backend._step_out_cache_loc(forward_batch, 0)
|
||||
self.assertEqual(result.tolist(), [0, 3, 6, 9])
|
||||
|
||||
def test_normal_case_step1(self):
|
||||
backend = self._make_backend(topk=2, speculative_num_steps=3)
|
||||
loc = torch.arange(12, dtype=torch.int32)
|
||||
forward_batch = SimpleNamespace(out_cache_loc=loc, batch_size=2)
|
||||
result = backend._step_out_cache_loc(forward_batch, 1)
|
||||
self.assertEqual(result.tolist(), [1, 4, 7, 10])
|
||||
|
||||
def test_normal_case_step2(self):
|
||||
backend = self._make_backend(topk=2, speculative_num_steps=3)
|
||||
loc = torch.arange(12, dtype=torch.int32)
|
||||
forward_batch = SimpleNamespace(out_cache_loc=loc, batch_size=2)
|
||||
result = backend._step_out_cache_loc(forward_batch, 2)
|
||||
self.assertEqual(result.tolist(), [2, 5, 8, 11])
|
||||
|
||||
def test_normal_case_different_dimensions(self):
|
||||
backend = self._make_backend(topk=3, speculative_num_steps=2)
|
||||
loc = torch.arange(24, dtype=torch.int32)
|
||||
forward_batch = SimpleNamespace(out_cache_loc=loc, batch_size=4)
|
||||
# view(4,3,2).permute(2,0,1).reshape(2,-1)
|
||||
result = backend._step_out_cache_loc(forward_batch, 0)
|
||||
self.assertEqual(result.tolist(), [0, 2, 4, 6, 8, 10, 12, 14, 16, 18, 20, 22])
|
||||
|
||||
def test_normal_case_returns_tensor(self):
|
||||
backend = self._make_backend(topk=2, speculative_num_steps=3)
|
||||
loc = torch.arange(12, dtype=torch.int32)
|
||||
forward_batch = SimpleNamespace(out_cache_loc=loc, batch_size=2)
|
||||
result = backend._step_out_cache_loc(forward_batch, 0)
|
||||
self.assertIsInstance(result, torch.Tensor)
|
||||
|
||||
|
||||
class TestCommonTemplate(unittest.TestCase):
|
||||
def _make_backend(self, speculative_num_steps):
|
||||
backend = object.__new__(DeepseekV4AscendMultiStepDraftBackend)
|
||||
backend.speculative_num_steps = speculative_num_steps
|
||||
return backend
|
||||
|
||||
def test_calls_call_fn_for_each_step(self):
|
||||
backend = self._make_backend(speculative_num_steps=4)
|
||||
forward_batch = SimpleNamespace(spec_info=object())
|
||||
call_fn = MagicMock()
|
||||
backend.common_template(forward_batch, call_fn)
|
||||
# range(speculative_num_steps - 1) = range(3) -> i=0,1,2
|
||||
self.assertEqual(call_fn.call_count, 3)
|
||||
for i, call in enumerate(call_fn.call_args_list):
|
||||
self.assertEqual(call.args[0], i)
|
||||
self.assertIs(call.args[1], forward_batch)
|
||||
|
||||
def test_single_step_no_calls(self):
|
||||
backend = self._make_backend(speculative_num_steps=1)
|
||||
forward_batch = SimpleNamespace(spec_info=object())
|
||||
call_fn = MagicMock()
|
||||
backend.common_template(forward_batch, call_fn)
|
||||
self.assertEqual(call_fn.call_count, 0)
|
||||
|
||||
def test_two_steps_one_call(self):
|
||||
backend = self._make_backend(speculative_num_steps=2)
|
||||
forward_batch = SimpleNamespace(spec_info=object())
|
||||
call_fn = MagicMock()
|
||||
backend.common_template(forward_batch, call_fn)
|
||||
self.assertEqual(call_fn.call_count, 1)
|
||||
self.assertEqual(call_fn.call_args_list[0].args[0], 0)
|
||||
|
||||
def test_asserts_spec_info_not_none(self):
|
||||
backend = self._make_backend(speculative_num_steps=3)
|
||||
forward_batch = SimpleNamespace(spec_info=None)
|
||||
call_fn = MagicMock()
|
||||
with self.assertRaises(AssertionError):
|
||||
backend.common_template(forward_batch, call_fn)
|
||||
self.assertEqual(call_fn.call_count, 0)
|
||||
|
||||
def test_call_fn_exception_propagates(self):
|
||||
backend = self._make_backend(speculative_num_steps=3)
|
||||
forward_batch = SimpleNamespace(spec_info=object())
|
||||
call_fn = MagicMock(side_effect=RuntimeError("boom"))
|
||||
with self.assertRaises(RuntimeError):
|
||||
backend.common_template(forward_batch, call_fn)
|
||||
self.assertEqual(call_fn.call_count, 1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -99,6 +99,7 @@ PER_COMMIT_SUITES = {
|
||||
],
|
||||
HWBackend.NPU: [
|
||||
"base-a-test-1-gpu-small",
|
||||
"stage-a-unit-test-npu",
|
||||
"stage-b-test-1-npu-a2",
|
||||
"stage-b-test-2-npu-a2",
|
||||
"stage-b-test-4-npu-a3",
|
||||
@@ -337,6 +338,7 @@ def run_a_suite(args):
|
||||
if not f.endswith("/conftest.py")
|
||||
and not f.endswith("/__init__.py")
|
||||
and not f.endswith("/cpu/utils.py")
|
||||
and not f.endswith("/run_tests.py")
|
||||
]
|
||||
|
||||
# Strict: all discovered files must have proper registration
|
||||
|
||||
Reference in New Issue
Block a user