[NPU] Add NPU arch35 support and enhance DSV4 processing in DeepSeek-V4 (#37373)
Co-authored-by: AndyLi429 <AndyLi429@noreply.gitcode.com> Co-authored-by: Kailong Lu <kelonlu@163.com> Co-authored-by: cx <chengxin65@huawei.com> Co-authored-by: ranjiewen <ranjiewen@huawei.com> Co-authored-by: HEX1A0A <1a0ahex@gmail.com> Co-authored-by: vstone-w <374330057@qq.com> Co-authored-by: Even Zhou <even.y.zhou@outlook.com> Co-authored-by: ClownBin <chaobin1993@126.com> Co-authored-by: sglang-npu-bot <sglangnpu@163.com>
This commit is contained in:
co-authored by
AndyLi429
Kailong Lu
cx
ranjiewen
HEX1A0A
vstone-w
Even Zhou
ClownBin
sglang-npu-bot
parent
df623d3cbd
commit
62a4a6ea0e
@@ -32,6 +32,7 @@ from sglang.srt.disaggregation.mooncake.conn import (
|
||||
)
|
||||
from sglang.srt.disaggregation.utils import (
|
||||
MetadataBuffers,
|
||||
get_dsv4_c4_state_indices,
|
||||
get_dsv4_c128_state_indices,
|
||||
setup_state_kv_args,
|
||||
)
|
||||
@@ -365,7 +366,12 @@ class TestEagleDsaSeedTransfer(unittest.TestCase):
|
||||
buffers.set_buf(self._make_req(seed))
|
||||
buffers.set_buf(self._make_req(None, metadata_buffer_index=1))
|
||||
|
||||
self.assertTrue(torch.equal(buffers.output_dsa_topk_indices[0], seed))
|
||||
self.assertTrue(
|
||||
torch.equal(
|
||||
buffers.output_dsa_topk_indices[0],
|
||||
seed.to(buffers.output_dsa_topk_indices.device),
|
||||
)
|
||||
)
|
||||
self.assertEqual(buffers.output_dsa_topk_indices[1].tolist(), [-1, -1, -1])
|
||||
ptrs, data_lens, item_lens = buffers.get_buf_infos()
|
||||
self.assertEqual(ptrs[-2], buffers.output_dsa_topk_indices.data_ptr())
|
||||
@@ -519,6 +525,39 @@ class TestEagleDsaSeedTransfer(unittest.TestCase):
|
||||
self.assertEqual(future_map.topk_index_buf.shape, (4, 3))
|
||||
|
||||
|
||||
class TestDSV4C4StateIndices(unittest.TestCase):
|
||||
def test_non_mtp_to_mtp_maps_the_same_logical_positions(self):
|
||||
# seq_len=13 keeps logical positions [8, 13) for the overlap C4 state.
|
||||
src = get_dsv4_c4_state_indices(2, 13, ring_size=8)
|
||||
dst = get_dsv4_c4_state_indices(2, 13, ring_size=16)
|
||||
|
||||
np.testing.assert_array_equal(src, np.array([16, 17, 18, 19, 20]))
|
||||
np.testing.assert_array_equal(dst, np.array([40, 41, 42, 43, 44]))
|
||||
self.assertEqual(src.size, dst.size)
|
||||
|
||||
def test_ring_wrap_preserves_position_order(self):
|
||||
np.testing.assert_array_equal(
|
||||
get_dsv4_c4_state_indices(0, 10, ring_size=8),
|
||||
np.array([4, 5, 6, 7, 0, 1], dtype=np.int32),
|
||||
)
|
||||
|
||||
def test_short_and_empty_sequences(self):
|
||||
np.testing.assert_array_equal(
|
||||
get_dsv4_c4_state_indices(3, 3, ring_size=8),
|
||||
np.array([24, 25, 26], dtype=np.int32),
|
||||
)
|
||||
np.testing.assert_array_equal(
|
||||
get_dsv4_c4_state_indices(3, 0, ring_size=8),
|
||||
np.empty((0,), dtype=np.int32),
|
||||
)
|
||||
|
||||
def test_invalid_ring_size_is_rejected(self):
|
||||
with self.assertRaises(ValueError):
|
||||
get_dsv4_c4_state_indices(0, 8, ring_size=4)
|
||||
with self.assertRaises(ValueError):
|
||||
get_dsv4_c4_state_indices(0, 8, ring_size=10)
|
||||
|
||||
|
||||
class TestDSV4C128StateIndices(unittest.TestCase):
|
||||
def test_online_aligned_boundary_has_no_partial_state(self):
|
||||
np.testing.assert_array_equal(
|
||||
|
||||
@@ -37,7 +37,10 @@ def _fake_group() -> SimpleNamespace:
|
||||
|
||||
|
||||
def _make_receiver(ps: ParallelState) -> SchedulerRequestReceiver:
|
||||
group = _fake_group()
|
||||
tp_group = _fake_group()
|
||||
attn_tp_group = _fake_group()
|
||||
attn_cp_group = _fake_group()
|
||||
world_group = _fake_group()
|
||||
return SchedulerRequestReceiver(
|
||||
recv_from_tokenizer=None,
|
||||
recv_from_rpc=None,
|
||||
@@ -45,13 +48,13 @@ def _make_receiver(ps: ParallelState) -> SchedulerRequestReceiver:
|
||||
input_blocker=None,
|
||||
mm_receiver=None,
|
||||
ps=ps,
|
||||
tp_group=group,
|
||||
tp_cpu_group=group,
|
||||
attn_tp_group=group,
|
||||
attn_tp_cpu_group=group,
|
||||
attn_cp_group=group,
|
||||
attn_cp_cpu_group=group,
|
||||
world_group=group,
|
||||
tp_group=tp_group,
|
||||
tp_cpu_group=tp_group,
|
||||
attn_tp_group=attn_tp_group,
|
||||
attn_tp_cpu_group=attn_tp_group,
|
||||
attn_cp_group=attn_cp_group,
|
||||
attn_cp_cpu_group=attn_cp_group,
|
||||
world_group=world_group,
|
||||
server_args=SimpleNamespace(
|
||||
enable_dp_attention=True,
|
||||
enable_dp_attention_local_control_broadcast=False,
|
||||
@@ -63,6 +66,94 @@ def _make_receiver(ps: ParallelState) -> SchedulerRequestReceiver:
|
||||
)
|
||||
|
||||
|
||||
class TestRequestReceiverBroadcast(unittest.TestCase):
|
||||
def test_local_control_skips_full_tp_broadcast_for_decode_dp(self):
|
||||
# Decode uses pure DP attention (attn_tp=attn_cp=1). The DP controller
|
||||
# sends control requests to every local leader, so no per-tick Gloo
|
||||
# broadcast should remain in SchedulerRequestReceiver.
|
||||
ps = SimpleNamespace(
|
||||
attn_tp_rank=0,
|
||||
attn_cp_rank=0,
|
||||
attn_tp_size=1,
|
||||
attn_cp_size=1,
|
||||
tp_size=32,
|
||||
)
|
||||
receiver = _make_receiver(ps)
|
||||
control_req = SimpleNamespace(kind="control")
|
||||
parallel = SimpleNamespace(
|
||||
enable_dp_attention=True,
|
||||
enable_dp_attention_local_control_broadcast=True,
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"sglang.srt.managers.scheduler_components.request_receiver."
|
||||
"get_parallel",
|
||||
return_value=parallel,
|
||||
),
|
||||
patch(
|
||||
"sglang.srt.managers.scheduler_components.request_receiver."
|
||||
"attn_cp_tp_broadcast_pyobj",
|
||||
side_effect=lambda requests: requests,
|
||||
),
|
||||
patch(
|
||||
"sglang.srt.managers.scheduler_components.request_receiver."
|
||||
"broadcast_pyobj"
|
||||
) as broadcast,
|
||||
):
|
||||
result = receiver._broadcast_reqs_across_ranks([control_req])
|
||||
|
||||
self.assertEqual(result, [control_req])
|
||||
broadcast.assert_not_called()
|
||||
|
||||
def test_default_control_uses_full_tp_broadcast(self):
|
||||
ps = SimpleNamespace(
|
||||
attn_tp_rank=0,
|
||||
attn_cp_rank=0,
|
||||
attn_tp_size=1,
|
||||
attn_cp_size=1,
|
||||
tp_size=32,
|
||||
)
|
||||
receiver = _make_receiver(ps)
|
||||
control_req = SimpleNamespace(kind="control")
|
||||
parallel = SimpleNamespace(
|
||||
enable_dp_attention=True,
|
||||
enable_dp_attention_local_control_broadcast=False,
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"sglang.srt.managers.scheduler_components.request_receiver."
|
||||
"get_parallel",
|
||||
return_value=parallel,
|
||||
),
|
||||
patch(
|
||||
"sglang.srt.managers.scheduler_components.request_receiver."
|
||||
"is_ep_scale_joiner",
|
||||
return_value=False,
|
||||
),
|
||||
patch(
|
||||
"sglang.srt.managers.scheduler_components.request_receiver."
|
||||
"attn_cp_tp_broadcast_pyobj",
|
||||
side_effect=lambda requests: requests,
|
||||
),
|
||||
patch(
|
||||
"sglang.srt.managers.scheduler_components.request_receiver."
|
||||
"broadcast_pyobj",
|
||||
side_effect=lambda requests, *_args, **_kwargs: requests,
|
||||
) as broadcast,
|
||||
):
|
||||
result = receiver._broadcast_reqs_across_ranks([control_req])
|
||||
|
||||
self.assertEqual(result, [control_req])
|
||||
broadcast.assert_called_once_with(
|
||||
[control_req],
|
||||
receiver.tp_group.rank,
|
||||
receiver.tp_cpu_group,
|
||||
src=receiver.tp_group.ranks[0],
|
||||
)
|
||||
|
||||
|
||||
class TestPPCPRankOffsets(unittest.TestCase):
|
||||
def test_request_receiver_uses_cp_size_for_pp_recv_rank(self):
|
||||
ps = _make_ps()
|
||||
|
||||
@@ -56,11 +56,328 @@ sys.modules.setdefault("sglang.srt.speculative", ModuleType("sglang.srt.speculat
|
||||
sys.modules.setdefault("sglang.srt.speculative.eagle_utils", _eagle_stub)
|
||||
|
||||
from sglang.srt.hardware_backend.npu.attention.ascend_dsv4_backend import (
|
||||
C4IndexerAscendBackendMixin,
|
||||
CompressorAscendBackendMixin,
|
||||
DeepseekV4AscendAttnBackend,
|
||||
DeepseekV4AscendMultiStepDraftBackend,
|
||||
_apply_hadamard,
|
||||
_build_cycle_state_block_table,
|
||||
_get_kv_indices,
|
||||
_sparse_attn_kv_quant_kwargs,
|
||||
_sparse_attn_ops,
|
||||
_walsh_hadamard_matrix,
|
||||
)
|
||||
from sglang.srt.hardware_backend.npu.dsv4.dsv4_common_hooks import (
|
||||
dsv4_state_payloads,
|
||||
)
|
||||
from sglang.srt.hardware_backend.npu.dsv4.dsv4_memory_pool import DSV4NPUTokenToKVPool
|
||||
from sglang.srt.hardware_backend.npu.dsv4.dsv4_req_to_token_pool import (
|
||||
DSV4ReqToTokenTablesMixin,
|
||||
)
|
||||
|
||||
|
||||
class TestVerifyCompressPositions(unittest.TestCase):
|
||||
@staticmethod
|
||||
def _backend():
|
||||
backend = DeepseekV4AscendAttnBackend.__new__(DeepseekV4AscendAttnBackend)
|
||||
backend._dsv4_compress_ratios = (4, 128)
|
||||
return backend
|
||||
|
||||
def _assert_device_path_matches_cpu_reference(
|
||||
self,
|
||||
*,
|
||||
positions,
|
||||
live_seq_lens,
|
||||
n_draft,
|
||||
ratio,
|
||||
dst_size,
|
||||
):
|
||||
backend = self._backend()
|
||||
positions = torch.tensor(positions, dtype=torch.int64)
|
||||
live_seq_lens = torch.tensor(live_seq_lens, dtype=torch.int32)
|
||||
final_seq_lens = torch.where(
|
||||
live_seq_lens > 0,
|
||||
live_seq_lens + int(n_draft),
|
||||
live_seq_lens,
|
||||
)
|
||||
expected = torch.full((dst_size,), -1, dtype=torch.int64)
|
||||
actual = torch.full((dst_size,), -2, dtype=torch.int64)
|
||||
|
||||
backend._fill_verify_positions_cmp_padding_one(
|
||||
positions,
|
||||
expected,
|
||||
ratio=ratio,
|
||||
seq_lens_cpu=final_seq_lens,
|
||||
n_draft=n_draft,
|
||||
)
|
||||
backend._fill_verify_positions_cmp_padding_one_device(
|
||||
positions,
|
||||
actual,
|
||||
ratio=ratio,
|
||||
live_seq_lens=live_seq_lens,
|
||||
n_draft=n_draft,
|
||||
)
|
||||
|
||||
self.assertEqual(actual.tolist(), expected.tolist())
|
||||
|
||||
def test_uses_group_start_rope_position(self):
|
||||
backend = DeepseekV4AscendAttnBackend.__new__(DeepseekV4AscendAttnBackend)
|
||||
backend._dsv4_compress_ratios = (4, 128)
|
||||
|
||||
# Two linear three-token verify trees. Their completed C4 groups end
|
||||
# at token positions 7 and 11, whose compressed RoPE positions are the
|
||||
# corresponding group starts 4 and 8.
|
||||
positions = torch.tensor([7, 8, 9, 10, 11, 12], dtype=torch.int64)
|
||||
final_seq_lens = torch.tensor([10, 13], dtype=torch.int32)
|
||||
dst = torch.full((4,), -1, dtype=torch.int64)
|
||||
|
||||
backend._fill_verify_positions_cmp_padding_one(
|
||||
positions,
|
||||
dst,
|
||||
ratio=4,
|
||||
seq_lens_cpu=final_seq_lens,
|
||||
n_draft=3,
|
||||
)
|
||||
|
||||
self.assertEqual(dst.tolist(), [4, 8, 0, 0])
|
||||
|
||||
def test_c128_uses_group_start_rope_position(self):
|
||||
backend = DeepseekV4AscendAttnBackend.__new__(DeepseekV4AscendAttnBackend)
|
||||
backend._dsv4_compress_ratios = (4, 128)
|
||||
dst = torch.full((2,), -1, dtype=torch.int64)
|
||||
|
||||
backend._fill_verify_positions_cmp_padding_one(
|
||||
torch.tensor([126, 127, 128], dtype=torch.int64),
|
||||
dst,
|
||||
ratio=128,
|
||||
seq_lens_cpu=torch.tensor([129], dtype=torch.int32),
|
||||
n_draft=3,
|
||||
)
|
||||
|
||||
self.assertEqual(dst.tolist(), [0, 0])
|
||||
|
||||
def test_no_completed_group_clears_destination(self):
|
||||
backend = DeepseekV4AscendAttnBackend.__new__(DeepseekV4AscendAttnBackend)
|
||||
backend._dsv4_compress_ratios = (4, 128)
|
||||
dst = torch.full((2,), -1, dtype=torch.int64)
|
||||
|
||||
backend._fill_verify_positions_cmp_padding_one(
|
||||
torch.tensor([5, 6], dtype=torch.int64),
|
||||
dst,
|
||||
ratio=4,
|
||||
seq_lens_cpu=torch.tensor([7], dtype=torch.int32),
|
||||
n_draft=2,
|
||||
)
|
||||
|
||||
self.assertEqual(dst.tolist(), [0, 0])
|
||||
|
||||
def test_zero_length_graph_padding_does_not_emit_a_position(self):
|
||||
backend = DeepseekV4AscendAttnBackend.__new__(DeepseekV4AscendAttnBackend)
|
||||
backend._dsv4_compress_ratios = (4, 128)
|
||||
dst = torch.full((2,), -1, dtype=torch.int64)
|
||||
|
||||
backend._fill_verify_positions_cmp_padding_one(
|
||||
torch.tensor([0, 0, 0, 10, 11, 12], dtype=torch.int64),
|
||||
dst,
|
||||
ratio=4,
|
||||
seq_lens_cpu=torch.tensor([0, 13], dtype=torch.int32),
|
||||
n_draft=3,
|
||||
)
|
||||
|
||||
self.assertEqual(dst.tolist(), [8, 0])
|
||||
|
||||
def test_device_path_matches_reference_across_boundaries_and_padding(self):
|
||||
cases = (
|
||||
# Existing C4 and C128 boundary examples.
|
||||
dict(
|
||||
positions=[7, 8, 9, 10, 11, 12],
|
||||
live_seq_lens=[7, 10],
|
||||
n_draft=3,
|
||||
ratio=4,
|
||||
dst_size=4,
|
||||
),
|
||||
dict(
|
||||
positions=[126, 127, 128],
|
||||
live_seq_lens=[126],
|
||||
n_draft=3,
|
||||
ratio=128,
|
||||
dst_size=2,
|
||||
),
|
||||
# A zero-length graph-padding row may appear before a live row.
|
||||
dict(
|
||||
positions=[0, 0, 0, 10, 11, 12],
|
||||
live_seq_lens=[0, 10],
|
||||
n_draft=3,
|
||||
ratio=4,
|
||||
dst_size=4,
|
||||
),
|
||||
# More than one boundary per request and destination truncation.
|
||||
dict(
|
||||
positions=list(range(2, 11)) + list(range(7, 16)),
|
||||
live_seq_lens=[2, 7],
|
||||
n_draft=9,
|
||||
ratio=4,
|
||||
dst_size=3,
|
||||
),
|
||||
# Preserve values from a non-linear tree position array rather than
|
||||
# reconstructing them arithmetically from sequence lengths.
|
||||
dict(
|
||||
positions=[50, 90, 51, 52, 70, 71, 120, 72],
|
||||
live_seq_lens=[3, 126],
|
||||
n_draft=4,
|
||||
ratio=4,
|
||||
dst_size=4,
|
||||
),
|
||||
# Long-context C128 boundaries around 128K.
|
||||
dict(
|
||||
positions=[131070, 131071, 131072, 131073],
|
||||
live_seq_lens=[131070],
|
||||
n_draft=4,
|
||||
ratio=128,
|
||||
dst_size=2,
|
||||
),
|
||||
)
|
||||
for case in cases:
|
||||
with self.subTest(case=case):
|
||||
self._assert_device_path_matches_cpu_reference(**case)
|
||||
|
||||
def test_stable_compact_preserves_boolean_index_order_and_zero_tail(self):
|
||||
dst = torch.full((4,), -1, dtype=torch.int64)
|
||||
values = torch.tensor([11, 22, 33, 44, 55, 66], dtype=torch.int64)
|
||||
keep = torch.tensor([False, True, False, True, True, False])
|
||||
|
||||
DeepseekV4AscendAttnBackend._stable_compact_1d(dst, values, keep)
|
||||
|
||||
self.assertEqual(dst.tolist(), [22, 44, 55, 0])
|
||||
|
||||
def test_stable_compact_matches_boolean_index_truncation(self):
|
||||
dst = torch.full((2,), -1, dtype=torch.int64)
|
||||
values = torch.tensor([5, 6, 7, 8, 9], dtype=torch.int64)
|
||||
keep = torch.tensor([True, False, True, True, True])
|
||||
|
||||
DeepseekV4AscendAttnBackend._stable_compact_1d(dst, values, keep)
|
||||
|
||||
self.assertEqual(dst.tolist(), values[keep][:2].tolist())
|
||||
|
||||
def test_stable_compact_matches_all_small_boolean_masks(self):
|
||||
values = torch.arange(1, 7, dtype=torch.int64)
|
||||
for mask_bits in range(1 << values.numel()):
|
||||
keep = torch.tensor(
|
||||
[(mask_bits >> index) & 1 for index in range(values.numel())],
|
||||
dtype=torch.bool,
|
||||
)
|
||||
for dst_size in range(1, values.numel() + 1):
|
||||
dst = torch.full((dst_size,), -1, dtype=torch.int64)
|
||||
DeepseekV4AscendAttnBackend._stable_compact_1d(dst, values, keep)
|
||||
selected = values[keep][:dst_size]
|
||||
expected = torch.zeros_like(dst)
|
||||
expected[: selected.numel()].copy_(selected)
|
||||
self.assertEqual(dst.tolist(), expected.tolist())
|
||||
|
||||
|
||||
class TestMultiStepDraftCompressedLocs(unittest.TestCase):
|
||||
def test_draft_steps_skip_compressed_locs_but_preserve_full_and_swa(self):
|
||||
backend = DeepseekV4AscendMultiStepDraftBackend.__new__(
|
||||
DeepseekV4AscendMultiStepDraftBackend
|
||||
)
|
||||
backend.topk = 1
|
||||
backend.speculative_num_steps = 3
|
||||
bundle = SimpleNamespace(
|
||||
out_full_loc=torch.arange(6, dtype=torch.int64),
|
||||
out_swa_loc=torch.arange(10, 16, dtype=torch.int64),
|
||||
out_c4_loc=torch.tensor([101, 102], dtype=torch.int64),
|
||||
out_c128_loc=torch.tensor([201], dtype=torch.int64),
|
||||
)
|
||||
forward_batch = SimpleNamespace(
|
||||
batch_size=2,
|
||||
out_cache_loc=bundle.out_full_loc,
|
||||
out_cache_loc_dsv4=bundle,
|
||||
seq_lens=torch.tensor([7, 11], dtype=torch.int32),
|
||||
)
|
||||
|
||||
with patch("torch.cumsum", side_effect=AssertionError("unexpected compaction")):
|
||||
step = backend._step_out_cache_loc_dsv4(forward_batch, step_id=1)
|
||||
|
||||
self.assertEqual(step.out_full_loc.tolist(), [1, 4])
|
||||
self.assertEqual(step.out_swa_loc.tolist(), [11, 14])
|
||||
self.assertEqual(step.out_c4_loc.numel(), 0)
|
||||
self.assertEqual(step.out_c128_loc.numel(), 0)
|
||||
self.assertEqual(step.out_c4_loc.dtype, bundle.out_c4_loc.dtype)
|
||||
self.assertEqual(step.out_c128_loc.dtype, bundle.out_c128_loc.dtype)
|
||||
|
||||
|
||||
class TestC4StateTransferLayout(unittest.TestCase):
|
||||
@patch(
|
||||
"sglang.srt.hardware_backend.npu.utils.is_npu_arch35",
|
||||
return_value=True,
|
||||
)
|
||||
def test_payload_uses_each_peers_private_ring_size(self, _):
|
||||
def state_rows(ring_size):
|
||||
req_pool = SimpleNamespace(
|
||||
c128_page_size=1,
|
||||
req_to_c128_sidecar=torch.zeros((4, 1), dtype=torch.int32),
|
||||
get_dsv4_c4_state_ring_size=lambda: ring_size,
|
||||
)
|
||||
payloads = dsv4_state_payloads(req_pool, 2, 13, page_size=1)
|
||||
return next(
|
||||
payload()
|
||||
for state_type, payload in payloads.items()
|
||||
if state_type.value == "dsv4_c4_state"
|
||||
)
|
||||
|
||||
self.assertEqual(state_rows(8).tolist(), [16, 17, 18, 19, 20])
|
||||
self.assertEqual(state_rows(16).tolist(), [40, 41, 42, 43, 44])
|
||||
|
||||
def test_req_pool_reads_ring_size_from_registered_kv_pool(self):
|
||||
req_pool = DSV4ReqToTokenTablesMixin.__new__(DSV4ReqToTokenTablesMixin)
|
||||
req_pool._dsv4_allocator = MagicMock()
|
||||
req_pool._dsv4_allocator.get_kvcache().get_ring_size.return_value = 16
|
||||
|
||||
self.assertEqual(req_pool.get_dsv4_c4_state_ring_size(), 16)
|
||||
req_pool._dsv4_allocator.get_kvcache().get_ring_size.assert_called_once_with(4)
|
||||
|
||||
def test_registers_single_rows_instead_of_request_banks(self):
|
||||
attn_state = torch.empty((32, 5), dtype=torch.float32)
|
||||
indexer_state = torch.empty((32, 7), dtype=torch.float32)
|
||||
pool = DSV4NPUTokenToKVPool.__new__(DSV4NPUTokenToKVPool)
|
||||
pool.compress_state_pools = [
|
||||
SimpleNamespace(
|
||||
ratio=4,
|
||||
ring_size=8,
|
||||
kv_score_buffer=SimpleNamespace(kv_score=attn_state),
|
||||
)
|
||||
]
|
||||
pool.indexer_compress_state_pools = [
|
||||
SimpleNamespace(
|
||||
ratio=4,
|
||||
ring_size=8,
|
||||
kv_score_buffer=SimpleNamespace(kv_score=indexer_state),
|
||||
)
|
||||
]
|
||||
|
||||
_, data_lens, item_lens = pool.get_c4_state_buf_infos()
|
||||
|
||||
self.assertEqual(data_lens, [attn_state.nbytes, indexer_state.nbytes])
|
||||
self.assertEqual(
|
||||
item_lens,
|
||||
[attn_state[0].nbytes, indexer_state[0].nbytes],
|
||||
)
|
||||
|
||||
|
||||
class TestC4IndexerInitialization(unittest.TestCase):
|
||||
@patch(
|
||||
"sglang.srt.hardware_backend.npu.attention.ascend_dsv4_backend.is_npu_arch35",
|
||||
return_value=True,
|
||||
)
|
||||
def test_arch35_indexer_uses_float8_kv(self, _):
|
||||
indexer = torch.nn.Module()
|
||||
indexer.head_dim = 8
|
||||
indexer.compressor = SimpleNamespace()
|
||||
backend = C4IndexerAscendBackendMixin.__new__(C4IndexerAscendBackendMixin)
|
||||
|
||||
backend._ensure_npu_c4_indexer(indexer, torch.device("cpu"))
|
||||
|
||||
self.assertEqual(indexer.compressor.li_kv_dtype, "float8")
|
||||
|
||||
|
||||
class TestWalshHadamardMatrix(unittest.TestCase):
|
||||
@@ -181,6 +498,318 @@ class TestApplyHadamard(unittest.TestCase):
|
||||
self.assertTrue(torch.equal(out, expected))
|
||||
|
||||
|
||||
class TestCompressorStateTableABI(unittest.TestCase):
|
||||
def test_arch35_cycle_table_is_one_bank_per_request(self):
|
||||
req_pool_indices = torch.tensor([7, 3], dtype=torch.int64)
|
||||
table = _build_cycle_state_block_table(req_pool_indices)
|
||||
self.assertEqual(tuple(table.shape), (2,))
|
||||
self.assertEqual(table.dtype, torch.int32)
|
||||
self.assertEqual(table.tolist(), [7, 3])
|
||||
|
||||
def test_arch35_cycle_table_rejects_explicit_shape(self):
|
||||
with self.assertRaises(ValueError):
|
||||
_build_cycle_state_block_table(torch.zeros((2, 8), dtype=torch.int32))
|
||||
|
||||
@patch(
|
||||
"sglang.srt.hardware_backend.npu.attention.ascend_dsv4_backend.is_npu_arch35",
|
||||
return_value=True,
|
||||
)
|
||||
def test_arch35_eager_metadata_builds_cycle_table(self, _):
|
||||
backend = CompressorAscendBackendMixin.__new__(CompressorAscendBackendMixin)
|
||||
backend.forward_metadata = SimpleNamespace()
|
||||
backend.token_to_kv_pool = MagicMock()
|
||||
backend.req_to_token = torch.empty((0, 0), dtype=torch.int32)
|
||||
backend.req_to_token_pool = MagicMock()
|
||||
backend._dsv4_compress_ratios = ()
|
||||
backend._compute_compress_locs = MagicMock(return_value={})
|
||||
|
||||
forward_mode = MagicMock()
|
||||
forward_mode.is_decode.return_value = True
|
||||
forward_mode.is_target_verify.return_value = False
|
||||
forward_batch = SimpleNamespace(
|
||||
forward_mode=forward_mode,
|
||||
req_pool_indices=torch.tensor([7, 3], dtype=torch.int64),
|
||||
seq_lens=torch.tensor([5, 9], dtype=torch.int32),
|
||||
out_cache_loc=torch.empty(0, dtype=torch.int64),
|
||||
out_cache_loc_dsv4=None,
|
||||
batch_size=2,
|
||||
)
|
||||
|
||||
backend._build_npu_compress_metadata(forward_batch)
|
||||
|
||||
table = getattr(backend.forward_metadata, "dsv4_cycle_state_block_table", None)
|
||||
self.assertIsNotNone(table)
|
||||
self.assertEqual(table.tolist(), [7, 3])
|
||||
self.assertEqual(table.dtype, torch.int32)
|
||||
|
||||
@patch(
|
||||
"sglang.srt.hardware_backend.npu.attention.ascend_dsv4_backend.is_npu_arch35",
|
||||
return_value=True,
|
||||
)
|
||||
def test_arch35_graph_replay_slices_static_req_pool_buffer_to_graph_bs(self, _):
|
||||
backend = DeepseekV4AscendAttnBackend.__new__(DeepseekV4AscendAttnBackend)
|
||||
table = torch.zeros(1, dtype=torch.int32)
|
||||
graph_mode = MagicMock()
|
||||
graph_mode.is_decode.return_value = False
|
||||
graph_mode.is_target_verify.return_value = False
|
||||
ctx = SimpleNamespace(
|
||||
fm=SimpleNamespace(dsv4_cycle_state_block_table=table),
|
||||
forward_batch=SimpleNamespace(
|
||||
req_pool_indices=torch.arange(7, 19, dtype=torch.int64)
|
||||
),
|
||||
graph_mode=graph_mode,
|
||||
bs=1,
|
||||
)
|
||||
backend._build_dsv4_graph_replay_ctx = MagicMock(return_value=ctx)
|
||||
for name in (
|
||||
"_refresh_graph_seq_metadata",
|
||||
"_refresh_graph_compress_page_tables_direct",
|
||||
"_refresh_graph_explicit_state_block_tables",
|
||||
"_refresh_graph_swa_metadata_direct",
|
||||
"_refresh_graph_dspark_sparse_metadata",
|
||||
"_refresh_graph_kernel_metadata",
|
||||
):
|
||||
setattr(backend, name, MagicMock())
|
||||
|
||||
backend._apply_dsv4_graph_metadata(SimpleNamespace())
|
||||
|
||||
self.assertIs(ctx.fm.dsv4_cycle_state_block_table, table)
|
||||
self.assertEqual(table.tolist(), [7])
|
||||
|
||||
@patch(
|
||||
"sglang.srt.hardware_backend.npu.attention.ascend_dsv4_backend.is_npu_arch35",
|
||||
return_value=True,
|
||||
)
|
||||
def test_arch35_graph_capture_allocates_cycle_table_buffer(self, _):
|
||||
backend = DeepseekV4AscendAttnBackend.__new__(DeepseekV4AscendAttnBackend)
|
||||
metadata = SimpleNamespace()
|
||||
backend.device = "cpu"
|
||||
backend.graph_metadata = {
|
||||
2: metadata,
|
||||
"swa_page_table": torch.full((2, 4), -1, dtype=torch.int32),
|
||||
"c4_page_table": torch.full((2, 4), -1, dtype=torch.int32),
|
||||
"c128_page_table": torch.full((2, 4), -1, dtype=torch.int32),
|
||||
"kernel_metadata_c1a": torch.zeros(1024, dtype=torch.int32),
|
||||
"kernel_metadata_c4a": torch.zeros(1024, dtype=torch.int32),
|
||||
"kernel_metadata_c128a": torch.zeros(1024, dtype=torch.int32),
|
||||
"kernel_metadata_li_quant": torch.zeros(1024, dtype=torch.int32),
|
||||
"c4_topk_indices": torch.full((2, 1), -1, dtype=torch.int32),
|
||||
}
|
||||
backend._dsv4_graph_tokens_per_req = 1
|
||||
backend._dsv4_index_topk = 1
|
||||
backend._dsv4_state_pools_by_ratio = {}
|
||||
backend._dsv4_sliding_window_size = 128
|
||||
backend._is_dspark_draft_worker = False
|
||||
forward_mode = MagicMock()
|
||||
forward_mode.is_target_verify.return_value = False
|
||||
forward_mode.is_draft_extend_v2.return_value = False
|
||||
|
||||
backend._init_dsv4_graph_metadata(2, forward_mode)
|
||||
|
||||
table = getattr(metadata, "dsv4_cycle_state_block_table", None)
|
||||
self.assertIsNotNone(table)
|
||||
self.assertEqual(tuple(table.shape), (2,))
|
||||
self.assertEqual(table.dtype, torch.int32)
|
||||
|
||||
@patch(
|
||||
"sglang.srt.hardware_backend.npu.attention.ascend_dsv4_backend.is_npu_arch35",
|
||||
return_value=True,
|
||||
)
|
||||
def test_arch35_forward_reuses_batch_cycle_table(self, _):
|
||||
table = torch.tensor([7, 3], dtype=torch.int32)
|
||||
backend = CompressorAscendBackendMixin.__new__(CompressorAscendBackendMixin)
|
||||
backend.graph_mode = False
|
||||
backend.forward_metadata = SimpleNamespace(
|
||||
dsv4_cycle_state_block_table=table,
|
||||
positions_cmp_padding_c128=torch.empty(0, dtype=torch.int64),
|
||||
actual_seq_lengths_q_pa=torch.tensor([0, 1, 2], dtype=torch.int32),
|
||||
seqused=torch.ones(2, dtype=torch.int32),
|
||||
start_pos=torch.zeros(2, dtype=torch.int32),
|
||||
c128_loc=None,
|
||||
)
|
||||
backend.token_to_kv_pool = MagicMock()
|
||||
backend.token_to_kv_pool._get_state_pool.return_value = SimpleNamespace(
|
||||
state_cache_3d=torch.empty(0)
|
||||
)
|
||||
backend._ensure_compressor_hadamard = MagicMock()
|
||||
backend._ensure_fused_caches = MagicMock()
|
||||
backend._compressor_epilog_npu = MagicMock()
|
||||
|
||||
compressor = SimpleNamespace(
|
||||
ratio=128,
|
||||
overlap=False,
|
||||
layer_id=0,
|
||||
is_in_indexer=False,
|
||||
freqs_cis=None,
|
||||
rotary_emb=None,
|
||||
_fused_wkv_w=torch.empty(0),
|
||||
_fused_wgate_w=torch.empty(0),
|
||||
ape=torch.empty(0),
|
||||
_fused_norm_weight_fp32=torch.empty(0),
|
||||
rope_head_dim=64,
|
||||
norm=SimpleNamespace(variance_epsilon=1e-6),
|
||||
rotate=False,
|
||||
)
|
||||
forward_mode = MagicMock()
|
||||
forward_mode.is_prefill.return_value = False
|
||||
forward_mode.is_target_verify.return_value = False
|
||||
forward_batch = SimpleNamespace(
|
||||
req_pool_indices=torch.tensor([7, 3], dtype=torch.int64),
|
||||
forward_mode=forward_mode,
|
||||
)
|
||||
rope = MagicMock()
|
||||
rope.get_cos_sin.return_value = (torch.empty(0), torch.empty(0))
|
||||
|
||||
with (
|
||||
patch(
|
||||
"sglang.srt.hardware_backend.npu.attention.ascend_dsv4_backend."
|
||||
"Dsv4NpuRoPE.for_freqs",
|
||||
return_value=rope,
|
||||
),
|
||||
patch.object(torch.ops, "custom", MagicMock(), create=True) as custom_ops,
|
||||
patch.object(torch.ops, "npu", MagicMock(), create=True) as npu_ops,
|
||||
):
|
||||
custom_ops.compressor.return_value = torch.empty((0, 1))
|
||||
backend.forward_compress(compressor, torch.empty((2, 1)), forward_batch)
|
||||
backend.forward_compress(compressor, torch.empty((2, 1)), forward_batch)
|
||||
|
||||
self.assertEqual(npu_ops.compressor.call_count, 0)
|
||||
self.assertIs(
|
||||
custom_ops.compressor.call_args_list[0].kwargs["state_block_table"], table
|
||||
)
|
||||
self.assertIs(
|
||||
custom_ops.compressor.call_args_list[1].kwargs["state_block_table"], table
|
||||
)
|
||||
|
||||
|
||||
class TestArch35SparseAttentionDispatch(unittest.TestCase):
|
||||
_ARCH35_PATCH_TARGET = (
|
||||
"sglang.srt.hardware_backend.npu.attention.ascend_dsv4_backend.is_npu_arch35"
|
||||
)
|
||||
|
||||
@patch(_ARCH35_PATCH_TARGET, return_value=True)
|
||||
def test_arch35_uses_kv_quant_ops_and_layout_kwargs(self, _):
|
||||
with patch("torch.ops.custom", MagicMock(), create=True) as custom_ops:
|
||||
metadata_op, attention_op = _sparse_attn_ops()
|
||||
kwargs = _sparse_attn_kv_quant_kwargs()
|
||||
|
||||
self.assertIs(
|
||||
metadata_op, custom_ops.npu_kv_quant_sparse_attn_sharedkv_metadata
|
||||
)
|
||||
self.assertIs(attention_op, custom_ops.npu_kv_quant_sparse_attn_sharedkv)
|
||||
self.assertEqual(
|
||||
kwargs,
|
||||
{"kv_quant_mode": 1, "tile_size": 64, "rope_head_dim": 64},
|
||||
)
|
||||
|
||||
@patch(_ARCH35_PATCH_TARGET, return_value=False)
|
||||
def test_pre_arch35_keeps_legacy_ops_without_quant_kwargs(self, _):
|
||||
with (
|
||||
patch("torch.ops.custom", MagicMock(), create=True) as custom_ops,
|
||||
patch("torch.ops.npu", MagicMock(), create=True) as npu_ops,
|
||||
):
|
||||
metadata_op, attention_op = _sparse_attn_ops()
|
||||
kwargs = _sparse_attn_kv_quant_kwargs()
|
||||
|
||||
self.assertIs(metadata_op, custom_ops.npu_sparse_attn_sharedkv_metadata)
|
||||
self.assertIs(attention_op, npu_ops.sparse_attn_sharedkv)
|
||||
self.assertEqual(kwargs, {})
|
||||
|
||||
|
||||
class TestSparseAttentionMetadata(unittest.TestCase):
|
||||
_ARCH35_PATCH_TARGET = (
|
||||
"sglang.srt.hardware_backend.npu.attention.ascend_dsv4_backend.is_npu_arch35"
|
||||
)
|
||||
|
||||
def test_device_metadata_receives_sequence_lengths(self):
|
||||
cu_seqlens_q = torch.tensor([0, 2, 3], dtype=torch.int32)
|
||||
seqused_kv = torch.tensor([8, 12], dtype=torch.int32)
|
||||
|
||||
for is_arch35, metadata_op_name in (
|
||||
(False, "npu_sparse_attn_sharedkv_metadata"),
|
||||
(True, "npu_kv_quant_sparse_attn_sharedkv_metadata"),
|
||||
):
|
||||
with (
|
||||
self.subTest(is_arch35=is_arch35),
|
||||
patch(self._ARCH35_PATCH_TARGET, return_value=is_arch35),
|
||||
patch("torch.ops.custom", MagicMock(), create=True) as custom_ops,
|
||||
patch("torch.ops.npu", MagicMock(), create=True),
|
||||
):
|
||||
backend = DeepseekV4AscendAttnBackend.__new__(
|
||||
DeepseekV4AscendAttnBackend
|
||||
)
|
||||
backend.forward_metadata = SimpleNamespace()
|
||||
backend._is_dspark_draft_worker = False
|
||||
backend._dsv4_sliding_window_size = 128
|
||||
backend._dsv4_q_head_num = 64
|
||||
backend._dsv4_kv_head_num = 1
|
||||
backend._dsv4_head_dim = 512
|
||||
backend._dsv4_has_c4 = True
|
||||
backend._dsv4_has_c128 = True
|
||||
backend._dsv4_index_topk = 512
|
||||
backend._dsv4_index_n_heads = 16
|
||||
backend._dsv4_index_head_dim = 128
|
||||
|
||||
backend._kernel_metadata_from_parts(
|
||||
bs=2,
|
||||
actual_seq_lengths_q_pa=cu_seqlens_q,
|
||||
actual_seq_lengths_kv=seqused_kv,
|
||||
block_tables=torch.zeros((2, 1), dtype=torch.int32),
|
||||
max_seqlen_q=2,
|
||||
is_nextn=False,
|
||||
)
|
||||
|
||||
metadata_op = getattr(custom_ops, metadata_op_name)
|
||||
self.assertEqual(metadata_op.call_count, 3)
|
||||
for call in metadata_op.call_args_list:
|
||||
self.assertIs(call.kwargs["cu_seqlens_q"], cu_seqlens_q)
|
||||
self.assertIs(call.kwargs["seqused_kv"], seqused_kv)
|
||||
|
||||
def test_dspark_host_metadata_receives_host_sequence_lengths(self):
|
||||
cu_seqlens_q = torch.tensor([0, 2, 3], dtype=torch.int32)
|
||||
seqused_kv = torch.tensor([8, 12], dtype=torch.int32)
|
||||
cu_seqlens_q_cpu = cu_seqlens_q.clone()
|
||||
seqused_kv_cpu = seqused_kv.clone()
|
||||
|
||||
with (
|
||||
patch("torch.ops.npu", MagicMock(), create=True) as npu_ops,
|
||||
patch(
|
||||
"sglang.srt.hardware_backend.npu.attention.ascend_dsv4_backend._sparse_attn_ops",
|
||||
return_value=(MagicMock(), MagicMock()),
|
||||
),
|
||||
):
|
||||
backend = DeepseekV4AscendAttnBackend.__new__(DeepseekV4AscendAttnBackend)
|
||||
backend.forward_metadata = SimpleNamespace(
|
||||
actual_seq_lengths_q_pa_cpu=cu_seqlens_q_cpu,
|
||||
seq_lens_cpu_int=seqused_kv_cpu,
|
||||
)
|
||||
backend._is_dspark_draft_worker = True
|
||||
backend._dsv4_sliding_window_size = 128
|
||||
backend._dsv4_q_head_num = 64
|
||||
backend._dsv4_kv_head_num = 1
|
||||
backend._dsv4_head_dim = 512
|
||||
backend._dsv4_has_c4 = False
|
||||
backend._dsv4_has_c128 = False
|
||||
|
||||
kernel_metadata = backend._kernel_metadata_from_parts(
|
||||
bs=2,
|
||||
actual_seq_lengths_q_pa=cu_seqlens_q,
|
||||
actual_seq_lengths_kv=seqused_kv,
|
||||
block_tables=torch.zeros((2, 1), dtype=torch.int32),
|
||||
max_seqlen_q=2,
|
||||
is_nextn=False,
|
||||
)
|
||||
|
||||
metadata_op = npu_ops.sparse_attn_sharedkv_metadata_host
|
||||
metadata_op.assert_called_once()
|
||||
self.assertIs(kernel_metadata["c1a_metadata"], metadata_op.return_value)
|
||||
self.assertIs(metadata_op.call_args.kwargs["cu_seqlens_q"], cu_seqlens_q_cpu)
|
||||
actual_seqused_kv = metadata_op.call_args.kwargs["seqused_kv"]
|
||||
torch.testing.assert_close(actual_seqused_kv, seqused_kv_cpu[:2])
|
||||
self.assertEqual(actual_seqused_kv.dtype, torch.int32)
|
||||
self.assertEqual(actual_seqused_kv.device.type, "cpu")
|
||||
|
||||
|
||||
class TestGetKvIndices(unittest.TestCase):
|
||||
_PATCH_TARGET = (
|
||||
"sglang.srt.hardware_backend.npu.attention.ascend_dsv4_backend.get_attn_backend"
|
||||
@@ -383,5 +1012,69 @@ class TestCommonTemplate(unittest.TestCase):
|
||||
self.assertEqual(call_fn.call_count, 1)
|
||||
|
||||
|
||||
class TestCompressorEpilogEmptyWrite(unittest.TestCase):
|
||||
@staticmethod
|
||||
def _backend(*, loc, graph_mode=False):
|
||||
backend = CompressorAscendBackendMixin.__new__(CompressorAscendBackendMixin)
|
||||
backend.graph_mode = graph_mode
|
||||
backend.token_to_kv_pool = MagicMock()
|
||||
backend.forward_metadata = SimpleNamespace(c4_loc=loc, c128_loc=loc)
|
||||
return backend
|
||||
|
||||
@staticmethod
|
||||
def _compressor(*, li_kv_dtype="bf16", is_in_indexer=False):
|
||||
return SimpleNamespace(
|
||||
ratio=128,
|
||||
layer_id=0,
|
||||
is_in_indexer=is_in_indexer,
|
||||
li_kv_dtype=li_kv_dtype,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _verify_batch():
|
||||
forward_mode = MagicMock()
|
||||
forward_mode.is_target_verify.return_value = True
|
||||
return SimpleNamespace(forward_mode=forward_mode)
|
||||
|
||||
def test_all_slots_masked_skips_compress_write(self):
|
||||
backend = self._backend(loc=torch.zeros(3, dtype=torch.int32))
|
||||
backend._compressor_epilog_npu(
|
||||
self._compressor(), torch.zeros(3, 512), self._verify_batch()
|
||||
)
|
||||
backend.token_to_kv_pool.set_compress_buffer.assert_not_called()
|
||||
|
||||
def test_partially_masked_slots_writes_surviving_rows(self):
|
||||
backend = self._backend(loc=torch.tensor([0, 7, 0], dtype=torch.int32))
|
||||
kv = torch.arange(12, dtype=torch.float32).view(3, 4)
|
||||
backend._compressor_epilog_npu(self._compressor(), kv, self._verify_batch())
|
||||
|
||||
backend.token_to_kv_pool.set_compress_buffer.assert_called_once()
|
||||
_, written_loc, written_kv, _, _ = (
|
||||
backend.token_to_kv_pool.set_compress_buffer.call_args.args
|
||||
)
|
||||
self.assertEqual(written_loc.tolist(), [7])
|
||||
self.assertEqual(written_kv.tolist(), [kv[1].tolist()])
|
||||
|
||||
def test_graph_mode_keeps_static_shape_write(self):
|
||||
backend = self._backend(loc=torch.zeros(3, dtype=torch.int32), graph_mode=True)
|
||||
backend._compressor_epilog_npu(
|
||||
self._compressor(), torch.ones(3, 4), self._verify_batch()
|
||||
)
|
||||
|
||||
backend.token_to_kv_pool.set_compress_buffer.assert_called_once()
|
||||
written_kv = backend.token_to_kv_pool.set_compress_buffer.call_args.args[2]
|
||||
self.assertEqual(written_kv.shape[0], 3)
|
||||
self.assertEqual(written_kv.abs().sum().item(), 0.0)
|
||||
|
||||
def test_all_slots_masked_skips_fused_indexer_write(self):
|
||||
backend = self._backend(loc=torch.zeros(3, dtype=torch.int32))
|
||||
compressor = self._compressor(li_kv_dtype="float8", is_in_indexer=True)
|
||||
with patch("torch.ops.custom", MagicMock(), create=True) as custom_ops:
|
||||
backend._compressor_epilog_npu(
|
||||
compressor, torch.zeros(3, 512), self._verify_batch()
|
||||
)
|
||||
custom_ops.indexer_compress_epilog.assert_not_called()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -0,0 +1,558 @@
|
||||
import inspect
|
||||
import os
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.test.ci.ci_register import register_npu_ci
|
||||
|
||||
register_npu_ci(est_time=1, suite="stage-a-unit-test-npu")
|
||||
|
||||
# Load the quantization package first so `base_config`, `moe_methods`, and
|
||||
# `linear_method_npu` initialize in dependency order. Importing `fp4_moe_methods`
|
||||
# (or `linear_method_npu`) directly from a cold process triggers a circular
|
||||
# import: linear_method_npu -> base_config -> quantization/__init__ ->
|
||||
# gguf/unquant/gptq_moe -> moe_methods -> linear_method_npu (partially
|
||||
# initialized, `_get_float8_e8m0fnu_dtype` not yet defined). Initializing the
|
||||
# package first mirrors how the engine loads quantization at model-config time.
|
||||
import sglang.srt.layers.quantization # noqa: F401
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.hardware_backend.npu.quantization import fp4_moe_methods
|
||||
from sglang.srt.hardware_backend.npu.quantization.fp4_moe_methods import (
|
||||
NPUW4A4Fp4MoEMethod,
|
||||
_apply_swiglu_limit_npu,
|
||||
_configure_dsv4_deepep_dispatcher,
|
||||
_pair_pack_mxfp_act_scale,
|
||||
_reshape_mxfp4_scale_for_npu,
|
||||
npu_apply_without_routing_weights_w4a4_mxfp,
|
||||
w4a8_mxfp_gmm,
|
||||
)
|
||||
from sglang.srt.layers.moe.fused_moe_triton import FusedMoE
|
||||
from sglang.srt.layers.moe.token_dispatcher import deepep
|
||||
from sglang.srt.layers.quantization.fp8 import Fp8Config, Fp8MoEMethod
|
||||
|
||||
_NOT_PASSED = object()
|
||||
|
||||
|
||||
class TestFP4MethodGate(unittest.TestCase):
|
||||
def test_pre_arch35_keeps_fp8_moe_method(self):
|
||||
config = Fp8Config(is_fp4_experts=True)
|
||||
layer = FusedMoE.__new__(FusedMoE)
|
||||
|
||||
with (
|
||||
patch("sglang.srt.layers.quantization.fp8.is_npu", return_value=True),
|
||||
patch(
|
||||
"sglang.srt.layers.quantization.fp8.is_npu_arch35",
|
||||
return_value=False,
|
||||
),
|
||||
):
|
||||
method = config.get_quant_method(layer, "model.layers.0.experts")
|
||||
|
||||
self.assertIsInstance(method, Fp8MoEMethod)
|
||||
|
||||
|
||||
class TestApplySwiGLULimitNpu(unittest.TestCase):
|
||||
def test_clamps_gate_and_up_asymmetrically(self):
|
||||
# DeepSeek-V4 clamps gate (first half) to <= limit but only the upper
|
||||
# bound, while up (second half) is clamped symmetrically to [-limit, limit].
|
||||
# A regression that swapped these would silently change expert activations.
|
||||
gate_up = torch.tensor([[8.0, -9.0, 9.0, -9.0]])
|
||||
_apply_swiglu_limit_npu(gate_up, 7.0)
|
||||
self.assertTrue(torch.equal(gate_up, torch.tensor([[7.0, -9.0, 7.0, -7.0]])))
|
||||
|
||||
def test_noop_when_limit_none(self):
|
||||
gate_up = torch.tensor([[8.0, -9.0]])
|
||||
_apply_swiglu_limit_npu(gate_up, None)
|
||||
self.assertTrue(torch.equal(gate_up, torch.tensor([[8.0, -9.0]])))
|
||||
|
||||
def test_noop_when_limit_nonpositive(self):
|
||||
gate_up = torch.tensor([[8.0, -9.0]])
|
||||
_apply_swiglu_limit_npu(gate_up, 0.0)
|
||||
self.assertTrue(torch.equal(gate_up, torch.tensor([[8.0, -9.0]])))
|
||||
|
||||
|
||||
class TestReshapeMxfp4ScaleForNpu(unittest.TestCase):
|
||||
def test_packs_scale_to_gmm_layout(self):
|
||||
# [E, N, K/32] -> [E, K/64, N, 2] is the packed-pair layout the GMM reads;
|
||||
# getting the transpose axis wrong silently dequantizes with the wrong scale.
|
||||
scale = torch.arange(8, dtype=torch.uint8).view(1, 2, 4)
|
||||
out = _reshape_mxfp4_scale_for_npu(scale)
|
||||
self.assertEqual(tuple(out.shape), (1, 2, 2, 2))
|
||||
self.assertTrue(torch.equal(out, scale.view(1, 2, 2, 2).transpose(1, 2)))
|
||||
|
||||
def test_rejects_odd_k_dim(self):
|
||||
with self.assertRaises(ValueError):
|
||||
_reshape_mxfp4_scale_for_npu(torch.zeros(1, 2, 3, dtype=torch.uint8))
|
||||
|
||||
|
||||
class TestMxfp4ScaleWeightLoader(unittest.TestCase):
|
||||
def test_reinterprets_e8m0_scale_as_raw_uint8(self):
|
||||
loaded = []
|
||||
|
||||
def weight_loader(param, loaded_weight, *args, **kwargs):
|
||||
loaded.append(loaded_weight.clone())
|
||||
|
||||
layer = torch.nn.Module()
|
||||
method = NPUW4A4Fp4MoEMethod(fp8_method=MagicMock(), prefix="test")
|
||||
method.create_weights(
|
||||
layer,
|
||||
num_experts=1,
|
||||
hidden_size=64,
|
||||
intermediate_size_per_partition=64,
|
||||
params_dtype=torch.bfloat16,
|
||||
weight_loader=weight_loader,
|
||||
)
|
||||
checkpoint_scale = torch.tensor([0.5, 0.25, 0.125], dtype=torch.float8_e8m0fnu)
|
||||
|
||||
layer.w13_weight_scale_inv.weight_loader(
|
||||
layer.w13_weight_scale_inv,
|
||||
checkpoint_scale,
|
||||
"model.layers.0.mlp.experts.0.gate_proj.weight_scale_inv",
|
||||
"w1",
|
||||
0,
|
||||
)
|
||||
|
||||
self.assertEqual(loaded[0].dtype, torch.uint8)
|
||||
self.assertTrue(torch.equal(loaded[0], checkpoint_scale.view(torch.uint8)))
|
||||
|
||||
|
||||
class TestPairPackMxfpActScale(unittest.TestCase):
|
||||
def test_packs_as_view(self):
|
||||
# The GMM expects a pair-packed *view* of the per-token scale, not a copy;
|
||||
# materializing a copy here would break the kernel's aliasing contract.
|
||||
flat = torch.arange(8).view(2, 4)
|
||||
packed = _pair_pack_mxfp_act_scale(flat)
|
||||
self.assertEqual(tuple(packed.shape), (2, 2, 2))
|
||||
self.assertEqual(packed.data_ptr(), flat.data_ptr())
|
||||
|
||||
def test_rejects_odd_scale_dim(self):
|
||||
with self.assertRaises(ValueError):
|
||||
_pair_pack_mxfp_act_scale(torch.zeros(2, 3))
|
||||
|
||||
def test_unflattens_low_latency_deepep_scale_as_view(self):
|
||||
# DeepEP returns one flat E8M0 scale per 32-element block. Passing
|
||||
# that flat buffer to GMM would use the wrong scale layout and either
|
||||
# fail or dequantize activations incorrectly.
|
||||
flat = torch.arange(4, dtype=torch.uint8)
|
||||
packed = _pair_pack_mxfp_act_scale(flat, input_shape=(2, 64))
|
||||
|
||||
self.assertEqual(tuple(packed.shape), (2, 1, 2))
|
||||
self.assertEqual(packed.data_ptr(), flat.data_ptr())
|
||||
self.assertTrue(torch.equal(packed, torch.tensor([[[0, 1]], [[2, 3]]])))
|
||||
|
||||
def test_rejects_low_latency_deepep_scale_with_wrong_length(self):
|
||||
with self.assertRaises(ValueError):
|
||||
_pair_pack_mxfp_act_scale(
|
||||
torch.zeros(3, dtype=torch.uint8), input_shape=(2, 64)
|
||||
)
|
||||
|
||||
|
||||
class TestDsv4DeepEPMxfp8DispatcherConfig(unittest.TestCase):
|
||||
@staticmethod
|
||||
def _deepep_backend():
|
||||
return SimpleNamespace(is_deepep=lambda: True)
|
||||
|
||||
def test_a5_deepep_defaults_low_latency_dispatch_to_mxfp8(self):
|
||||
dispatcher = MagicMock()
|
||||
layer = SimpleNamespace(dispatcher=dispatcher)
|
||||
|
||||
with (
|
||||
patch.object(fp4_moe_methods, "is_npu_arch35", return_value=True),
|
||||
patch(
|
||||
"sglang.srt.layers.moe.get_moe_a2a_backend",
|
||||
return_value=self._deepep_backend(),
|
||||
),
|
||||
patch.dict(os.environ, {}, clear=True),
|
||||
):
|
||||
_configure_dsv4_deepep_dispatcher(layer)
|
||||
|
||||
dispatcher.set_quant_config.assert_called_once_with(
|
||||
{
|
||||
"normal_dispatcher_output_dtype": "bf16",
|
||||
"low_latency_dispatcher_output_dtype": "mxfp8",
|
||||
}
|
||||
)
|
||||
|
||||
def test_non_deepep_ignores_the_low_latency_quant_environment(self):
|
||||
dispatcher = MagicMock()
|
||||
layer = SimpleNamespace(dispatcher=dispatcher)
|
||||
|
||||
with (
|
||||
patch.object(fp4_moe_methods, "is_npu_arch35", return_value=True),
|
||||
patch(
|
||||
"sglang.srt.layers.moe.get_moe_a2a_backend",
|
||||
return_value=SimpleNamespace(is_deepep=lambda: False),
|
||||
),
|
||||
envs.SGLANG_NPU_DSV4_DEEPEP_LL_DISPATCH_QUANT_MODE.override("invalid"),
|
||||
):
|
||||
_configure_dsv4_deepep_dispatcher(layer)
|
||||
|
||||
dispatcher.set_quant_config.assert_called_once_with(
|
||||
{"dispatcher_output_dtype": "bf16"}
|
||||
)
|
||||
|
||||
def test_a5_deepep_allows_bf16_low_latency_fallback(self):
|
||||
dispatcher = MagicMock()
|
||||
layer = SimpleNamespace(dispatcher=dispatcher)
|
||||
|
||||
with (
|
||||
patch.object(fp4_moe_methods, "is_npu_arch35", return_value=True),
|
||||
patch(
|
||||
"sglang.srt.layers.moe.get_moe_a2a_backend",
|
||||
return_value=self._deepep_backend(),
|
||||
),
|
||||
envs.SGLANG_NPU_DSV4_DEEPEP_LL_DISPATCH_QUANT_MODE.override("bf16"),
|
||||
):
|
||||
_configure_dsv4_deepep_dispatcher(layer)
|
||||
|
||||
dispatcher.set_quant_config.assert_called_once_with(
|
||||
{
|
||||
"normal_dispatcher_output_dtype": "bf16",
|
||||
"low_latency_dispatcher_output_dtype": "bf16",
|
||||
}
|
||||
)
|
||||
|
||||
def test_a5_deepep_rejects_an_invalid_low_latency_quant_mode(self):
|
||||
layer = SimpleNamespace(dispatcher=MagicMock())
|
||||
|
||||
with (
|
||||
patch.object(fp4_moe_methods, "is_npu_arch35", return_value=True),
|
||||
patch(
|
||||
"sglang.srt.layers.moe.get_moe_a2a_backend",
|
||||
return_value=self._deepep_backend(),
|
||||
),
|
||||
envs.SGLANG_NPU_DSV4_DEEPEP_LL_DISPATCH_QUANT_MODE.override("invalid"),
|
||||
self.assertRaisesRegex(ValueError, "SGLANG_NPU_DSV4"),
|
||||
):
|
||||
_configure_dsv4_deepep_dispatcher(layer)
|
||||
|
||||
def test_non_a5_ignores_the_low_latency_quant_environment(self):
|
||||
dispatcher = MagicMock()
|
||||
layer = SimpleNamespace(dispatcher=dispatcher)
|
||||
|
||||
with (
|
||||
patch.object(fp4_moe_methods, "is_npu_arch35", return_value=False),
|
||||
patch(
|
||||
"sglang.srt.layers.moe.get_moe_a2a_backend",
|
||||
return_value=self._deepep_backend(),
|
||||
),
|
||||
envs.SGLANG_NPU_DSV4_DEEPEP_LL_DISPATCH_QUANT_MODE.override("invalid"),
|
||||
):
|
||||
_configure_dsv4_deepep_dispatcher(layer)
|
||||
|
||||
dispatcher.set_quant_config.assert_called_once_with(
|
||||
{"dispatcher_output_dtype": "bf16"}
|
||||
)
|
||||
|
||||
|
||||
class _LowLatencyBuffer:
|
||||
def __init__(self):
|
||||
self.kwargs = None
|
||||
|
||||
def low_latency_dispatch(
|
||||
self,
|
||||
hidden_states,
|
||||
topk_ids,
|
||||
num_max_dispatch_tokens_per_rank,
|
||||
num_experts,
|
||||
*,
|
||||
use_fp8,
|
||||
quant_mode=_NOT_PASSED,
|
||||
**kwargs,
|
||||
):
|
||||
self.kwargs = {"use_fp8": use_fp8, "quant_mode": quant_mode, **kwargs}
|
||||
return torch.empty(0), torch.empty(0), object(), object(), object()
|
||||
|
||||
|
||||
class _LegacyLowLatencyBuffer:
|
||||
def low_latency_dispatch(
|
||||
self,
|
||||
hidden_states,
|
||||
topk_ids,
|
||||
num_max_dispatch_tokens_per_rank,
|
||||
num_experts,
|
||||
*,
|
||||
use_fp8,
|
||||
**kwargs,
|
||||
):
|
||||
return torch.empty(0), torch.empty(0), object(), object(), object()
|
||||
|
||||
|
||||
class TestDeepEPLowLatencyMxfp8Dispatch(unittest.TestCase):
|
||||
@staticmethod
|
||||
def _dispatcher(quant_mode, buffer):
|
||||
dispatcher = object.__new__(deepep._DeepEPDispatcherImplLowLatency)
|
||||
dispatcher.quant_config = {}
|
||||
dispatcher.use_fp8 = False
|
||||
dispatcher.use_nvfp4 = False
|
||||
dispatcher.low_latency_quant_mode = quant_mode
|
||||
dispatcher._low_latency_quant_mode_runtime_checked = False
|
||||
dispatcher.num_max_dispatch_tokens_per_rank = 2
|
||||
dispatcher.num_experts = 2
|
||||
dispatcher.return_recv_hook = False
|
||||
dispatcher._get_buffer = lambda: buffer
|
||||
return dispatcher
|
||||
|
||||
def test_mxfp8_passes_the_kernel_quant_mode(self):
|
||||
buffer = _LowLatencyBuffer()
|
||||
dispatcher = self._dispatcher("mx_fp8_e4m3", buffer)
|
||||
|
||||
with (
|
||||
patch.dict(os.environ, {}, clear=True),
|
||||
patch.object(deepep, "_deepep_precompile_tp_barrier"),
|
||||
):
|
||||
dispatcher._dispatch_core(
|
||||
torch.zeros(1, 64),
|
||||
torch.zeros(1, 1, dtype=torch.int64),
|
||||
torch.ones(1, 1),
|
||||
)
|
||||
|
||||
self.assertEqual(buffer.kwargs["quant_mode"], "mx_fp8_e4m3")
|
||||
|
||||
def test_mxfp8_ops_strategy_uses_legacy_mxfp8_flags(self):
|
||||
buffer = _LowLatencyBuffer()
|
||||
dispatcher = self._dispatcher("mx_fp8_e4m3", buffer)
|
||||
|
||||
with (
|
||||
patch.dict(os.environ, {"DEEP_USE_MODE": "ops"}, clear=True),
|
||||
patch.object(deepep, "_deepep_precompile_tp_barrier"),
|
||||
):
|
||||
dispatcher._dispatch_core(
|
||||
torch.zeros(1, 64),
|
||||
torch.zeros(1, 1, dtype=torch.int64),
|
||||
torch.ones(1, 1),
|
||||
)
|
||||
|
||||
self.assertTrue(buffer.kwargs["use_fp8"])
|
||||
self.assertTrue(buffer.kwargs["use_ue8m0"])
|
||||
self.assertEqual(buffer.kwargs["quant_mode"], "mx_fp8_e4m3")
|
||||
|
||||
def test_mxfp8_rejects_an_unsupported_low_latency_strategy(self):
|
||||
dispatcher = self._dispatcher("mx_fp8_e4m3", _LowLatencyBuffer())
|
||||
|
||||
with (
|
||||
patch.dict(os.environ, {"DEEP_USE_MODE": "alltoall"}, clear=True),
|
||||
self.assertRaisesRegex(RuntimeError, "DEEP_USE_MODE"),
|
||||
):
|
||||
dispatcher._dispatch_core(
|
||||
torch.zeros(1, 64),
|
||||
torch.zeros(1, 1, dtype=torch.int64),
|
||||
torch.ones(1, 1),
|
||||
)
|
||||
|
||||
def test_mxfp8_checks_runtime_interface_once_per_dispatcher(self):
|
||||
buffer = _LowLatencyBuffer()
|
||||
dispatcher = self._dispatcher("mx_fp8_e4m3", buffer)
|
||||
|
||||
with (
|
||||
patch.dict(os.environ, {}, clear=True),
|
||||
patch.object(deepep, "_deepep_precompile_tp_barrier"),
|
||||
patch.object(
|
||||
deepep.inspect, "signature", wraps=inspect.signature
|
||||
) as signature,
|
||||
):
|
||||
dispatcher._dispatch_core(
|
||||
torch.zeros(1, 64),
|
||||
torch.zeros(1, 1, dtype=torch.int64),
|
||||
torch.ones(1, 1),
|
||||
)
|
||||
dispatcher._dispatch_core(
|
||||
torch.zeros(1, 64),
|
||||
torch.zeros(1, 1, dtype=torch.int64),
|
||||
torch.ones(1, 1),
|
||||
)
|
||||
|
||||
self.assertEqual(signature.call_count, 1)
|
||||
|
||||
def test_bf16_does_not_pass_a_quant_mode(self):
|
||||
buffer = _LowLatencyBuffer()
|
||||
dispatcher = self._dispatcher(None, buffer)
|
||||
|
||||
with patch.object(deepep, "_deepep_precompile_tp_barrier"):
|
||||
dispatcher._dispatch_core(
|
||||
torch.zeros(1, 64),
|
||||
torch.zeros(1, 1, dtype=torch.int64),
|
||||
torch.ones(1, 1),
|
||||
)
|
||||
|
||||
self.assertIs(buffer.kwargs["quant_mode"], _NOT_PASSED)
|
||||
|
||||
def test_mxfp8_rejects_legacy_runtime_without_quant_mode(self):
|
||||
dispatcher = self._dispatcher("mx_fp8_e4m3", _LegacyLowLatencyBuffer())
|
||||
|
||||
with self.assertRaisesRegex(RuntimeError, "quant_mode"):
|
||||
dispatcher._dispatch_core(
|
||||
torch.zeros(1, 64),
|
||||
torch.zeros(1, 1, dtype=torch.int64),
|
||||
torch.ones(1, 1),
|
||||
)
|
||||
|
||||
|
||||
class TestW4A8MxfpGmmInputScale(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.input = torch.randn(2, 64)
|
||||
self.input_scale = torch.ones(2, 1, 2)
|
||||
self.weight = torch.empty(2, 64, 32, dtype=torch.uint8)
|
||||
self.weight_scale = torch.ones(2, 1, 32, 2, dtype=torch.uint8)
|
||||
self.group_list = torch.tensor([1, 1], dtype=torch.int32)
|
||||
|
||||
def _call_gmm(self, input_scale):
|
||||
return w4a8_mxfp_gmm(
|
||||
input=self.input,
|
||||
input_scale=input_scale,
|
||||
weight=self.weight,
|
||||
weight_scale=self.weight_scale,
|
||||
group_list_type=1,
|
||||
group_list=self.group_list,
|
||||
output_dtype=torch.bfloat16,
|
||||
)
|
||||
|
||||
def test_supplied_scale_skips_dynamic_quant(self):
|
||||
expected = torch.randn(2, 32)
|
||||
with (
|
||||
patch.object(
|
||||
torch.ops.npu, "npu_dynamic_mx_quant", create=True
|
||||
) as dynamic_quant,
|
||||
patch.object(
|
||||
torch.ops.npu,
|
||||
"npu_grouped_matmul",
|
||||
return_value=[expected],
|
||||
create=True,
|
||||
) as grouped_matmul,
|
||||
):
|
||||
output = self._call_gmm(self.input_scale)
|
||||
|
||||
dynamic_quant.assert_not_called()
|
||||
self.assertIs(output, expected)
|
||||
call_kwargs = grouped_matmul.call_args.kwargs
|
||||
self.assertIs(call_kwargs["per_token_scale"][0], self.input_scale)
|
||||
self.assertEqual(call_kwargs["group_list"].dtype, torch.int64)
|
||||
self.assertTrue(torch.equal(call_kwargs["group_list"], self.group_list))
|
||||
|
||||
def test_flat_deepep_scale_skips_dynamic_quant_after_layout_adaptation(self):
|
||||
flat_scale = torch.arange(4, dtype=torch.uint8)
|
||||
expected = torch.randn(2, 32)
|
||||
with (
|
||||
patch.object(
|
||||
torch.ops.npu, "npu_dynamic_mx_quant", create=True
|
||||
) as dynamic_quant,
|
||||
patch.object(
|
||||
torch.ops.npu,
|
||||
"npu_grouped_matmul",
|
||||
return_value=[expected],
|
||||
create=True,
|
||||
) as grouped_matmul,
|
||||
):
|
||||
output = self._call_gmm(flat_scale)
|
||||
|
||||
dynamic_quant.assert_not_called()
|
||||
self.assertIs(output, expected)
|
||||
packed_scale = grouped_matmul.call_args.kwargs["per_token_scale"][0]
|
||||
self.assertEqual(tuple(packed_scale.shape), (2, 1, 2))
|
||||
self.assertEqual(packed_scale.data_ptr(), flat_scale.data_ptr())
|
||||
|
||||
def test_missing_scale_uses_dynamic_quant(self):
|
||||
quantized = torch.empty(2, 64, dtype=torch.float8_e4m3fn)
|
||||
quantized_scale = torch.ones(2, 1, 2)
|
||||
expected = torch.randn(2, 32)
|
||||
with (
|
||||
patch.object(
|
||||
torch.ops.npu,
|
||||
"npu_dynamic_mx_quant",
|
||||
return_value=(quantized, quantized_scale),
|
||||
create=True,
|
||||
) as dynamic_quant,
|
||||
patch.object(
|
||||
torch.ops.npu,
|
||||
"npu_grouped_matmul",
|
||||
return_value=[expected],
|
||||
create=True,
|
||||
) as grouped_matmul,
|
||||
):
|
||||
output = self._call_gmm(None)
|
||||
|
||||
dynamic_quant.assert_called_once()
|
||||
self.assertIs(output, expected)
|
||||
self.assertIs(
|
||||
grouped_matmul.call_args.kwargs["per_token_scale"][0], quantized_scale
|
||||
)
|
||||
|
||||
|
||||
class TestW4A8MxfpGmmChain(unittest.TestCase):
|
||||
def test_passes_swiglu_limit_to_quant(self):
|
||||
gate_up = torch.randn(1, 64)
|
||||
activated = torch.randn(1, 32)
|
||||
activated_scale = torch.randn(1, 1)
|
||||
expected = torch.randn(1, 32)
|
||||
layer = SimpleNamespace(
|
||||
w13_weight=MagicMock(),
|
||||
w13_weight_scale_inv=MagicMock(),
|
||||
w2_weight=MagicMock(),
|
||||
w2_weight_scale_inv=MagicMock(),
|
||||
moe_runner_config=SimpleNamespace(swiglu_limit=7.0),
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
fp4_moe_methods, "w4a8_mxfp_gmm", side_effect=[gate_up, expected]
|
||||
) as gmm,
|
||||
patch.object(
|
||||
fp4_moe_methods,
|
||||
"swiglu_quant",
|
||||
return_value=(activated, activated_scale),
|
||||
) as swiglu,
|
||||
):
|
||||
output = npu_apply_without_routing_weights_w4a4_mxfp(
|
||||
layer,
|
||||
torch.randn(1, 4),
|
||||
torch.ones(1, 1, 2),
|
||||
group_list_type=1,
|
||||
group_list=torch.tensor([1], dtype=torch.int64),
|
||||
output_dtype=torch.bfloat16,
|
||||
)
|
||||
|
||||
self.assertIs(output, expected)
|
||||
self.assertTrue(torch.equal(swiglu.call_args.args[0], gate_up))
|
||||
self.assertTrue(swiglu.call_args.kwargs["do_limit"])
|
||||
self.assertEqual(swiglu.call_args.kwargs["limit"], 7.0)
|
||||
self.assertIs(gmm.call_args_list[1].kwargs["input"], activated)
|
||||
self.assertIs(gmm.call_args_list[1].kwargs["input_scale"], activated_scale)
|
||||
|
||||
|
||||
class TestProcessWeightsAfterLoadingZeroScale(unittest.TestCase):
|
||||
@staticmethod
|
||||
def _method():
|
||||
return NPUW4A4Fp4MoEMethod(fp8_method=MagicMock(), prefix="test")
|
||||
|
||||
def test_raises_when_w13_scales_never_loaded(self):
|
||||
# An all-zero scale is the signature of a checkpoint whose scale names
|
||||
# never matched; without this guard every routed expert computes silently
|
||||
# as zero instead of failing loudly.
|
||||
layer = SimpleNamespace(
|
||||
w13_weight_scale_inv=torch.nn.Parameter(
|
||||
torch.zeros(2, 2, 4, dtype=torch.uint8), requires_grad=False
|
||||
),
|
||||
w2_weight_scale_inv=torch.nn.Parameter(
|
||||
torch.zeros(2, 2, 4, dtype=torch.uint8), requires_grad=False
|
||||
),
|
||||
)
|
||||
with self.assertRaises(RuntimeError):
|
||||
self._method().process_weights_after_loading(layer)
|
||||
|
||||
def test_raises_when_w2_scales_never_loaded(self):
|
||||
layer = SimpleNamespace(
|
||||
w13_weight_scale_inv=torch.nn.Parameter(
|
||||
torch.ones(2, 2, 4, dtype=torch.uint8), requires_grad=False
|
||||
),
|
||||
w2_weight_scale_inv=torch.nn.Parameter(
|
||||
torch.zeros(2, 2, 4, dtype=torch.uint8), requires_grad=False
|
||||
),
|
||||
)
|
||||
with self.assertRaises(RuntimeError):
|
||||
self._method().process_weights_after_loading(layer)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,171 @@
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.test.ci.ci_register import register_npu_ci
|
||||
|
||||
register_npu_ci(est_time=1, suite="stage-a-unit-test-npu")
|
||||
|
||||
# Load the quantization package first so `base_config`, `moe_methods`, and
|
||||
# `linear_method_npu` initialize in dependency order. Importing
|
||||
# `linear_method_npu` directly from a cold process triggers a circular import:
|
||||
# linear_method_npu -> base_config -> quantization/__init__ ->
|
||||
# gguf/unquant/gptq_moe -> moe_methods -> linear_method_npu (partially
|
||||
# initialized, `_get_float8_e8m0fnu_dtype` not yet defined). Initializing the
|
||||
# package first mirrors how the engine loads quantization at model-config time.
|
||||
import sglang.srt.layers.quantization # noqa: F401
|
||||
from sglang.srt.hardware_backend.npu.quantization.linear_method_npu import (
|
||||
npu_w8a8_mxfp8_linear,
|
||||
)
|
||||
from sglang.srt.hardware_backend.npu.quantization.w8a8_mxfp8 import (
|
||||
process_npu_arch35_mxfp8_linear_weights,
|
||||
)
|
||||
from sglang.srt.layers.quantization.fp8 import Fp8Config
|
||||
|
||||
|
||||
class TestNPUW8A8BlockFP8Linear(unittest.TestCase):
|
||||
def test_fp8_config_preserves_ue8m0_scale_format(self):
|
||||
quant_config = Fp8Config.from_config(
|
||||
{
|
||||
"quant_method": "fp8",
|
||||
"activation_scheme": "dynamic",
|
||||
"weight_block_size": [128, 128],
|
||||
"scale_fmt": "ue8m0",
|
||||
}
|
||||
)
|
||||
|
||||
self.assertEqual(quant_config.scale_fmt, "ue8m0")
|
||||
|
||||
def test_layout_only_ue8m0_conversion_preserves_fp8_weight(self):
|
||||
original_weight = torch.randint(1, 255, (128, 64), dtype=torch.uint8).view(
|
||||
torch.float8_e4m3fn
|
||||
)
|
||||
layer = SimpleNamespace(
|
||||
weight=torch.nn.Parameter(original_weight.clone(), requires_grad=False),
|
||||
weight_scale_inv=torch.nn.Parameter(
|
||||
torch.tensor([[2**-12]], dtype=torch.float32), requires_grad=False
|
||||
),
|
||||
)
|
||||
|
||||
process_npu_arch35_mxfp8_linear_weights(layer, [128, 128], scale_fmt="ue8m0")
|
||||
|
||||
self.assertEqual(layer.weight.shape, (64, 128))
|
||||
torch.testing.assert_close(
|
||||
layer.weight.data.T.contiguous().view(torch.uint8),
|
||||
original_weight.view(torch.uint8),
|
||||
)
|
||||
self.assertEqual(layer.weight_scale_inv.shape, (1, 128, 2))
|
||||
self.assertTrue(
|
||||
torch.equal(
|
||||
layer.weight_scale_inv.data,
|
||||
torch.full((1, 128, 2), 0x73, dtype=torch.uint8),
|
||||
)
|
||||
)
|
||||
self.assertTrue(layer.weight_scale_inv.format_ue8m0)
|
||||
|
||||
def test_rejects_non_ue8m0_scale_format(self):
|
||||
with self.assertRaisesRegex(ValueError, "scale_fmt='ue8m0'"):
|
||||
process_npu_arch35_mxfp8_linear_weights(
|
||||
SimpleNamespace(), [128, 128], scale_fmt="float32"
|
||||
)
|
||||
|
||||
def test_rejects_non_fp8_weight(self):
|
||||
with self.assertRaisesRegex(ValueError, "expects float8_e4m3fn weights"):
|
||||
npu_w8a8_mxfp8_linear(
|
||||
torch.empty(1, 128, dtype=torch.bfloat16),
|
||||
torch.empty(128, 64, dtype=torch.bfloat16),
|
||||
[128, 128],
|
||||
torch.empty(1),
|
||||
)
|
||||
|
||||
def test_quantizes_flattened_input_and_restores_batch_shape(self):
|
||||
input_tensor = torch.randn(2, 3, 128, dtype=torch.bfloat16)
|
||||
weight = torch.empty(128, 64, dtype=torch.float8_e4m3fn)
|
||||
weight_scale = torch.empty(2, 64, 2, dtype=torch.uint8)
|
||||
bias = torch.randn(64, dtype=torch.float32)
|
||||
quantized = torch.empty(6, 128, dtype=torch.float8_e4m3fn)
|
||||
input_scale = torch.empty(6, 2, 2, dtype=torch.uint8)
|
||||
matmul_output = torch.randn(6, 64, dtype=torch.bfloat16)
|
||||
|
||||
npu_ops = MagicMock()
|
||||
npu_ops.npu_dynamic_mx_quant.return_value = (quantized, input_scale)
|
||||
npu_ops.npu_quant_matmul.return_value = matmul_output
|
||||
with patch.object(torch.ops, "npu", npu_ops, create=True):
|
||||
output = npu_w8a8_mxfp8_linear(
|
||||
input_tensor,
|
||||
weight,
|
||||
[64, 128],
|
||||
weight_scale,
|
||||
bias=bias,
|
||||
)
|
||||
|
||||
self.assertEqual(output.shape, (2, 3, 64))
|
||||
quant_call = npu_ops.npu_dynamic_mx_quant.call_args
|
||||
self.assertEqual(quant_call.args[0].shape, (6, 128))
|
||||
self.assertEqual(quant_call.kwargs["dst_type"], torch.float8_e4m3fn)
|
||||
|
||||
matmul_call = npu_ops.npu_quant_matmul.call_args
|
||||
self.assertIs(matmul_call.args[0], quantized)
|
||||
self.assertIs(matmul_call.args[1], weight)
|
||||
self.assertIs(matmul_call.kwargs["scale"], weight_scale)
|
||||
self.assertIs(matmul_call.kwargs["pertoken_scale"], input_scale)
|
||||
self.assertIs(matmul_call.kwargs["bias"], bias)
|
||||
self.assertEqual(matmul_call.kwargs["group_sizes"], (1, 1, 32))
|
||||
|
||||
def test_rejects_noncontiguous_input(self):
|
||||
input_tensor = torch.randn(2, 3, 128, dtype=torch.bfloat16).transpose(0, 1)
|
||||
weight = torch.empty(128, 64, dtype=torch.float8_e4m3fn)
|
||||
weight_scale = torch.empty(2, 64, 2, dtype=torch.uint8)
|
||||
|
||||
with self.assertRaisesRegex(RuntimeError, "view size is not compatible"):
|
||||
npu_w8a8_mxfp8_linear(input_tensor, weight, [64, 128], weight_scale)
|
||||
|
||||
def test_preserves_supported_input_dtype(self):
|
||||
input_tensor = torch.randn(2, 128, dtype=torch.float16)
|
||||
weight = torch.empty(128, 64, dtype=torch.float8_e4m3fn)
|
||||
weight_scale = torch.empty(2, 64, 2, dtype=torch.uint8)
|
||||
npu_ops = MagicMock()
|
||||
npu_ops.npu_dynamic_mx_quant.return_value = (
|
||||
torch.empty(2, 128, dtype=torch.float8_e4m3fn),
|
||||
torch.empty(2, 2, 2, dtype=torch.uint8),
|
||||
)
|
||||
npu_ops.npu_quant_matmul.return_value = torch.empty(2, 64)
|
||||
|
||||
with patch.object(torch.ops, "npu", npu_ops, create=True):
|
||||
npu_w8a8_mxfp8_linear(input_tensor, weight, [128, 128], weight_scale)
|
||||
|
||||
self.assertEqual(
|
||||
npu_ops.npu_quant_matmul.call_args.kwargs["output_dtype"],
|
||||
torch.float16,
|
||||
)
|
||||
|
||||
def test_converts_bias_to_float32(self):
|
||||
input_tensor = torch.randn(2, 128, dtype=torch.bfloat16)
|
||||
weight = torch.empty(128, 64, dtype=torch.float8_e4m3fn)
|
||||
weight_scale = torch.empty(2, 64, 2, dtype=torch.uint8)
|
||||
bias = torch.randn(64, dtype=torch.bfloat16)
|
||||
npu_ops = MagicMock()
|
||||
npu_ops.npu_dynamic_mx_quant.return_value = (
|
||||
torch.empty(2, 128, dtype=torch.float8_e4m3fn),
|
||||
torch.empty(2, 2, 2, dtype=torch.uint8),
|
||||
)
|
||||
npu_ops.npu_quant_matmul.return_value = torch.empty(2, 64)
|
||||
|
||||
with patch.object(torch.ops, "npu", npu_ops, create=True):
|
||||
npu_w8a8_mxfp8_linear(
|
||||
input_tensor,
|
||||
weight,
|
||||
[128, 128],
|
||||
weight_scale,
|
||||
bias=bias,
|
||||
)
|
||||
|
||||
quant_bias = npu_ops.npu_quant_matmul.call_args.kwargs["bias"]
|
||||
self.assertEqual(quant_bias.dtype, torch.float32)
|
||||
torch.testing.assert_close(quant_bias, bias.float())
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,50 @@
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
from sglang.test.ci.ci_register import register_npu_ci
|
||||
|
||||
register_npu_ci(est_time=1, suite="stage-a-unit-test-npu")
|
||||
|
||||
from sglang.srt.hardware_backend.npu.utils import is_npu_arch35
|
||||
|
||||
|
||||
class TestArch35Capability(unittest.TestCase):
|
||||
def tearDown(self):
|
||||
is_npu_arch35.cache_clear()
|
||||
|
||||
def test_arch35_is_supported_from_acl_device_info(self):
|
||||
with (
|
||||
patch("sglang.srt.hardware_backend.npu.utils.is_npu", return_value=True),
|
||||
patch.dict(
|
||||
"sys.modules",
|
||||
{
|
||||
"acl": SimpleNamespace(
|
||||
rt=SimpleNamespace(get_device_info=lambda *_: (3510, 0))
|
||||
)
|
||||
},
|
||||
),
|
||||
):
|
||||
self.assertTrue(is_npu_arch35())
|
||||
|
||||
def test_non_arch35_npu_is_not_supported(self):
|
||||
with (
|
||||
patch("sglang.srt.hardware_backend.npu.utils.is_npu", return_value=True),
|
||||
patch.dict(
|
||||
"sys.modules",
|
||||
{
|
||||
"acl": SimpleNamespace(
|
||||
rt=SimpleNamespace(get_device_info=lambda *_: (2901, 0))
|
||||
)
|
||||
},
|
||||
),
|
||||
):
|
||||
self.assertFalse(is_npu_arch35())
|
||||
|
||||
@patch("sglang.srt.hardware_backend.npu.utils.is_npu", return_value=False)
|
||||
def test_non_npu_is_not_supported(self, _):
|
||||
self.assertFalse(is_npu_arch35())
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user