diff --git a/.github/workflows/pr-test-npu.yml b/.github/workflows/pr-test-npu.yml index d0186b204..8e3b1510d 100644 --- a/.github/workflows/pr-test-npu.yml +++ b/.github/workflows/pr-test-npu.yml @@ -60,6 +60,7 @@ jobs: - "python/pyproject_npu.toml" - "scripts/ci/npu/npu_ci_install_dependency.sh" - "test/registered/ascend/**" + - "test/registered/unit/npu/**" - ".github/workflows/pr-test-npu.yml" multimodal_gen: - "python/sglang/multimodal_gen/**/!(*.md|*.ipynb)" @@ -89,6 +90,46 @@ jobs: echo "CANN_image_a3=swr.cn-southwest-2.myhuaweicloud.com/base_image/ascend-ci/cann:9.0.0-a3-ubuntu22.04-py3.11" >> $GITHUB_OUTPUT echo "CANN_image_910b=swr.cn-southwest-2.myhuaweicloud.com/base_image/ascend-ci/cann:9.0.0-910b-ubuntu22.04-py3.11" >> $GITHUB_OUTPUT + stage-a-unit-test-npu: + needs: [check-changes, pr-gate, set-image-config] + if: needs.check-changes.outputs.main_package == 'true' + runs-on: linux-aarch64-a2-1 + container: + image: ${{ needs.set-image-config.outputs.CANN_image_910b }} + steps: + - name: Checkout code + uses: actions/checkout@v4 + with: + ref: ${{ inputs.ref || github.ref }} + + - name: Mark repository safe + run: | + git config --system --add safe.directory ${GITHUB_WORKSPACE} + + - name: Install dependencies + env: + TORCH_CACHE_URL: "http://cache-service.nginx-pypi-cache.svc.cluster.local/whl/cpu" + PYPI_CACHE_URL: "http://cache-service.nginx-pypi-cache.svc.cluster.local/pypi/simple" + UV_INDEX_URL: "http://cache-service.nginx-pypi-cache.svc.cluster.local/pypi/simple" + GITHUB_PROXY_URL: "https://gh-proxy.test.osinfra.cn/" + RUSTUP_CACHE_URL: "http://cache-service.nginx-pypi-cache.svc.cluster.local:8082" + run: | + # speed up by using infra cache services + CACHING_URL="cache-service.nginx-pypi-cache.svc.cluster.local" + sed -Ei "s@(ports|archive).ubuntu.com@${CACHING_URL}:8081@g" /etc/apt/sources.list + pip config set global.index-url http://${CACHING_URL}/pypi/simple + pip config set global.trusted-host "${CACHING_URL}" + + bash scripts/ci/npu/npu_ci_install_dependency.sh 910b + + - name: Run test + timeout-minutes: 15 + env: + SGLANG_IS_IN_CI: true + run: | + cd test + python3 run_suite.py --hw npu --suite stage-a-unit-test-npu + stage-b-test-1-npu-a2: needs: [check-changes, pr-gate, set-image-config] if: needs.check-changes.outputs.main_package == 'true' @@ -439,6 +480,7 @@ jobs: [ check-changes, + stage-a-unit-test-npu, stage-b-test-1-npu-a2, stage-b-test-2-npu-a2, stage-b-test-4-npu-a3, diff --git a/test/registered/unit/npu/attention/test_npu_ascend_backend.py b/test/registered/unit/npu/attention/test_npu_ascend_backend.py new file mode 100644 index 000000000..e9955cd77 --- /dev/null +++ b/test/registered/unit/npu/attention/test_npu_ascend_backend.py @@ -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() diff --git a/test/registered/unit/npu/attention/test_npu_ascend_dsv4_backend.py b/test/registered/unit/npu/attention/test_npu_ascend_dsv4_backend.py new file mode 100644 index 000000000..5cf5f1e16 --- /dev/null +++ b/test/registered/unit/npu/attention/test_npu_ascend_dsv4_backend.py @@ -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() diff --git a/test/run_suite.py b/test/run_suite.py index a6907ac1f..7d0a73c9b 100644 --- a/test/run_suite.py +++ b/test/run_suite.py @@ -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