From 0f35f896262989edd22763d933b1ecc1886195d4 Mon Sep 17 00:00:00 2001 From: "ext.zhaoxingchen5" Date: Thu, 23 Jul 2026 11:04:03 +0800 Subject: [PATCH] feat: rewrite attention tests in python aclnn, add metadata & mc2 dispatch ops --- test/python_test/RegisterOps.cpp | 422 +++++++++++++ test/python_test/compressor_gen.py | 151 ----- test/python_test/custom_ops.py | 157 +++++ test/python_test/pytorch_npu_helper.hpp | 24 +- .../quant_lightning_indexer_gen.py | 131 ---- test/python_test/sparse_attn_sharedkv_gen.py | 144 ----- test/python_test/test_compressor.cpp | 193 ------ test/python_test/test_compressor.py | 126 ++++ test/python_test/test_dispatch_ffn_combine.py | 229 +++++++ .../test_dispatch_gmm_combine_decode.py | 589 ++++++++++++++++++ test/python_test/test_index_group_matmul.py | 5 + .../test_lightning_indexer_quant_metadata.py | 124 ++++ .../test_quant_lightning_indexer.cpp | 166 ----- .../test_quant_lightning_indexer.py | 230 +++++++ .../test_quant_lightning_indexer_metadata.py | 119 ++++ .../python_test/test_sparse_attn_sharedkv.cpp | 213 ------- test/python_test/test_sparse_attn_sharedkv.py | 127 ++++ .../test_sparse_attn_sharedkv_metadata.py | 134 ++++ 18 files changed, 2283 insertions(+), 1001 deletions(-) delete mode 100644 test/python_test/compressor_gen.py delete mode 100644 test/python_test/quant_lightning_indexer_gen.py delete mode 100644 test/python_test/sparse_attn_sharedkv_gen.py delete mode 100644 test/python_test/test_compressor.cpp create mode 100644 test/python_test/test_compressor.py create mode 100644 test/python_test/test_dispatch_ffn_combine.py create mode 100644 test/python_test/test_dispatch_gmm_combine_decode.py create mode 100644 test/python_test/test_lightning_indexer_quant_metadata.py delete mode 100644 test/python_test/test_quant_lightning_indexer.cpp create mode 100644 test/python_test/test_quant_lightning_indexer.py create mode 100644 test/python_test/test_quant_lightning_indexer_metadata.py delete mode 100644 test/python_test/test_sparse_attn_sharedkv.cpp create mode 100644 test/python_test/test_sparse_attn_sharedkv.py create mode 100644 test/python_test/test_sparse_attn_sharedkv_metadata.py diff --git a/test/python_test/RegisterOps.cpp b/test/python_test/RegisterOps.cpp index ab0300c..dc3dea8 100644 --- a/test/python_test/RegisterOps.cpp +++ b/test/python_test/RegisterOps.cpp @@ -1430,6 +1430,420 @@ std::tuple moe_gating_top_k_hash_impl_npu( return std::make_tuple(y, expert_idx, out); } +// ==================== compressor ==================== +// Single-step op (25 params). kv_state/score_state are in-place (Ref). +// When enable_grad=false, wkvProj/softmaxRes/normX/normRstd are {1} placeholders. +at::Tensor compressor_impl_npu( + const at::Tensor& x, + const at::Tensor& wkv, + const at::Tensor& wgate, + at::Tensor& kv_state, + at::Tensor& score_state, + const at::Tensor& ape, + const at::Tensor& norm_weight, + const at::Tensor& rope_sin, + const at::Tensor& rope_cos, + const at::Tensor& kv_block_table, + const at::Tensor& score_block_table, + int64_t rope_head_dim, + int64_t cmp_ratio, + int64_t coff, + double norm_eps, + int64_t rotary_mode, + bool enable_grad) { + auto x_sizes = x.sizes().vec(); + const int64_t B = x_sizes[0]; + const int64_t S = x_sizes[1]; + const int64_t HEAD_DIM = wkv.sizes()[0] / coff; // coff_d / coff + const int64_t SR = (S + cmp_ratio - 1) / cmp_ratio; + at::Tensor cmp_kv = at::empty({B, SR, HEAD_DIM}, x.options()); + at::Tensor wkv_proj = at::empty({1}, x.options()); + at::Tensor softmax_res = at::empty({1}, x.options()); + at::Tensor norm_x = at::empty({1}, x.options()); + at::Tensor norm_rstd = at::empty({1}, x.options()); + const c10::optional null_opt; + EXEC_NPU_CMD(aclnnCompressor, + x, wkv, wgate, kv_state, score_state, ape, norm_weight, + rope_sin, rope_cos, kv_block_table, score_block_table, + null_opt, null_opt, null_opt, + rope_head_dim, cmp_ratio, coff, norm_eps, rotary_mode, enable_grad, + cmp_kv, wkv_proj, softmax_res, norm_x, norm_rstd); + return cmp_kv; +} + +// ==================== sparse_attn_sharedkv (two-stage) ==================== +// Helper: metadata op (AICPU). Must be in separate function due to +// EXEC_NPU_CMD static variable binding. +static void sparse_attn_sharedkv_metadata_helper( + const at::Tensor& cu_seq_q, + const at::Tensor& cu_seq_ori_kv, + const at::Tensor& cu_seq_cmp_kv, + const at::Tensor& seqused_q, + const at::Tensor& seqused_kv, + int64_t n1, int64_t kv_n, int64_t d, + int64_t batch, int64_t s1, int64_t s2, + int64_t ori_top_k, int64_t cmp_top_k, int64_t cmp_ratio, + int64_t ori_mask_mode, int64_t cmp_mask_mode, + int64_t ori_win_left, int64_t ori_win_right, + char* layout_q_c, char* layout_kv_c, + bool is_prefill, bool return_metadata_only, + at::Tensor& meta_t) { + EXEC_NPU_CMD(aclnnSparseAttnSharedkvMetadata, + cu_seq_q, cu_seq_ori_kv, cu_seq_cmp_kv, seqused_q, seqused_kv, + n1, kv_n, d, batch, s1, s2, + ori_top_k, cmp_top_k, cmp_ratio, + ori_mask_mode, cmp_mask_mode, ori_win_left, ori_win_right, + layout_q_c, layout_kv_c, is_prefill, return_metadata_only, + meta_t); +} + +// Helper: main op. +static void sparse_attn_sharedkv_main_helper( + const at::Tensor& query, + const at::Tensor& ori_kv, + const c10::optional& cmp_kv, + const c10::optional& ori_sparse_idx, + const c10::optional& cmp_sparse_idx, + const at::Tensor& ori_bt, + const c10::optional& cmp_bt, + const c10::optional& cu_seq_q, + const c10::optional& cu_seq_ori_kv, + const c10::optional& cu_seq_cmp_kv, + const c10::optional& seqused_q, + const at::Tensor& seqused_kv, + const at::Tensor& sinks, + const at::Tensor& meta_t, + double softmax_scale, int64_t cmp_ratio, + int64_t ori_mask_mode, int64_t cmp_mask_mode, + int64_t ori_kv_stride, int64_t cmp_kv_stride, + int64_t ori_win_left, int64_t ori_win_right, + char* layout_q_c, char* layout_kv_c, + bool return_softmax_lse, + at::Tensor& attn_out, at::Tensor& lse_out) { + EXEC_NPU_CMD(aclnnSparseAttnSharedkv, + query, ori_kv, cmp_kv, ori_sparse_idx, cmp_sparse_idx, + ori_bt, cmp_bt, + cu_seq_q, cu_seq_ori_kv, cu_seq_cmp_kv, seqused_q, + seqused_kv, sinks, meta_t, + softmax_scale, cmp_ratio, ori_mask_mode, cmp_mask_mode, + ori_kv_stride, cmp_kv_stride, + ori_win_left, ori_win_right, + layout_q_c, layout_kv_c, return_softmax_lse, + attn_out, lse_out); +} + +// Public impl: runs metadata then main, returns attn_out. +at::Tensor sparse_attn_sharedkv_impl_npu( + const at::Tensor& query, + const at::Tensor& ori_kv, + const at::Tensor& ori_bt, + const at::Tensor& seqused_kv, + const at::Tensor& sinks, + int64_t n1, int64_t kv_n, int64_t d, + int64_t s1, int64_t s2, + int64_t ori_mask_mode, int64_t cmp_mask_mode, + int64_t ori_win_left, int64_t ori_win_right, + double softmax_scale, int64_t cmp_ratio, + std::string layout_q, std::string layout_kv) { + auto q_sizes = query.sizes().vec(); + const int64_t B = q_sizes[0]; + // Metadata inputs: cu_seqlens (prefix-sum of lengths), seqused_q + auto opts_i32 = query.options().dtype(at::kInt); + std::vector cu_q_h(B + 1, 0), cu_ori_h(B + 1, 0), cu_cmp_h(B + 1, 0); + std::vector sused_q_h(B, (int32_t)s1); + for (int64_t i = 0; i < B; ++i) { + cu_q_h[i + 1] = cu_q_h[i] + (int32_t)s1; + cu_ori_h[i + 1] = cu_ori_h[i] + (int32_t)s2; + cu_cmp_h[i + 1] = cu_cmp_h[i] + (int32_t)s2; + } + at::Tensor cu_seq_q = at::from_blob(cu_q_h.data(), {B + 1}, at::kInt).to(query.device()); + at::Tensor cu_seq_ori_kv = at::from_blob(cu_ori_h.data(), {B + 1}, at::kInt).to(query.device()); + at::Tensor cu_seq_cmp_kv = at::from_blob(cu_cmp_h.data(), {B + 1}, at::kInt).to(query.device()); + at::Tensor seqused_q = at::from_blob(sused_q_h.data(), {B}, at::kInt).to(query.device()); + const int64_t SAS_META_SIZE = 1024; + at::Tensor meta_t = at::empty({SAS_META_SIZE}, opts_i32.device(query.device())); + char* lq = const_cast(layout_q.c_str()); + char* lk = const_cast(layout_kv.c_str()); + // Step 1: metadata + sparse_attn_sharedkv_metadata_helper( + cu_seq_q, cu_seq_ori_kv, cu_seq_cmp_kv, seqused_q, seqused_kv, + n1, kv_n, d, B, s1, s2, + 512, 512, cmp_ratio, // ori_top_k, cmp_top_k + ori_mask_mode, cmp_mask_mode, ori_win_left, ori_win_right, + lq, lk, true, false, meta_t); + // Step 2: main + at::Tensor attn_out = at::empty({B, s1, n1, d}, query.options()); + at::Tensor lse_out = at::empty({B, s1, n1}, query.options().dtype(at::kFloat)); + const c10::optional null_opt; + sparse_attn_sharedkv_main_helper( + query, ori_kv, null_opt, null_opt, null_opt, ori_bt, null_opt, + null_opt, null_opt, null_opt, null_opt, + seqused_kv, sinks, meta_t, + softmax_scale, cmp_ratio, ori_mask_mode, cmp_mask_mode, 0, 0, + ori_win_left, ori_win_right, lq, lk, false, + attn_out, lse_out); + return attn_out; +} + +// ==================== quant_lightning_indexer (two-stage) ==================== +// Helper: metadata op. +static void quant_lightning_indexer_metadata_helper( + const at::Tensor& aslq, const at::Tensor& aslk, + int64_t num_heads_q, int64_t num_heads_k, int64_t head_dim, + int64_t query_quant_mode, int64_t key_quant_mode, + int64_t batch_size, int64_t max_seq_q, int64_t max_seq_k, + char* layout_query_c, char* layout_key_c, + int64_t sparse_count, int64_t sparse_mode, + int64_t pre_token, int64_t next_token, int64_t cmp_ratio, + at::Tensor& meta_t) { + EXEC_NPU_CMD(aclnnQuantLightningIndexerMetadata, + aslq, aslk, num_heads_q, num_heads_k, head_dim, + query_quant_mode, key_quant_mode, + batch_size, max_seq_q, max_seq_k, + layout_query_c, layout_key_c, + sparse_count, sparse_mode, pre_token, next_token, cmp_ratio, + meta_t); +} + +// Helper: main op. +static void quant_lightning_indexer_main_helper( + const at::Tensor& query, const at::Tensor& key, + const at::Tensor& weights, const at::Tensor& q_scale, const at::Tensor& k_scale, + const at::Tensor& aslq, const at::Tensor& aslk, + const at::Tensor& bt, const at::Tensor& meta_t, + int64_t query_quant_mode, int64_t key_quant_mode, + char* layout_query_c, char* layout_key_c, + int64_t sparse_count, int64_t sparse_mode, + int64_t pre_token, int64_t next_token, int64_t cmp_ratio, + bool return_values, int64_t stride, int64_t scale_stride, + at::Tensor& idx_out, at::Tensor& val_out) { + EXEC_NPU_CMD(aclnnQuantLightningIndexer, + query, key, weights, q_scale, k_scale, + aslq, aslk, bt, meta_t, + query_quant_mode, key_quant_mode, + layout_query_c, layout_key_c, + sparse_count, sparse_mode, pre_token, next_token, cmp_ratio, + return_values, stride, scale_stride, + idx_out, val_out); +} + +// Public impl: runs metadata then main, returns idx_out. +at::Tensor quant_lightning_indexer_impl_npu( + const at::Tensor& query, const at::Tensor& key, + const at::Tensor& weights, const at::Tensor& q_scale, const at::Tensor& k_scale, + const at::Tensor& aslq, const at::Tensor& aslk, + const at::Tensor& block_table, + int64_t num_heads_q, int64_t num_heads_k, int64_t head_dim, + int64_t query_quant_mode, int64_t key_quant_mode, + std::string layout_query, std::string layout_key, + int64_t sparse_count, int64_t sparse_mode, + int64_t pre_token, int64_t next_token, int64_t cmp_ratio) { + auto q_sizes = query.sizes().vec(); + const int64_t batch_size = q_sizes[0]; + const int64_t max_seq_q = q_sizes[1]; + const int64_t max_seq_k = aslk.max().item(); + const int64_t QLI_META_SIZE = 1024; + auto opts_i32 = query.options().dtype(at::kInt); + at::Tensor meta_t = at::empty({QLI_META_SIZE}, opts_i32.device(query.device())); + char* lq = const_cast(layout_query.c_str()); + char* lk = const_cast(layout_key.c_str()); + // Step 1: metadata + quant_lightning_indexer_metadata_helper( + aslq, aslk, num_heads_q, num_heads_k, head_dim, + query_quant_mode, key_quant_mode, + batch_size, max_seq_q, max_seq_k, lq, lk, + sparse_count, sparse_mode, pre_token, next_token, cmp_ratio, meta_t); + // Step 2: main + at::Tensor idx_out = at::empty({batch_size, max_seq_q, num_heads_k, sparse_count}, opts_i32.device(query.device())); + at::Tensor val_out = at::empty({batch_size, max_seq_q, num_heads_k, sparse_count}, + query.options().dtype(at::kFloat)); + quant_lightning_indexer_main_helper( + query, key, weights, q_scale, k_scale, aslq, aslk, block_table, meta_t, + query_quant_mode, key_quant_mode, lq, lk, + sparse_count, sparse_mode, pre_token, next_token, cmp_ratio, + false, 1, 1, idx_out, val_out); + return idx_out; +} + +// dispatch_ffn_combine (int8 path: aclnnDispatchFFNCombine) +std::tuple dispatch_ffn_combine_impl_npu( + const at::Tensor& x, + std::vector weight1, + std::vector weight2, + const at::Tensor& expert_idx, + std::vector scale1, + std::vector scale2, + const at::Tensor& probs, + std::string group, + int64_t max_output_size, + at::Tensor& out, + at::Tensor& expert_token_nums, + const c10::optional& x_active_mask, + double swiglu_limit) { + at::TensorList weight1_list = at::TensorList(weight1); + at::TensorList weight2_list = at::TensorList(weight2); + at::TensorList scale1_list = at::TensorList(scale1); + at::TensorList scale2_list = at::TensorList(scale2); + char* group_ep_ptr = const_cast(group.c_str()); + at::Tensor mask_tensor = x_active_mask.has_value() ? x_active_mask.value() : at::Tensor(); + EXEC_NPU_CMD(aclnnDispatchFFNCombine, + x, + weight1_list, + weight2_list, + expert_idx, + scale1_list, + scale2_list, + probs, + mask_tensor, + group_ep_ptr, + max_output_size, + swiglu_limit, + out, + expert_token_nums); + return std::make_tuple(out, expert_token_nums); +} + +// dispatch_gmm_combine_decode (aclnnDispatchGmmCombineDecode) +std::tuple dispatch_gmm_combine_decode_impl_npu( + const at::Tensor& x, + const at::Tensor& expert_ids, + std::vector gmm1_permuted_weight, + std::vector gmm1_permuted_weight_scale, + std::vector gmm2_weight, + std::vector gmm2_weight_scale, + const at::Tensor& expert_scales, + const c10::optional& expert_smooth_scales, + const c10::optional& x_active_mask, + std::string group_ep, + int64_t ep_rank_size, + int64_t ep_rank_id, + int64_t moe_expert_num, + int64_t shared_expert_num, + int64_t shared_expert_rank_num, + int64_t quant_mode, + int64_t global_bs) { + auto x_shape = x.sizes(); + int bs = x_shape[0]; + int h = x_shape[1]; + at::Tensor output = at::empty({bs, h}, x.options()); + + bool is_shared_expert = (ep_rank_id < shared_expert_rank_num); + int64_t num_local_experts = is_shared_expert ? 1 : moe_expert_num / (ep_rank_size - shared_expert_rank_num); + at::Tensor expert_token_nums = at::empty({num_local_experts}, expert_ids.options().dtype(at::kLong)); + + at::TensorList gmm1_w_list = at::TensorList(gmm1_permuted_weight); + at::TensorList gmm1_ws_list = at::TensorList(gmm1_permuted_weight_scale); + at::TensorList gmm2_w_list = at::TensorList(gmm2_weight); + at::TensorList gmm2_ws_list = at::TensorList(gmm2_weight_scale); + + std::vector group_ep_chrs(group_ep.begin(), group_ep.end()); + group_ep_chrs.push_back('\0'); + char* group_ep_ptr = &group_ep_chrs[0]; + + EXEC_NPU_CMD(aclnnDispatchGmmCombineDecode, + x, + expert_ids, + gmm1_w_list, + gmm1_ws_list, + gmm2_w_list, + gmm2_ws_list, + expert_scales, + expert_smooth_scales, + x_active_mask, + group_ep_ptr, + ep_rank_size, + ep_rank_id, + moe_expert_num, + shared_expert_num, + shared_expert_rank_num, + quant_mode, + global_bs, + output, + expert_token_nums); + return std::make_tuple(output, expert_token_nums); +} + +// ==================== standalone metadata ops ==================== + +// sparse_attn_sharedkv_metadata (standalone entry) +at::Tensor sparse_attn_sharedkv_metadata_impl_npu( + const at::Tensor& cu_seq_q, + const at::Tensor& cu_seq_ori_kv, + const at::Tensor& cu_seq_cmp_kv, + const at::Tensor& seqused_q, + const at::Tensor& seqused_kv, + int64_t num_heads_q, int64_t num_heads_kv, int64_t head_dim, + int64_t batch_size, int64_t max_seq_q, int64_t max_seq_kv, + int64_t ori_top_k, int64_t cmp_top_k, int64_t cmp_ratio, + int64_t ori_mask_mode, int64_t cmp_mask_mode, + int64_t ori_win_left, int64_t ori_win_right, + std::string layout_q, std::string layout_kv, + bool has_ori_kv, bool has_cmp_kv) { + const int64_t META_SIZE = 1024; + auto opts_i32 = cu_seq_q.options().dtype(at::kInt); + at::Tensor meta_t = at::empty({META_SIZE}, opts_i32.device(cu_seq_q.device())); + char* lq = const_cast(layout_q.c_str()); + char* lkv = const_cast(layout_kv.c_str()); + EXEC_NPU_CMD(aclnnSparseAttnSharedkvMetadata, + cu_seq_q, cu_seq_ori_kv, cu_seq_cmp_kv, seqused_q, seqused_kv, + num_heads_q, num_heads_kv, head_dim, + batch_size, max_seq_q, max_seq_kv, + ori_top_k, cmp_top_k, cmp_ratio, + ori_mask_mode, cmp_mask_mode, ori_win_left, ori_win_right, + lq, lkv, has_ori_kv, has_cmp_kv, + meta_t); + return meta_t; +} + +// quant_lightning_indexer_metadata (standalone entry) +at::Tensor quant_lightning_indexer_metadata_impl_npu( + const at::Tensor& aslq, const at::Tensor& aslk, + int64_t num_heads_q, int64_t num_heads_k, int64_t head_dim, + int64_t query_quant_mode, int64_t key_quant_mode, + int64_t batch_size, int64_t max_seq_q, int64_t max_seq_k, + std::string layout_query, std::string layout_key, + int64_t sparse_count, int64_t sparse_mode, + int64_t pre_token, int64_t next_token, int64_t cmp_ratio) { + const int64_t QLI_META_SIZE = 1024; + auto opts_i32 = aslq.options().dtype(at::kInt); + at::Tensor meta_t = at::empty({QLI_META_SIZE}, opts_i32.device(aslq.device())); + char* lq = const_cast(layout_query.c_str()); + char* lk = const_cast(layout_key.c_str()); + EXEC_NPU_CMD(aclnnQuantLightningIndexerMetadata, + aslq, aslk, num_heads_q, num_heads_k, head_dim, + query_quant_mode, key_quant_mode, + batch_size, max_seq_q, max_seq_k, + lq, lk, + sparse_count, sparse_mode, pre_token, next_token, cmp_ratio, + meta_t); + return meta_t; +} + +// lightning_indexer_quant_metadata (standalone entry) +at::Tensor lightning_indexer_quant_metadata_impl_npu( + const at::Tensor& aslq, const at::Tensor& aslk, + int64_t num_heads_q, int64_t num_heads_k, int64_t head_dim, + int64_t query_quant_mode, int64_t key_quant_mode, + int64_t batch_size, int64_t max_seq_q, int64_t max_seq_k, + std::string layout_query, std::string layout_key, + int64_t sparse_count, int64_t sparse_mode, + bool is_fd, + int64_t pre_token, int64_t next_token, int64_t cmp_ratio) { + const int64_t LIQ_META_SIZE = 1024; + auto opts_i32 = aslq.options().dtype(at::kInt); + at::Tensor meta_t = at::empty({LIQ_META_SIZE}, opts_i32.device(aslq.device())); + char* lq = const_cast(layout_query.c_str()); + char* lk = const_cast(layout_key.c_str()); + EXEC_NPU_CMD(aclnnLightningIndexerQuantMetadata, + aslq, aslk, num_heads_q, num_heads_k, head_dim, + query_quant_mode, key_quant_mode, + batch_size, max_seq_q, max_seq_k, + lq, lk, + sparse_count, sparse_mode, is_fd, pre_token, next_token, cmp_ratio, + meta_t); + return meta_t; +} + PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("select_unshared_kv", &select_unshared_kv_impl_npu, "select_unshared_kv"); m.def("cache_unshared_kv", &cache_unshared_kv_impl_npu, "cache_unshared_kv"); @@ -1481,4 +1895,12 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("multi_latent_attention", &multi_latent_attention_impl_npu, "multi_latent_attention"); m.def("rms_norm_dynamic_quant", &rms_norm_dynamic_quant_impl_npu, "rms_norm_dynamic_quant"); m.def("lightning_indexer_quant", &lightning_indexer_quant_impl_npu, "lightning_indexer_quant"); + m.def("compressor", &compressor_impl_npu, "compressor"); + m.def("sparse_attn_sharedkv", &sparse_attn_sharedkv_impl_npu, "sparse_attn_sharedkv"); + m.def("quant_lightning_indexer", &quant_lightning_indexer_impl_npu, "quant_lightning_indexer"); + m.def("dispatch_ffn_combine", &dispatch_ffn_combine_impl_npu, "dispatch_ffn_combine"); + m.def("dispatch_gmm_combine_decode", &dispatch_gmm_combine_decode_impl_npu, "dispatch_gmm_combine_decode"); + m.def("sparse_attn_sharedkv_metadata", &sparse_attn_sharedkv_metadata_impl_npu, "sparse_attn_sharedkv_metadata"); + m.def("quant_lightning_indexer_metadata", &quant_lightning_indexer_metadata_impl_npu, "quant_lightning_indexer_metadata"); + m.def("lightning_indexer_quant_metadata", &lightning_indexer_quant_metadata_impl_npu, "lightning_indexer_quant_metadata"); } diff --git a/test/python_test/compressor_gen.py b/test/python_test/compressor_gen.py deleted file mode 100644 index 052975d..0000000 --- a/test/python_test/compressor_gen.py +++ /dev/null @@ -1,151 +0,0 @@ -#!/usr/bin/env python3 -# -*- coding: utf-8 -*- -""" -compressor(AICORE MIX_AIC) fused-op golden + BSH data generation. - -Scenario (priority: non-OVERLAP, direct semantics): - {cmp_ratio=128, coff=1, head_dim=512} (one of the 3 legal combos of CheckScenarioConsistency) - B=1, S=128 -> Sr=ceil(S/cmp_ratio)=1, hidden_size=1024 (512-aligned, in [1024,10240]), - start_pos=0, block_size>=S -> no state history dependency (head/tail holders all 0, no r/w on kv/score_state). - -Fused formula (precisely locked from compressor_block_cube.h ComputeMm1 + block_vec.h Vec1/Vec2): - coffD = coff*head_dim = 512 - kv = x @ wkv^T # (S,hidden) @ (coffD,hidden)^T -> (S, coffD) - score = x @ wgate^T # (S, coffD) - score = score + ape # ape=(cmp_ratio, coffD), add positional encoding per row within group - # softmax within group (cmp_ratio rows) along "row/group" dim (each column independent); coff=1 => ReduceSize=cmp_ratio - p = softmax(score, axis=row-within-group) # (cmp_ratio, coffD) per group - cmp = sum_within_group(p * kv) # (Sr, coffD) weighted-sum compression within group - norm = RmsNorm(cmp, norm_weight(head_dim), eps) # per-row rmsnorm - out = HALF-rope (on last rope_head_dim=64 dims), first head_dim-64 dims unchanged - cmp_kv = out.reshape(B, Sr, head_dim) # BSH -""" -import os -import sys -import numpy as np - -DATA_DIR = "/tmp/compressor_data" - -# ------------------ case params (non-OVERLAP minimal scenario) ------------------ -B = 1 -S = 128 -CMP_RATIO = 128 -COFF = 1 -HEAD_DIM = 512 -COFF_D = COFF * HEAD_DIM # 512 -HIDDEN = 1024 # in [1024,10240] and 512-aligned -ROPE_HEAD_DIM = 64 -NORM_EPS = 1e-6 -BLOCK_SIZE = 128 # >= S, single block holds all -SR = (S + CMP_RATIO - 1) // CMP_RATIO # 1 - - -def rms_norm(x, weight, eps): - """x:(rows, head_dim) weight:(head_dim,) per-row rmsnorm""" - var = np.mean(x.astype(np.float32) ** 2, axis=-1, keepdims=True) - return (x / np.sqrt(var + eps)) * weight - - -def half_rope(x, cos, sin): - """ - HALF mode (rope.h RotaryPosEmb MODE==HALF): only acts on last rope_head_dim dims per row. - x:(rows, head_dim); cos/sin:(rows, rope_head_dim) - rotate_half(v) = [-v[half:], v[:half]] (half = rope_head_dim/2) - out_rope = v*cos + rotate_half(v)*sin - first head_dim-rope_head_dim dims unchanged. - """ - out = x.astype(np.float32).copy() - r = ROPE_HEAD_DIM - base = HEAD_DIM - r - v = out[:, base:base + r] - half = r // 2 - rot = np.concatenate([-v[:, half:], v[:, :half]], axis=-1) - out[:, base:base + r] = v * cos + rot * sin - return out - - -def golden(x, wkv, wgate, ape, norm_weight, rope_cos, rope_sin): - """ - x:(B,S,hidden) wkv/wgate:(coffD,hidden) ape:(cmp_ratio,coffD) - norm_weight:(head_dim) rope_cos/sin:(B,Sr,rope_head_dim) - return cmp_kv:(B,Sr,head_dim) float32 - """ - cmp_kv = np.zeros((B, SR, HEAD_DIM), dtype=np.float32) - for b in range(B): - kv = x[b] @ wkv.T # (S, coffD) - score = x[b] @ wgate.T # (S, coffD) - for g in range(SR): - r0 = g * CMP_RATIO - r1 = min(r0 + CMP_RATIO, S) - n = r1 - r0 - sc = score[r0:r1] + ape[:n] # (n, coffD) - kvg = kv[r0:r1] # (n, coffD) - # softmax within group (along row dim, each column independent) - m = np.max(sc, axis=0, keepdims=True) - e = np.exp(sc - m) - p = e / np.sum(e, axis=0, keepdims=True) # (n, coffD) - cmp = np.sum(p * kvg, axis=0, keepdims=True) # (1, coffD) - norm = rms_norm(cmp, norm_weight, NORM_EPS) # (1, head_dim) coffD==head_dim - out = half_rope(norm, rope_cos[b, g:g + 1], rope_sin[b, g:g + 1]) - cmp_kv[b, g] = out[0] - return cmp_kv - - -def main(): - dtype_str = sys.argv[1] if len(sys.argv) > 1 else "fp16" - os.makedirs(DATA_DIR, exist_ok=True) - rng = np.random.default_rng(2025) - - x = (rng.standard_normal((B, S, HIDDEN)) * 0.1).astype(np.float32) - wkv = (rng.standard_normal((COFF_D, HIDDEN)) * 0.05).astype(np.float32) - wgate = (rng.standard_normal((COFF_D, HIDDEN)) * 0.05).astype(np.float32) - ape = (rng.standard_normal((CMP_RATIO, COFF_D)) * 0.1).astype(np.float32) - norm_weight = (rng.standard_normal((HEAD_DIM,)) * 0.1 + 1.0).astype(np.float32) - rope_cos = (rng.standard_normal((B, SR, ROPE_HEAD_DIM)) * 0.1).astype(np.float32) - rope_sin = (rng.standard_normal((B, SR, ROPE_HEAD_DIM)) * 0.1).astype(np.float32) - - # state inputs (zero-initialized in-place; not used in minimal scenario) - # kv_state/score_state shape: (block_num, block_size, coffD) paged; block_num=1 in minimal scenario - block_num = 1 - kv_state = np.zeros((block_num, BLOCK_SIZE, COFF_D), dtype=np.float32) - score_state = np.zeros((block_num, BLOCK_SIZE, COFF_D), dtype=np.float32) - - # block_table: (batchSize, maxBlockNumPerBatch) int32, points to physical block index - # maxBlockNumPerBatch = ceil(S/BLOCK_SIZE) = 1; batch 0 uses block 0 - max_block = (S + BLOCK_SIZE - 1) // BLOCK_SIZE - kv_block_table = np.zeros((B, max_block), dtype=np.int32) - score_block_table = np.zeros((B, max_block), dtype=np.int32) - - gold = golden(x, wkv, wgate, ape, norm_weight, rope_cos, rope_sin) - - def cast(a): - if dtype_str == "fp16": - return a.astype(np.float16) - u32 = a.astype(np.float32).view(np.uint32) - return ((u32 + 0x8000) >> 16).astype(np.uint16) - - cast(x).tofile(os.path.join(DATA_DIR, "x.bin")) - cast(wkv).tofile(os.path.join(DATA_DIR, "wkv.bin")) - cast(wgate).tofile(os.path.join(DATA_DIR, "wgate.bin")) - ape.astype(np.float32).tofile(os.path.join(DATA_DIR, "ape.bin")) - kv_state.tofile(os.path.join(DATA_DIR, "kv_state.bin")) - score_state.tofile(os.path.join(DATA_DIR, "score_state.bin")) - kv_block_table.tofile(os.path.join(DATA_DIR, "kv_block_table.bin")) - score_block_table.tofile(os.path.join(DATA_DIR, "score_block_table.bin")) - cast(norm_weight).tofile(os.path.join(DATA_DIR, "norm_weight.bin")) - cast(rope_cos).tofile(os.path.join(DATA_DIR, "rope_cos.bin")) - cast(rope_sin).tofile(os.path.join(DATA_DIR, "rope_sin.bin")) - cast(gold).tofile(os.path.join(DATA_DIR, "golden_cmp_kv.bin")) - - with open(os.path.join(DATA_DIR, "meta.txt"), "w") as f: - f.write(f"dtype={dtype_str}\n") - f.write(f"B={B} S={S} SR={SR} HIDDEN={HIDDEN} HEAD_DIM={HEAD_DIM} COFF_D={COFF_D}\n") - f.write(f"CMP_RATIO={CMP_RATIO} COFF={COFF} ROPE_HEAD_DIM={ROPE_HEAD_DIM}\n") - f.write(f"BLOCK_SIZE={BLOCK_SIZE} block_num={block_num} NORM_EPS={NORM_EPS}\n") - - print(f"[gen] dtype={dtype_str} x{x.shape} wkv{wkv.shape} ape{ape.shape}") - print(f"[gen] cmp_kv(golden){gold.shape} -> {DATA_DIR}") - - -if __name__ == "__main__": - main() \ No newline at end of file diff --git a/test/python_test/custom_ops.py b/test/python_test/custom_ops.py index 74d1d9d..7be792e 100644 --- a/test/python_test/custom_ops.py +++ b/test/python_test/custom_ops.py @@ -529,3 +529,160 @@ def lightning_indexer_quant_npu(query, key, weights, query_scale, key_scale, query, key, weights, query_scale, key_scale, query_quant_mode, key_quant_mode, layout_query, layout_key, sparse_count, sparse_mode) + + +# compressor: single-step fusion (matmul + gate + softmax + compressKV + RmsNorm + RoPE). +# x:(B,S,HIDDEN), wkv/wgate:(COFF_D,HIDDEN), kv_state/score_state:(blockNum,blockSize,COFF_D) fp32 in-place, +# ape:(CMP_RATIO,COFF_D) fp32, norm_weight:(HEAD_DIM), rope_sin/cos:(B,SR,ROPE_HEAD_DIM), +# kv/score_block_table:(B,maxBlock) int32. +# Output: cmp_kv (B, SR, HEAD_DIM). enable_grad=False -> auxiliary outputs are placeholder. +def compressor_npu(x, wkv, wgate, kv_state, score_state, ape, norm_weight, + rope_sin, rope_cos, kv_block_table, score_block_table, + rope_head_dim=64, cmp_ratio=128, coff=1, + norm_eps=1e-6, rotary_mode=1, enable_grad=False): + return custom_ops_lib.compressor( + x, wkv, wgate, kv_state, score_state, ape, norm_weight, + rope_sin, rope_cos, kv_block_table, score_block_table, + rope_head_dim, cmp_ratio, coff, norm_eps, rotary_mode, enable_grad) + + +# sparse_attn_sharedkv: two-stage (metadata + main) flash-attention with sliding window + sink. +# query:(B,S1,N1,D), ori_kv:(blockNum,blockSize,KV_N,D) paged, ori_bt:(B,maxBlocks) int32, +# seqused_kv:(B,) int32, sinks:(N1,) fp32. +# Output: attn_out (B, S1, N1, D). +def sparse_attn_sharedkv_npu(query, ori_kv, ori_bt, seqused_kv, sinks, + n1=64, kv_n=1, d=512, s1=4, s2=16, + ori_mask_mode=4, cmp_mask_mode=3, + ori_win_left=127, ori_win_right=0, + softmax_scale=None, cmp_ratio=1, + layout_q="BSND", layout_kv="PA_ND"): + import math + if softmax_scale is None: + softmax_scale = 1.0 / math.sqrt(d) + return custom_ops_lib.sparse_attn_sharedkv( + query, ori_kv, ori_bt, seqused_kv, sinks, + n1, kv_n, d, s1, s2, + ori_mask_mode, cmp_mask_mode, ori_win_left, ori_win_right, + softmax_scale, cmp_ratio, layout_q, layout_kv) + + +# quant_lightning_indexer: two-stage (metadata + main) fp8 quantized indexer with PA. +# query:(B,qSeq,numHeadsQ,headDim) fp8_e4m3, key:(blockNum,blockSize,numHeadsK,headDim) fp8_e4m3, +# weights:(B,qSeq,numHeadsQ) fp32, q_scale:(B,qSeq,numHeadsQ) fp32, +# k_scale:(blockNum,blockSize,numHeadsK) fp32, aslq/aslk:(B,) int32, block_table:(B,maxBlocks) int32. +# Output: idx_out (B, qSeq, numHeadsK, sparseCount) int32. +def quant_lightning_indexer_npu(query, key, weights, q_scale, k_scale, + aslq, aslk, block_table, + num_heads_q=64, num_heads_k=1, head_dim=128, + query_quant_mode=0, key_quant_mode=0, + layout_query="BSND", layout_key="PA_BSND", + sparse_count=8, sparse_mode=3, + pre_token=9223372036854775807, + next_token=9223372036854775807, + cmp_ratio=1): + return custom_ops_lib.quant_lightning_indexer( + query, key, weights, q_scale, k_scale, aslq, aslk, block_table, + num_heads_q, num_heads_k, head_dim, + query_quant_mode, key_quant_mode, + layout_query, layout_key, + sparse_count, sparse_mode, pre_token, next_token, cmp_ratio) + + +# dispatch_ffn_combine (multi-card MoE dispatch+FFN+combine fusion, int8 path) +# weight1/weight2: list of tensors (per-expert NZ int8), scale1/scale2: list of tensors, +# probs:(M,topk) fp32, group: HCCL comm name string. +# out/expert_token_nums: pre-allocated output tensors (in-place write). +def dispatch_ffn_combine_npu(x, weight1, weight2, expert_idx, + scale1, scale2, probs, group, + max_output_size, out, expert_token_nums, + x_active_mask=None, swiglu_limit=0.0): + return custom_ops_lib.dispatch_ffn_combine( + x, weight1, weight2, expert_idx, scale1, scale2, probs, + group, max_output_size, out, expert_token_nums, + x_active_mask, swiglu_limit) + + +# dispatch_gmm_combine_decode (multi-card MoE dispatch+GMM+combine decode fusion) +# gmm1/gmm2 weight/scale: list of tensors. +# group_ep: HCCL comm name string. Returns (output, expert_token_nums). +def dispatch_gmm_combine_decode_npu(x, expert_ids, gmm1_permuted_weight, + gmm1_permuted_weight_scale, + gmm2_weight, gmm2_weight_scale, + expert_scales, group_ep, + ep_rank_size, ep_rank_id, + moe_expert_num, shared_expert_num=1, + shared_expert_rank_num=0, quant_mode=0, + global_bs=0, expert_smooth_scales=None, + x_active_mask=None): + return custom_ops_lib.dispatch_gmm_combine_decode( + x, expert_ids, gmm1_permuted_weight, gmm1_permuted_weight_scale, + gmm2_weight, gmm2_weight_scale, expert_scales, + expert_smooth_scales, x_active_mask, + group_ep, ep_rank_size, ep_rank_id, moe_expert_num, + shared_expert_num, shared_expert_rank_num, quant_mode, global_bs) + + +# ==================== standalone metadata ops ==================== + +# sparse_attn_sharedkv_metadata: AICPU metadata op for sparse_attn_sharedkv. +# Computes tiling/scheduling metadata for the two-stage sparse attention. +# Output: int32 tensor of size 1024 (opaque scheduling metadata). +def sparse_attn_sharedkv_metadata_npu(cu_seq_q, cu_seq_ori_kv, cu_seq_cmp_kv, + seqused_q, seqused_kv, + num_heads_q=64, num_heads_kv=1, head_dim=512, + batch_size=1, max_seq_q=4, max_seq_kv=16, + ori_top_k=0, cmp_top_k=0, cmp_ratio=1, + ori_mask_mode=4, cmp_mask_mode=3, + ori_win_left=127, ori_win_right=0, + layout_q="BSND", layout_kv="PA_ND", + has_ori_kv=True, has_cmp_kv=False): + return custom_ops_lib.sparse_attn_sharedkv_metadata( + cu_seq_q, cu_seq_ori_kv, cu_seq_cmp_kv, seqused_q, seqused_kv, + num_heads_q, num_heads_kv, head_dim, + batch_size, max_seq_q, max_seq_kv, + ori_top_k, cmp_top_k, cmp_ratio, + ori_mask_mode, cmp_mask_mode, ori_win_left, ori_win_right, + layout_q, layout_kv, has_ori_kv, has_cmp_kv) + + +# quant_lightning_indexer_metadata: AICPU metadata op for quant_lightning_indexer. +# Computes tiling/scheduling metadata for two-stage fp8 quantized indexer. +# Output: int32 tensor of size 1024 (opaque scheduling metadata). +def quant_lightning_indexer_metadata_npu(aslq, aslk, + num_heads_q=64, num_heads_k=1, head_dim=128, + query_quant_mode=0, key_quant_mode=0, + batch_size=1, max_seq_q=1, max_seq_k=128, + layout_query="BSND", layout_key="PA_BSND", + sparse_count=8, sparse_mode=3, + pre_token=9223372036854775807, + next_token=9223372036854775807, + cmp_ratio=1): + return custom_ops_lib.quant_lightning_indexer_metadata( + aslq, aslk, + num_heads_q, num_heads_k, head_dim, + query_quant_mode, key_quant_mode, + batch_size, max_seq_q, max_seq_k, + layout_query, layout_key, + sparse_count, sparse_mode, pre_token, next_token, cmp_ratio) + + +# lightning_indexer_quant_metadata: AICPU metadata op for lightning_indexer_quant. +# Similar to quant_lightning_indexer_metadata but with additional isFd parameter. +# Output: int32 tensor of size 1024 (opaque scheduling metadata). +def lightning_indexer_quant_metadata_npu(aslq, aslk, + num_heads_q=64, num_heads_k=1, head_dim=128, + query_quant_mode=0, key_quant_mode=0, + batch_size=1, max_seq_q=1, max_seq_k=128, + layout_query="BSND", layout_key="BSND", + sparse_count=128, sparse_mode=0, + is_fd=False, + pre_token=9223372036854775807, + next_token=9223372036854775807, + cmp_ratio=1): + return custom_ops_lib.lightning_indexer_quant_metadata( + aslq, aslk, + num_heads_q, num_heads_k, head_dim, + query_quant_mode, key_quant_mode, + batch_size, max_seq_q, max_seq_k, + layout_query, layout_key, + sparse_count, sparse_mode, is_fd, pre_token, next_token, cmp_ratio) diff --git a/test/python_test/pytorch_npu_helper.hpp b/test/python_test/pytorch_npu_helper.hpp index 9974475..6925b12 100644 --- a/test/python_test/pytorch_npu_helper.hpp +++ b/test/python_test/pytorch_npu_helper.hpp @@ -357,8 +357,16 @@ limitations under the License. return nullptr; } at::ScalarType scalar_data_type = at_tensor.scalar_type(); - aclDataType acl_data_type = - kATenScalarTypeToAclDataTypeTable[static_cast(scalar_data_type)]; + aclDataType acl_data_type; + if (scalar_data_type == at::ScalarType::Float8_e5m2) { + acl_data_type = static_cast(35); + } else if (scalar_data_type == at::ScalarType::Float8_e4m3fn) { + acl_data_type = static_cast(36); + } else { + auto idx = static_cast(scalar_data_type); + acl_data_type = (idx >= 0 && idx <= static_cast(at::ScalarType::NumOptions)) + ? kATenScalarTypeToAclDataTypeTable[idx] : ACL_DT_UNDEFINED; + } TORCH_CHECK( acl_data_type != ACL_DT_UNDEFINED, std::string(c10::toString(scalar_data_type)) + " has not been supported") @@ -516,7 +524,17 @@ limitations under the License. } inline aclDataType ConvertType(const at::ScalarType scalarType) { - return kATenScalarTypeToAclDataTypeTable[static_cast(scalarType)]; + // Float8 types: enum values beyond old table coverage + if (scalarType == at::ScalarType::Float8_e5m2) { + return static_cast(35); // ACL_FLOAT8_E5M2 + } else if (scalarType == at::ScalarType::Float8_e4m3fn) { + return static_cast(36); // ACL_FLOAT8_E4M3FN + } + auto idx = static_cast(scalarType); + if (idx < 0 || idx > static_cast(at::ScalarType::NumOptions)) { + return ACL_DT_UNDEFINED; + } + return kATenScalarTypeToAclDataTypeTable[idx]; } template diff --git a/test/python_test/quant_lightning_indexer_gen.py b/test/python_test/quant_lightning_indexer_gen.py deleted file mode 100644 index b780703..0000000 --- a/test/python_test/quant_lightning_indexer_gen.py +++ /dev/null @@ -1,131 +0,0 @@ -# -*- coding: utf-8 -*- -# Generate quant_lightning_indexer(PA_BSND) inputs + CPU golden(sparse_indices) written to .bin. -# Kernel hard constraint: layout_key only supports PA_BSND (paged KV). -# golden.forward uses unpaged key_bnsd; NPU uses paged key[block_num,block_size,kH,hd]+block_table. -import sys, types, os, math -for name in ("test", "custom_ops"): - sys.modules[name] = types.ModuleType(name) - -GOLDEN_DIR = "/export/home/weinan5/zhaoxingcheng/xllm_ops/xllm_ops/attention/quant_lightning_indexer/tests/pytest" -sys.path.insert(0, GOLDEN_DIR) - -import numpy as np -import torch -import quant_lightning_indexer_golden as G - -OUT = "/tmp/quant_lightning_indexer_data" -os.makedirs(OUT, exist_ok=True) -np.random.seed(2026); torch.manual_seed(2026) - -# ---- Case (PA_BSND): B=1, q_seq=4, k_seq=128(=1 block), block_size=128 ---- -batch_size = 1 -q_seq = 4 -block_size = 128 -k_seq = 128 -q_head_num = 64 -k_head_num = 1 -head_dim = 128 -block_num = 1 -qk_dtype = torch.int8 -dequant_dtype = torch.float16 -actual_seq_dtype = torch.int32 -cmp_ratio = 1 -act_seq_q = [q_seq] -act_seq_k = [k_seq * cmp_ratio] -query_quant_mode = 0 -key_quant_mode = 0 -layout_query = "BSND" -layout_key = "PA_BSND" -sparse_count = 8 -sparse_mode = 3 -q_t_size = q_seq; k_t_size = k_seq - -qr = [-100, 100]; kr = [-100, 100]; wr = [-25, 25] -qsr = [0, 255]; ksr = [0, 65504] - -indexer_op = G.GeneralizedQLI(batch_size, q_seq, k_seq, q_t_size, k_t_size, - q_head_num, k_head_num, head_dim, block_size, block_num, - qk_dtype, dequant_dtype, actual_seq_dtype, act_seq_q, act_seq_k, - query_quant_mode, key_quant_mode, layout_query, layout_key, - sparse_count, sparse_mode, cmp_ratio) - -# query / weights / q_scale (BSND) -query = torch.tensor(np.random.uniform(qr[0], qr[1], (batch_size, q_seq, q_head_num, head_dim))).to(qk_dtype) -weights = torch.tensor(np.random.uniform(wr[0], wr[1], (batch_size, q_seq, q_head_num))).to(dequant_dtype) -q_scale = torch.tensor(np.random.uniform(qsr[0], qsr[1], (batch_size, q_seq, q_head_num))).to(dequant_dtype) - -# key_bnsd (for golden) : [B, kH, k_max_s2, hd] -k_max_s2 = math.floor(max(act_seq_k) / cmp_ratio) -key_bnsd = torch.tensor(np.random.uniform(kr[0], kr[1], (batch_size, k_head_num, k_max_s2, head_dim))).to(qk_dtype) -k_scale_bns = torch.tensor(np.random.uniform(ksr[0], ksr[1], (batch_size, k_head_num, k_max_s2))).to(dequant_dtype) - -aslq = torch.tensor(act_seq_q).to(actual_seq_dtype) -aslk = torch.tensor(act_seq_k).to(actual_seq_dtype) - -# ---- Build block_table + paged key / k_scale (for NPU) ---- -k_max_block_num_per_batch = math.ceil(k_max_s2 / block_size) -key_block_num_per_batch = [] -key_block_num_sum = 0 -for cur_act_k in act_seq_k: - cur_cmp = math.floor(cur_act_k / cmp_ratio) - n = math.ceil(cur_cmp / block_size) - key_block_num_per_batch.append(n); key_block_num_sum += n -assert block_num >= key_block_num_sum, "block_num too small" - -block_id_list = np.arange(block_num).astype(np.int32) # no permutation, keep deterministic -block_table = np.full((batch_size, k_max_block_num_per_batch), -1, dtype=np.int32) -cur = 0 -for bi, thr in enumerate(key_block_num_per_batch): - for ib in range(thr): - block_table[bi][ib] = block_id_list[cur]; cur += 1 - -# paged key: [block_num, block_size, kH, hd] -key_expand = torch.zeros((batch_size, k_head_num, k_max_block_num_per_batch * block_size, head_dim), dtype=qk_dtype) -key_expand[:, :, :k_max_s2, :] = key_bnsd -key_pa = torch.zeros((block_num, block_size, k_head_num, head_dim), dtype=qk_dtype) -for ib_ in range(batch_size): - for i_block, cbid in enumerate(block_table[ib_]): - if cbid == -1: continue - sp = i_block * block_size - for i_n in range(k_head_num): - key_pa[cbid, :, i_n, :] = key_expand[ib_, i_n, sp:sp + block_size, :] - -ks_expand = torch.zeros((batch_size, k_head_num,k_max_block_num_per_batch * block_size), dtype=dequant_dtype) -ks_expand[:, :, :k_max_s2] = k_scale_bns -kscale_pa = torch.zeros((block_num, block_size, k_head_num), dtype=dequant_dtype) -for ib_ in range(batch_size): - for i_block, cbid in enumerate(block_table[ib_]): - if cbid == -1: continue - sp = i_block * block_size - for i_n in range(k_head_num): - kscale_pa[cbid, :, i_n] = ks_expand[ib_, i_n, sp:sp + block_size] - -def save(name, t, np_dtype): - arr = np.ascontiguousarray(t.cpu().numpy().astype(np_dtype)) - arr.tofile(os.path.join(OUT, name)) - print(f"{name}: shape={arr.shape} dtype={arr.dtype} bytes={arr.nbytes}") - -save("query.bin", query, np.int8) -save("key.bin", key_pa, np.int8) # paged key -save("weights.bin", weights, np.float16) -save("q_scale.bin", q_scale, np.float16) -save("k_scale.bin", kscale_pa, np.float16) # paged k_scale -save("aslq.bin", aslq, np.int32) -save("aslk.bin", aslk, np.int32) -save("block_table.bin", torch.from_numpy(block_table), np.int32) - -# CPU golden (uses unpaged key_bnsd / k_scale_bns) -y, y_value = indexer_op.forward(query, key_bnsd, weights, q_scale, k_scale_bns, aslq, aslk, torch.from_numpy(block_table)) -save("golden_indices.bin", y, np.int32) -print("golden y shape:", tuple(y.shape)) -for s in range(q_seq): - print(f"golden y[0,{s},0,:]=", y[0, s, 0, :].tolist()) - -with open(os.path.join(OUT, "meta.txt"), "w") as f: - f.write(f"batch_size={batch_size}\nq_seq={q_seq}\nk_seq={k_seq}\n") - f.write(f"q_head_num={q_head_num}\nk_head_num={k_head_num}\nhead_dim={head_dim}\n") - f.write(f"block_size={block_size}\nblock_num={block_num}\n") - f.write(f"k_max_block_num_per_batch={k_max_block_num_per_batch}\n") - f.write(f"sparse_count={sparse_count}\nsparse_mode={sparse_mode}\ncmp_ratio={cmp_ratio}\n") - f.write(f"query_quant_mode={query_quant_mode}\nkey_quant_mode={key_quant_mode}\n") -print("meta written. DONE.") \ No newline at end of file diff --git a/test/python_test/sparse_attn_sharedkv_gen.py b/test/python_test/sparse_attn_sharedkv_gen.py deleted file mode 100644 index e81a150..0000000 --- a/test/python_test/sparse_attn_sharedkv_gen.py +++ /dev/null @@ -1,144 +0,0 @@ -#!/usr/bin/env python3 -# -*- coding: utf-8 -*- -""" -sparse_attn_sharedkv SWA (Sliding Window Attention + attention sink) single-path golden + PA_ND data generation. - -Kernel semantics (precisely locked by swa_kernel.h CalcParams / tiling.cpp CheckFeature): - For each batch b, each query absolute position s1, each q head n1 (total N1=64): - diag = actKvS2 - actS1 + s1 (causal alignment) - maskRight = diag + oriWinRight (closed-interval upper bound, default 0) - maskLeft = max(diag - oriWinLeft, 0) (closed-interval lower bound, default 127) - j in [maskLeft, maskRight] - logit[j] = (Q[s1,n1,:] . K[j,:]) * softmax_scale - m = max(max_j logit[j], sinks[n1]) - denom = sum_j exp(logit[j]-m) + exp(sinks[n1]-m) (sink as an extra softmax term) - out[s1,n1,:] = sum_j (exp(logit[j]-m)/denom) * V[j,:] (sink value=0, no contribution to numerator) - -Hard constraints: N1(qHead)=64; D=512; KV_N(kvHead)=1; layout_q=BSND; layout_kv=PA_ND; - ori_mask_mode=4; cmp_mask_mode=3; ori_win_left=127; ori_win_right=0; qType==oriKvType in {fp16,bf16}. - -PA_ND: golden uses unpaged kv_bnsd as reference; NPU uses paged kv_pa[block_num,block_size,kvH=1,d] + block_table, mapped from the same source by slicing. -""" -import os -import sys -import numpy as np - -DATA_DIR = "/tmp/sparse_attn_sharedkv_data" - -# ------------------ Case parameters (SWA single-path, minimal scale) ------------------ -B = 1 # batch -S1 = 4 # query sequence length (per batch) -S2 = 16 # kv sequence length (per batch) -N1 = 64 # q head count (hard constraint =64) -KV_N = 1 # kv head count (hard constraint =1) -D = 512 # head dim (hard constraint =512) -BLOCK_SIZE = 16 # tokens per PA_ND block -ORI_WIN_LEFT = 127 -ORI_WIN_RIGHT = 0 -SOFTMAX_SCALE = 1.0 / np.sqrt(D) - - -def golden_swa(q, kv, sinks, actS1, actS2, win_left, win_right, scale): - """ - q: [B, S1, N1, D] float32 - kv: [B, S2, KV_N, D] float32 (unpaged, KV_N=1 broadcast to N1) - sinks:[N1] float32 - return: out [B, S1, N1, D] float32, lse [B, S1, N1] float32 - """ - out = np.zeros((B, S1, N1, D), dtype=np.float32) - lse = np.zeros((B, S1, N1), dtype=np.float32) - for b in range(B): - for s1 in range(actS1): - diag = actS2 - actS1 + s1 - mask_right = diag + win_right - mask_left = max(diag - win_left, 0) - if mask_right < mask_left: - # no valid kv, sink only - for n1 in range(N1): - out[b, s1, n1, :] = 0.0 - lse[b, s1, n1] = sinks[n1] # log(exp(sink-sink))+sink = sink - continue - j_range = np.arange(mask_left, mask_right + 1) # closed interval - k = kv[b, j_range, 0, :] # [J, D] (kvHead=0 broadcast) - for n1 in range(N1): - logit = (q[b, s1, n1, :] @ k.T) * scale # [J] - sink = sinks[n1] - m = max(float(logit.max()), float(sink)) - exp_l = np.exp(logit - m) # [J] - exp_s = np.exp(sink - m) - denom = float(exp_l.sum()) + float(exp_s) - p = exp_l / denom # [J] - out[b, s1, n1, :] = p @ k # sink value=0 - lse[b, s1, n1] = np.log(denom) + m - return out, lse - - -def to_paged(kv, block_size): - """ - kv: [B, S2, KV_N, D] -> paged kv[block_num, block_size, KV_N, D] + block_table[B, max_blocks] - Same-source slice mapping: token (b, t) -> global block = b*blocks_per_b + t//block_size - """ - blocks_per_b = (S2 + block_size - 1) // block_size - max_blocks = blocks_per_b - block_num = B * blocks_per_b - kv_pa = np.zeros((block_num, block_size, KV_N, D), dtype=kv.dtype) - block_table = np.zeros((B, max_blocks), dtype=np.int32) - for b in range(B): - for blk in range(blocks_per_b): - gblk = b * blocks_per_b + blk - block_table[b, blk] = gblk - for tok in range(block_size): - t = blk * block_size + tok - if t < S2: - kv_pa[gblk, tok, :, :] = kv[b, t, :, :] - return kv_pa, block_table - - -def main(): - dtype_str = sys.argv[1] if len(sys.argv) > 1 else "fp16" - np_dtype = np.float16 if dtype_str == "fp16" else np.uint16 # bf16 stored as uint16 bits - os.makedirs(DATA_DIR, exist_ok=True) - rng = np.random.default_rng(1234) - - # ---- generate data (float32 reference) ---- - q = rng.standard_normal((B, S1, N1, D)).astype(np.float32) * 0.1 - kv = rng.standard_normal((B, S2, KV_N, D)).astype(np.float32) * 0.1 - sinks = rng.standard_normal((N1,)).astype(np.float32) * 0.1 - seqused_kv = np.array([S2] * B, dtype=np.int32) - - # ---- golden (unpaged) ---- - out, lse = golden_swa(q, kv, sinks, S1, S2, ORI_WIN_LEFT, ORI_WIN_RIGHT, SOFTMAX_SCALE) - - # ---- paged (NPU input) ---- - kv_pa_f32, block_table = to_paged(kv, BLOCK_SIZE) - - # ---- dtype conversion and write to disk ---- - def cast(x): - if dtype_str == "fp16": - return x.astype(np.float16) - # bf16: take the high 16 bits of float32 - u32 = x.astype(np.float32).view(np.uint32) - return ((u32 + 0x8000) >> 16).astype(np.uint16) - - cast(q).tofile(os.path.join(DATA_DIR, "q.bin")) - cast(kv_pa_f32).tofile(os.path.join(DATA_DIR, "kv_pa.bin")) - sinks.astype(np.float32).tofile(os.path.join(DATA_DIR, "sinks.bin")) - block_table.tofile(os.path.join(DATA_DIR, "block_table.bin")) - seqused_kv.tofile(os.path.join(DATA_DIR, "seqused_kv.bin")) - cast(out).tofile(os.path.join(DATA_DIR, "golden_out.bin")) - lse.astype(np.float32).tofile(os.path.join(DATA_DIR, "golden_lse.bin")) - - # ---- meta info ---- - with open(os.path.join(DATA_DIR, "meta.txt"), "w") as f: - f.write(f"dtype={dtype_str}\n") - f.write(f"B={B} S1={S1} S2={S2} N1={N1} KV_N={KV_N} D={D}\n") - f.write(f"BLOCK_SIZE={BLOCK_SIZE} block_num={kv_pa_f32.shape[0]} max_blocks={block_table.shape[1]}\n") - f.write(f"softmax_scale={SOFTMAX_SCALE}\n") - f.write(f"ori_win_left={ORI_WIN_LEFT} ori_win_right={ORI_WIN_RIGHT}\n") - - print(f"[gen] dtype={dtype_str} q{q.shape} kv_pa{kv_pa_f32.shape} block_table{block_table.shape}") - print(f"[gen] out{out.shape} lse{lse.shape} -> {DATA_DIR}") - - -if __name__ == "__main__": - main() \ No newline at end of file diff --git a/test/python_test/test_compressor.cpp b/test/python_test/test_compressor.cpp deleted file mode 100644 index 46f2e5d..0000000 --- a/test/python_test/test_compressor.cpp +++ /dev/null @@ -1,193 +0,0 @@ -#include -#include -#include -#include -#include -#include -#include -#include -#include "acl/acl.h" -#include "aclnn_compressor.h" - -// ---- Case (non-OVERLAP {128,1,512}, consistent with compressor_gen.py) ---- -static const int64_t B = 1; -static const int64_t S = 128; -static const int64_t SR = 1; // ceil(S/CMP_RATIO) -static const int64_t HIDDEN = 1024; -static const int64_t HEAD_DIM = 512; -static const int64_t COFF_D = 512; // coff*head_dim -static const int64_t CMP_RATIO = 128; -static const int64_t COFF = 1; -static const int64_t ROPE_HEAD_DIM = 64; -static const int64_t BLOCK_SIZE = 128; -static const int64_t BLOCK_NUM = 1; -static const double NORM_EPS = 1e-6; -static const int64_t ROTARY_MODE = 1; // HALF -static const bool ENABLE_GRAD = false; -static const char* DATA_DIR = "/tmp/compressor_data/"; - -static aclDataType g_dtype = ACL_FLOAT16; - -static std::vector ReadBin(const std::string &name) { - std::ifstream f(std::string(DATA_DIR) + name, std::ios::binary | std::ios::ate); - if (!f) { printf("open %s FAILED\n", name.c_str()); return {}; } - std::streamsize sz = f.tellg(); f.seekg(0, std::ios::beg); - std::vector buf(sz); f.read(buf.data(), sz); return buf; -} - -std::tuple CreateTensor(size_t size, std::vector shape, - aclDataType dType, const void* hostData = nullptr) { - std::vector strides(shape.size(), 1); - for (int i = (int)shape.size() - 2; i >= 0; --i) strides[i] = strides[i + 1] * shape[i + 1]; - void* dev = nullptr; - auto ret = aclrtMalloc(&dev, size, ACL_MEM_MALLOC_HUGE_FIRST); - if (ret != ACL_SUCCESS) { printf("aclrtMalloc %d\n", ret); return {nullptr, nullptr}; } - aclTensor* t = aclCreateTensor(shape.data(), shape.size(), dType, strides.data(), 0, - aclFormat::ACL_FORMAT_ND, shape.data(), shape.size(), dev); - if (t == nullptr) { aclrtFree(dev); return {nullptr, nullptr}; } - if (hostData) aclrtMemcpy(dev, size, hostData, size, ACL_MEMCPY_HOST_TO_DEVICE); - return {t, dev}; -} - -// 16-bit(fp16/bf16) bit pattern -> float32 -static float ToF32(uint16_t v, bool isBf16) { - uint32_t u; - if (isBf16) { - u = static_cast(v) << 16; - } else { - uint32_t sign = (v >> 15) & 0x1; - uint32_t exp = (v >> 10) & 0x1f; - uint32_t man = v & 0x3ff; - if (exp == 0) { - if (man == 0) { u = sign << 31; } - else { - exp = 127 - 15 + 1; - while ((man & 0x400) == 0) { man <<= 1; exp--; } - man &= 0x3ff; - u = (sign << 31) | (exp << 23) | (man << 13); - } - } else if (exp == 0x1f) { - u = (sign << 31) | (0xff << 23) | (man << 13); - } else { - u = (sign << 31) | ((exp - 15 + 127) << 23) | (man << 13); - } - } - float f; memcpy(&f, &u, 4); return f; -} - -int main(int argc, char** argv) { - bool isBf16 = (argc > 1 && std::string(argv[1]) == "bf16"); - g_dtype = isBf16 ? ACL_BF16 : ACL_FLOAT16; - size_t elemSz = 2; - - int32_t deviceId = 0; aclrtStream stream; - aclError ret = aclInit(nullptr); - if (ret != ACL_SUCCESS) { printf("aclInit %d\n", ret); return -1; } - ret = aclrtSetDevice(deviceId); - if (ret != ACL_SUCCESS) { printf("aclrtSetDevice %d\n", ret); return -1; } - ret = aclrtCreateStream(&stream); - if (ret != ACL_SUCCESS) { printf("aclrtCreateStream %d\n", ret); return -1; } - - auto xBuf = ReadBin("x.bin"); - auto wkvBuf = ReadBin("wkv.bin"); - auto wgateBuf = ReadBin("wgate.bin"); - auto apeBuf = ReadBin("ape.bin"); - auto kvStateBuf = ReadBin("kv_state.bin"); - auto scoreStateBuf = ReadBin("score_state.bin"); - auto kvBlockTableBuf = ReadBin("kv_block_table.bin"); - auto scoreBlockTableBuf = ReadBin("score_block_table.bin"); - auto normWBuf = ReadBin("norm_weight.bin"); - auto ropeCosBuf = ReadBin("rope_cos.bin"); - auto ropeSinBuf = ReadBin("rope_sin.bin"); - auto goldenBuf = ReadBin("golden_cmp_kv.bin"); - if (xBuf.empty() || wkvBuf.empty() || goldenBuf.empty() || apeBuf.empty()) { - printf("read bin FAILED\n"); return -1; - } - - // ---- input tensors ---- - aclTensor *x, *wkv, *wgate, *kvState, *scoreState, *ape, *normW, *ropeSin, *ropeCos; - void *xDev, *wkvDev, *wgateDev, *kvStateDev, *scoreStateDev, *apeDev, *normWDev, *ropeSinDev, *ropeCosDev; - std::tie(x, xDev) = CreateTensor(xBuf.size(), {B, S, HIDDEN}, g_dtype, xBuf.data()); - std::tie(wkv, wkvDev) = CreateTensor(wkvBuf.size(), {COFF_D, HIDDEN}, g_dtype, wkvBuf.data()); - std::tie(wgate, wgateDev) = CreateTensor(wgateBuf.size(), {COFF_D, HIDDEN}, g_dtype, wgateBuf.data()); - // kv_state/score_state: float32, in-place, PageAttention: [blockNum, blockSize, coffD] - std::tie(kvState, kvStateDev) = CreateTensor(kvStateBuf.size(), {BLOCK_NUM, BLOCK_SIZE, COFF_D}, ACL_FLOAT, kvStateBuf.data()); - std::tie(scoreState, scoreStateDev) = CreateTensor(scoreStateBuf.size(), {BLOCK_NUM, BLOCK_SIZE, COFF_D}, ACL_FLOAT, scoreStateBuf.data()); - std::tie(ape, apeDev) = CreateTensor(apeBuf.size(), {CMP_RATIO, COFF_D}, ACL_FLOAT, apeBuf.data()); - std::tie(normW, normWDev) = CreateTensor(normWBuf.size(), {HEAD_DIM}, g_dtype, normWBuf.data()); - std::tie(ropeSin, ropeSinDev) = CreateTensor(ropeSinBuf.size(), {B, SR, ROPE_HEAD_DIM}, g_dtype, ropeSinBuf.data()); - std::tie(ropeCos, ropeCosDev) = CreateTensor(ropeCosBuf.size(), {B, SR, ROPE_HEAD_DIM}, g_dtype, ropeCosBuf.data()); - - // kv/score block_table: int32, shape (batchSize, maxBlockNumPerBatch) - const int64_t MAX_BLOCK = (S + BLOCK_SIZE - 1) / BLOCK_SIZE; // 1 - aclTensor *kvBlockTable, *scoreBlockTable; - void *kvBlockTableDev, *scoreBlockTableDev; - std::tie(kvBlockTable, kvBlockTableDev) = CreateTensor(kvBlockTableBuf.size(), {B, MAX_BLOCK}, ACL_INT32, kvBlockTableBuf.data()); - std::tie(scoreBlockTable, scoreBlockTableDev) = CreateTensor(scoreBlockTableBuf.size(), {B, MAX_BLOCK}, ACL_INT32, scoreBlockTableBuf.data()); - - // ---- output tensors ---- - // cmp_kv: BSH -> (B, SR, HEAD_DIM); when enable_grad=false the 4 auxiliary outputs have shape=0, use {1} as placeholder - aclTensor *cmpKv, *wkvProj, *softmaxRes, *normX, *normRstd; - void *cmpKvDev, *wkvProjDev, *softmaxResDev, *normXDev, *normRstdDev; - int64_t cmpKvElems = B * SR * HEAD_DIM; - std::tie(cmpKv, cmpKvDev) = CreateTensor(cmpKvElems * elemSz, {B, SR, HEAD_DIM}, g_dtype); - std::tie(wkvProj, wkvProjDev) = CreateTensor(elemSz, {1}, g_dtype); - std::tie(softmaxRes, softmaxResDev) = CreateTensor(elemSz, {1}, g_dtype); - std::tie(normX, normXDev) = CreateTensor(elemSz, {1}, g_dtype); - std::tie(normRstd, normRstdDev) = CreateTensor(elemSz, {1}, g_dtype); - - // ---- aclnnCompressor (25 params: kv_state/score_state are in-place Ref, passed only once) ---- - aclOpExecutor* executor = nullptr; uint64_t wsSize = 0; void* ws = nullptr; - ret = aclnnCompressorGetWorkspaceSize( - x, wkv, wgate, kvState, scoreState, ape, normW, ropeSin, ropeCos, - kvBlockTable, scoreBlockTable, nullptr, nullptr, nullptr, // kvBT, scoreBT, cuSeqlens, seqused, startPos - ROPE_HEAD_DIM, CMP_RATIO, COFF, NORM_EPS, ROTARY_MODE, ENABLE_GRAD, - cmpKv, wkvProj, softmaxRes, normX, normRstd, - &wsSize, &executor); - if (ret != ACL_SUCCESS) { printf("Compressor GetWorkspaceSize FAILED ret=%d\n", ret); return -1; } - printf("Compressor GetWorkspaceSize OK, ws=%lu\n", wsSize); - if (wsSize > 0) { ret = aclrtMalloc(&ws, wsSize, ACL_MEM_MALLOC_HUGE_FIRST); - if (ret != ACL_SUCCESS) { printf("malloc ws %d\n", ret); return -1; } } - ret = aclnnCompressor(ws, wsSize, executor, stream); - if (ret != ACL_SUCCESS) { printf("Compressor Execute FAILED ret=%d\n", ret); return -1; } - ret = aclrtSynchronizeStream(stream); - if (ret != ACL_SUCCESS) { printf("Compressor Sync %d\n", ret); return -1; } - printf("Compressor Execute OK\n"); - - // ---- read back cmp_kv and compare ---- - std::vector outHost(cmpKvElems); - ret = aclrtMemcpy(outHost.data(), cmpKvElems * elemSz, cmpKvDev, cmpKvElems * elemSz, ACL_MEMCPY_DEVICE_TO_HOST); - if (ret != ACL_SUCCESS) { printf("memcpy out %d\n", ret); return -1; } - - const uint16_t* golden = reinterpret_cast(goldenBuf.data()); - double maxAbs = 0.0, maxRel = 0.0; - int failCnt = 0; - double atol = 2e-2, rtol = 2e-2; - for (int64_t i = 0; i < cmpKvElems; ++i) { - float g = ToF32(golden[i], isBf16); - float o = ToF32(outHost[i], isBf16); - double ad = std::fabs(g - o); - double rd = ad / (std::fabs(g) + 1e-6); - if (ad > maxAbs) maxAbs = ad; - if (rd > maxRel) maxRel = rd; - if (ad > atol && rd > rtol) failCnt++; - } - printf("==== cmpKvElems=%ld maxAbs=%.5f maxRel=%.5f fail=%d ====\n", (long)cmpKvElems, maxAbs, maxRel, failCnt); - printf(failCnt == 0 ? "RESULT: PASS\n" : "RESULT: FAIL\n"); - - aclDestroyTensor(x); aclDestroyTensor(wkv); aclDestroyTensor(wgate); - aclDestroyTensor(kvState); aclDestroyTensor(scoreState); aclDestroyTensor(ape); - aclDestroyTensor(normW); aclDestroyTensor(ropeSin); aclDestroyTensor(ropeCos); - aclDestroyTensor(cmpKv); aclDestroyTensor(wkvProj); aclDestroyTensor(softmaxRes); - aclDestroyTensor(normX); aclDestroyTensor(normRstd); - aclDestroyTensor(kvBlockTable); aclDestroyTensor(scoreBlockTable); - aclrtFree(xDev); aclrtFree(wkvDev); aclrtFree(wgateDev); - aclrtFree(kvStateDev); aclrtFree(scoreStateDev); aclrtFree(apeDev); - aclrtFree(normWDev); aclrtFree(ropeSinDev); aclrtFree(ropeCosDev); - aclrtFree(cmpKvDev); aclrtFree(wkvProjDev); aclrtFree(softmaxResDev); - aclrtFree(normXDev); aclrtFree(normRstdDev); - aclrtFree(kvBlockTableDev); aclrtFree(scoreBlockTableDev); - if (ws) aclrtFree(ws); - aclrtDestroyStream(stream); aclrtResetDevice(deviceId); aclFinalize(); - return 0; -} \ No newline at end of file diff --git a/test/python_test/test_compressor.py b/test/python_test/test_compressor.py new file mode 100644 index 0000000..53a84f6 --- /dev/null +++ b/test/python_test/test_compressor.py @@ -0,0 +1,126 @@ +import pytest +import numpy as np +import torch + +torch_npu = pytest.importorskip("torch_npu") +custom_ops = pytest.importorskip("custom_ops") + + +# ---- Golden (CPU fp32 reference, inlined from compressor_gen.py) ---- +def rms_norm(x, weight, eps): + """x:(rows, head_dim) weight:(head_dim,) per-row rmsnorm""" + var = np.mean(x.astype(np.float32) ** 2, axis=-1, keepdims=True) + return (x / np.sqrt(var + eps)) * weight + + +def half_rope(x, cos, sin, head_dim, rope_head_dim): + """HALF mode: only acts on last rope_head_dim dims per row.""" + out = x.astype(np.float32).copy() + r = rope_head_dim + base = head_dim - r + v = out[:, base:base + r] + half = r // 2 + rot = np.concatenate([-v[:, half:], v[:, :half]], axis=-1) + out[:, base:base + r] = v * cos + rot * sin + return out + + +def compressor_golden(x, wkv, wgate, ape, norm_weight, rope_cos, rope_sin, + cmp_ratio, head_dim, rope_head_dim, norm_eps): + """ + x:(B,S,hidden) wkv/wgate:(coffD,hidden) ape:(cmp_ratio,coffD) + norm_weight:(head_dim) rope_cos/sin:(B,Sr,rope_head_dim) + return cmp_kv:(B,Sr,head_dim) float32 + """ + B, S, _ = x.shape + SR = (S + cmp_ratio - 1) // cmp_ratio + cmp_kv = np.zeros((B, SR, head_dim), dtype=np.float32) + for b in range(B): + kv = x[b] @ wkv.T + score = x[b] @ wgate.T + for g in range(SR): + r0 = g * cmp_ratio + r1 = min(r0 + cmp_ratio, S) + n = r1 - r0 + sc = score[r0:r1] + ape[:n] + kvg = kv[r0:r1] + m = np.max(sc, axis=0, keepdims=True) + e = np.exp(sc - m) + p = e / np.sum(e, axis=0, keepdims=True) + cmp = np.sum(p * kvg, axis=0, keepdims=True) + norm = rms_norm(cmp, norm_weight, norm_eps) + out = half_rope(norm, rope_cos[b, g:g + 1], rope_sin[b, g:g + 1], + head_dim, rope_head_dim) + cmp_kv[b, g] = out[0] + return cmp_kv + + +# ---- Test cases ---- +# Using minimal scenario: S=cmp_ratio so SR=1; single block holds all. +CASES = [ + # (B, S, HIDDEN, HEAD_DIM, CMP_RATIO, COFF, ROPE_HEAD_DIM, BLOCK_SIZE) + (1, 128, 1024, 512, 128, 1, 64, 128), +] + + +@pytest.mark.parametrize("B,S,HIDDEN,HEAD_DIM,CMP_RATIO,COFF,ROPE_HD,BLOCK_SIZE", CASES) +def test_compressor(B, S, HIDDEN, HEAD_DIM, CMP_RATIO, COFF, ROPE_HD, BLOCK_SIZE): + torch.manual_seed(2025) + np.random.seed(2025) + COFF_D = COFF * HEAD_DIM + SR = (S + CMP_RATIO - 1) // CMP_RATIO + NORM_EPS = 1e-6 + ROTARY_MODE = 1 # HALF + + # Generate inputs (fp32 reference) + x_np = (np.random.randn(B, S, HIDDEN) * 0.1).astype(np.float32) + wkv_np = (np.random.randn(COFF_D, HIDDEN) * 0.05).astype(np.float32) + wgate_np = (np.random.randn(COFF_D, HIDDEN) * 0.05).astype(np.float32) + ape_np = (np.random.randn(CMP_RATIO, COFF_D) * 0.1).astype(np.float32) + norm_w_np = (np.random.randn(HEAD_DIM) * 0.1 + 1.0).astype(np.float32) + rope_cos_np = (np.random.randn(B, SR, ROPE_HD) * 0.1).astype(np.float32) + rope_sin_np = (np.random.randn(B, SR, ROPE_HD) * 0.1).astype(np.float32) + + # Golden + ref = compressor_golden(x_np, wkv_np, wgate_np, ape_np, norm_w_np, + rope_cos_np, rope_sin_np, + CMP_RATIO, HEAD_DIM, ROPE_HD, NORM_EPS) + + # To half-precision tensors for NPU + x_t = torch.from_numpy(x_np).half().npu() + wkv_t = torch.from_numpy(wkv_np).half().npu() + wgate_t = torch.from_numpy(wgate_np).half().npu() + ape_t = torch.from_numpy(ape_np).float().npu() + norm_w_t = torch.from_numpy(norm_w_np).half().npu() + rope_sin_t = torch.from_numpy(rope_sin_np).half().npu() + rope_cos_t = torch.from_numpy(rope_cos_np).half().npu() + + # In-place state (float32, paged) + block_num = (S + BLOCK_SIZE - 1) // BLOCK_SIZE + kv_state_t = torch.zeros(block_num, BLOCK_SIZE, COFF_D, dtype=torch.float32).npu() + score_state_t = torch.zeros(block_num, BLOCK_SIZE, COFF_D, dtype=torch.float32).npu() + + # Block tables: (B, maxBlock) int32 + max_block = block_num + kv_bt = torch.zeros(B, max_block, dtype=torch.int32).npu() + score_bt = torch.zeros(B, max_block, dtype=torch.int32).npu() + + out = custom_ops.compressor_npu( + x_t, wkv_t, wgate_t, kv_state_t, score_state_t, + ape_t, norm_w_t, rope_sin_t, rope_cos_t, + kv_bt, score_bt, + rope_head_dim=ROPE_HD, cmp_ratio=CMP_RATIO, coff=COFF, + norm_eps=NORM_EPS, rotary_mode=ROTARY_MODE, enable_grad=False) + + out_np = out.cpu().float().numpy() + ref_fp16 = ref.astype(np.float16).astype(np.float32) + assert out_np.shape == ref_fp16.shape, f"{out_np.shape} vs {ref_fp16.shape}" + + atol, rtol = 2e-2, 2e-2 + diff = np.abs(out_np - ref_fp16) + rel = diff / (np.abs(ref_fp16) + 1e-6) + fail_count = int(np.sum((diff > atol) & (rel > rtol))) + total = out_np.size + assert fail_count == 0, ( + f"compressor mismatch: {fail_count}/{total} elems exceed tol, " + f"maxAbs={diff.max():.5f} maxRel={rel.max():.5f}") \ No newline at end of file diff --git a/test/python_test/test_dispatch_ffn_combine.py b/test/python_test/test_dispatch_ffn_combine.py new file mode 100644 index 0000000..0aa5976 --- /dev/null +++ b/test/python_test/test_dispatch_ffn_combine.py @@ -0,0 +1,229 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- +""" +pytest test for dispatch_ffn_combine (multi-card MoE dispatch+FFN+combine fusion op). +Migrated from vllm-ascend multicard_ops_a3. + +Dependencies: + - custom_ops_lib (pybind11 module built from RegisterOps.cpp) + - HCCL multi-card (world_size=2) + +NOTE: This test requires: + 1. custom_ops_lib.so compiled with aclnnDispatchFFNCombine support + 2. At least 2 NPU cards available + 3. HCCL communication environment +""" +import random + +import torch +import torch.distributed as dist +import torch.multiprocessing as mp +import torch_npu +from torch.distributed.distributed_c10d import _get_default_group + +import custom_ops_lib + + +class TestDispatchFFNCombine: + def __init__(self, rank, world_size, port): + self.rank = rank + self.world_size = world_size + self.master_ip = "127.0.0.1" + self.port = port + + def get_hcomm(self, comm_group): + hcomm_info = None + if torch.__version__ > "2.0.1": + hcomm_info = comm_group._get_backend(torch.device("npu")).get_hccl_comm_name(self.rank) + else: + hcomm_info = comm_group.get_hccl_comm_name(self.rank) + return hcomm_info + + def setup_ep_tp( + self, + rank, + tp_size, + ep_size, + backend_type, + ep_ranks_list=None, + tp_ranks_list=None, + ): + for i in range(tp_size): + if ep_ranks_list: + ep_ranks = ep_ranks_list[i] + else: + ep_ranks = [x + ep_size * i for x in range(ep_size)] + ep_group = dist.new_group(backend=backend_type, ranks=ep_ranks) + if rank in ep_ranks: + ep_group_tmp = ep_group + for i in range(ep_size): + if tp_ranks_list: + tp_ranks = tp_ranks_list[i] + else: + tp_ranks = [x * ep_size + i for x in range(tp_size)] + tp_group = dist.new_group(backend=backend_type, ranks=tp_ranks) + if rank in tp_ranks: + tp_group_tmp = tp_group + return ep_group_tmp, tp_group_tmp + + def generate_hcom(self): + torch_npu.npu.set_device(self.rank) + dist.init_process_group( + backend="hccl", + rank=self.rank, + world_size=self.world_size, + init_method=f"tcp://127.0.0.1:{self.port}", + ) + + ep_size = 0 + tp_size = self.world_size + hcomm_info_dist = { + "default_pg_info": None, + "ep_hcomm_info": None, + "group_ep": None, + "tp_hcomm_info": None, + "group_tp": None, + } + if ep_size and tp_size: + group_ep, group_tp = self.setup_ep_tp(self.rank, tp_size, ep_size, "hccl", None, None) + hcomm_info_dist["ep_hcomm_info"] = self.get_hcomm(group_ep) + hcomm_info_dist["tp_hcomm_info"] = self.get_hcomm(group_tp) + hcomm_info_dist["group_ep"] = group_ep + hcomm_info_dist["group_tp"] = group_tp + else: + if dist.is_available(): + default_pg = _get_default_group() + hcomm_info_dist["default_pg_info"] = self.get_hcomm(default_pg) + hcomm_info = hcomm_info_dist["default_pg_info"] + self.hcomm_info = hcomm_info + + def run_tensor_list(self) -> bool: + torch_npu.npu.set_device(self.rank) + m = 64 + k = 1024 + n = 1024 + topk = 8 + e = 8 + k2 = n // 2 + n2 = k + + torch_npu.npu.config.allow_internal_format = True + x = self.generate_random_tensor((m, k), dtype=torch.bfloat16).npu() + weight1 = self.generate_random_tensor((e, k, n), dtype=torch.int8).npu() + weight1 = torch_npu.npu_format_cast(weight1, 29) + weight2 = self.generate_random_tensor((e, k2, n2), dtype=torch.int8).npu() + weight2 = torch_npu.npu_format_cast(weight2, 29) + + expert_idx = torch.randint(0, self.world_size * e, (m, topk), dtype=torch.int32).npu() + scale1 = torch.randint(0, 1, (e, n), dtype=torch.int64).npu() + scale2 = torch.randint(0, 1, (e, n2), dtype=torch.int64).npu() + probs = torch.randn(size=(m, topk), dtype=torch.float32).npu() + xactmask = torch.randint(0, 2, (m,), dtype=torch.bool).npu() + + weight1_nz_npu = [] + weight2_nz_npu = [] + scale1_npu = [] + scale2_npu = [] + for i in range(e): + weight1_nz_npu.append(torch_npu.npu_format_cast(weight1[i].npu(), 29)) + scale1_npu.append(scale1[i].npu()) + weight2_nz_npu.append(torch_npu.npu_format_cast(weight2[i].npu(), 29)) + scale2_npu.append(scale2[i].npu()) + + out = self.generate_random_tensor((m, k), dtype=torch.bfloat16).npu() + expert_token_nums = self.generate_random_tensor((1, e), dtype=torch.int32).npu() + + custom_ops_lib.dispatch_ffn_combine( + x, weight1_nz_npu, weight2_nz_npu, expert_idx, + scale1_npu, scale2_npu, probs, + self.hcomm_info, 512, out, expert_token_nums, + xactmask, 0.0, + ) + return True + + def run_normal(self) -> bool: + torch_npu.npu.set_device(self.rank) + m = 64 + k = 1024 + n = 1024 + topk = 8 + e = 8 + k2 = n // 2 + n2 = k + + torch_npu.npu.config.allow_internal_format = True + x = self.generate_random_tensor((m, k), dtype=torch.bfloat16).npu() + weight1 = self.generate_random_tensor((e, k, n), dtype=torch.int8).npu() + weight1 = torch_npu.npu_format_cast(weight1, 29) + weight2 = self.generate_random_tensor((e, k2, n2), dtype=torch.int8).npu() + weight2 = torch_npu.npu_format_cast(weight2, 29) + + expert_idx = torch.randint(0, self.world_size * e, (m, topk), dtype=torch.int32).npu() + scale1 = torch.randint(0, 1, (e, n), dtype=torch.int64).npu() + scale2 = torch.randint(0, 1, (e, n2), dtype=torch.int64).npu() + probs = torch.randn(size=(m, topk), dtype=torch.float32).npu() + + weight1_nz_npu = [] + weight2_nz_npu = [] + scale1_npu = [] + scale2_npu = [] + weight1_nz_npu.append(torch_npu.npu_format_cast(weight1.npu(), 29)) + scale1_npu.append(scale1.npu()) + weight2_nz_npu.append(torch_npu.npu_format_cast(weight2.npu(), 29)) + scale2_npu.append(scale2.npu()) + + out = self.generate_random_tensor((m, k), dtype=torch.bfloat16).npu() + expert_token_nums = self.generate_random_tensor((1, e), dtype=torch.int32).npu() + + custom_ops_lib.dispatch_ffn_combine( + x, weight1_nz_npu, weight2_nz_npu, expert_idx, + scale1_npu, scale2_npu, probs, + self.hcomm_info, 512, out, expert_token_nums, + None, 0.0, + ) + return True + + def generate_random_tensor(self, size, dtype): + if dtype in [torch.float16, torch.bfloat16, torch.float32]: + return torch.randn(size=size, dtype=dtype) + elif dtype is torch.int8: + return torch.randint(-16, 16, size=size, dtype=dtype) + elif dtype is torch.int32: + return torch.randint(-1024, 1024, size=size, dtype=dtype) + else: + raise ValueError(f"Invalid dtype: {dtype}") + + +def worker(rank: int, world_size: int, port: int, q: mp.SimpleQueue): + op = TestDispatchFFNCombine(rank, world_size, port) + op.generate_hcom() + out1 = op.run_tensor_list() + q.put(out1) + out2 = op.run_normal() + q.put(out2) + + +@torch.inference_mode() +def test_dispatch_ffn_combine_kernel(): + world_size = 2 + mp.set_start_method("fork", force=True) + + q = mp.SimpleQueue() + p_list = [] + port = 29501 + random.randint(0, 10000) + + for rank in range(world_size): + p = mp.Process(target=worker, args=(rank, world_size, port, q)) + p.start() + p_list.append(p) + + results = [q.get() for _ in range(world_size)] + + for p in p_list: + p.join() + + assert all(results) + + +if __name__ == "__main__": + test_dispatch_ffn_combine_kernel() \ No newline at end of file diff --git a/test/python_test/test_dispatch_gmm_combine_decode.py b/test/python_test/test_dispatch_gmm_combine_decode.py new file mode 100644 index 0000000..3537040 --- /dev/null +++ b/test/python_test/test_dispatch_gmm_combine_decode.py @@ -0,0 +1,589 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- +""" +pytest test for dispatch_gmm_combine_decode (multi-card MoE dispatch+GMM+combine decode). +Migrated from vllm-ascend multicard_ops_a3. + +Compares SmallOps (torch_npu native ops chain) vs FusionOp (single fused op). + +Dependencies: + - custom_ops_lib (pybind11 module built from RegisterOps.cpp) + - torch_npu native MoE ops (npu_moe_distribute_dispatch_v2, npu_grouped_matmul, etc.) + - HCCL multi-card (ep_world_size=2) + +NOTE: This test requires: + 1. custom_ops_lib.so compiled with aclnnDispatchGmmCombineDecode support + 2. At least 2 NPU cards available + 3. HCCL communication environment +""" +import gc +import os +import sys +import time +from pathlib import Path + +import numpy as np +import torch +import torch.distributed as dist +import torch.multiprocessing as mp +import torch_npu + +import custom_ops_lib + +torch.manual_seed(42) +torch_npu.npu.config.allow_internal_format = True +LOG_NAME = "dispatch_gmm_combine_decode_test_logs" +BASE_KWARGS = { + "batch_size": 64, + "token_hidden_size": 7168, + "moe_intermediate_size": 2048, + "ep_world_size": 2, + "moe_expert_num": 8, + "shared_expert_rank_num": 0, + "top_k": 4, + "test_bfloat16": True, + "enable_dynamic_bs": False, + "test_graph": False, + "with_mc2_mask": False, + "dynamic_eplb": False, + "w8a8_dynamic": True, + "is_nz": True, +} + + +def redirect_output(log_file_path): + log_path = Path(LOG_NAME) / log_file_path + log_path.parent.mkdir(parents=True, exist_ok=True) + f = open(LOG_NAME + "/" + log_file_path, "w") # noqa: SIM115 + os.dup2(f.fileno(), sys.stdout.fileno()) + os.dup2(f.fileno(), sys.stderr.fileno()) + return f + + +def permute_weight(w: torch.Tensor, tile_n): + *dims, n = w.shape + order = list(range(len(dims))) + [-2, -3, -1] + return w.reshape(*dims, 2, n // tile_n, tile_n // 2).permute(order).reshape(*dims, n).contiguous() + + +def output_to_file(rank_id): + return rank_id > 0 + + +class DecodeMoeOps(torch.nn.Module): + def __init__( + self, + gmm1_weight, + gmm1_weight_scale, + gmm2_weight, + gmm2_weight_scale, + ep_hcomm_info, + batch_size, + token_hidden_size, + moe_intermediate_size, + ep_world_size, + moe_expert_num, + global_rank_id, + shared_expert_rank_num=0, + dynamic_eplb=False, + w8a8_dynamic=True, + is_nz=True, + ): + super().__init__() + if w8a8_dynamic: + assert gmm1_weight_scale is not None and gmm2_weight_scale is not None, ( + "gmm1_weight_scale and gmm2_weight_scale must be provided for w8a8_dynamic" + ) + else: + assert gmm1_weight_scale is None and gmm2_weight_scale is None, ( + "gmm1_weight_scale and gmm2_weight_scale must be None for w8a8_dynamic" + ) + self.ep_hcomm_info = ep_hcomm_info + self.batch_size = batch_size + self.token_hidden_size = token_hidden_size + self.moe_intermediate_size = moe_intermediate_size + self.ep_world_size = ep_world_size + self.moe_expert_num = moe_expert_num + self.global_rank_id = global_rank_id + self.shared_expert_rank_num = shared_expert_rank_num + is_shared_expert = global_rank_id < shared_expert_rank_num + moe_expert_num_per_rank = moe_expert_num // (ep_world_size - shared_expert_rank_num) + self.local_expert_num = 1 if is_shared_expert else moe_expert_num_per_rank + self.ep_recv_count_size = self.local_expert_num * ep_world_size + self.dynamic_eplb = dynamic_eplb + self.w8a8_dynamic = w8a8_dynamic + self.is_nz = is_nz + self.gmm1_weight = torch.empty([self.local_expert_num, self.token_hidden_size, self.moe_intermediate_size * 2]) + self.gmm2_weight = torch.empty([self.local_expert_num, self.moe_intermediate_size, self.token_hidden_size]) + if self.w8a8_dynamic: + self.gmm1_weight_scale = torch.empty([self.local_expert_num, self.moe_intermediate_size * 2]) + self.gmm2_weight_scale = torch.empty([self.local_expert_num, self.token_hidden_size]) + else: + self.gmm1_weight_scale = None + self.gmm2_weight_scale = None + self.gmm1_weight_scale_fp32 = None + self.gmm2_weight_scale_fp32 = None + self._process_weights_after_loading(gmm1_weight, gmm1_weight_scale, gmm2_weight, gmm2_weight_scale) + + def _process_weights_after_loading(self, gmm1_weight, gmm1_weight_scale, gmm2_weight, gmm2_weight_scale): + if self.w8a8_dynamic: + gmm1_weight = torch_npu.npu_format_cast(gmm1_weight, torch_npu.Format.FRACTAL_NZ) + gmm2_weight = torch_npu.npu_format_cast(gmm2_weight, torch_npu.Format.FRACTAL_NZ) + self.gmm1_weight = torch.nn.Parameter(gmm1_weight, requires_grad=False) + self.gmm2_weight = torch.nn.Parameter(gmm2_weight, requires_grad=False) + if self.w8a8_dynamic: + self.gmm1_weight_scale = torch.nn.Parameter(gmm1_weight_scale, requires_grad=False) + self.gmm2_weight_scale = torch.nn.Parameter(gmm2_weight_scale, requires_grad=False) + self.gmm1_weight_scale_fp32 = torch.nn.Parameter(gmm1_weight_scale.float(), requires_grad=False) + self.gmm2_weight_scale_fp32 = torch.nn.Parameter(gmm2_weight_scale.float(), requires_grad=False) + + def _apply_ops(self, x, expert_ids, smooth_scales, expert_scales, x_active_mask): + raise NotImplementedError("To be implemented in subclass") + + def forward(self, x, expert_ids, smooth_scales, expert_scales, x_active_mask): + return self._apply_ops(x, expert_ids, smooth_scales, expert_scales, x_active_mask) + + +class SmallOps(DecodeMoeOps): + def __init__( + self, + gmm1_weight, + gmm1_weight_scale, + gmm2_weight, + gmm2_weight_scale, + ep_hcomm_info, + batch_size, + token_hidden_size, + moe_intermediate_size, + ep_world_size, + moe_expert_num, + global_rank_id, + shared_expert_rank_num=0, + dynamic_eplb=False, + w8a8_dynamic=True, + is_nz=True, + ): + super().__init__( + gmm1_weight, + gmm1_weight_scale, + gmm2_weight, + gmm2_weight_scale, + ep_hcomm_info, + batch_size, + token_hidden_size, + moe_intermediate_size, + ep_world_size, + moe_expert_num, + global_rank_id, + shared_expert_rank_num, + dynamic_eplb, + w8a8_dynamic, + is_nz, + ) + self.tp_hcomm_info = "" + + def _apply_ops(self, x, expert_ids, smooth_scales, expert_scales, x_active_mask): + outputs = torch_npu.npu_moe_distribute_dispatch_v2( + x=x, + expert_ids=expert_ids, + expert_scales=expert_scales, + x_active_mask=x_active_mask, + group_ep=self.ep_hcomm_info, + ep_world_size=self.ep_world_size, + ep_rank_id=self.global_rank_id, + moe_expert_num=self.moe_expert_num, + group_tp=self.tp_hcomm_info, + tp_world_size=1, + tp_rank_id=0, + expert_shard_type=0, + shared_expert_num=1, + shared_expert_rank_num=self.shared_expert_rank_num, + quant_mode=2 if self.w8a8_dynamic else 0, + global_bs=self.batch_size * self.ep_world_size, + expert_token_nums_type=1, + ) + ( + expand_x, + dynamic_scales, + assist_info_for_combine, + expert_token_nums, + ep_send_counts, + tp_send_counts, + expand_scales, + ) = outputs + output_dtype = x.dtype + + y1_int32 = torch_npu.npu_grouped_matmul( + x=[expand_x], + weight=[self.gmm1_weight], + split_item=3, + group_list_type=1, + group_type=0, + group_list=expert_token_nums, + output_dtype=torch.int32 if self.w8a8_dynamic else output_dtype, + )[0] + y1_scale = None + if self.w8a8_dynamic: + y1, y1_scale = torch_npu.npu_dequant_swiglu_quant( + x=y1_int32, + weight_scale=self.gmm1_weight_scale.to(torch.float32), + activation_scale=dynamic_scales, + bias=None, + quant_scale=None, + quant_offset=None, + group_index=expert_token_nums, + activate_left=True, + quant_mode=1, + ) + else: + y1 = torch_npu.npu_swiglu(y1_int32) + y2 = torch_npu.npu_grouped_matmul( + x=[y1], + weight=[self.gmm2_weight], + scale=[self.gmm2_weight_scale] if self.w8a8_dynamic else None, + per_token_scale=[y1_scale] if self.w8a8_dynamic else None, + split_item=2, + group_list_type=1, + group_type=0, + group_list=expert_token_nums, + output_dtype=output_dtype, + )[0] + combine_output = torch_npu.npu_moe_distribute_combine_v2( + expand_x=y2, + expert_ids=expert_ids, + assist_info_for_combine=assist_info_for_combine, + ep_send_counts=ep_send_counts, + expert_scales=expert_scales, + x_active_mask=x_active_mask, + group_ep=self.ep_hcomm_info, + ep_world_size=self.ep_world_size, + ep_rank_id=self.global_rank_id, + moe_expert_num=self.moe_expert_num, + tp_send_counts=tp_send_counts, + expand_scales=expand_scales, + group_tp=self.tp_hcomm_info, + tp_world_size=1, + tp_rank_id=0, + expert_shard_type=0, + shared_expert_num=1, + shared_expert_rank_num=self.shared_expert_rank_num, + global_bs=self.batch_size * self.ep_world_size, + ) + return (combine_output, expert_token_nums) + + +class FusionOp(DecodeMoeOps): + def __init__( + self, + gmm1_weight, + gmm1_weight_scale, + gmm2_weight, + gmm2_weight_scale, + ep_hcomm_info, + batch_size, + token_hidden_size, + moe_intermediate_size, + ep_world_size, + moe_expert_num, + global_rank_id, + shared_expert_rank_num=0, + dynamic_eplb=False, + w8a8_dynamic=True, + is_nz=True, + ): + super().__init__( + gmm1_weight, + gmm1_weight_scale, + gmm2_weight, + gmm2_weight_scale, + ep_hcomm_info, + batch_size, + token_hidden_size, + moe_intermediate_size, + ep_world_size, + moe_expert_num, + global_rank_id, + shared_expert_rank_num, + dynamic_eplb, + w8a8_dynamic, + is_nz, + ) + + def _apply_ops(self, x, expert_ids, smooth_scales, expert_scales, x_active_mask): + smooth_scales = torch.zeros(128 * 1024 * 1024).npu() + output, expert_token_nums = custom_ops_lib.dispatch_gmm_combine_decode( + x, + expert_ids, + self.gmm1_weight, + self.gmm1_weight_scale_fp32, + self.gmm2_weight, + self.gmm2_weight_scale_fp32, + expert_scales, + smooth_scales, + x_active_mask, + self.ep_hcomm_info, + self.ep_world_size, + self.global_rank_id, + self.moe_expert_num, + 1, # shared_expert_num + self.shared_expert_rank_num, + 0, # quant_mode + self.batch_size * self.ep_world_size, # global_bs + ) + return (output, expert_token_nums) + + def _process_weights_after_loading(self, gmm1_weight, gmm1_weight_scale, gmm2_weight, gmm2_weight_scale): + if self.is_nz: + gmm1_weight = torch_npu.npu_format_cast(gmm1_weight, torch_npu.Format.FRACTAL_NZ) + gmm2_weight = torch_npu.npu_format_cast(gmm2_weight, torch_npu.Format.FRACTAL_NZ) + + if self.dynamic_eplb: + self.gmm1_weight = [weight.clone() for weight in gmm1_weight.unbind(dim=0)] + self.gmm2_weight = [weight.clone() for weight in gmm2_weight.unbind(dim=0)] + if self.w8a8_dynamic: + self.gmm1_weight_scale_fp32 = [weight.clone() for weight in gmm1_weight_scale.unbind(dim=0)] + self.gmm2_weight_scale_fp32 = [weight.clone() for weight in gmm2_weight_scale.unbind(dim=0)] + else: + self.gmm1_weight_scale_fp32 = [torch.ones(1).npu().to(gmm1_weight.dtype)] + self.gmm2_weight_scale_fp32 = [torch.ones(1).npu().to(gmm2_weight.dtype)] + else: + self.gmm1_weight = [gmm1_weight.clone()] + self.gmm2_weight = [gmm2_weight.clone()] + if self.w8a8_dynamic: + self.gmm1_weight_scale_fp32 = [gmm1_weight_scale.clone()] + self.gmm2_weight_scale_fp32 = [gmm2_weight_scale.clone()] + else: + self.gmm1_weight_scale_fp32 = [torch.ones(1).npu().to(gmm1_weight.dtype)] + self.gmm2_weight_scale_fp32 = [torch.ones(1).npu().to(gmm2_weight.dtype)] + + +def generate_datas( + batch_size, + token_hidden_size, + moe_intermediate_size, + ep_world_size, + moe_expert_num, + global_rank_id, + shared_expert_rank_num=0, + top_k=8, + test_bfloat16=True, + enable_dynamic_bs=False, + with_mc2_mask=False, + w8a8_dynamic=True, +): + is_shared_expert = global_rank_id < shared_expert_rank_num + moe_expert_num_per_rank = moe_expert_num // (ep_world_size - shared_expert_rank_num) + actual_bs = int( + torch.randint(2 if with_mc2_mask else 1, batch_size, [1]).item() if enable_dynamic_bs else batch_size + ) + local_expert_num = 1 if is_shared_expert else moe_expert_num_per_rank + gmm1_input_dim = token_hidden_size + gmm1_output_dim = moe_intermediate_size * 2 + gmm2_input_dim = moe_intermediate_size + gmm2_output_dim = token_hidden_size + x = torch.rand([actual_bs, token_hidden_size]) * 0.5 - 0.5 + expert_ids = ( + torch.arange(global_rank_id * batch_size * top_k, global_rank_id * batch_size * top_k + actual_bs * top_k) + .to(torch.int32) + .view(actual_bs, top_k) + ) + expert_ids = expert_ids % moe_expert_num + gmm1_weight_scale = None + gmm2_weight_scale = None + if w8a8_dynamic: + if is_shared_expert: + gmm1_weight = torch.ones([local_expert_num, gmm1_input_dim, gmm1_output_dim]).to(torch.int8) * 4 + gmm2_weight = torch.ones([local_expert_num, gmm2_input_dim, gmm2_output_dim]).to(torch.int8) * 4 + gmm1_weight[:, :, ::2] = gmm1_weight[:, :, ::2] * -1 + gmm2_weight[:, :, ::2] = gmm2_weight[:, :, ::2] * -1 + gmm1_weight_scale = torch.ones([local_expert_num, gmm1_output_dim]) * 0.0015 + gmm2_weight_scale = torch.ones([local_expert_num, gmm2_output_dim]) * 0.0015 + else: + gmm1_weight = torch.randint(-16, 16, [local_expert_num, gmm1_input_dim, gmm1_output_dim]).to(torch.int8) + gmm2_weight = torch.randint(-16, 16, [local_expert_num, gmm2_input_dim, gmm2_output_dim]).to(torch.int8) + gmm1_weight_scale = torch.rand([local_expert_num, gmm1_output_dim]) * 0.003 + 0.0015 + gmm2_weight_scale = torch.rand([local_expert_num, gmm2_output_dim]) * 0.003 + 0.0015 + else: + if is_shared_expert: + gmm1_weight = ( + torch.ones([local_expert_num, gmm1_input_dim, gmm1_output_dim]).to( + torch.bfloat16 if test_bfloat16 else torch.float16 + ) + * 0.5 + ) + gmm2_weight = ( + torch.ones([local_expert_num, gmm2_input_dim, gmm2_output_dim]).to( + torch.bfloat16 if test_bfloat16 else torch.float16 + ) + * 0.5 + ) + else: + gmm1_weight = ( + torch.rand([local_expert_num, gmm1_input_dim, gmm1_output_dim]).to( + torch.bfloat16 if test_bfloat16 else torch.float16 + ) + * 0.25 + ) + gmm2_weight = ( + torch.rand([local_expert_num, gmm2_input_dim, gmm2_output_dim]).to( + torch.bfloat16 if test_bfloat16 else torch.float16 + ) + * 0.25 + ) + gmm1_weight[:, ::2, :] = gmm1_weight[:, ::2, :] * -1 + gmm2_weight[:, ::2, :] = gmm2_weight[:, ::2, :] * -1 + expert_scales = torch.rand(actual_bs, top_k) + if test_bfloat16: + x = x.bfloat16() + if w8a8_dynamic: + assert gmm1_weight_scale is not None and gmm2_weight_scale is not None, ( + "gmm1_weight_scale and gmm2_weight_scale must be provided for w8a8_dynamic" + ) + gmm1_weight_scale = gmm1_weight_scale.bfloat16() + gmm2_weight_scale = gmm2_weight_scale.bfloat16() + else: + x = x.half() + smooth_sales = None + x_active_mask = None + valid_token_num = actual_bs + if with_mc2_mask: + valid_token_num = int(torch.randint(1, actual_bs, [1]).item()) + x_active_mask = torch.cat((torch.ones(valid_token_num), torch.zeros(actual_bs - valid_token_num))).bool() + return ( + (x, expert_ids, smooth_sales, expert_scales, x_active_mask), + (gmm1_weight, gmm1_weight_scale, gmm2_weight, gmm2_weight_scale), + actual_bs, + valid_token_num, + ) + + +def run_once( + local_rank_id, + batch_size, + token_hidden_size, + moe_intermediate_size, + ep_world_size, + moe_expert_num, + shared_expert_rank_num=0, + top_k=8, + test_bfloat16=True, + enable_dynamic_bs=False, + test_graph=False, + with_mc2_mask=False, + dynamic_eplb=False, + w8a8_dynamic=True, + is_nz=True, +): + log_file = redirect_output(f"local_rank_{local_rank_id}.log") if output_to_file(local_rank_id) else None + global_rank_id = local_rank_id # single node + device_id = local_rank_id % 8 + torch_npu.npu.set_device(device_id) + + # Initialize distributed environment + os.environ["MASTER_ADDR"] = "127.0.0.1" + os.environ["MASTER_PORT"] = "29500" + dist.init_process_group(backend="hccl", rank=local_rank_id, world_size=ep_world_size) + ep_ranks_list = list(np.arange(0, ep_world_size)) + ep_group = dist.new_group(backend="hccl", ranks=ep_ranks_list) + ep_group_small = dist.new_group(backend="hccl", ranks=ep_ranks_list) + + ep_hcomm_info_fused = ep_group._get_backend(torch.device("npu")).get_hccl_comm_name(local_rank_id) + ep_hcomm_info_small = ep_group_small._get_backend(torch.device("npu")).get_hccl_comm_name(local_rank_id) + torch_npu.npu.synchronize(device_id) + + parameter = ( + batch_size, + token_hidden_size, + moe_intermediate_size, + ep_world_size, + moe_expert_num, + global_rank_id, + shared_expert_rank_num, + ) + input_datas, weight_datas, actual_bs, valid_token_num = generate_datas( + *parameter, top_k, test_bfloat16, enable_dynamic_bs, with_mc2_mask, w8a8_dynamic + ) + input_datas = [data.npu() if data is not None else None for data in input_datas] + weight_datas = [data.npu() if data is not None else None for data in weight_datas] + small_ops = SmallOps(*weight_datas, ep_hcomm_info_small, *parameter, dynamic_eplb, w8a8_dynamic, is_nz).npu() + fused_ops = FusionOp(*weight_datas, ep_hcomm_info_fused, *parameter, dynamic_eplb, w8a8_dynamic, is_nz).npu() + + # test performance + start_time = time.perf_counter() + for _ in range(100): + small_op_token_output, small_op_count_output = small_ops(*input_datas) + torch_npu.npu.synchronize(device_id) + end_time = time.perf_counter() + elapsed_time = end_time - start_time + elapsed_time_us = elapsed_time * 1000000 + print(f"rank-{global_rank_id} small {elapsed_time_us} us") + start_time = time.perf_counter() + for _ in range(100): + fused_op_token_output, fused_op_count_output = fused_ops(*input_datas) + torch_npu.npu.synchronize(device_id) + end_time = time.perf_counter() + elapsed_time = end_time - start_time + elapsed_time_us = elapsed_time * 1000000 + print(f"rank-{global_rank_id} fused {elapsed_time_us} us") + small_op_token_output, small_op_count_output = small_ops(*input_datas) + torch_npu.npu.synchronize(device_id) + print(f"rank-{global_rank_id} Small op End") + fused_op_token_output, fused_op_count_output = fused_ops(*input_datas) + torch_npu.npu.synchronize(device_id) + print(f"rank-{global_rank_id} Fused op End") + dist.destroy_process_group() + if log_file is not None: + log_file.close() + try: + torch.testing.assert_close( + small_op_token_output[0:valid_token_num].cpu(), + fused_op_token_output[0:valid_token_num].cpu(), + atol=2.0, + rtol=0.02, + ) + torch.testing.assert_close(small_op_count_output.cpu(), fused_op_count_output.cpu()) + except Exception as e: + print(f"rank-{global_rank_id} Assert close Failed: {e}") + else: + print(f"rank-{global_rank_id} Assert close Pass") + gc.collect() + torch.npu.empty_cache() + torch.npu.reset_peak_memory_stats() + + +@torch.inference_mode() +def test_dispatch_gmm_combine_decode_base(): + custom_kwargs = BASE_KWARGS.copy() + custom_kwargs["batch_size"] = 32 + custom_kwargs["ep_world_size"] = 2 + custom_kwargs["moe_expert_num"] = 8 + custom_kwargs["top_k"] = 4 + custom_kwargs["w8a8_dynamic"] = True + custom_kwargs["is_nz"] = True + ep_world_size = custom_kwargs["ep_world_size"] + custom_args = tuple(custom_kwargs.values()) + print(f"{custom_kwargs=}") + mp.spawn(run_once, args=custom_args, nprocs=ep_world_size, join=True) + print(f"{custom_kwargs=}") + + +@torch.inference_mode() +def test_dispatch_gmm_combine_decode_with_mc2_mask(): + custom_kwargs = BASE_KWARGS.copy() + custom_kwargs["with_mc2_mask"] = True + ep_world_size = custom_kwargs["ep_world_size"] + custom_args = tuple(custom_kwargs.values()) + mp.spawn(run_once, args=custom_args, nprocs=ep_world_size, join=True) + + +@torch.inference_mode() +def test_dispatch_gmm_combine_decode_dynamic_eplb(): + custom_kwargs = BASE_KWARGS.copy() + custom_kwargs["dynamic_eplb"] = True + ep_world_size = custom_kwargs["ep_world_size"] + custom_args = tuple(custom_kwargs.values()) + mp.spawn(run_once, args=custom_args, nprocs=ep_world_size, join=True) + + +if __name__ == "__main__": + test_dispatch_gmm_combine_decode_base() \ No newline at end of file diff --git a/test/python_test/test_index_group_matmul.py b/test/python_test/test_index_group_matmul.py index c0783fd..6bc8a7e 100644 --- a/test/python_test/test_index_group_matmul.py +++ b/test/python_test/test_index_group_matmul.py @@ -35,6 +35,11 @@ torch_npu = pytest.importorskip("torch_npu") custom_ops = pytest.importorskip("custom_ops") +try: + torch_npu.npu.config.allow_internal_format = True # allow FRACTAL_NZ tensors +except Exception: + pass + def index_group_matmul_golden(a, b, scale, per_token_scale, group_list): """Reference int8 grouped matmul + dequant, computed in fp32 -> bf16.""" diff --git a/test/python_test/test_lightning_indexer_quant_metadata.py b/test/python_test/test_lightning_indexer_quant_metadata.py new file mode 100644 index 0000000..e4f62b0 --- /dev/null +++ b/test/python_test/test_lightning_indexer_quant_metadata.py @@ -0,0 +1,124 @@ +# Copyright 2025 The xLLM Authors. All Rights Reserved. +# AICPU metadata op test for lightning_indexer_quant_metadata. +# Verifies: no crash, correct output shape/dtype, output self-consistency. + +from dataclasses import dataclass + +import pytest +import torch +import torch_npu # noqa: F401 +import custom_ops + +META_SIZE = 1024 + + +@dataclass(frozen=True) +class LightningIndexerQuantMetadataCase: + name: str + batch_size: int + max_seq_q: int + max_seq_k: int + num_heads_q: int + num_heads_k: int + head_dim: int + query_quant_mode: int = 0 + key_quant_mode: int = 0 + layout_query: str = "BSND" + layout_key: str = "BSND" + sparse_count: int = 128 + sparse_mode: int = 0 + is_fd: bool = False + pre_token: int = 9223372036854775807 + next_token: int = 9223372036854775807 + cmp_ratio: int = 1 + + +CASES = [ + LightningIndexerQuantMetadataCase( + name="decode_b1_sq1_sk128_n4_nk1_d128", + batch_size=1, max_seq_q=1, max_seq_k=128, + num_heads_q=4, num_heads_k=1, head_dim=128, + ), + LightningIndexerQuantMetadataCase( + name="decode_b4_sq1_sk256_n4_nk1_d128", + batch_size=4, max_seq_q=1, max_seq_k=256, + num_heads_q=4, num_heads_k=1, head_dim=128, + ), + LightningIndexerQuantMetadataCase( + name="decode_b1_sq1_sk512_n4_nk1_d128_sparse64", + batch_size=1, max_seq_q=1, max_seq_k=512, + num_heads_q=4, num_heads_k=1, head_dim=128, + sparse_count=64, + ), + LightningIndexerQuantMetadataCase( + name="decode_b2_sq1_sk1024_n4_nk1_d128_causal", + batch_size=2, max_seq_q=1, max_seq_k=1024, + num_heads_q=4, num_heads_k=1, head_dim=128, + sparse_mode=3, + ), + LightningIndexerQuantMetadataCase( + name="decode_b1_sq4_sk128_n4_nk1_d128_is_fd", + batch_size=1, max_seq_q=4, max_seq_k=128, + num_heads_q=4, num_heads_k=1, head_dim=128, + is_fd=True, + ), +] + + +def _make_inputs(case: LightningIndexerQuantMetadataCase): + """Create actual_seq_lengths tensors on NPU.""" + B = case.batch_size + device = "npu" + # actual_seq_lengths_query: (B,) int32, each = max_seq_q + aslq = torch.full((B,), case.max_seq_q, dtype=torch.int32, device=device) + # actual_seq_lengths_key: (B,) int32, each = max_seq_k + aslk = torch.full((B,), case.max_seq_k, dtype=torch.int32, device=device) + return aslq, aslk + + +@pytest.mark.parametrize("case", CASES, ids=[c.name for c in CASES]) +def test_lightning_indexer_quant_metadata(case: LightningIndexerQuantMetadataCase): + torch.npu.set_device(0) + aslq, aslk = _make_inputs(case) + + meta_t = custom_ops.lightning_indexer_quant_metadata_npu( + aslq, aslk, + num_heads_q=case.num_heads_q, num_heads_k=case.num_heads_k, + head_dim=case.head_dim, + query_quant_mode=case.query_quant_mode, + key_quant_mode=case.key_quant_mode, + batch_size=case.batch_size, + max_seq_q=case.max_seq_q, max_seq_k=case.max_seq_k, + layout_query=case.layout_query, layout_key=case.layout_key, + sparse_count=case.sparse_count, sparse_mode=case.sparse_mode, + is_fd=case.is_fd, + pre_token=case.pre_token, next_token=case.next_token, + cmp_ratio=case.cmp_ratio, + ) + torch.npu.synchronize() + + # Verify output shape and dtype + assert meta_t.shape == (META_SIZE,), f"Expectedshape ({META_SIZE},), got {meta_t.shape}" + assert meta_t.dtype == torch.int32, f"Expected dtype int32, got {meta_t.dtype}" + + # Verify determinism: call again and compare + meta_t2 = custom_ops.lightning_indexer_quant_metadata_npu( + aslq, aslk, + num_heads_q=case.num_heads_q, num_heads_k=case.num_heads_k, + head_dim=case.head_dim, + query_quant_mode=case.query_quant_mode, + key_quant_mode=case.key_quant_mode, + batch_size=case.batch_size, + max_seq_q=case.max_seq_q, max_seq_k=case.max_seq_k, + layout_query=case.layout_query, layout_key=case.layout_key, + sparse_count=case.sparse_count, sparse_mode=case.sparse_mode, + is_fd=case.is_fd, + pre_token=case.pre_token, next_token=case.next_token, + cmp_ratio=case.cmp_ratio, + ) + torch.npu.synchronize() + assert torch.equal(meta_t.cpu(), meta_t2.cpu()), "Metadata output is not deterministic" + + +if __name__ == "__main__": + pytest.main([__file__, "-v", "-x"]) \ No newline at end of file diff --git a/test/python_test/test_quant_lightning_indexer.cpp b/test/python_test/test_quant_lightning_indexer.cpp deleted file mode 100644 index c873b63..0000000 --- a/test/python_test/test_quant_lightning_indexer.cpp +++ /dev/null @@ -1,166 +0,0 @@ -#include -#include -#include -#include -#include -#include -#include "acl/acl.h" -#include "aclnn_quant_lightning_indexer.h" -#include "aclnn_quant_lightning_indexer_metadata.h" - -// ---- Case (PA_BSND, consistent with quant_lightning_indexer_gen.py) ---- -static const int64_t batchSize = 1; -static const int64_t qSeq = 4; -static const int64_t kSeq = 128; // = block_size (1 block) -static const int64_t blockSize = 128; -static const int64_t blockNum = 1; -static const int64_t kMaxBlockPerBatch = 1; -static const int64_t numHeadsQ = 64; -static const int64_t numHeadsK = 1; -static const int64_t headDim = 128; -static const int64_t queryQuantMode = 0; -static const int64_t keyQuantMode = 0; -static const int64_t sparseCount = 8; -static const int64_t sparseMode = 3; -static const int64_t preToken = INT64_MAX; -static const int64_t nextToken = INT64_MAX; -static const int64_t cmpRatio = 1; -static const int64_t stride = 1; -static const int64_t scaleStride = 1; -static const bool returnValues = false; -static std::string layoutQuery = "BSND"; -static std::string layoutKey = "PA_BSND"; -static const int64_t QLI_META_SIZE = 1024; -static const char* DATA_DIR = "/tmp/quant_lightning_indexer_data/"; - -static std::vector ReadBin(const std::string &name) { - std::ifstream f(std::string(DATA_DIR) + name, std::ios::binary | std::ios::ate); - if (!f) { printf("open %s FAILED\n", name.c_str()); return {}; } - std::streamsize sz = f.tellg(); f.seekg(0, std::ios::beg); - std::vector buf(sz); f.read(buf.data(), sz); return buf; -} - -std::tuple CreateTensor(size_t size, std::vector shape, - aclDataType dType, const void* hostData = nullptr) { - std::vector strides(shape.size(), 1); - for (int i = (int)shape.size() - 2; i >= 0; --i) strides[i] = strides[i + 1] * shape[i + 1]; - void* dev = nullptr; - auto ret = aclrtMalloc(&dev, size, ACL_MEM_MALLOC_HUGE_FIRST); - if (ret != ACL_SUCCESS) { printf("aclrtMalloc %d\n", ret); return {nullptr, nullptr}; } - aclTensor* t = aclCreateTensor(shape.data(), shape.size(), dType, strides.data(), 0, - aclFormat::ACL_FORMAT_ND, shape.data(), shape.size(), dev); - if (t == nullptr) { aclrtFree(dev); return {nullptr, nullptr}; } - if (hostData) aclrtMemcpy(dev, size, hostData, size, ACL_MEMCPY_HOST_TO_DEVICE); - return {t, dev}; -} - -int main() { - int32_t deviceId = 0; aclrtStream stream; - aclError ret = aclInit(nullptr); - if (ret != ACL_SUCCESS) { printf("aclInit %d\n", ret); return -1; } - ret = aclrtSetDevice(deviceId); - if (ret != ACL_SUCCESS) { printf("aclrtSetDevice %d\n", ret); return -1; } - ret = aclrtCreateStream(&stream); - if (ret != ACL_SUCCESS) { printf("aclrtCreateStream %d\n", ret); return -1; } - - auto qBuf = ReadBin("query.bin"); - auto kBuf = ReadBin("key.bin"); - auto wBuf = ReadBin("weights.bin"); - auto qsBuf = ReadBin("q_scale.bin"); - auto ksBuf = ReadBin("k_scale.bin"); - auto aslqBuf = ReadBin("aslq.bin"); - auto aslkBuf = ReadBin("aslk.bin"); - auto btBuf = ReadBin("block_table.bin"); - auto goldenBuf = ReadBin("golden_indices.bin"); - if (qBuf.empty() || kBuf.empty() || goldenBuf.empty() || btBuf.empty()) { printf("read bin FAILED\n"); return -1; } - - // ================= Step 1: metadata op ================= - aclTensor *aslqT, *aslkT, *metaT; - void *aslqDev, *aslkDev, *metaDev; - std::tie(aslqT, aslqDev) = CreateTensor(aslqBuf.size(), {batchSize}, ACL_INT32, aslqBuf.data()); - std::tie(aslkT, aslkDev) = CreateTensor(aslkBuf.size(), {batchSize}, ACL_INT32, aslkBuf.data()); - std::tie(metaT, metaDev) = CreateTensor(QLI_META_SIZE * sizeof(int32_t), {QLI_META_SIZE}, ACL_INT32); - - int64_t maxSeqQ = qSeq, maxSeqK = kSeq * cmpRatio; - aclOpExecutor* mExec = nullptr; uint64_t mWs = 0; void* mWsPtr = nullptr; - ret = aclnnQuantLightningIndexerMetadataGetWorkspaceSize( - aslqT, aslkT, numHeadsQ, numHeadsK, headDim, queryQuantMode, keyQuantMode, - batchSize, maxSeqQ, maxSeqK, &layoutQuery[0], &layoutKey[0], - sparseCount, sparseMode, preToken, nextToken, cmpRatio, - metaT, &mWs, &mExec); - if (ret != ACL_SUCCESS) { printf("Metadata GetWorkspaceSize FAILED ret=%d\n", ret); return -1; } - if (mWs > 0) { ret = aclrtMalloc(&mWsPtr, mWs, ACL_MEM_MALLOC_HUGE_FIRST); - if (ret != ACL_SUCCESS) { printf("malloc mWs %d\n", ret); return -1; } } - ret = aclnnQuantLightningIndexerMetadata(mWsPtr, mWs, mExec, stream); - if (ret != ACL_SUCCESS) { printf("Metadata Execute FAILED ret=%d\n", ret); return -1; } - ret = aclrtSynchronizeStream(stream); - if (ret != ACL_SUCCESS) { printf("Metadata Sync %d\n", ret); return -1; } - printf("Metadata OK\n"); - - // ================= Step 2: main op ================= - aclTensor *query, *key, *weights, *qScale, *kScale, *btT, *idxOut, *valOut; - void *qDev, *kDev, *wDev, *qsDev, *ksDev, *btDev, *idxDev, *valDev; - std::tie(query, qDev) = CreateTensor(qBuf.size(), {batchSize, qSeq, numHeadsQ, headDim}, ACL_INT8, qBuf.data()); - // paged key: [blockNum, blockSize, kH, hd] - std::tie(key, kDev) = CreateTensor(kBuf.size(), {blockNum, blockSize, numHeadsK, headDim}, ACL_INT8, kBuf.data()); - std::tie(weights, wDev) = CreateTensor(wBuf.size(), {batchSize, qSeq, numHeadsQ}, ACL_FLOAT16, wBuf.data()); - std::tie(qScale, qsDev) = CreateTensor(qsBuf.size(), {batchSize, qSeq, numHeadsQ}, ACL_FLOAT16, qsBuf.data()); - // paged k_scale: [blockNum, blockSize, kH] - std::tie(kScale, ksDev) = CreateTensor(ksBuf.size(), {blockNum, blockSize, numHeadsK}, ACL_FLOAT16, ksBuf.data()); - std::tie(btT, btDev) = CreateTensor(btBuf.size(), {batchSize, kMaxBlockPerBatch}, ACL_INT32, btBuf.data()); - - int64_t idxCount = batchSize * qSeq * numHeadsK * sparseCount; - std::tie(idxOut, idxDev) = CreateTensor(idxCount * sizeof(int32_t), - {batchSize, qSeq, numHeadsK, sparseCount}, ACL_INT32); - std::tie(valOut, valDev) = CreateTensor(idxCount * sizeof(float), - {batchSize, qSeq, numHeadsK, sparseCount}, ACL_FLOAT); - - aclOpExecutor* executor = nullptr; uint64_t wsSize = 0; void* ws = nullptr; - ret = aclnnQuantLightningIndexerGetWorkspaceSize( - query, key, weights, qScale, kScale, - aslqT, aslkT, btT, metaT, - queryQuantMode, keyQuantMode, &layoutQuery[0], &layoutKey[0], - sparseCount, sparseMode, preToken, nextToken, cmpRatio, - returnValues, stride, scaleStride, - idxOut, valOut, &wsSize, &executor); - if (ret != ACL_SUCCESS) { printf("Main GetWorkspaceSize FAILED ret=%d\n", ret); return -1; } - printf("Main GetWorkspaceSize OK, ws=%lu\n", wsSize); - if (wsSize > 0) { ret = aclrtMalloc(&ws, wsSize, ACL_MEM_MALLOC_HUGE_FIRST); - if (ret != ACL_SUCCESS) { printf("malloc ws %d\n", ret); return -1; } } - ret = aclnnQuantLightningIndexer(ws, wsSize, executor, stream); - if (ret != ACL_SUCCESS) { printf("Main Execute FAILED ret=%d\n", ret); return -1; } - ret = aclrtSynchronizeStream(stream); - if (ret != ACL_SUCCESS) { printf("Main Sync %d\n", ret); return -1; } - - std::vector idxHost(idxCount); - ret = aclrtMemcpy(idxHost.data(), idxCount * sizeof(int32_t), idxDev, - idxCount * sizeof(int32_t),ACL_MEMCPY_DEVICE_TO_HOST); - if (ret != ACL_SUCCESS) { printf("memcpy out %d\n", ret); return -1; } - - const int32_t* golden = reinterpret_cast(goldenBuf.data()); - int total = 0, setMatch = 0; - for (int s = 0; s < qSeq; ++s) { - const int32_t* npuRow = &idxHost[s * sparseCount]; - const int32_t* gRow = &golden[s * sparseCount]; - printf("s=%d NPU =", s); for (int j = 0; j < sparseCount; ++j) printf(" %d", npuRow[j]); printf("\n"); - printf("s=%d golden=", s); for (int j = 0; j < sparseCount; ++j) printf(" %d", gRow[j]); printf("\n"); - for (int j = 0; j < sparseCount; ++j) { - total++; - for (int k = 0; k < sparseCount; ++k) if (npuRow[j] == gRow[k]) { setMatch++; break; } - } - } - printf("==== set-match %d/%d ====\n", setMatch, total); - printf(setMatch == total ? "RESULT: PASS\n" : "RESULT: FAIL\n"); - - aclDestroyTensor(query); aclDestroyTensor(key); aclDestroyTensor(weights); - aclDestroyTensor(qScale); aclDestroyTensor(kScale); aclDestroyTensor(btT); - aclDestroyTensor(idxOut); aclDestroyTensor(valOut); - aclDestroyTensor(aslqT); aclDestroyTensor(aslkT); aclDestroyTensor(metaT); - aclrtFree(qDev); aclrtFree(kDev); aclrtFree(wDev); aclrtFree(qsDev); aclrtFree(ksDev); - aclrtFree(btDev); aclrtFree(idxDev); aclrtFree(valDev); - aclrtFree(aslqDev); aclrtFree(aslkDev); aclrtFree(metaDev); - if (mWsPtr) aclrtFree(mWsPtr); - if (ws) aclrtFree(ws); - aclrtDestroyStream(stream); aclrtResetDevice(deviceId); aclFinalize(); - return 0; -} \ No newline at end of file diff --git a/test/python_test/test_quant_lightning_indexer.py b/test/python_test/test_quant_lightning_indexer.py new file mode 100644 index 0000000..7e5688b --- /dev/null +++ b/test/python_test/test_quant_lightning_indexer.py @@ -0,0 +1,230 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- +""" +pytest test for quant_lightning_indexer (PA_BSND paged indexer, two-stage metadata). +Golden: weighted dot-product score -> topk indices (set comparison). +NPU: metadata + main two-stage with paged key. + +Platform adaptation: + - A5 (ascend950): query/key fp8_e4m3fn, weights/scale fp32 + - 910C (Ascend910C): query/key int8, weights/scale fp16 +""" +import math +import os + +import numpy as np +import pytest +import torch +import torch_npu + +from custom_ops import quant_lightning_indexer_npu + +INVALID_IDX = -1 + + +def _get_platform(): + """Detect SOC platform: 'A5' or '910C' or 'unknown'.""" + soc = torch_npu._C._npu_get_soc_version() + # ascend950 -> A5, Ascend910C/Ascend910B -> 910C family + if soc in (200, ): # SOC_VERSION_ASCEND950 == 200 + return "A5" + # 910C soc version is typically 220 / 221 + return "910C" + + +PLATFORM = _get_platform() + + +# ============ Golden reference (CPU fp32 approximate) ============ + +def golden_quant_lightning_indexer(query_fp8_raw, key_bnsd_fp8_raw, + weights, q_scale, k_scale_bns, + act_seq_q, act_seq_k, + q_head_num, k_head_num, head_dim, + sparse_count, sparse_mode, cmp_ratio): + """ + Golden for quant_lightning_indexer (simplified from GeneralizedQLI.forward). + query_fp8_raw: [B, qSeq, qHead, hd] float32 (dequantized from fp8) + key_bnsd_fp8_raw: [B, kHead, kSeq, hd] float32 (dequantized from fp8) + weights: [B, qSeq, qHead]float32 + q_scale: [B, qSeq, qHead] float32 + k_scale_bns: [B, kHead, kSeq] float32 + act_seq_q: list[int] per-batch actual query length + act_seq_k: list[int] per-batch actual key length (before cmp_ratio) + Returns: indices [B, qSeq, kHead, sparse_count] int32 + """ + B = query_fp8_raw.shape[0] + q_seq = query_fp8_raw.shape[1] + k_seq = key_bnsd_fp8_raw.shape[2] + + out = np.full((B, q_seq, k_head_num, sparse_count), INVALID_IDX, dtype=np.int32) + + for b in range(B): + actual_q = act_seq_q[b] + actual_k = int(math.floor(act_seq_k[b] / cmp_ratio)) + for s1 in range(actual_q): + for kh in range(k_head_num): + # Score = sum over qHeads of: + # weights[b,s1,n] * q_scale[b,s1,n] * (q[b,s1,n,:] . k[b,kh,s2,:]) * k_scale[b,kh,s2] + # sparse_mode=3: causal -> valid s2 in [0, s1] (but since q_seq < k_seq, use all actual_k) + if sparse_mode == 3: + valid_k = actual_k # decode: all cached keys visible + else: + valid_k = actual_k + + if valid_k <= 0: + continue + + # q: [qHead, hd], k: [valid_k, hd] + q = query_fp8_raw[b, s1, :, :] # [qHead, hd] + k = key_bnsd_fp8_raw[b, kh, :valid_k, :] # [valid_k, hd] + + # dot product: [qHead, valid_k] + dots = q @ k.T # [qHead, valid_k] + + # per-head weighting + w = weights[b, s1, :] # [qHead] + qs = q_scale[b, s1, :] # [qHead] + ks = k_scale_bns[b, kh, :valid_k] # [valid_k] + + # score[s2] = sum_n(w[n] * qs[n] * dots[n, s2]) * ks[s2] + coeff = w * qs # [qHead] + score = (coeff[:, None] * dots).sum(axis=0) * ks # [valid_k] + + # topk descending + order = np.argsort(-score) + take = min(sparse_count, valid_k) + out[b, s1, kh, :take] = order[:take].astype(np.int32) + + return out + + +def to_paged_key(key_bnsd, k_scale_bns, B, k_head_num, k_seq, head_dim, block_size): + """ + key_bnsd: [B, kHead, kSeq, hd] -> key_pa[block_num, block_size, kHead, hd] + block_table + k_scale_pa + """ + blocks_per_b = (k_seq + block_size - 1) // block_size + block_num = B * blocks_per_b + key_pa = np.zeros((block_num, block_size, k_head_num, head_dim), dtype=key_bnsd.dtype) + kscale_pa = np.zeros((block_num, block_size, k_head_num), dtype=k_scale_bns.dtype) + block_table = np.zeros((B, blocks_per_b), dtype=np.int32) + + for b in range(B): + for blk in range(blocks_per_b): + gblk = b * blocks_per_b + blk + block_table[b, blk] = gblk + for tok in range(block_size): + t = blk * block_size + tok + if t < k_seq: + for kh in range(k_head_num): + key_pa[gblk, tok, kh, :] = key_bnsd[b, kh, t, :] + kscale_pa[gblk, tok, kh] = k_scale_bns[b, kh, t] + return key_pa, kscale_pa, block_table + + +def _valid_set(row): + return set(int(x) for x in row if int(x) != INVALID_IDX) + + +# ============ Test cases ============ + +CASES = [ + # (B, q_seq, k_seq, q_head, k_head, head_dim, block_size, sparse_count, sparse_mode, cmp_ratio) + (1, 4, 128, 64, 1, 128, 128, 8, 3, 1), +] + + +@pytest.mark.parametrize("B,Q_SEQ,K_SEQ,Q_HEAD,K_HEAD,HD,BLOCK_SIZE,SPARSE_COUNT,SPARSE_MODE,CMP_RATIO", CASES) +def test_quant_lightning_indexer(B, Q_SEQ, K_SEQ, Q_HEAD, K_HEAD, HD, BLOCK_SIZE, SPARSE_COUNT, SPARSE_MODE, CMP_RATIO): + np.random.seed(2026) + torch.manual_seed(2026) + + # Platform-dependent dtype selection + if PLATFORM == "A5": + qk_dtype = torch.float8_e4m3fn # A5: fp8 query/key + scale_dtype = torch.float32 # A5: fp32 weights/scale + else: + qk_dtype = torch.int8 # 910C: int8 query/key + scale_dtype = torch.float16 # 910C: fp16 weights/scale + + # Generate data via float -> quantize -> dequantize for golden + query_f32 = np.random.uniform(-100, 100, (B, Q_SEQ, Q_HEAD, HD)).astype(np.float32) + key_bnsd_f32 = np.random.uniform(-100, 100, (B, K_HEAD, K_SEQ, HD)).astype(np.float32) + weights_f32 = np.random.uniform(-25, 25, (B, Q_SEQ, Q_HEAD)).astype(np.float32) + q_scale_f32 = np.random.uniform(0, 255, (B, Q_SEQ, Q_HEAD)).astype(np.float32) + k_scale_f32 = np.random.uniform(0, 65504, (B, K_HEAD, K_SEQ)).astype(np.float32) + + if qk_dtype == torch.float8_e4m3fn: + # fp8: quantize-dequantize to get actual representable values + query_quant = torch.from_numpy(query_f32).to(torch.float8_e4m3fn) + key_quant = torch.from_numpy(key_bnsd_f32).to(torch.float8_e4m3fn) + query_deq = query_quant.float().numpy() + key_deq = key_quant.float().numpy() + else: + # int8: clamp + round to [-127, 127] + query_i8 = np.clip(np.round(query_f32), -127, 127).astype(np.int8) + key_i8 = np.clip(np.round(key_bnsd_f32), -127, 127).astype(np.int8) + query_quant = torch.from_numpy(query_i8) + key_quant = torch.from_numpy(key_i8) # [B, K_HEAD, K_SEQ, HD] + query_deq = query_i8.astype(np.float32) + key_deq = key_i8.astype(np.float32) + + # Scale dtype cast for golden (both paths use fp32 golden math) + if scale_dtype == torch.float16: + # Simulate fp16 precision loss in weights/scale + weights_f32 = torch.from_numpy(weights_f32).to(torch.float16).float().numpy() + q_scale_f32 = torch.from_numpy(q_scale_f32).to(torch.float16).float().numpy() + k_scale_f32 = torch.from_numpy(k_scale_f32).to(torch.float16).float().numpy() + + act_seq_q = [Q_SEQ] * B + act_seq_k = [K_SEQ * CMP_RATIO] * B + + # Golden (CPU) + golden_idx = golden_quant_lightning_indexer( + query_deq, key_deq, weights_f32, q_scale_f32, k_scale_f32, + act_seq_q, act_seq_k, Q_HEAD, K_HEAD, HD, + SPARSE_COUNT, SPARSE_MODE, CMP_RATIO) + + # Page the key and k_scale for NPU + key_pa, kscale_pa, block_table = to_paged_key( + key_deq, k_scale_f32, B, K_HEAD, K_SEQ, HD, BLOCK_SIZE) + + # Prepare NPU tensors + query_npu = query_quant.npu() + if qk_dtype == torch.float8_e4m3fn: + key_npu = torch.from_numpy(key_pa).to(torch.float8_e4m3fn).npu() + else: + key_npu = torch.from_numpy(key_pa.astype(np.int8)).npu() + weights_npu = torch.from_numpy(weights_f32).to(scale_dtype).npu() + q_scale_npu = torch.from_numpy(q_scale_f32).to(scale_dtype).npu() + k_scale_npu = torch.from_numpy(kscale_pa).to(scale_dtype).npu() + aslq_npu = torch.tensor(act_seq_q, dtype=torch.int32).npu() + aslk_npu = torch.tensor(act_seq_k, dtype=torch.int32).npu() + bt_npu = torch.from_numpy(block_table).int().npu() + + # Run NPU + npu_out = quant_lightning_indexer_npu( + query_npu, key_npu, weights_npu, q_scale_npu, k_scale_npu, + aslq_npu, aslk_npu, bt_npu, + num_heads_q=Q_HEAD, num_heads_k=K_HEAD, head_dim=HD, + query_quant_mode=0, key_quant_mode=0, + layout_query="BSND", layout_key="PA_BSND", + sparse_count=SPARSE_COUNT, sparse_mode=SPARSE_MODE, + cmp_ratio=CMP_RATIO) + npu_result = npu_out.cpu().numpy() + + # Compare using set overlap (low precision causes topk tie-break differences) + assert npu_result.shape == golden_idx.shape, f"Shape mismatch: {npu_result.shape} vs {golden_idx.shape}" + for b in range(B): + for s1 in range(Q_SEQ): + for kh in range(K_HEAD): + got = _valid_set(npu_result[b, s1, kh]) + exp = _valid_set(golden_idx[b, s1, kh]) + overlap = len(got & exp) + min_match = max(1, int(len(exp) * 0.5)) + assert overlap >= min_match, ( + f"b={b} s1={s1} kh={kh}: got {sorted(got)} exp {sorted(exp)}") + + +if __name__ == "__main__": + pytest.main([__file__, "-v", "-s"]) \ No newline at end of file diff --git a/test/python_test/test_quant_lightning_indexer_metadata.py b/test/python_test/test_quant_lightning_indexer_metadata.py new file mode 100644 index 0000000..ca03f03 --- /dev/null +++ b/test/python_test/test_quant_lightning_indexer_metadata.py @@ -0,0 +1,119 @@ +# Copyright 2025 The xLLM Authors. All Rights Reserved. +# AICPU metadata op test for quant_lightning_indexer_metadata. +# Verifies: no crash, correct output shape/dtype, output self-consistency. + +from dataclasses import dataclass + +import pytest +import torch +import torch_npu # noqa: F401 +import custom_ops + +META_SIZE = 1024 + + +@dataclass(frozen=True) +class QuantLightningIndexerMetadataCase: + name: str + batch_size: int + max_seq_q: int + max_seq_k: int + num_heads_q: int + num_heads_k: int + head_dim: int + query_quant_mode: int = 0 + key_quant_mode: int = 0 + layout_query: str = "BSND" + layout_key: str = "PA_BSND" + sparse_count: int = 8 + sparse_mode: int = 3 + pre_token: int = 9223372036854775807 + next_token: int = 9223372036854775807 + cmp_ratio: int = 1 + + +CASES = [ + QuantLightningIndexerMetadataCase( + name="decode_b1_sq1_sk128_n64_nk1_d128", + batch_size=1, max_seq_q=1, max_seq_k=128, + num_heads_q=64, num_heads_k=1, head_dim=128, + ), + QuantLightningIndexerMetadataCase( + name="decode_b4_sq1_sk256_n64_nk1_d128", + batch_size=4, max_seq_q=1, max_seq_k=256, + num_heads_q=64, num_heads_k=1, head_dim=128, + ), + QuantLightningIndexerMetadataCase( + name="decode_b1_sq1_sk512_n64_nk1_d128_sparse16", + batch_size=1, max_seq_q=1, max_seq_k=512, + num_heads_q=64, num_heads_k=1, head_dim=128, + sparse_count=16, + ), + QuantLightningIndexerMetadataCase( + name="decode_b2_sq1_sk1024_n64_nk1_d128", + batch_size=2, max_seq_q=1, max_seq_k=1024, + num_heads_q=64, num_heads_k=1, head_dim=128, + ), + QuantLightningIndexerMetadataCase( + name="decode_b1_sq4_sk128_n64_nk1_d128_prefill", + batch_size=1, max_seq_q=4, max_seq_k=128, + num_heads_q=64, num_heads_k=1, head_dim=128, + ), +] + + +def _make_inputs(case: QuantLightningIndexerMetadataCase): + """Create actual_seq_lengths tensors on NPU.""" + B = case.batch_size + device = "npu" + # actual_seq_lengths_query: (B,) int32, each = max_seq_q + aslq = torch.full((B,), case.max_seq_q, dtype=torch.int32, device=device) + # actual_seq_lengths_key: (B,) int32, each = max_seq_k + aslk = torch.full((B,), case.max_seq_k, dtype=torch.int32, device=device) + return aslq, aslk + + +@pytest.mark.parametrize("case", CASES, ids=[c.name for c in CASES]) +def test_quant_lightning_indexer_metadata(case: QuantLightningIndexerMetadataCase): + torch.npu.set_device(0) + aslq, aslk = _make_inputs(case) + + meta_t = custom_ops.quant_lightning_indexer_metadata_npu( + aslq, aslk, + num_heads_q=case.num_heads_q, num_heads_k=case.num_heads_k, + head_dim=case.head_dim, + query_quant_mode=case.query_quant_mode, + key_quant_mode=case.key_quant_mode, + batch_size=case.batch_size, + max_seq_q=case.max_seq_q, max_seq_k=case.max_seq_k, + layout_query=case.layout_query, layout_key=case.layout_key, + sparse_count=case.sparse_count, sparse_mode=case.sparse_mode, + pre_token=case.pre_token, next_token=case.next_token, + cmp_ratio=case.cmp_ratio, + ) + torch.npu.synchronize() + + # Verify output shape and dtype + assert meta_t.shape == (META_SIZE,), f"Expected shape ({META_SIZE},), got {meta_t.shape}" + assert meta_t.dtype == torch.int32, f"Expected dtype int32, got {meta_t.dtype}" + + # Verify determinism: call again and compare + meta_t2 = custom_ops.quant_lightning_indexer_metadata_npu( + aslq, aslk, + num_heads_q=case.num_heads_q, num_heads_k=case.num_heads_k, + head_dim=case.head_dim, + query_quant_mode=case.query_quant_mode, + key_quant_mode=case.key_quant_mode, + batch_size=case.batch_size, + max_seq_q=case.max_seq_q, max_seq_k=case.max_seq_k, + layout_query=case.layout_query, layout_key=case.layout_key, + sparse_count=case.sparse_count, sparse_mode=case.sparse_mode, + pre_token=case.pre_token, next_token=case.next_token, + cmp_ratio=case.cmp_ratio, + ) + torch.npu.synchronize() + assert torch.equal(meta_t.cpu(), meta_t2.cpu()), "Metadata output is not deterministic" + + +if __name__ == "__main__": + pytest.main([__file__, "-v", "-x"]) \ No newline at end of file diff --git a/test/python_test/test_sparse_attn_sharedkv.cpp b/test/python_test/test_sparse_attn_sharedkv.cpp deleted file mode 100644 index 1e43b9a..0000000 --- a/test/python_test/test_sparse_attn_sharedkv.cpp +++ /dev/null @@ -1,213 +0,0 @@ -#include -#include -#include -#include -#include -#include -#include -#include -#include "acl/acl.h" -#include "aclnn_sparse_attn_sharedkv.h" -#include "aclnn_sparse_attn_sharedkv_metadata.h" - -// ---- Case (SWA single-path, consistent with sparse_attn_sharedkv_gen.py) ---- -static const int64_t B = 1; -static const int64_t S1 = 4; -static const int64_t S2 = 16; -static const int64_t N1 = 64; // qHead hard constraint =64 -static const int64_t KV_N = 1; // kvHead hard constraint =1 -static const int64_t D = 512; // headDim hard constraint =512 -static const int64_t BLOCK_SIZE = 16; -static const int64_t BLOCK_NUM = 1; // B*ceil(S2/BLOCK_SIZE)=1 -static const int64_t MAX_BLOCKS = 1; -static const int64_t ORI_MASK_MODE = 4; -static const int64_t CMP_MASK_MODE = 3; -static const int64_t ORI_WIN_LEFT = 127; -static const int64_t ORI_WIN_RIGHT = 0; -static const int64_t CMP_RATIO = 1; -static const int64_t ORI_KV_STRIDE = 0; -static const int64_t CMP_KV_STRIDE = 0; -static const int64_t ORI_TOP_K = 512; -static const int64_t CMP_TOP_K = 512; -static const double SOFTMAX_SCALE = 1.0 / 22.627416997969522; // 1/sqrt(512) -static std::string layoutQ = "BSND"; -static std::string layoutKv = "PA_ND"; -static const int64_t SAS_META_SIZE = 1024; -static const char* DATA_DIR = "/tmp/sparse_attn_sharedkv_data/"; - -// dtype: 0=fp16, 1=bf16 (determined by command line argv[1]; default fp16) -static aclDataType g_dtype = ACL_FLOAT16; - -static std::vector ReadBin(const std::string &name) { - std::ifstream f(std::string(DATA_DIR) + name, std::ios::binary | std::ios::ate); - if (!f) { printf("open %s FAILED\n", name.c_str()); return {}; } - std::streamsize sz = f.tellg(); f.seekg(0, std::ios::beg); - std::vector buf(sz); f.read(buf.data(), sz); return buf; -} - -std::tuple CreateTensor(size_t size, std::vector shape, - aclDataType dType, const void* hostData = nullptr) { - std::vector strides(shape.size(), 1); - for (int i = (int)shape.size() - 2; i >= 0; --i) strides[i] = strides[i + 1] * shape[i + 1]; - void* dev = nullptr; - auto ret = aclrtMalloc(&dev, size, ACL_MEM_MALLOC_HUGE_FIRST); - if (ret != ACL_SUCCESS) { printf("aclrtMalloc %d\n", ret); return {nullptr, nullptr}; } - aclTensor* t = aclCreateTensor(shape.data(), shape.size(), dType, strides.data(), 0, - aclFormat::ACL_FORMAT_ND, shape.data(), shape.size(), dev); - if (t == nullptr) { aclrtFree(dev); return {nullptr, nullptr}; } - if (hostData) aclrtMemcpy(dev, size, hostData, size, ACL_MEMCPY_HOST_TO_DEVICE); - return {t, dev}; -} - -// 16-bit(fp16/bf16) bit pattern -> float32 -static float ToF32(uint16_t v, bool isBf16) { - uint32_t u; - if (isBf16) { - u = static_cast(v) << 16; - } else { - uint32_t sign = (v >> 15) & 0x1; - uint32_t exp = (v >> 10) & 0x1f; - uint32_t man = v & 0x3ff; - if (exp == 0) { - if (man == 0) { u = sign << 31; } - else { - exp = 127 - 15 + 1; - while ((man & 0x400) == 0) { man <<= 1; exp--; } - man &= 0x3ff; - u = (sign << 31) | (exp << 23) | (man << 13); - } - } else if (exp == 0x1f) { - u = (sign << 31) | (0xff << 23) | (man << 13); - } else { - u = (sign << 31) | ((exp - 15 + 127) << 23) | (man << 13); - } - } - float f; memcpy(&f, &u, 4); return f; -} - -int main(int argc, char** argv) { - bool isBf16 = (argc > 1 && std::string(argv[1]) == "bf16"); - g_dtype = isBf16 ? ACL_BF16 : ACL_FLOAT16; - size_t elemSz = 2; // fp16/bf16 are both 2 bytes - - int32_t deviceId = 0; aclrtStream stream; - aclError ret = aclInit(nullptr); - if (ret != ACL_SUCCESS) { printf("aclInit %d\n", ret); return -1; } - ret = aclrtSetDevice(deviceId); - if (ret != ACL_SUCCESS) { printf("aclrtSetDevice %d\n", ret); return -1; } - ret = aclrtCreateStream(&stream); - if (ret != ACL_SUCCESS) { printf("aclrtCreateStream %d\n", ret); return -1; } - - auto qBuf = ReadBin("q.bin"); - auto kvBuf = ReadBin("kv_pa.bin"); - auto sinksBuf = ReadBin("sinks.bin"); - auto btBuf = ReadBin("block_table.bin"); - auto susedBuf = ReadBin("seqused_kv.bin"); - auto goldenBuf = ReadBin("golden_out.bin"); - if (qBuf.empty() || kvBuf.empty() || goldenBuf.empty() || btBuf.empty() || susedBuf.empty()) { - printf("read bin FAILED\n"); return -1; - } - - // ================= Step 1: metadata op (25 params) ================= - // wrapper unconditionally applies Contiguous to optional inputs, nullptr will crash; pass non-null cu_seqlens/seqused tensors - // cu_seqlens prefix sum: length B+1, [0, S1](q), [0, S2](kv); seqused: length B - std::vector cuSeqQHost(B + 1, 0), cuSeqOriKvHost(B + 1, 0), cuSeqCmpKvHost(B + 1, 0); - for (int64_t i = 0; i < B; ++i) { cuSeqQHost[i + 1] = cuSeqQHost[i] + (int32_t)S1; - cuSeqOriKvHost[i + 1] = cuSeqOriKvHost[i] + (int32_t)S2; - cuSeqCmpKvHost[i + 1] = cuSeqCmpKvHost[i] + (int32_t)S2; } - std::vector susedQHost(B, (int32_t)S1); - aclTensor *cuSeqQT_m, *cuSeqOriKvT_m, *cuSeqCmpKvT_m, *susedQT_m, *susedKvT_m, *metaT; - void *cuSeqQDev_m, *cuSeqOriKvDev_m, *cuSeqCmpKvDev_m, *susedQDev_m, *susedKvDev_m, *metaDev; - std::tie(cuSeqQT_m, cuSeqQDev_m) = CreateTensor((B + 1) * sizeof(int32_t), {B + 1}, ACL_INT32, cuSeqQHost.data()); - std::tie(cuSeqOriKvT_m, cuSeqOriKvDev_m) = CreateTensor((B + 1) * sizeof(int32_t), {B + 1}, ACL_INT32, cuSeqOriKvHost.data()); - std::tie(cuSeqCmpKvT_m, cuSeqCmpKvDev_m) = CreateTensor((B + 1) * sizeof(int32_t), {B + 1}, ACL_INT32, cuSeqCmpKvHost.data()); - std::tie(susedQT_m, susedQDev_m) = CreateTensor(B * sizeof(int32_t), {B}, ACL_INT32, susedQHost.data()); - std::tie(susedKvT_m, susedKvDev_m) = CreateTensor(susedBuf.size(), {B}, ACL_INT32, susedBuf.data()); - std::tie(metaT, metaDev) = CreateTensor(SAS_META_SIZE * sizeof(int32_t), {SAS_META_SIZE}, ACL_INT32); - - aclOpExecutor* mExec = nullptr; uint64_t mWs = 0; void* mWsPtr = nullptr; - ret = aclnnSparseAttnSharedkvMetadataGetWorkspaceSize( - cuSeqQT_m, cuSeqOriKvT_m, cuSeqCmpKvT_m, susedQT_m, susedKvT_m, - N1, KV_N, D, B, S1, S2, - ORI_TOP_K, CMP_TOP_K, CMP_RATIO, - ORI_MASK_MODE, CMP_MASK_MODE, ORI_WIN_LEFT, ORI_WIN_RIGHT, - &layoutQ[0], &layoutKv[0], true, false, - metaT, &mWs, &mExec); - if (ret != ACL_SUCCESS) { printf("Metadata GetWorkspaceSize FAILED ret=%d\n", ret); return -1; } - if (mWs > 0) { ret = aclrtMalloc(&mWsPtr, mWs, ACL_MEM_MALLOC_HUGE_FIRST); - if (ret != ACL_SUCCESS) { printf("malloc mWs %d\n", ret); return -1; } } - ret = aclnnSparseAttnSharedkvMetadata(mWsPtr, mWs, mExec, stream); - if (ret != ACL_SUCCESS) { printf("Metadata Execute FAILED ret=%d\n", ret); return -1; } - ret = aclrtSynchronizeStream(stream); - if (ret != ACL_SUCCESS) { printf("Metadata Sync %d\n", ret); return -1; } - printf("Metadata OK\n"); - - // ================= Step 2: main op (27 params) ================= - aclTensor *query, *oriKv, *oriBt, *susedKvT, *sinksT, *attnOut, *lseOut; - void *qDev, *kvDev, *btDev, *susedKvDev, *sinksDev, *outDev, *lseDev; - std::tie(query, qDev) = CreateTensor(qBuf.size(), {B, S1, N1, D}, g_dtype, qBuf.data()); - // paged ori_kv: [blockNum, blockSize, kvHead=1, D] - std::tie(oriKv, kvDev) = CreateTensor(kvBuf.size(), {BLOCK_NUM, BLOCK_SIZE, KV_N, D}, g_dtype, kvBuf.data()); - std::tie(oriBt, btDev) = CreateTensor(btBuf.size(), {B, MAX_BLOCKS}, ACL_INT32, btBuf.data()); - std::tie(susedKvT, susedKvDev) = CreateTensor(susedBuf.size(), {B}, ACL_INT32, susedBuf.data()); - std::tie(sinksT, sinksDev) = CreateTensor(sinksBuf.size(), {N1}, ACL_FLOAT, sinksBuf.data()); - - int64_t outElems = B * S1 * N1 * D; - std::tie(attnOut, outDev) = CreateTensor(outElems * elemSz, {B, S1, N1, D}, g_dtype); - int64_t lseElems = B * S1 * N1; - std::tie(lseOut, lseDev) = CreateTensor(lseElems * sizeof(float), {B, S1, N1}, ACL_FLOAT); - - aclOpExecutor* executor = nullptr; uint64_t wsSize = 0; void* ws = nullptr; - // 27 params: q, oriKv, cmpKv, oriSparseIdx, cmpSparseIdx, oriBt, cmpBt, - // cuSeqQ, cuSeqOriKv, cuSeqCmpKv, sequsedQ, sequsedKv, sinks, metadata, - // softmaxScale, cmpRatio, oriMaskMode, cmpMaskMode, oriKvStride, cmpKvStride, - // oriWinLeft, oriWinRight, layoutQ, layoutKv, returnSoftmaxLse, attnOut, lseOut - ret = aclnnSparseAttnSharedkvGetWorkspaceSize( - query, oriKv, nullptr, nullptr, nullptr, oriBt, nullptr, - nullptr, nullptr, nullptr, nullptr, susedKvT, sinksT, metaT, - SOFTMAX_SCALE, CMP_RATIO, ORI_MASK_MODE, CMP_MASK_MODE, ORI_KV_STRIDE, CMP_KV_STRIDE, - ORI_WIN_LEFT, ORI_WIN_RIGHT, &layoutQ[0], &layoutKv[0], false, - attnOut, lseOut, &wsSize, &executor); - if (ret != ACL_SUCCESS) { printf("Main GetWorkspaceSize FAILED ret=%d\n", ret); return -1; } - printf("Main GetWorkspaceSize OK, ws=%lu\n", wsSize); - if (wsSize > 0) { ret = aclrtMalloc(&ws, wsSize, ACL_MEM_MALLOC_HUGE_FIRST); - if (ret != ACL_SUCCESS) { printf("malloc ws %d\n", ret); return -1; } } - ret = aclnnSparseAttnSharedkv(ws, wsSize, executor, stream); - if (ret != ACL_SUCCESS) { printf("Main Execute FAILED ret=%d\n", ret); return -1; } - ret = aclrtSynchronizeStream(stream); - if (ret != ACL_SUCCESS) { printf("Main Sync %d\n", ret); return -1; } - - std::vector outHost(outElems); - ret = aclrtMemcpy(outHost.data(), outElems * elemSz, outDev, outElems * elemSz, ACL_MEMCPY_DEVICE_TO_HOST); - if (ret != ACL_SUCCESS) { printf("memcpy out %d\n", ret); return -1; } - - const uint16_t* golden = reinterpret_cast(goldenBuf.data()); - double maxAbs = 0.0, maxRel = 0.0; - int failCnt = 0; - double atol = 2e-2, rtol = 2e-2; - for (int64_t i = 0; i < outElems; ++i) { - float g = ToF32(golden[i], isBf16); - float o = ToF32(outHost[i], isBf16); - double ad = std::fabs(g - o); - double rd = ad / (std::fabs(g) + 1e-6); - if (ad > maxAbs) maxAbs = ad; - if (rd > maxRel) maxRel = rd; - if (ad > atol && rd > rtol) failCnt++; - } - printf("==== outElems=%ld maxAbs=%.5f maxRel=%.5f fail=%d ====\n", (long)outElems, maxAbs, maxRel, failCnt); - printf(failCnt == 0 ? "RESULT: PASS\n" : "RESULT: FAIL\n"); - - aclDestroyTensor(query); aclDestroyTensor(oriKv); aclDestroyTensor(oriBt); - aclDestroyTensor(susedKvT); aclDestroyTensor(sinksT); aclDestroyTensor(attnOut); aclDestroyTensor(lseOut); - aclDestroyTensor(susedKvT_m); aclDestroyTensor(metaT); - aclDestroyTensor(cuSeqQT_m); aclDestroyTensor(cuSeqOriKvT_m); - aclDestroyTensor(cuSeqCmpKvT_m); aclDestroyTensor(susedQT_m); - aclrtFree(qDev); aclrtFree(kvDev); aclrtFree(btDev); aclrtFree(susedKvDev); - aclrtFree(sinksDev); aclrtFree(outDev); aclrtFree(lseDev); - aclrtFree(susedKvDev_m); aclrtFree(metaDev); - aclrtFree(cuSeqQDev_m); aclrtFree(cuSeqOriKvDev_m); - if (mWsPtr) aclrtFree(mWsPtr); - if (ws) aclrtFree(ws); - aclrtDestroyStream(stream); aclrtResetDevice(deviceId); aclFinalize(); - return 0; -} \ No newline at end of file diff --git a/test/python_test/test_sparse_attn_sharedkv.py b/test/python_test/test_sparse_attn_sharedkv.py new file mode 100644 index 0000000..fb61794 --- /dev/null +++ b/test/python_test/test_sparse_attn_sharedkv.py @@ -0,0 +1,127 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- +""" +pytest test for sparse_attn_sharedkv (SWA single-path + attention sink). +Golden: SWA sliding-window attention + sink term, unpaged reference. +NPU: PA_ND paged kv + metadata two-stage. +""" +import numpy as np +import pytest +import torch + +from custom_ops import sparse_attn_sharedkv_npu + +# ============ Golden reference (SWA + sink, numpy CPU float32) ============ + +def golden_swa(q, kv, sinks, B, S1, S2, N1, KV_N, D, win_left, win_right, scale): + """ + q: [B, S1, N1, D] float32 + kv: [B, S2, KV_N, D] float32 (KV_N=1 broadcast) + sinks:[N1] float32 + return: out [B, S1, N1, D] float32 + """ + out = np.zeros((B, S1, N1, D), dtype=np.float32) + for b in range(B): + for s1_idx in range(S1): + diag = S2 - S1 + s1_idx + mask_right = diag + win_right + mask_left = max(diag - win_left, 0) + if mask_right < mask_left: + continue + j_range = np.arange(mask_left, mask_right + 1) + k = kv[b, j_range, 0, :] # [J, D] + for n1 in range(N1): + logit = (q[b, s1_idx, n1, :] @ k.T) * scale # [J] + sink = sinks[n1] + m = max(float(logit.max()), float(sink)) + exp_l = np.exp(logit - m) + exp_s = np.exp(sink - m) + denom = float(exp_l.sum()) + float(exp_s) + p = exp_l / denom + out[b, s1_idx, n1, :] = p @ k + return out + + +def to_paged(kv, B, S2, KV_N, D, block_size): + """ + kv: [B, S2, KV_N, D] -> kv_pa[block_num, block_size, KV_N, D] + block_table[B, max_blocks] + """ + blocks_per_b = (S2 + block_size - 1) // block_size + block_num = B * blocks_per_b + kv_pa = np.zeros((block_num, block_size, KV_N, D), dtype=kv.dtype) + block_table = np.zeros((B, blocks_per_b), dtype=np.int32) + for b in range(B): + for blk in range(blocks_per_b): + gblk = b * blocks_per_b + blk + block_table[b, blk] = gblk + for tok in range(block_size): + t = blk * block_size + tok + if t < S2: + kv_pa[gblk, tok, :, :] = kv[b, t, :, :] + return kv_pa, block_table + + +# ============ Test cases ============ + +CASES = [ + # (B, S1, S2, N1, KV_N, D, block_size, win_left, win_right, dtype_str) + (1, 4, 16, 64, 1, 512, 16, 127, 0, "fp16"), +] + + +@pytest.mark.parametrize("B,S1,S2,N1,KV_N,D,BLOCK_SIZE,WIN_LEFT,WIN_RIGHT,dtype_str", CASES) +def test_sparse_attn_sharedkv(B, S1, S2, N1, KV_N, D, BLOCK_SIZE, WIN_LEFT, WIN_RIGHT, dtype_str): + rng = np.random.default_rng(1234) + scale = 1.0 / np.sqrt(D) + + # Generate data in float32 + q_f32 = rng.standard_normal((B, S1, N1, D)).astype(np.float32) * 0.1 + kv_f32 = rng.standard_normal((B, S2, KV_N, D)).astype(np.float32) * 0.1 + sinks_f32 = rng.standard_normal((N1,)).astype(np.float32) * 0.1 + seqused_kv = np.array([S2] * B, dtype=np.int32) + + # Cast to target dtype for quantization alignment + torch_dtype = torch.float16 if dtype_str == "fp16" else torch.bfloat16 + q_t = torch.from_numpy(q_f32).to(torch_dtype) + kv_t = torch.from_numpy(kv_f32).to(torch_dtype) + + # Golden uses the quantized values (fp16/bf16 -> fp32) + q_ref = q_t.float().numpy() + kv_ref = kv_t.float().numpy() + golden_out = golden_swa(q_ref, kv_ref, sinks_f32, B, S1, S2, N1, KV_N, D, + WIN_LEFT, WIN_RIGHT, scale) + # Cast golden to target dtype for comparison + golden_out_t = torch.from_numpy(golden_out).to(torch_dtype).float().numpy() + + # Page the kv + kv_np = kv_t.float().numpy() + kv_pa, block_table = to_paged(kv_np, B, S2, KV_N, D, BLOCK_SIZE) + kv_pa_t = torch.from_numpy(kv_pa).to(torch_dtype) + + # Prepare NPU inputs + query_npu = q_t.npu() + ori_kv_npu = kv_pa_t.npu() + ori_bt_npu = torch.from_numpy(block_table).int().npu() + seqused_kv_npu = torch.from_numpy(seqused_kv).int().npu() + sinks_npu = torch.from_numpy(sinks_f32).float().npu() + + # Run NPU + npu_out = sparse_attn_sharedkv_npu( + query_npu, ori_kv_npu, ori_bt_npu, seqused_kv_npu, sinks_npu, + n1=N1, kv_n=KV_N, d=D, s1=S1, s2=S2, + ori_mask_mode=4, cmp_mask_mode=3, + ori_win_left=WIN_LEFT, ori_win_right=WIN_RIGHT, + softmax_scale=scale, cmp_ratio=1, + layout_q="BSND", layout_kv="PA_ND" + ) + npu_result = npu_out.cpu().float().numpy() + + # Compare + atol = 2e-2 + rtol = 2e-2 + np.testing.assert_allclose(npu_result, golden_out_t, atol=atol, rtol=rtol, + err_msg="sparse_attn_sharedkv NPU vs golden mismatch") + + +if __name__ == "__main__": + pytest.main([__file__, "-v", "-s"]) \ No newline at end of file diff --git a/test/python_test/test_sparse_attn_sharedkv_metadata.py b/test/python_test/test_sparse_attn_sharedkv_metadata.py new file mode 100644 index 0000000..d570296 --- /dev/null +++ b/test/python_test/test_sparse_attn_sharedkv_metadata.py @@ -0,0 +1,134 @@ +# Copyright 2025 The xLLM Authors. All Rights Reserved. +# AICPU metadata op test for sparse_attn_sharedkv_metadata. +# Verifies: no crash, correct output shape/dtype, output self-consistency. + +from dataclasses import dataclass + +import pytest +import torch +import torch_npu # noqa: F401 +import custom_ops + +META_SIZE = 1024 + + +@dataclass(frozen=True) +class SparseAttnSharedkvMetadataCase: + name: str + batch_size: int + max_seq_q: int + max_seq_kv: int + num_heads_q: int + num_heads_kv: int + head_dim: int + ori_mask_mode: int = 4 + cmp_mask_mode: int = 3 + ori_top_k: int = 0 + cmp_top_k: int = 0 + cmp_ratio: int = 1 + ori_win_left: int = 127 + ori_win_right: int = 0 + layout_q: str = "BSND" + layout_kv: str = "PA_ND" + has_ori_kv: bool = True + has_cmp_kv: bool = False + + +CASES = [ + SparseAttnSharedkvMetadataCase( + name="decode_b1_s1_kv16_n64_kvn1_d512", + batch_size=1, max_seq_q=1, max_seq_kv=16, + num_heads_q=64, num_heads_kv=1, head_dim=512, + ), + SparseAttnSharedkvMetadataCase( + name="decode_b4_s1_kv128_n64_kvn1_d512", + batch_size=4, max_seq_q=1, max_seq_kv=128, + num_heads_q=64, num_heads_kv=1, head_dim=512, + ), + SparseAttnSharedkvMetadataCase( + name="prefill_b1_s4_kv16_n64_kvn1_d512", + batch_size=1, max_seq_q=4, max_seq_kv=16, + num_heads_q=64, num_heads_kv=1, head_dim=512, + ), + SparseAttnSharedkvMetadataCase( + name="decode_b2_s1_kv256_n64_kvn1_d512_swa", + batch_size=2, max_seq_q=1, max_seq_kv=256, + num_heads_q=64, num_heads_kv=1, head_dim=512, + ori_mask_mode=4, ori_win_left=127, ori_win_right=0, + ), + SparseAttnSharedkvMetadataCase( + name="decode_b1_s1_kv64_n64_kvn1_d512_has_cmp", + batch_size=1, max_seq_q=1, max_seq_kv=64, + num_heads_q=64, num_heads_kv=1, head_dim=512, + has_ori_kv=True, has_cmp_kv=True, cmp_ratio=4, + ), +] + + +def _make_inputs(case: SparseAttnSharedkvMetadataCase): + """Create input tensors for metadata op on NPU.""" + B = case.batch_size + device = "npu" + # cu_seq_lens: cumulative sequence lengths, shape (B+1,) + seq_q_lens = torch.full((B,), case.max_seq_q, dtype=torch.int32) + cu_seq_q = torch.zeros(B + 1, dtype=torch.int32) + cu_seq_q[1:] = torch.cumsum(seq_q_lens, dim=0) + + seq_kv_lens = torch.full((B,), case.max_seq_kv, dtype=torch.int32) + cu_seq_ori_kv = torch.zeros(B + 1, dtype=torch.int32) + cu_seq_ori_kv[1:] = torch.cumsum(seq_kv_lens, dim=0) + + cu_seq_cmp_kv = torch.zeros(B + 1, dtype=torch.int32) + if case.has_cmp_kv and case.cmp_ratio > 0: + cmp_kv_lens = torch.full((B,), case.max_seq_kv // case.cmp_ratio, dtype=torch.int32) + cu_seq_cmp_kv[1:] = torch.cumsum(cmp_kv_lens, dim=0) + + seqused_q = seq_q_lens.clone() + seqused_kv = seq_kv_lens.clone() + + return ( + cu_seq_q.to(device), cu_seq_ori_kv.to(device), cu_seq_cmp_kv.to(device), + seqused_q.to(device), seqused_kv.to(device), + ) + + +@pytest.mark.parametrize("case", CASES, ids=[c.name for c in CASES]) +def test_sparse_attn_sharedkv_metadata(case: SparseAttnSharedkvMetadataCase): + torch.npu.set_device(0) + cu_seq_q, cu_seq_ori_kv, cu_seq_cmp_kv, seqused_q, seqused_kv = _make_inputs(case) + + meta_t = custom_ops.sparse_attn_sharedkv_metadata_npu( + cu_seq_q, cu_seq_ori_kv, cu_seq_cmp_kv, seqused_q, seqused_kv, + num_heads_q=case.num_heads_q, num_heads_kv=case.num_heads_kv, + head_dim=case.head_dim, + batch_size=case.batch_size, max_seq_q=case.max_seq_q, max_seq_kv=case.max_seq_kv, + ori_top_k=case.ori_top_k, cmp_top_k=case.cmp_top_k, cmp_ratio=case.cmp_ratio, + ori_mask_mode=case.ori_mask_mode, cmp_mask_mode=case.cmp_mask_mode, + ori_win_left=case.ori_win_left, ori_win_right=case.ori_win_right, + layout_q=case.layout_q, layout_kv=case.layout_kv, + has_ori_kv=case.has_ori_kv, has_cmp_kv=case.has_cmp_kv, + ) + torch.npu.synchronize() + + # Verify output shape and dtype + assert meta_t.shape == (META_SIZE,), f"Expected shape ({META_SIZE},), got {meta_t.shape}" + assert meta_t.dtype == torch.int32, f"Expected dtype int32, got {meta_t.dtype}" + + # Verify determinism: call again and compare + meta_t2 = custom_ops.sparse_attn_sharedkv_metadata_npu( + cu_seq_q, cu_seq_ori_kv, cu_seq_cmp_kv, seqused_q, seqused_kv, + num_heads_q=case.num_heads_q, num_heads_kv=case.num_heads_kv, + head_dim=case.head_dim, + batch_size=case.batch_size, max_seq_q=case.max_seq_q, max_seq_kv=case.max_seq_kv, + ori_top_k=case.ori_top_k, cmp_top_k=case.cmp_top_k, cmp_ratio=case.cmp_ratio, + ori_mask_mode=case.ori_mask_mode, cmp_mask_mode=case.cmp_mask_mode, + ori_win_left=case.ori_win_left, ori_win_right=case.ori_win_right, + layout_q=case.layout_q, layout_kv=case.layout_kv, + has_ori_kv=case.has_ori_kv, has_cmp_kv=case.has_cmp_kv, + ) + torch.npu.synchronize() + assert torch.equal(meta_t.cpu(), meta_t2.cpu()), "Metadata output is not deterministic" + + +if __name__ == "__main__": + pytest.main([__file__, "-v", "-x"]) \ No newline at end of file