diff --git a/benchmark/one-level-arch/kernels/fa/qsmla_config_pto.hpp b/benchmark/one-level-arch/kernels/fa/qsmla_config_pto.hpp new file mode 100644 index 00000000..35b65b10 --- /dev/null +++ b/benchmark/one-level-arch/kernels/fa/qsmla_config_pto.hpp @@ -0,0 +1,170 @@ +#ifndef QSMLA_CONFIG_PTO_HPP +#define QSMLA_CONFIG_PTO_HPP + +#include + +static constexpr int qsmla_pe_row_begin( + int row_count, int pe_count, int pe_id) +{ + const int rows_per_pe = row_count / pe_count; + const int extra_rows = row_count % pe_count; + return pe_id * rows_per_pe + (pe_id < extra_rows ? pe_id : extra_rows); +} + +static constexpr int qsmla_pe_row_end( + int row_count, int pe_count, int pe_id) +{ + return qsmla_pe_row_begin(row_count, pe_count, pe_id + 1); +} + +static constexpr int qsmla_pe_row_offset( + int row_count, int pe_count, int pe_id, int row_stride) +{ + return qsmla_pe_row_begin(row_count, pe_count, pe_id) * row_stride; +} + +struct QsmlaSwaRange { + int begin; + int end; +}; + +static constexpr QsmlaSwaRange qsmla_swa_range( + int kv_sequence_length, int q_sequence_length, int q_position, + int win_left, int win_right) +{ + const int diagonal = kv_sequence_length - q_sequence_length + q_position; + const int unclipped_begin = win_left < 0 ? 0 : diagonal - win_left; + const int unclipped_end = + win_right < 0 ? kv_sequence_length : diagonal + win_right + 1; + const int begin = unclipped_begin < 0 ? 0 : + (unclipped_begin > kv_sequence_length ? kv_sequence_length + : unclipped_begin); + const int end = unclipped_end < 0 ? 0 : + (unclipped_end > kv_sequence_length ? kv_sequence_length + : unclipped_end); + return {begin, end < begin ? begin : end}; +} + +static constexpr QsmlaSwaRange qsmla_swa_block_range( + const QsmlaSwaRange& token_range, int tile_k) +{ + return {token_range.begin / tile_k, + (token_range.end + tile_k - 1) / tile_k}; +} + +static constexpr float qsmla_swa_mask_value( + int kv_position, const QsmlaSwaRange& token_range) +{ + return kv_position >= token_range.begin && kv_position < token_range.end + ? 0.0f + : -1.0e30f; +} + +static inline void qsmla_build_shared_swa_masks( + float* first_mask, float* last_mask, float* zero_mask, + int mask_rows, int tile_k, int first_kv_block, int kv_block_count, + int kv_sequence_length, int q_sequence_length, int q_position, + int win_left, int win_right) +{ + const QsmlaSwaRange range = qsmla_swa_range( + kv_sequence_length, q_sequence_length, q_position, + win_left, win_right); + const int first_block_begin = first_kv_block * tile_k; + const int last_block_begin = + (first_kv_block + kv_block_count - 1) * tile_k; + for (int row = 0; row < mask_rows; ++row) { + for (int column = 0; column < tile_k; ++column) { + const int offset = row * tile_k + column; + first_mask[offset] = qsmla_swa_mask_value( + first_block_begin + column, range); + last_mask[offset] = qsmla_swa_mask_value( + last_block_begin + column, range); + zero_mask[offset] = 0.0f; + } + } +} + +struct QsmlaWorkItem { + int batch; + int q_token; + int kv_head; + int g_slice; + int q_head_begin; + int m_real; +}; + +template +struct QsmlaConfig { + static_assert(B_ > 0 && S1_ > 0 && S2_ > 0, "B/S1/S2 must be positive"); + static_assert(N1_ > 0 && N2_ > 0 && N1_ % N2_ == 0, + "N1 must be positive and divisible by N2"); + static_assert(D_ > 0 && K_ >= 0, "D must be positive and K non-negative"); + static_assert(Tm_ > 0 && Tk_ > 0 && Td_ > 0, "tile sizes must be positive"); + static_assert(GSliceMax_ > 0 && GSliceMax_ <= 64, + "gSlice must be positive and no larger than the MM1 M block"); + + static constexpr int B = B_; + static constexpr int S1 = S1_; + static constexpr int S2 = S2_; + static constexpr int N1 = N1_; + static constexpr int N2 = N2_; + static constexpr int D = D_; + static constexpr int K = K_; + static constexpr int TileM = Tm_; + static constexpr int TileK = Tk_; + static constexpr int TileD = Td_; + static constexpr int GSliceMax = GSliceMax_; + + static constexpr int G = N1 / N2; + static constexpr int GSliceCount = (G + GSliceMax - 1) / GSliceMax; + static constexpr int WorkCount = B * S1 * N2 * GSliceCount; + static constexpr int DBlockCount = (D + TileD - 1) / TileD; + static constexpr int KvBlockCount = (S2 + TileK - 1) / TileK; + + static constexpr int g_head_begin(int kv_head, int g_slice) { + return kv_head * G + g_slice * GSliceMax; + } + + static constexpr int g_slice_size(int g_slice) { + const int remaining = G - g_slice * GSliceMax; + return remaining < GSliceMax ? remaining : GSliceMax; + } + + static constexpr std::size_t q_offset(int batch, int token, int q_head, int dim) { + return (((static_cast(batch) * S1 + token) * N1 + q_head) * D + dim); + } + + static constexpr std::size_t kv_offset(int batch, int token, int kv_head, int dim) { + return (((static_cast(batch) * S2 + token) * N2 + kv_head) * D + dim); + } + + static constexpr std::size_t out_offset(int batch, int token, int q_head, int dim) { + return q_offset(batch, token, q_head, dim); + } + + static constexpr std::size_t q_work_offset(const QsmlaWorkItem& work) { + return q_offset(work.batch, work.q_token, work.q_head_begin, 0); + } + + static constexpr std::size_t kv_work_offset(const QsmlaWorkItem& work) { + return kv_offset(work.batch, 0, work.kv_head, 0); + } + + static constexpr std::size_t out_work_offset(const QsmlaWorkItem& work) { + return out_offset(work.batch, work.q_token, work.q_head_begin, 0); + } + + static constexpr QsmlaWorkItem decode_work(int work_id) { + const int g_slice = work_id % GSliceCount; + work_id /= GSliceCount; + const int kv_head = work_id % N2; + work_id /= N2; + const int q_token = work_id % S1; + const int batch = work_id / S1; + return {batch, q_token, kv_head, g_slice, + g_head_begin(kv_head, g_slice), g_slice_size(g_slice)}; + } +}; + +#endif diff --git a/benchmark/one-level-arch/kernels/fa/quant_sparse_flash_mla_onepass_pto.hpp b/benchmark/one-level-arch/kernels/fa/quant_sparse_flash_mla_onepass_pto.hpp new file mode 100644 index 00000000..a9e37fd6 --- /dev/null +++ b/benchmark/one-level-arch/kernels/fa/quant_sparse_flash_mla_onepass_pto.hpp @@ -0,0 +1,461 @@ +#ifndef QUANT_SPARSE_FLASH_MLA_ONEPASS_PTO_HPP +#define QUANT_SPARSE_FLASH_MLA_ONEPASS_PTO_HPP + +// ============================================================================= +// quant_sparse_flash_mla_onepass_pto.hpp +// Quant Sparse Flash MLA (SWA mode) — One-pass (fused) variant +// +// 一遍式 (fused online softmax) 实现, 将归约 (m,l) 和 P@V 合并为单遍. +// QK^T 只算 1 次 (两遍式算 2 次). +// +// 【D=512 支持】 +// QK^T 沿 D 累加得到 score [kTm, kTk] (不依赖 D 分块) +// PV 按 D 分块: O[Db] 数组, 每个 O[dd] = [kTm, kTd] +// 每个 j 迭代中: 算完 score → rescale 所有 O[dd] → 对每个 dd 做 P@V +// 同时存活 tile: O[0..Db-1] + score/mask/m/l/scale/V/P ≈ Db+10 +// Db=8 时峰值约 18 个 tile, 每个 O tile 8KB, 总 64KB +// +// 【与两遍式的差异】 +// - QK^T 只算 1 次, 计算量减半 +// - online softmax 中同时 rescale 旧 O 并累加新 PV +// - tile 寄存器压力更大 (O[Db] 数组跨 j 循环存活) +// +// 【mask 方式】 +// BSND 共享 Q token 路径只预生成 first/last/zero 三个 [TileM,TileK] +// mask,内部整块直接复用 zero mask;旧 2D 路径保留 [s1,s2] mask。 +// mask 写入仍放在 Tile/CUBE 流程前,避免当前编码问题。 +// +// 【切换方法】 +// test 文件中: +// #include "fa/quant_sparse_flash_mla_onepass_pto.hpp" +// quant_sparse_flash_mla_swa_onepass_pto<...>(...) +// ============================================================================= + +#include +#include "template_asm.h" +#include "qsmla_config_pto.hpp" + +using namespace pto; + +static inline void build_swa_mask_onepass( + float* mask, int s1, int s2, int win_left, int win_right, + int q_position = -1, int q_sequence_length = -1) +{ + for (int q = 0; q < s1; ++q) { + const int logical_q = q_position >= 0 ? q_position : q; + const int logical_s1 = q_position >= 0 ? q_sequence_length : s1; + const QsmlaSwaRange range = qsmla_swa_range( + s2, logical_s1, logical_q, win_left, win_right); + for (int kv = 0; kv < s2; ++kv) { + mask[q * s2 + kv] = qsmla_swa_mask_value(kv, range); + } + } +} + +template +void quant_sparse_flash_mla_swa_onepass_config_pto( + odttype* out_ptr, + qdtype* q_ptr, + kvdtype* ori_kv_ptr, + float softmax_scale, + int ori_win_left, + int ori_win_right, + float* q_descale, + float* ori_kv_descale, + int* ori_sparse_indices, + int* ori_block_table, + int* cu_seqlens_q, + int* cu_seqlens_ori_kv, + int* seqused_q, + int* seqused_ori_kv, + float* sinks, + int* metadata, + float* softmax_lse, + int q_position = -1, + int q_sequence_length = -1) +{ + constexpr int s1 = Config::S1; + constexpr int s2 = Config::S2; + constexpr int D = Config::D; + constexpr int kTm = Config::TileM; + constexpr int kTk = Config::TileK; + constexpr int kTd = Config::TileD; + static_assert(D % kTd == 0, + "one-pass D-tail support is implemented in the next shape-generalization step"); + static_assert(!SharedSwaMask || s1 == kTm, + "shared BSND SWA mask expects one full M tile"); + constexpr int Db = D / kTd; + + constexpr int MaskTileElements = kTm * kTk; + constexpr int MaskBufferElements = + SharedSwaMask ? 3 * MaskTileElements : s1 * s2; + float mask_buf[MaskBufferElements]; + if constexpr (!SharedSwaMask) { + build_swa_mask_onepass( + mask_buf, s1, s2, ori_win_left, ori_win_right, + q_position, q_sequence_length); + } + + using gmQ = global_tensor>; + using gmKV = global_tensor>; + using gmO = global_tensor>; + + using tileQ = TileLeft; + using tileKSrc = Tile; + using tileKRight = TileRight; + using tileW = Tile; + using tileMask = Tile; + using tileW_left = TileLeft; + + using tileO = Tile; + using tileO_cast = Tile; + + using tileV = TileRight; + using tileMax = Tile; + using tileSum = Tile; + + using itQ = global_iterator; + using itKSrc = global_iterator; + using itV = global_iterator; + using itO = global_iterator; + + itQ gIterQ(q_ptr); + itO gIterO(out_ptr); + + const int Qb = (s1 + kTm - 1) / kTm; + const float scale = softmax_scale; + + for (int i = 0; i < Qb; ++i) { + const int q_row_begin = i * kTm; + constexpr bool shared_q_position = SharedSwaMask; + const int logical_q_sequence_length = + shared_q_position ? q_sequence_length : s1; + const int first_logical_q = + shared_q_position ? q_position : q_row_begin; + const int last_q_row = + q_row_begin + kTm < s1 ? q_row_begin + kTm - 1 : s1 - 1; + const int last_logical_q = + shared_q_position ? q_position : last_q_row; + const QsmlaSwaRange first_range = qsmla_swa_range( + s2, logical_q_sequence_length, first_logical_q, + ori_win_left, ori_win_right); + const QsmlaSwaRange last_range = qsmla_swa_range( + s2, logical_q_sequence_length, last_logical_q, + ori_win_left, ori_win_right); + const QsmlaSwaRange kv_range = {first_range.begin, last_range.end}; + const QsmlaSwaRange kv_blocks = qsmla_swa_block_range(kv_range, kTk); + const int kv_block_count = kv_blocks.end - kv_blocks.begin; + if constexpr (SharedSwaMask) { + qsmla_build_shared_swa_masks( + mask_buf, + mask_buf + MaskTileElements, + mask_buf + 2 * MaskTileElements, + kTm, kTk, kv_blocks.begin, kv_block_count, + s2, q_sequence_length, q_position, + ori_win_left, ori_win_right); + } + kvdtype* clipped_kv_ptr = + ori_kv_ptr + static_cast(kv_blocks.begin) * kTk * D; + itKSrc gIterKSrc(clipped_kv_ptr); + itV gIterV(clipped_kv_ptr); + + // ============================================================ + // 一遍式 fused online softmax + PV + // + // O[Db] 数组跨 j 循环存活: + // O[0..Db-1] 每个 [kTm, kTd] float = 8KB, 共 Db*8KB + // Db=8 → 64KB (128 tile tags × 32KB = 4MB, 硬件足够) + // + // 每个 j 迭代: + // 1. QK^T 沿全 D 累加 → score [kTm, kTk] + // 2. score *= scale + mask + // 3. online softmax: m_new, scale_old, l_new + // 4. rescale 所有 O[dd] *= scale_old + // 5. p = exp(score - m_new) + // 6. 对每个 dd: PV = p @ V[dd], O[dd] += PV + // + // 最终: O[dd] /= l, 写回 + // ============================================================ + + // 初始化 m, l + tileMax tMax; TEXPANDS(tMax, -1e30f); + tileSum tSum; TEXPANDS(tSum, 0.0f); + + // 初始化 O[Db] 数组 + tileO tO[Db]; + #pragma clang loop unroll(full) + for (int dd = 0; dd < Db; ++dd) { + TEXPANDS(tO[dd], 0.0f); + } + + for (int j = 0; j < kv_block_count; ++j) { + + // --- Step 1: QK^T 沿全 D 累加 --- + // SuperScalarModel main does not preserve the input ACC correctly + // across TMATMUL_ACC calls. Convert each independent partial from + // ACC NZ to Vec ND immediately and accumulate in the Vec tile. + tileW tW; + TEXPANDS(tW, 0.0f); + #pragma clang loop unroll(full) + for (int dd = 0; dd < Db; ++dd) { + tileQ tQ; + auto gQ = gIterQ(i, dd); + TLOAD(tQ, gQ); // RowMajor Q -> Left tile + + tileKSrc tKSrc; + auto gK = gIterKSrc(j, dd); + TLOAD(tKSrc, gK); // Original K block [Tk, Td] + + tileKRight tK; + TTRANS(tK, tKSrc); // [Tk, Td] -> [Td, Tk] for MM1 SrcR + + tileW tW_partial; + TMATMUL(tW_partial, tQ, tK); + TADD(tW, tW, tW_partial); + } + + // --- Step 2: scale + mask --- + TMULS(tW, tW, scale); + + tileMask tMask; + if constexpr (SharedSwaMask) { + float* selected_mask = mask_buf + 2 * MaskTileElements; + if (j == 0) { + selected_mask = mask_buf; + } + if (j + 1 == kv_block_count) { + selected_mask = mask_buf + MaskTileElements; + } + using gmSharedMask = + global_tensor>; + using itSharedMask = global_iterator; + itSharedMask gIterMask(selected_mask); + auto gMask = gIterMask(0, 0); + TLOAD(tMask, gMask); + } else { + using gmFullMask = global_tensor>; + using itFullMask = global_iterator; + itFullMask gIterMask(mask_buf); + auto gMask = gIterMask(i, kv_blocks.begin + j); + TLOAD(tMask, gMask); + } + TADD(tW, tW, tMask); + + // --- Step 3: online softmax --- + tileMax tLocalMax; + TROWMAX(tLocalMax, tW); + tileMax tNewMax; + TMAX(tNewMax, tMax, tLocalMax); + + // rescale factor: exp(m_old - m_new) + tileMax tScale; + TSUB(tScale, tMax, tNewMax); + TEXP(tScale, tScale); + + // rescale old sum + tileSum tScaledOldSum; + TMUL(tScaledOldSum, tSum, tScale); + + // p = exp(score - m_new) + TROWEXPANDSUB(tW, tW, tNewMax); + TEXP(tW, tW); + + // local sum + tileSum tLocalSum; + TROWSUM(tLocalSum, tW); + + // l_new = l_old' + local_sum + tileSum tNewSum; + TADD(tNewSum, tScaledOldSum, tLocalSum); + + // --- Step 4: rescale all O[dd] --- + #pragma clang loop unroll(full) + for (int dd = 0; dd < Db; ++dd) { + TROWEXPANDMUL(tO[dd], tO[dd], tScale); + } + + // --- Step 5: P@V for each D block --- + tileW_left tW_left; + // PTO v0.58 Local CUBE reads Left payloads as NORM row-major and + // has no NZ dependency, so convert/copy FP32 probabilities directly + // into the qdtype Left tile without an intermediate Vec tile. + TCVT(tW_left, tW); + + #pragma clang loop unroll(full) + for (int dd = 0; dd < Db; ++dd) { + tileV tV; + auto gV = gIterV(j, dd); + TLOAD(tV, gV); // RowMajor V -> Right tile + + tileO tPV; + TMATMUL(tPV, tW_left, tV); + + TADD(tO[dd], tO[dd], tPV); + } + + // --- commit state --- + tMax = tNewMax; + tSum = tNewSum; + } + + // --- 最终归一化: O[dd] /= l, 写回 --- + tileSum tInvSum; + TRECIP(tInvSum, tSum); + + #pragma clang loop unroll(full) + for (int dd = 0; dd < Db; ++dd) { + TROWEXPANDMUL(tO[dd], tO[dd], tInvSum); + + tileO_cast tO_cast; + TCVT(tO_cast, tO[dd]); + auto gO = gIterO(i, dd); + TSTORE(gO, tO_cast); + } + } +} + +// Compatibility entry for the existing fixed two-dimensional smoke. New BSND +// dispatch code should instantiate QsmlaConfig explicitly and call the Config +// entry above. +template +void quant_sparse_flash_mla_swa_onepass_pto( + odttype* out_ptr, + qdtype* q_ptr, + kvdtype* ori_kv_ptr, + float softmax_scale, + int ori_win_left, + int ori_win_right, + float* q_descale, + float* ori_kv_descale, + int* ori_sparse_indices, + int* ori_block_table, + int* cu_seqlens_q, + int* cu_seqlens_ori_kv, + int* seqused_q, + int* seqused_ori_kv, + float* sinks, + int* metadata, + float* softmax_lse) +{ + using Config = QsmlaConfig<1, s1, s2, 1, 1, D, 0, kTm, kTk, kTd>; + quant_sparse_flash_mla_swa_onepass_config_pto( + out_ptr, q_ptr, ori_kv_ptr, softmax_scale, ori_win_left, ori_win_right, + q_descale, ori_kv_descale, ori_sparse_indices, ori_block_table, + cu_seqlens_q, cu_seqlens_ori_kv, seqused_q, seqused_ori_kv, + sinks, metadata, softmax_lse, -1, -1); +} + +// Stage-1 BSND dispatcher. A work item owns all G rows for one +// (batch, qToken, kvHead, gSlice), so its online-softmax state never crosses +// work-item boundaries. Each slice is further divided into full TileM chunks. +// The final short chunk is zero-padded in a kernel-local buffer because the +// current Local TEPL reduction chain cannot consistently consume a partial-M +// ValidRow tile; only its valid rows are copied back. N2>1 needs a strided KV +// view and is intentionally deferred instead of treating non-contiguous +// [S2,N2,D] storage as [S2,D]. +template +void quant_sparse_flash_mla_swa_onepass_bsnd_pto( + odttype* out_ptr, + qdtype* q_ptr, + kvdtype* ori_kv_ptr, + float softmax_scale, + int ori_win_left, + int ori_win_right, + float* q_descale, + float* ori_kv_descale, + int* ori_sparse_indices, + int* ori_block_table, + int* cu_seqlens_q, + int* cu_seqlens_ori_kv, + int* seqused_q, + int* seqused_ori_kv, + float* sinks, + int* metadata, + float* softmax_lse) +{ + static_assert(Config::N2 == 1, + "Stage-1 BSND dispatcher currently requires contiguous N2=1 KV"); + + auto run_full_rows = [&](int row_offset, const QsmlaWorkItem& work) { + using WorkConfig = QsmlaConfig< + 1, Config::TileM, Config::S2, 1, 1, Config::D, Config::K, + Config::TileM, Config::TileK, Config::TileD, Config::TileM>; + + quant_sparse_flash_mla_swa_onepass_config_pto< + qdtype, kvdtype, odttype, WorkConfig, true>( + out_ptr + Config::out_work_offset(work) + row_offset * Config::D, + q_ptr + Config::q_work_offset(work) + row_offset * Config::D, + ori_kv_ptr + Config::kv_work_offset(work), + softmax_scale, ori_win_left, ori_win_right, + q_descale, ori_kv_descale, ori_sparse_indices, ori_block_table, + cu_seqlens_q, cu_seqlens_ori_kv, seqused_q, seqused_ori_kv, + sinks, metadata, softmax_lse, work.q_token, Config::S1); + }; + + auto run_tail_rows = [&](int row_offset, const QsmlaWorkItem& work) { + static_assert(Rows > 0 && Rows < Config::TileM); + qdtype padded_q[Config::TileM * Config::D]; + odttype padded_out[Config::TileM * Config::D]; + qdtype* work_q = q_ptr + Config::q_work_offset(work) + row_offset * Config::D; + odttype* work_out = out_ptr + Config::out_work_offset(work) + row_offset * Config::D; + + for (int row = 0; row < Config::TileM; ++row) { + for (int dim = 0; dim < Config::D; ++dim) { + padded_q[row * Config::D + dim] = + row < Rows ? work_q[row * Config::D + dim] + : static_cast(0.0f); + } + } + + using TailConfig = QsmlaConfig< + 1, Config::TileM, Config::S2, 1, 1, Config::D, Config::K, + Config::TileM, Config::TileK, Config::TileD, Config::TileM>; + quant_sparse_flash_mla_swa_onepass_config_pto< + qdtype, kvdtype, odttype, TailConfig, true>( + padded_out, padded_q, + ori_kv_ptr + Config::kv_work_offset(work), + softmax_scale, ori_win_left, ori_win_right, + q_descale, ori_kv_descale, ori_sparse_indices, ori_block_table, + cu_seqlens_q, cu_seqlens_ori_kv, seqused_q, seqused_ori_kv, + sinks, metadata, softmax_lse, work.q_token, Config::S1); + + for (int row = 0; row < Rows; ++row) { + for (int dim = 0; dim < Config::D; ++dim) { + work_out[row * Config::D + dim] = + padded_out[row * Config::D + dim]; + } + } + }; + + constexpr int kFullSliceChunks = Config::GSliceMax / Config::TileM; + constexpr int kFullSliceTail = Config::GSliceMax % Config::TileM; + constexpr int kLastSliceRows = Config::G % Config::GSliceMax; + constexpr int kLastSliceChunks = kLastSliceRows / Config::TileM; + constexpr int kLastSliceTail = kLastSliceRows % Config::TileM; + + for (int work_id = 0; work_id < Config::WorkCount; ++work_id) { + const QsmlaWorkItem work = Config::decode_work(work_id); + if (work.m_real == Config::GSliceMax) { + for (int chunk = 0; chunk < kFullSliceChunks; ++chunk) { + run_full_rows(chunk * Config::TileM, work); + } + if constexpr (kFullSliceTail != 0) { + run_tail_rows.template operator()( + kFullSliceChunks * Config::TileM, work); + } + } else { + for (int chunk = 0; chunk < kLastSliceChunks; ++chunk) { + run_full_rows(chunk * Config::TileM, work); + } + if constexpr (kLastSliceTail != 0) { + run_tail_rows.template operator()( + kLastSliceChunks * Config::TileM, work); + } + } + } +} + +#endif diff --git a/benchmark/one-level-arch/kernels/fa/quant_sparse_flash_mla_pto.hpp b/benchmark/one-level-arch/kernels/fa/quant_sparse_flash_mla_pto.hpp new file mode 100644 index 00000000..6a8c6354 --- /dev/null +++ b/benchmark/one-level-arch/kernels/fa/quant_sparse_flash_mla_pto.hpp @@ -0,0 +1,285 @@ +#ifndef QUANT_SPARSE_FLASH_MLA_PTO_HPP +#define QUANT_SPARSE_FLASH_MLA_PTO_HPP + +// ============================================================================= +// quant_sparse_flash_mla_pto.hpp +// Quant Sparse Flash MLA (SWA mode) on PTO Tile-OP +// +// 【计算语义】 +// O = softmax(Q @ K^T * softmax_scale, mask=invalid) @ V +// Q: [s1, D], KV: [s2, D] (MLA shared K=V=ori_kv), O: [s1, D] +// +// 【SWA 滑动窗口 — kernel 内部 token 级 mask, 使用 TSEL】 +// mask 为 UINT32 条件矩阵 (1 表示无效/被 mask, 0 表示有效) +// TSELECT_Impl(dst, mask, neg_inf, score): mask=1 → dst=-1e30 +// 窗口范围 (对第 q 个 Q token, 0-indexed): +// diagonal = (s2 - s1) + q +// valid kv: [diagonal - win_left, diagonal + win_right] (闭区间) +// mask 在 kernel 函数内部根据 win_left/win_right 计算, 放在 stack 上. +// 与原算子入参完全一致: win_left/win_right 为标量属性, 无额外 mask 入参. +// +// 【入参说明】 +// win_left / win_right : 滑动窗口参数, kernel 内部用于计算 mask +// ori_sparse_indices : SWA 模式下不使用 (为 nullptr), 保留入参签名 +// ori_block_table : 非 PA 场景下不使用, 保留入参签名 +// 其余可选入参均为 nullptr 占位, 保留签名方便后续扩展 +// +// 【D=512 分块】 +// D 超出单 tile 上限, 沿 D 维切分为 Db 块 (kTd) +// QK^T: 每个 D 块独立 TMATMUL, 转为 Vec 后用 TADD 累加 +// PV: 每个 D 分块独立计算并存储 +// +// 【两遍式】 +// Pass 1: online softmax 归约 (m, l), 含 mask +// Pass 2: 归一化 P, 计算 P@V, 含 mask +// +// 当前 template_asm.hpp 的 packed TSEL 包装与模型三输入契约不一致, +// 因此使用类型安全的四参数 TSELECT_Impl 和完整 UINT32 条件矩阵。 +// ============================================================================= + +#include +#include "template_asm.h" + +using namespace pto; + +// CPU 侧 mask 预计算 (在 kernel 内部调用, 放在 stack 上) +// 生成 UINT32 条件矩阵:窗口外为 1,窗口内为 0。 +static inline void build_swa_mask_select( + uint32_t* maskBuf, int s1, int s2, int win_left, int win_right) +{ + const int causal_offset = s2 - s1; + for (int q = 0; q < s1; ++q) { + int diagonal = causal_offset + q; + int lo = diagonal - win_left; + int hi = diagonal + win_right; + for (int kv = 0; kv < s2; ++kv) { + bool valid = (kv >= lo) && (kv <= hi); + maskBuf[q * s2 + kv] = valid ? 0u : 1u; + } + } +} + +template +void quant_sparse_flash_mla_swa_pto( + odttype* out_ptr, + qdtype* q_ptr, + kvdtype* ori_kv_ptr, + float softmax_scale, + int ori_win_left, + int ori_win_right, + float* q_descale, + float* ori_kv_descale, + int* ori_sparse_indices, + int* ori_block_table, + int* cu_seqlens_q, + int* cu_seqlens_ori_kv, + int* seqused_q, + int* seqused_ori_kv, + float* sinks, + int* metadata, + float* softmax_lse) +{ + constexpr int Db = D / kTd; + + uint32_t maskBuf[s1 * s2]; + build_swa_mask_select(maskBuf, s1, s2, ori_win_left, ori_win_right); + + using gmQ = global_tensor>; + using gmKV = global_tensor>; + // 与 gmKV 共用同一块内存,逻辑上表示 K^T;TCOPYIN 将 DN 转为 ZN。 + using gmKT = global_tensor>; + using gmO = global_tensor>; + using gmMask = global_tensor>; + + using tileQ = TileLeft; + using tileK = TileRight; + using tileW_out = TileAcc; + + // score tile: RowMajor float + using tileW = Tile; + + using tileMask = Tile; + + using tileW_cast = Tile; + using tileW_left = TileLeft; + + using tileO_out = TileAcc; + using tileO = Tile; + using tileO_cast = Tile; + + using tileV = TileRight; + using tileMax = Tile; + using tileSum = Tile; + + using itQ = global_iterator; + using itK = global_iterator; + using itV = global_iterator; + using itO = global_iterator; + using itMask = global_iterator; + + itQ gIterQ(q_ptr); + itK gIterK(ori_kv_ptr); + itV gIterV(ori_kv_ptr); + itO gIterO(out_ptr); + itMask gIterMask(maskBuf); + + const int Qb = (s1 + kTm - 1) / kTm; + const int Kb = (s2 + kTk - 1) / kTk; + + const float scale = softmax_scale; + + for (int i = 0; i < Qb; ++i) { + + // ============================================================ + // Pass 1: online softmax 归约 row max (m) 与 row sum (l) + // 遍历全部 KV 块, 用 mask 屏蔽窗口外 token + // ============================================================ + tileMax tMax; TEXPANDS(tMax, -1e30f); + tileSum tSum; TEXPANDS(tSum, 0.0f); + + // TSEL 的 true-value: -1e30 (mask 命中时取此值) + tileW tNegInf; TEXPANDS(tNegInf, -1e30f); + + for (int j = 0; j < Kb; ++j) { + + // QK^T 沿 D 维累加 + tileW tW; + TEXPANDS(tW, 0.0f); + #pragma clang loop unroll(full) + for (int dd = 0; dd < Db; ++dd) { + tileQ tQ; + auto gQ = gIterQ(i, dd); + TCOPYIN(tQ, gQ); + + tileK tK; + auto gK = gIterK(dd, j); + TCOPYIN(tK, gK); + + tileW_out tW_out; + TMATMUL(tW_out, tQ, tK); + tileW tW_partial; + TCVT_Impl(tW_partial, tW_out); + TADD(tW, tW, tW_partial); + } + + TMULS(tW, tW, scale); + + // mask=1 选择 neg_inf,mask=0 保留 score。 + { + tileMask tMask; + auto gMask = gIterMask(i, j); + TLOAD(tMask, gMask); + tileW tMasked; + TSELECT_Impl(tMasked, tMask, tNegInf, tW); + tW = tMasked; + } + + // m_new = max(m_old, rowmax(score)) + tileMax tLocalMax; + TROWMAX(tLocalMax, tW); + tileMax tNewMax; + TMAX(tNewMax, tMax, tLocalMax); + + // rescale = exp(m_old - m_new); l_old' = l_old * rescale + tileMax tScale; + TSUB(tScale, tMax, tNewMax); + TEXP(tScale, tScale); + tileSum tScaledOldSum; + TMUL(tScaledOldSum, tSum, tScale); + + // local_sum = rowsum(exp(score - m_new)) + TROWEXPANDSUB(tW, tW, tNewMax); + TEXP(tW, tW); + tileSum tLocalSum; + TROWSUM(tLocalSum, tW); + + // l_new = l_old' + local_sum + tileSum tNewSum; + TADD(tNewSum, tScaledOldSum, tLocalSum); + + tMax = tNewMax; + tSum = tNewSum; + } + + // ============================================================ + // Pass 2: p = exp(QK*scale, mask - m) / l, O = Σ p·V + // QK^T 用全 D 累加, PV 按 D 分块计算 + // ============================================================ + tileSum tInvSum; + TRECIP(tInvSum, tSum); + + for (int dd = 0; dd < Db; ++dd) { + tileO tO; + TEXPANDS(tO, 0.0f); + + for (int j = 0; j < Kb; ++j) { + + // 计算完整 QK^T (沿 D 累加, 与 Pass 1 一致) + tileW tW; + TEXPANDS(tW, 0.0f); + #pragma clang loop unroll(full) + for (int dd2 = 0; dd2 < Db; ++dd2) { + tileQ tQ; + auto gQ = gIterQ(i, dd2); + TCOPYIN(tQ, gQ); + + tileK tK; + auto gK = gIterK(dd2, j); + TCOPYIN(tK, gK); + + tileW_out tW_out; + TMATMUL(tW_out, tQ, tK); + tileW tW_partial; + TCVT_Impl(tW_partial, tW_out); + TADD(tW, tW, tW_partial); + } + + TMULS(tW, tW, scale); + + // 应用 token 级 mask: TSEL(score, mask, neg_inf) + { + tileMask tMask; + auto gMask = gIterMask(i, j); + TLOAD(tMask, gMask); + tileW tMasked; + TSELECT_Impl(tMasked, tMask, tNegInf, tW); + tW = tMasked; + } + + // p = exp(score - m) / l + TROWEXPANDSUB(tW, tW, tMax); + TEXP(tW, tW); + TROWEXPANDMUL(tW, tW, tInvSum); + + // cast p -> qdtype Left tile for TMATMUL + tileW_cast tExpW; + TCVT(tExpW, tW); + tileW_left tW_left; + TMOV_ND2NZ(tW_left, tExpW); + + // PV = p * V (当前 D 分块) + tileV tV; + auto gV = gIterV(j, dd); + TCOPYIN(tV, gV); + + tileO_out tPV_out; + TMATMUL(tPV_out, tW_left, tV); + tileO tPV; + TCVT_Impl(tPV, tPV_out); + + TADD(tO, tO, tPV); + } + + // 写回 O 分块 [kTm, kTd] + tileO_cast tO_cast; + TCVT(tO_cast, tO); + auto gO = gIterO(i, dd); + TCOPYOUT(gO, tO_cast); + } + } +} + +#endif diff --git a/benchmark/one-level-arch/kernels/fa/quant_sparse_flash_mla_tadd_pto.hpp b/benchmark/one-level-arch/kernels/fa/quant_sparse_flash_mla_tadd_pto.hpp new file mode 100644 index 00000000..21c61820 --- /dev/null +++ b/benchmark/one-level-arch/kernels/fa/quant_sparse_flash_mla_tadd_pto.hpp @@ -0,0 +1,778 @@ +#ifndef QUANT_SPARSE_FLASH_MLA_TADD_PTO_HPP +#define QUANT_SPARSE_FLASH_MLA_TADD_PTO_HPP + +// ============================================================================= +// quant_sparse_flash_mla_tadd_pto.hpp +// Quant Sparse Flash MLA (SWA mode) — TADD mask variant +// +// 本文件是 quant_sparse_flash_mla_pto.hpp (TSEL 版) 的备选实现, +// 使用 TADD 替代 TSEL 施加 mask, 代码更简单, 兼容旧版模拟器. +// +// 切换方法: 在 test 文件中修改 include 和函数名: +// #include "fa/quant_sparse_flash_mla_tadd_pto.hpp" +// quant_sparse_flash_mla_swa_tadd_pto<...>(...) +// +// 【与 TSEL 版的差异】 +// | 项目 | TADD 版 (本文件) | TSEL 版 (quant_sparse_flash_mla_pto.hpp) | +// |-------------|------------------------------|----------------------------------------------| +// | mask 格式 | float[s1*s2], 0.0/-1e30 | uint32_t 位打包[s1*ceil(s2/32)] | +// | mask 构建 | build_swa_mask 直接赋值 | build_swa_mask_bitpacked 位操作 | +// | mask 加载 | global_iterator 直接 TLOAD | 手动提取 block bit 重打包后 TLOAD | +// | mask 应用 | TADD(score, score, mask) | TSEL(score, mask, neg_inf) | +// | 额外 tile | 无 | tNegInf (TEXPANDS -1e30) | +// | 模拟器要求 | 无特殊要求 | 需更新后模拟器 (o_main 分支, ExecuteTSEL fix)| +// +// 【计算语义】 +// O = softmax(Q @ K^T * softmax_scale + mask) @ V +// Q: [s1, D], KV: [s2, D] (MLA shared K=V=ori_kv), O: [s1, D] +// +// 【SWA 滑动窗口 — kernel 内部 token 级 mask, 使用 TADD】 +// mask[q_idx, kv_idx] = 0.0f if kv_idx 在窗口内 (有效, score 不变) +// = -1e30f if kv_idx 在窗口外 (无效, score → -1e30) +// 窗口范围 (对第 q 个 Q token, 0-indexed): +// diagonal = (s2 - s1) + q +// valid kv: [diagonal - win_left, diagonal + win_right] (闭区间) +// mask 在 kernel 函数内部根据 win_left/win_right 计算, 放在 stack 上. +// 与原算子入参完全一致: win_left/win_right 为标量属性, 无额外 mask 入参. +// +// 【入参说明】 +// win_left / win_right : 滑动窗口参数, kernel 内部用于计算 mask +// ori_sparse_indices : SWA 模式下不使用 (为 nullptr), 保留入参签名 +// ori_block_table : 非 PA 场景下不使用, 保留入参签名 +// 其余可选入参均为 nullptr 占位, 保留签名方便后续扩展 +// +// 【D=512 分块】 +// D 超出单 tile 上限, 沿 D 维切分为 Db 块 (kTd) +// QK^T: 每个 D 块独立 TMATMUL, 转为 Vec 后用 TADD 累加 +// PV: 每个 D 分块独立计算并存储 +// +// 【两遍式】 +// Pass 1: online softmax 归约 (m, l), 含 mask +// Pass 2: 归一化 P, 计算 P@V, 含 mask +// ============================================================================= + +#include +#include "template_asm.h" +#include "qsmla_config_pto.hpp" + +using namespace pto; + +// CPU 侧 mask 预计算 (在 kernel 内部调用, 放在 stack 上) +// mask[q_idx * s2 + kv_idx] = 0.0f if valid, -1e30f if invalid +// valid: diagonal - win_left <= kv_idx <= diagonal + win_right +// where diagonal = (s2 - s1) + q_idx (causal offset + q position) +// kernel 中用 TADD: score += mask (0 保持原值, -1e30 屏蔽) +static inline void build_swa_mask_tadd( + float* mask, int s1, int s2, int win_left, int win_right, + int q_position = -1, int q_sequence_length = -1) +{ + for (int q = 0; q < s1; ++q) { + const int logical_q = q_position >= 0 ? q_position : q; + const int logical_s1 = q_position >= 0 ? q_sequence_length : s1; + const QsmlaSwaRange range = qsmla_swa_range( + s2, logical_s1, logical_q, win_left, win_right); + for (int kv = 0; kv < s2; ++kv) { + mask[q * s2 + kv] = qsmla_swa_mask_value(kv, range); + } + } +} + +template +void quant_sparse_flash_mla_swa_tadd_config_pto( + odttype* out_ptr, + qdtype* q_ptr, + kvdtype* ori_kv_ptr, + float softmax_scale, + int ori_win_left, + int ori_win_right, + float* q_descale, + float* ori_kv_descale, + int* ori_sparse_indices, + int* ori_block_table, + int* cu_seqlens_q, + int* cu_seqlens_ori_kv, + int* seqused_q, + int* seqused_ori_kv, + float* sinks, + int* metadata, + float* softmax_lse, + int q_position = -1, + int q_sequence_length = -1) +{ + constexpr int s1 = Config::S1; + constexpr int s2 = Config::S2; + constexpr int D = Config::D; + constexpr int kTm = Config::TileM; + constexpr int kTk = Config::TileK; + constexpr int kTd = Config::TileD; + static_assert(D % kTd == 0, + "tadd D-tail support is intentionally deferred"); + static_assert(!SharedSwaMask || s1 == kTm, + "shared BSND SWA mask expects one full M tile"); + constexpr int Db = D / kTd; + + constexpr int MaskTileElements = kTm * kTk; + constexpr int MaskBufferElements = + SharedSwaMask ? 3 * MaskTileElements : s1 * s2; + // BSND 的 TileM 行属于同一 Q token,因此只需要首块、尾块和 + // 全有效中间块三个小 mask;旧 2D 路径仍需要逐行 mask。 + float mask_buf[MaskBufferElements]; + if constexpr (!SharedSwaMask) { + build_swa_mask_tadd( + mask_buf, s1, s2, ori_win_left, ori_win_right, + q_position, q_sequence_length); + } + + using gmQ = global_tensor>; + using gmKV = global_tensor>; + using gmO = global_tensor>; + + using tileQ = TileLeft; + using tileKSrc = Tile; + using tileKRight = TileRight; + + // score tile 与 mask tile 都用 RowMajor, 保证 TADD 类型一致 + using tileW = Tile; + using tileMask = Tile; + using tileW_left = TileLeft; + + using tileO = Tile; + using tileO_cast = Tile; + + using tileV = TileRight; + using tileMax = Tile; + using tileSum = Tile; + + using itQ = global_iterator; + using itKSrc = global_iterator; + using itV = global_iterator; + using itO = global_iterator; + + itQ gIterQ(q_ptr); + itO gIterO(out_ptr); + + const int Qb = (s1 + kTm - 1) / kTm; + const float scale = softmax_scale; + + for (int i = 0; i < Qb; ++i) { + const int q_row_begin = i * kTm; + constexpr bool shared_q_position = SharedSwaMask; + const int logical_q_sequence_length = + shared_q_position ? q_sequence_length : s1; + const int first_logical_q = + shared_q_position ? q_position : q_row_begin; + const int last_q_row = + q_row_begin + kTm < s1 ? q_row_begin + kTm - 1 : s1 - 1; + const int last_logical_q = + shared_q_position ? q_position : last_q_row; + const QsmlaSwaRange first_range = qsmla_swa_range( + s2, logical_q_sequence_length, first_logical_q, + ori_win_left, ori_win_right); + const QsmlaSwaRange last_range = qsmla_swa_range( + s2, logical_q_sequence_length, last_logical_q, + ori_win_left, ori_win_right); + const QsmlaSwaRange kv_range = {first_range.begin, last_range.end}; + const QsmlaSwaRange kv_blocks = qsmla_swa_block_range(kv_range, kTk); + const int kv_block_count = kv_blocks.end - kv_blocks.begin; + if constexpr (SharedSwaMask) { + qsmla_build_shared_swa_masks( + mask_buf, + mask_buf + MaskTileElements, + mask_buf + 2 * MaskTileElements, + kTm, kTk, kv_blocks.begin, kv_block_count, + s2, q_sequence_length, q_position, + ori_win_left, ori_win_right); + } + kvdtype* clipped_kv_ptr = + ori_kv_ptr + static_cast(kv_blocks.begin) * kTk * D; + itKSrc gIterKSrc(clipped_kv_ptr); + itV gIterV(clipped_kv_ptr); + + // ============================================================ + // Pass 1: online softmax 归约 row max (m) 与 row sum (l) + // 只遍历与 SWA 有效区间相交的 KV 块 + // ============================================================ + tileMax tMax; TEXPANDS(tMax, -1e30f); + tileSum tSum; TEXPANDS(tSum, 0.0f); + + for (int j = 0; j < kv_block_count; ++j) { + + // QK^T 沿 D 维累加 + tileW tW; + TEXPANDS(tW, 0.0f); + #pragma clang loop unroll(full) + for (int dd = 0; dd < Db; ++dd) { + tileQ tQ; + auto gQ = gIterQ(i, dd); + TLOAD(tQ, gQ); + tileKSrc tKSrc; + auto gK = gIterKSrc(j, dd); + TLOAD(tKSrc, gK); + tileKRight tK; + TTRANS(tK, tKSrc); + tileW tW_partial; + TMATMUL(tW_partial, tQ, tK); + TADD(tW, tW, tW_partial); + } + + TMULS(tW, tW, scale); + + tileMask tMask; + if constexpr (SharedSwaMask) { + float* selected_mask = mask_buf + 2 * MaskTileElements; + if (j == 0) { + selected_mask = mask_buf; + } + if (j + 1 == kv_block_count) { + selected_mask = mask_buf + MaskTileElements; + } + using gmSharedMask = + global_tensor>; + using itSharedMask = global_iterator; + itSharedMask gIterMask(selected_mask); + auto gMask = gIterMask(0, 0); + TLOAD(tMask, gMask); + } else { + using gmFullMask = global_tensor>; + using itFullMask = global_iterator; + itFullMask gIterMask(mask_buf); + auto gMask = gIterMask(i, kv_blocks.begin + j); + TLOAD(tMask, gMask); + } + TADD(tW, tW, tMask); + + // m_new = max(m_old, rowmax(score)) + tileMax tLocalMax; + TROWMAX(tLocalMax, tW); + tileMax tNewMax; + TMAX(tNewMax, tMax, tLocalMax); + + // rescale = exp(m_old - m_new); l_old' = l_old * rescale + tileMax tScale; + TSUB(tScale, tMax, tNewMax); + TEXP(tScale, tScale); + tileSum tScaledOldSum; + TMUL(tScaledOldSum, tSum, tScale); + + // local_sum = rowsum(exp(score - m_new)) + TROWEXPANDSUB(tW, tW, tNewMax); + TEXP(tW, tW); + tileSum tLocalSum; + TROWSUM(tLocalSum, tW); + + // l_new = l_old' + local_sum + tileSum tNewSum; + TADD(tNewSum, tScaledOldSum, tLocalSum); + + tMax = tNewMax; + tSum = tNewSum; + } + + // ============================================================ + // Pass 2: p = exp(QK*scale + mask - m) / l, O = Σ p·V + // QK^T 用全 D 累加, PV 按 D 分块计算 + // ============================================================ + tileSum tInvSum; + TRECIP(tInvSum, tSum); + + for (int dd = 0; dd < Db; ++dd) { + tileO tO; + TEXPANDS(tO, 0.0f); + + for (int j = 0; j < kv_block_count; ++j) { + + // 计算完整 QK^T (沿 D 累加, 与 Pass 1 一致) + tileW tW; + TEXPANDS(tW, 0.0f); + #pragma clang loop unroll(full) + for (int dd2 = 0; dd2 < Db; ++dd2) { + tileQ tQ; + auto gQ = gIterQ(i, dd2); + TLOAD(tQ, gQ); + tileKSrc tKSrc; + auto gK = gIterKSrc(j, dd2); + TLOAD(tKSrc, gK); + tileKRight tK; + TTRANS(tK, tKSrc); + tileW tW_partial; + TMATMUL(tW_partial, tQ, tK); + TADD(tW, tW, tW_partial); + } + + TMULS(tW, tW, scale); + + tileMask tMask; + if constexpr (SharedSwaMask) { + float* selected_mask = mask_buf + 2 * MaskTileElements; + if (j == 0) { + selected_mask = mask_buf; + } + if (j + 1 == kv_block_count) { + selected_mask = mask_buf + MaskTileElements; + } + using gmSharedMask = + global_tensor>; + using itSharedMask = global_iterator; + itSharedMask gIterMask(selected_mask); + auto gMask = gIterMask(0, 0); + TLOAD(tMask, gMask); + } else { + using gmFullMask = global_tensor>; + using itFullMask = global_iterator; + itFullMask gIterMask(mask_buf); + auto gMask = gIterMask(i, kv_blocks.begin + j); + TLOAD(tMask, gMask); + } + TADD(tW, tW, tMask); + + // p = exp(score - m) / l + TROWEXPANDSUB(tW, tW, tMax); + TEXP(tW, tW); + TROWEXPANDMUL(tW, tW, tInvSum); + + tileW_left tW_left; + // PTO v0.58 Local CUBE reads Left payloads as NORM row-major + // and has no NZ dependency, so convert/copy probabilities + // directly into the qdtype Left tile. + TCVT(tW_left, tW); + + // PV = p * V (当前 D 分块) + auto gV = gIterV(j, dd); + tileV tV; + TLOAD(tV, gV); + tileO tPV; + TMATMUL(tPV, tW_left, tV); + + TADD(tO, tO, tPV); + } + + // 写回 O 分块 [kTm, kTd] + tileO_cast tO_cast; + TCVT(tO_cast, tO); + auto gO = gIterO(i, dd); + TSTORE(gO, tO_cast); + } + } +} + +// First-stage four-PE BSND path. A complete 64-head G slice remains one +// work item, while PE tid owns a contiguous 16-row Q/O slice. All PEs execute +// the same work_id and the same cooperative TMATMUL sequence; get_thread_idx() +// is a PE id, not a multi-core work distributor. +template +void quant_sparse_flash_mla_swa_tadd_4pe_bsnd_pto( + odttype* out_ptr, + qdtype* q_ptr, + kvdtype* ori_kv_ptr, + float softmax_scale, + int ori_win_left, + int ori_win_right, + float* q_descale, + float* ori_kv_descale, + int* ori_sparse_indices, + int* ori_block_table, + int* cu_seqlens_q, + int* cu_seqlens_ori_kv, + int* seqused_q, + int* seqused_ori_kv, + float* sinks, + int* metadata, + float* softmax_lse, + float* score_scratch, + qdtype* prob_scratch, + float* pv_scratch) +{ + constexpr int kPeNum = 4; + constexpr int kGroupM = Config::TileM; + constexpr int kTk = Config::TileK; + constexpr int kTd = Config::TileD; + constexpr int kDb = Config::D / kTd; + static_assert(Config::N2 == 1, + "the first four-PE BSND path requires contiguous N2=1 KV"); + static_assert(Config::GSliceMax == 64 && Config::G % 64 == 0, + "the first four-PE path supports complete 64-head G slices"); + static_assert(kGroupM == 64, + "the collective four-PE M tile is fixed at 64 rows"); + static_assert(Config::D % kTd == 0, + "four-PE D-tail support is intentionally deferred"); + constexpr int kPeRows = kGroupM / kPeNum; + const int pe_id = static_cast(get_thread_idx()); + + // The current cooperative CUBE contract mirrors fa_2d_unroll_gmma: + // PE0 stages complete Q/K/V/P matrices in SharedTReg, while TMATMUL maps + // one contiguous kPeRows-row accumulator shard to each PE. Left/Right + // describe operand roles; the shared payload itself remains row-major ND. + using tileQMatrix = SharedMatrixLeft; + using tileKMatrix = SharedMatrixRight; + using tilePMatrix = SharedMatrixLeft; + using tileVMatrix = SharedMatrixRight; + using tileQShared = SharedTile; + using tileKShared = SharedTile; + using tilePShared = SharedTile; + using tileVShared = SharedTile; + + using tileScoreCube = CubeAccumulatorM16; + using tilePVCube = CubeAccumulatorM16; + using tileW = + Tile; + using tileMask = tileW; + using tilePShard = + Tile; + using tileKSrc = + Tile; + using tileO = + Tile; + using tileOCast = + Tile; + // PTO v0.58.4 row reductions publish a dense one-column result. Keep + // the online-softmax state on the same physical descriptor so the + // following TMAX/TADD binary TEPL operations consume matching tiles. + using tileMax = + Tile; + using tileSum = tileMax; + + using gmQ = global_tensor>; + using gmKV = global_tensor>; + using gmV = global_tensor>; + using gmO = global_tensor>; + using itQ = global_iterator; + using itKSrc = global_iterator; + using itV = global_iterator; + using itO = global_iterator; + + constexpr int kMaskElements = kPeRows * kTk; + // mask_buf is PE-local. The three scratch pointers are supplied by the + // caller from shared GM. score/pv use disjoint PE-local shards within + // that region, while prob_scratch gathers all four PE shards into the + // complete shared P operand used by P@V. + float mask_buf[3 * kMaskElements]; + + using gmScoreScratch = + global_tensor>; + using gmProbScratch = + global_tensor>; + using gmPVScratch = + global_tensor>; + using itProbShard = global_iterator; + using itPShared = global_iterator; + gmScoreScratch gScore(score_scratch + pe_id * kPeRows * kTk); + gmPVScratch gPV(pv_scratch + pe_id * kPeRows * kTd); + itProbShard gIterProb(prob_scratch); + itPShared gIterP(prob_scratch); + + for (int work_id = 0; work_id < Config::WorkCount; ++work_id) { + const QsmlaWorkItem work = Config::decode_work(work_id); + qdtype* work_q = q_ptr + Config::q_work_offset(work); + kvdtype* work_kv = ori_kv_ptr + Config::kv_work_offset(work); + const std::size_t work_out_offset = Config::out_work_offset(work); + itQ gIterQ(work_q); + itKSrc gIterKSrc(work_kv); + itV gIterV(work_kv); + + const QsmlaSwaRange range = qsmla_swa_range( + Config::S2, Config::S1, work.q_token, + ori_win_left, ori_win_right); + const QsmlaSwaRange kv_blocks = + qsmla_swa_block_range(range, kTk); + const int kv_block_count = kv_blocks.end - kv_blocks.begin; + qsmla_build_shared_swa_masks( + mask_buf, mask_buf + kMaskElements, + mask_buf + 2 * kMaskElements, + kPeRows, kTk, kv_blocks.begin, kv_block_count, + Config::S2, Config::S1, work.q_token, + ori_win_left, ori_win_right); + + tileMax tMax; + tileSum tSum; + TEXPANDS(tMax, -1e30f); + TEXPANDS(tSum, 0.0f); + + // Pass 1: keep QK's D reduction in the native FP32 CUBE accumulator, + // then cross the explicit CUBE->GM->Vec boundary for online softmax. + for (int j = 0; j < kv_block_count; ++j) { + tileScoreCube tScoreCube; +#pragma clang loop unroll(full) + for (int dd = 0; dd < kDb; ++dd) { + tileQShared tQShared; + tileKSrc tKSrc; + tileKMatrix tKLocal; + tileKShared tKShared; + auto gQ = gIterQ(0, dd); + auto gK = gIterKSrc(kv_blocks.begin + j, dd); + TLOAD(tQShared, gQ); + TLOAD(tKSrc, gK); + TTRANS(tKLocal, tKSrc); + TMOV_L2S_PUBLISH(tKShared, tKLocal); + if (dd == 0) { + TMATMUL(tScoreCube, tQShared, tKShared, + fixp::keep_acc()); + } else { + TMATMUL_ACC(tScoreCube, tScoreCube, + tQShared, tKShared, fixp::keep_acc()); + } + } + + TSTORE_CUBE(gScore, tScoreCube); + tileW tW; + TLOAD(tW, gScore); + TMULS(tW, tW, softmax_scale); + + float* selected_mask = mask_buf + 2 * kMaskElements; + if (j == 0) { + selected_mask = mask_buf; + } + if (j + 1 == kv_block_count) { + selected_mask = mask_buf + kMaskElements; + } + using gmMask = + global_tensor>; + using itMask = global_iterator; + itMask gIterMask(selected_mask); + tileMask tMask; + auto gMask = gIterMask(0, 0); + TLOAD(tMask, gMask); + TADD(tW, tW, tMask); + + tileMax tLocalMax; + tileMax tNewMax; + TROWMAX(tLocalMax, tW); + TMAX(tNewMax, tMax, tLocalMax); + tileMax tScale; + TSUB(tScale, tMax, tNewMax); + TEXP(tScale, tScale); + tileSum tScaledOldSum; + TMUL(tScaledOldSum, tSum, tScale); + TROWEXPANDSUB(tW, tW, tNewMax); + TEXP(tW, tW); + tileSum tLocalSum; + TROWSUM(tLocalSum, tW); + TADD(tSum, tScaledOldSum, tLocalSum); + tMax = tNewMax; + } + + tileSum tInvSum; + TRECIP(tInvSum, tSum); + + // Pass 2: regenerate probabilities, gather PE shards into shared P, + // and let the collective P@V produce one kPeRows-row result per PE. + for (int out_dd = 0; out_dd < kDb; ++out_dd) { + tileO tO; + TEXPANDS(tO, 0.0f); + for (int j = 0; j < kv_block_count; ++j) { + tileScoreCube tScoreCube; +#pragma clang loop unroll(full) + for (int dd = 0; dd < kDb; ++dd) { + tileQShared tQShared; + tileKSrc tKSrc; + tileKMatrix tKLocal; + tileKShared tKShared; + auto gQ = gIterQ(0, dd); + auto gK = gIterKSrc(kv_blocks.begin + j, dd); + TLOAD(tQShared, gQ); + TLOAD(tKSrc, gK); + TTRANS(tKLocal, tKSrc); + TMOV_L2S_PUBLISH(tKShared, tKLocal); + if (dd == 0) { + TMATMUL(tScoreCube, tQShared, tKShared, + fixp::keep_acc()); + } else { + TMATMUL_ACC(tScoreCube, tScoreCube, + tQShared, tKShared, fixp::keep_acc()); + } + } + + TSTORE_CUBE(gScore, tScoreCube); + tileW tW; + TLOAD(tW, gScore); + TMULS(tW, tW, softmax_scale); + float* selected_mask = mask_buf + 2 * kMaskElements; + if (j == 0) { + selected_mask = mask_buf; + } + if (j + 1 == kv_block_count) { + selected_mask = mask_buf + kMaskElements; + } + using gmMask = + global_tensor>; + using itMask = global_iterator; + itMask gIterMask(selected_mask); + tileMask tMask; + auto gMask = gIterMask(0, 0); + TLOAD(tMask, gMask); + TADD(tW, tW, tMask); + TROWEXPANDSUB(tW, tW, tMax); + TEXP(tW, tW); + TROWEXPANDMUL(tW, tW, tInvSum); + + tilePShard tPShard; + TCVT(tPShard, tW); + auto gProbShard = gIterProb(pe_id, 0); + TSTORE(gProbShard, tPShard); + + tilePShared tPShared; + tileVShared tVShared; + auto gP = gIterP(0, 0); + auto gV = gIterV(kv_blocks.begin + j, out_dd); + TLOAD(tPShared, gP); + TLOAD(tVShared, gV); + tilePVCube tPVCube; + TMATMUL(tPVCube, tPShared, tVShared, fixp::keep_acc()); + TSTORE_CUBE(gPV, tPVCube); + tileO tPV; + TLOAD(tPV, gPV); + TADD(tO, tO, tPV); + } + + tileOCast tOCast; + TCVT(tOCast, tO); + itO gIterO(out_ptr + work_out_offset + + pe_id * kPeRows * Config::D); + auto gO = gIterO(0, out_dd); + TSTORE(gO, tOCast); + } + } +} + +// Compatibility entry for the existing fixed two-dimensional smoke. +template +void quant_sparse_flash_mla_swa_tadd_pto( + odttype* out_ptr, + qdtype* q_ptr, + kvdtype* ori_kv_ptr, + float softmax_scale, + int ori_win_left, + int ori_win_right, + float* q_descale, + float* ori_kv_descale, + int* ori_sparse_indices, + int* ori_block_table, + int* cu_seqlens_q, + int* cu_seqlens_ori_kv, + int* seqused_q, + int* seqused_ori_kv, + float* sinks, + int* metadata, + float* softmax_lse) +{ + using Config = QsmlaConfig<1, s1, s2, 1, 1, D, 0, kTm, kTk, kTd>; + quant_sparse_flash_mla_swa_tadd_config_pto< + qdtype, kvdtype, odttype, Config>( + out_ptr, q_ptr, ori_kv_ptr, softmax_scale, + ori_win_left, ori_win_right, + q_descale, ori_kv_descale, ori_sparse_indices, ori_block_table, + cu_seqlens_q, cu_seqlens_ori_kv, seqused_q, seqused_ori_kv, + sinks, metadata, softmax_lse, -1, -1); +} + +// BSND dispatcher aligned with the one-pass address model. One work item owns +// all G rows of one (batch, qToken, kvHead, gSlice). N2>1 remains deferred +// because [S2,N2,D] KV storage needs a strided view. +template +void quant_sparse_flash_mla_swa_tadd_bsnd_pto( + odttype* out_ptr, + qdtype* q_ptr, + kvdtype* ori_kv_ptr, + float softmax_scale, + int ori_win_left, + int ori_win_right, + float* q_descale, + float* ori_kv_descale, + int* ori_sparse_indices, + int* ori_block_table, + int* cu_seqlens_q, + int* cu_seqlens_ori_kv, + int* seqused_q, + int* seqused_ori_kv, + float* sinks, + int* metadata, + float* softmax_lse) +{ + static_assert(Config::N2 == 1, + "Stage-1 BSND dispatcher currently requires contiguous N2=1 KV"); + + auto run_full_rows = [&](int row_offset, const QsmlaWorkItem& work) { + using WorkConfig = QsmlaConfig< + 1, Config::TileM, Config::S2, 1, 1, Config::D, Config::K, + Config::TileM, Config::TileK, Config::TileD, Config::TileM>; + + quant_sparse_flash_mla_swa_tadd_config_pto< + qdtype, kvdtype, odttype, WorkConfig, true>( + out_ptr + Config::out_work_offset(work) + row_offset * Config::D, + q_ptr + Config::q_work_offset(work) + row_offset * Config::D, + ori_kv_ptr + Config::kv_work_offset(work), + softmax_scale, ori_win_left, ori_win_right, + q_descale, ori_kv_descale, ori_sparse_indices, ori_block_table, + cu_seqlens_q, cu_seqlens_ori_kv, seqused_q, seqused_ori_kv, + sinks, metadata, softmax_lse, work.q_token, Config::S1); + }; + + auto run_tail_rows = [&](int row_offset, const QsmlaWorkItem& work) { + static_assert(Rows > 0 && Rows < Config::TileM); + qdtype padded_q[Config::TileM * Config::D]; + odttype padded_out[Config::TileM * Config::D]; + qdtype* work_q = + q_ptr + Config::q_work_offset(work) + row_offset * Config::D; + odttype* work_out = + out_ptr + Config::out_work_offset(work) + row_offset * Config::D; + + for (int row = 0; row < Config::TileM; ++row) { + for (int dim = 0; dim < Config::D; ++dim) { + padded_q[row * Config::D + dim] = + row < Rows ? work_q[row * Config::D + dim] + : static_cast(0.0f); + } + } + + using TailConfig = QsmlaConfig< + 1, Config::TileM, Config::S2, 1, 1, Config::D, Config::K, + Config::TileM, Config::TileK, Config::TileD, Config::TileM>; + quant_sparse_flash_mla_swa_tadd_config_pto< + qdtype, kvdtype, odttype, TailConfig, true>( + padded_out, padded_q, + ori_kv_ptr + Config::kv_work_offset(work), + softmax_scale, ori_win_left, ori_win_right, + q_descale, ori_kv_descale, ori_sparse_indices, ori_block_table, + cu_seqlens_q, cu_seqlens_ori_kv, seqused_q, seqused_ori_kv, + sinks, metadata, softmax_lse, work.q_token, Config::S1); + + for (int row = 0; row < Rows; ++row) { + for (int dim = 0; dim < Config::D; ++dim) { + work_out[row * Config::D + dim] = + padded_out[row * Config::D + dim]; + } + } + }; + + constexpr int kFullSliceChunks = Config::GSliceMax / Config::TileM; + constexpr int kFullSliceTail = Config::GSliceMax % Config::TileM; + constexpr int kLastSliceRows = Config::G % Config::GSliceMax; + constexpr int kLastSliceChunks = kLastSliceRows / Config::TileM; + constexpr int kLastSliceTail = kLastSliceRows % Config::TileM; + + for (int work_id = 0; work_id < Config::WorkCount; ++work_id) { + const QsmlaWorkItem work = Config::decode_work(work_id); + if (work.m_real == Config::GSliceMax) { + for (int chunk = 0; chunk < kFullSliceChunks; ++chunk) { + run_full_rows(chunk * Config::TileM, work); + } + if constexpr (kFullSliceTail != 0) { + run_tail_rows.template operator()( + kFullSliceChunks * Config::TileM, work); + } + } else { + for (int chunk = 0; chunk < kLastSliceChunks; ++chunk) { + run_full_rows(chunk * Config::TileM, work); + } + if constexpr (kLastSliceTail != 0) { + run_tail_rows.template operator()( + kLastSliceChunks * Config::TileM, work); + } + } + } +} + +#endif diff --git a/benchmark/one-level-arch/test/kernel/fa/Makefile b/benchmark/one-level-arch/test/kernel/fa/Makefile index 86e099d6..1c44d7a5 100644 --- a/benchmark/one-level-arch/test/kernel/fa/Makefile +++ b/benchmark/one-level-arch/test/kernel/fa/Makefile @@ -84,6 +84,44 @@ ifeq ($(TESTCASE), fa_2d_unroll_gmma) TARGET = $(ELF_HEAD)/$(TESTCASE)_Sq$(Sq)_Skv$(Skv)_Tm$(Tm)_Tk$(Tk)_X$(X)_Y$(Y).elf endif +ifeq ($(TESTCASE), quant_sparse_flash_mla) + SRC_FILE += $(TEST_ROOT)/$(CASE_SRC_DIR)/quant_sparse_flash_mla.cpp + B ?= 1 + s1 ?= 64 + s2 ?= 128 + N1 ?= 1 + N2 ?= 1 + D ?= 512 + Tm ?= 32 + Tk ?= 32 + Td_block ?= 64 + softmax_scale ?= 0.125 + wleft ?= 1 + wright ?= 1 + IMPL ?= onepass + DEFINES += -DTbatch=$(B) + DEFINES += -DTs1=$(s1) + DEFINES += -DTs2=$(s2) + DEFINES += -DTn1=$(N1) + DEFINES += -DTn2=$(N2) + DEFINES += -DTd=$(D) + DEFINES += -DTm=$(Tm) + DEFINES += -DTk=$(Tk) + DEFINES += -DTd_block=$(Td_block) + DEFINES += -DTsoftmax_scale=$(softmax_scale) + DEFINES += -DTwleft=$(wleft) + DEFINES += -DTwright=$(wright) + ifeq ($(IMPL), tadd) + DEFINES += -DQSMLA_USE_TADD + endif + ifeq ($(IMPL), tadd_4pe) + DEFINES += -DQSMLA_USE_TADD_4PE + TARGET = $(ELF_HEAD)/$(TESTCASE)_B$(B)_s1$(s1)_s2$(s2)_N1$(N1)_N2$(N2)_D$(D)_Tm$(Tm)_Tk$(Tk)_Td$(Td_block)_IMPL$(IMPL).elf + else + TARGET = $(ELF_HEAD)/$(TESTCASE)_B$(B)_s1$(s1)_s2$(s2)_N1$(N1)_N2$(N2)_D$(D)_Tm$(Tm)_Tk$(Tk)_Td$(Td_block).elf + endif +endif + include ../../common/Makefile.common clean: diff --git a/benchmark/one-level-arch/test/kernel/fa/qsmla_stage0_cases.py b/benchmark/one-level-arch/test/kernel/fa/qsmla_stage0_cases.py new file mode 100644 index 00000000..25508d4d --- /dev/null +++ b/benchmark/one-level-arch/test/kernel/fa/qsmla_stage0_cases.py @@ -0,0 +1,232 @@ +#!/usr/bin/env python3 + +from dataclasses import dataclass +from pathlib import Path +from typing import FrozenSet, Optional, Tuple + + +@dataclass(frozen=True) +class QsmlaCase: + name: str + b: int + s1: int + s2: int + n1: int + n2: int + d: int + k: int + k1: Optional[int] = None + tm: int = 32 + tk: int = 32 + td: int = 64 + win_left: int = 1 + win_right: int = 1 + softmax_scale: float = 0.125 + source: str = "supernpubench:stage0-smoke" + mode: str = "SWA" + q_layout: str = "BSND" + kv_layout: str = "BSND" + logical_dtype: str = "fp16" + source_storage_dtype: str = "fp16" + stage0_compute_dtype: str = "fp16" + enable_stage: str = "stage0_reference" + coverage: FrozenSet[str] = frozenset() + reference_feasible: bool = True + note: str = "" + + +QSMLA_STAGE0_CASES: Tuple[QsmlaCase, ...] = ( + QsmlaCase( + name="baseline_swa", + b=1, s1=64, s2=128, n1=1, n2=1, d=512, k=128, + source="supernpubench:existing", + coverage=frozenset({"baseline", "k128", "s1_ne_s2"}), + ), + QsmlaCase( + name="transformer_decode_first", + b=1, s1=1, s2=8192, n1=64, n2=1, d=512, k=512, + win_left=127, win_right=0, softmax_scale=0.04419417, + source="transformer:decode_first", + mode="CSA", q_layout="TND", kv_layout="PA_BBND", + logical_dtype="hifp8", source_storage_dtype="uint8", + enable_stage="future_semantics", + coverage=frozenset({"q_tail", "batch_gqa", "k512", "s1_ne_s2"}), + reference_feasible=False, + note="Exact production shape; retained as a Stage 1 compile target.", + ), + QsmlaCase( + name="transformer_prefill_first", + b=1, s1=8192, s2=8192, n1=64, n2=1, d=512, k=512, + win_left=127, win_right=0, softmax_scale=0.04419417, + source="transformer:prefill_first", + mode="CSA", q_layout="TND", kv_layout="PA_BBND", + logical_dtype="hifp8", source_storage_dtype="uint8", + enable_stage="future_semantics", + coverage=frozenset({"batch_gqa", "k512"}), + reference_feasible=False, + note="Exact production prefill shape; Stage 1 compile target only.", + ), + QsmlaCase( + name="transformer_csa_small_prefill", + b=1, s1=16, s2=1024, n1=64, n2=1, d=512, k=512, + win_left=127, win_right=0, softmax_scale=0.04419417, + source="transformer:csa_small_prefill", + mode="CSA", q_layout="TND", kv_layout="PA_BBND", + logical_dtype="hifp8", source_storage_dtype="uint8", + enable_stage="future_semantics", + coverage=frozenset({"q_tail", "batch_gqa", "k512", "s1_ne_s2"}), + reference_feasible=False, + note="Shape-only coverage in Stage 0; semantics remain continuous SWA.", + ), + QsmlaCase( + name="transformer_ori_sparse_tnd_pa", + b=4, s1=4, s2=8192, n1=128, n2=1, d=512, k=0, k1=64, + source="transformer:ori_sparse_tnd_pa", + mode="ORI_SPARSE", q_layout="TND", kv_layout="PA_BBND", + logical_dtype="hifp8", source_storage_dtype="uint8", + enable_stage="future_semantics", + coverage=frozenset({"batch_gqa", "q_tail", "s1_ne_s2"}), + reference_feasible=False, + note="K1 is the ORI sparse candidate width; requires sparse, TND, PA and HIFLOAT8 semantics.", + ), + QsmlaCase( + name="transformer_ori_sparse_bsnd_pa", + b=4, s1=4, s2=8192, n1=128, n2=1, d=512, k=0, k1=128, + source="transformer:ori_sparse_bsnd_pa", + mode="ORI_SPARSE", q_layout="BSND", kv_layout="PA_BBND", + logical_dtype="hifp8", source_storage_dtype="uint8", + enable_stage="future_semantics", + coverage=frozenset({"batch_gqa", "q_tail", "k128", "s1_ne_s2"}), + reference_feasible=False, + note="BSND query plus ORI sparse PA path; K1=128.", + ), + QsmlaCase( + name="transformer_ori_cmp_sparse_tnd_pa", + b=4, s1=4, s2=8192, n1=128, n2=1, d=512, k=512, k1=128, + source="transformer:ori_cmp_sparse_tnd_pa", + mode="ORI_CMP_SPARSE", q_layout="TND", kv_layout="PA_BBND", + logical_dtype="hifp8", source_storage_dtype="uint8", + enable_stage="future_semantics", + coverage=frozenset({"batch_gqa", "q_tail", "k128", "k512", "s1_ne_s2"}), + reference_feasible=False, + note="Dual sparse ORI/CMP path with TND query and PA KV.", + ), + QsmlaCase( + name="transformer_ori_cmp_sparse_bsnd_pa", + b=4, s1=4, s2=8192, n1=128, n2=1, d=512, k=512, k1=128, + source="transformer:ori_cmp_sparse_bsnd_pa", + mode="ORI_CMP_SPARSE", q_layout="BSND", kv_layout="PA_BBND", + logical_dtype="hifp8", source_storage_dtype="uint8", + enable_stage="future_semantics", + coverage=frozenset({"batch_gqa", "q_tail", "k128", "k512", "s1_ne_s2"}), + reference_feasible=False, + note="Same numeric shape as TND variant but a distinct BSND semantic case.", + ), + QsmlaCase( + name="typical_bsnd_swa_1", + b=8, s1=4, s2=131072, n1=128, n2=1, d=512, k=0, + tm=64, win_left=128, win_right=576, + source="user:typical_case_1", + enable_stage="stage1_compile", + coverage=frozenset({"bsnd_layout", "batch_gqa", "g128_split", "s1_ne_s2", "typical_large_s2"}), + reference_feasible=False, + note="User-provided production BSND SWA shape; compile-only because S2 is 131072.", + ), + QsmlaCase( + name="typical_bsnd_swa_2", + b=8, s1=4, s2=131072, n1=128, n2=1, d=512, k=0, + tm=64, win_left=128, win_right=576, + source="user:typical_case_2", + enable_stage="stage1_compile", + coverage=frozenset({"bsnd_layout", "batch_gqa", "g128_split", "s1_ne_s2", "typical_large_s2"}), + reference_feasible=False, + note="User-provided production BSND SWA shape; compile-only because S2 is 131072.", + ), + QsmlaCase( + name="typical_bsnd_swa_3", + b=8, s1=4, s2=131072, n1=128, n2=1, d=512, k=0, + tm=64, win_left=128, win_right=576, + source="user:typical_case_3", + enable_stage="stage1_compile", + coverage=frozenset({"bsnd_layout", "batch_gqa", "g128_split", "s1_ne_s2", "typical_large_s2"}), + reference_feasible=False, + note="User-provided production BSND SWA shape; compile-only because S2 is 131072.", + ), + QsmlaCase( + name="typical_bsnd_swa_s2_128", + b=8, s1=4, s2=128, n1=128, n2=1, d=512, k=0, + tm=64, win_left=128, win_right=576, + source="user:typical_reduced_s2", + enable_stage="stage1_compile", + coverage=frozenset({"bsnd_layout", "batch_gqa", "g128_split", "s1_ne_s2", "typical_reduced_s2"}), + reference_feasible=False, + note="Production-like BSND SWA shape with S2 reduced to 128.", + ), + QsmlaCase( + name="typical_bsnd_swa_s2_1024", + b=8, s1=4, s2=1024, n1=128, n2=1, d=512, k=0, + tm=64, win_left=128, win_right=576, + source="user:typical_reduced_s2", + enable_stage="stage1_compile", + coverage=frozenset({"bsnd_layout", "batch_gqa", "g128_split", "s1_ne_s2", "typical_reduced_s2"}), + reference_feasible=False, + note="Production-like BSND SWA shape with S2 reduced to 1024.", + ), + QsmlaCase( + name="typical_bsnd_swa_s2_4096", + b=8, s1=4, s2=4096, n1=128, n2=1, d=512, k=0, + tm=64, win_left=128, win_right=576, + source="user:typical_reduced_s2", + enable_stage="stage1_compile", + coverage=frozenset({"bsnd_layout", "batch_gqa", "g128_split", "s1_ne_s2", "typical_reduced_s2"}), + reference_feasible=False, + note="Production-like BSND SWA shape with S2 reduced to 4096.", + ), + QsmlaCase( + name="smoke_batch_gqa", + b=2, s1=3, s2=5, n1=4, n2=2, d=8, k=128, + tm=2, tk=3, td=4, + coverage=frozenset({"batch_gqa", "q_tail", "kv_tail", "k128", "s1_ne_s2"}), + ), + QsmlaCase( + name="smoke_all_tails_k512", + b=1, s1=5, s2=7, n1=2, n2=1, d=10, k=512, + tm=4, tk=4, td=6, + coverage=frozenset({"batch_gqa", "q_tail", "kv_tail", "d_tail", "k512", "s1_ne_s2"}), + ), + QsmlaCase( + name="smoke_bsnd_g64", + b=1, s1=2, s2=3, n1=64, n2=1, d=8, k=128, + tm=64, tk=2, td=4, + q_layout="BSND", kv_layout="BSND", + coverage=frozenset({"bsnd_layout", "g64", "kv_tail", "k128", "s1_ne_s2"}), + note="Stage 1 direct [G,D] MM1 smoke with G=64.", + ), + QsmlaCase( + name="smoke_bsnd_g128", + b=1, s1=1, s2=3, n1=128, n2=1, d=8, k=512, + tm=64, tk=2, td=4, + q_layout="BSND", kv_layout="BSND", + coverage=frozenset({"bsnd_layout", "g128_split", "q_tail", "kv_tail", "k512", "s1_ne_s2"}), + note="Stage 1 split-G smoke: G=128 must be processed as two M=64 slices.", + ), +) + + +def validate_case(case: QsmlaCase) -> Optional[str]: + positive = (case.b, case.n1, case.n2, case.d, case.tm, case.tk, case.td) + if any(value <= 0 for value in positive): + return "B/N1/N2/D/Tm/Tk/Td must be positive" + if case.s1 < 0 or case.s2 < 0 or case.k < 0: + return "S1/S2/K must be non-negative" + if case.k1 is not None and case.k1 < 0: + return "K1 must be non-negative when present" + if case.n1 % case.n2 != 0: + return "N1 must be divisible by N2" + if case.win_left < -1 or case.win_right < -1: + return "window bounds must be -1 or non-negative" + return None + + +def case_output_dir(case: QsmlaCase, root: Path) -> Path: + return Path(root) / case.name diff --git a/benchmark/one-level-arch/test/kernel/fa/run_qsmla_stage0.py b/benchmark/one-level-arch/test/kernel/fa/run_qsmla_stage0.py new file mode 100644 index 00000000..cc6b7805 --- /dev/null +++ b/benchmark/one-level-arch/test/kernel/fa/run_qsmla_stage0.py @@ -0,0 +1,103 @@ +#!/usr/bin/env python3 + +import argparse +import subprocess +import sys +import tempfile +from pathlib import Path + +from qsmla_stage0_cases import QSMLA_STAGE0_CASES, QsmlaCase, validate_case + + +HERE = Path(__file__).resolve().parent +REFERENCE_SOURCE = HERE / "src" / "qsmla_cpu_ref.cpp" + + +def print_cases() -> None: + print("NAME B S1 S2 N1 N2 D K K1 MODE Q/KV LAYOUT LOGICAL/STORAGE -> STAGE0 ENABLE") + for case in QSMLA_STAGE0_CASES: + print( + f"{case.name:32} {case.b:2} {case.s1:5} {case.s2:5} " + f"{case.n1:4} {case.n2:3} {case.d:4} {case.k:4} " + f"{('-' if case.k1 is None else str(case.k1)):>4} {case.mode:16} " + f"{case.q_layout}/{case.kv_layout:15} " + f"{case.logical_dtype}/{case.source_storage_dtype} -> {case.stage0_compute_dtype:5} {case.enable_stage}" + ) + + +def compile_defines(case: QsmlaCase, output_root: Path): + values = { + "QB": case.b, "QS1": case.s1, "QS2": case.s2, + "QN1": case.n1, "QN2": case.n2, "QD": case.d, "QK": case.k, + "QK1": 0 if case.k1 is None else case.k1, + "QTM": case.tm, "QTK": case.tk, "QTD": case.td, + "QWIN_LEFT": case.win_left, "QWIN_RIGHT": case.win_right, + "QSOFTMAX_SCALE": f"{case.softmax_scale:.9g}f", + "QCASE_NAME": f'\"{case.name}\"', + "QOUTPUT_ROOT": f'\"{output_root.resolve()}\"', + "QLAYOUT_BSND": 1 if case.q_layout == "BSND" else 0, + "QKV_LAYOUT_BSND": 1 if case.kv_layout == "BSND" else 0, + } + return [f"-D{name}={value}" for name, value in values.items()] + + +def run_reference(case: QsmlaCase, output_root: Path) -> None: + error = validate_case(case) + if error is not None: + raise ValueError(f"invalid QSMLA case {case.name}: {error}") + if not case.reference_feasible: + raise ValueError(f"case {case.name} is a compile-shape target, not a Stage 0 reference target") + + output_root.mkdir(parents=True, exist_ok=True) + with tempfile.TemporaryDirectory(prefix="qsmla-stage0-build-") as build_dir: + binary = Path(build_dir) / case.name + command = [ + "g++", "-std=c++17", "-O2", str(REFERENCE_SOURCE), "-o", str(binary), + *compile_defines(case, output_root), + ] + print("compile:", " ".join(command)) + subprocess.run(command, check=True) + subprocess.run([str(binary)], check=True) + + +def parse_args(): + parser = argparse.ArgumentParser(description="List and generate QSMLA Stage 0 FP16 SWA references") + action = parser.add_mutually_exclusive_group(required=True) + action.add_argument("--list", action="store_true", help="list the fixed Stage 0 shape matrix") + action.add_argument("--case", metavar="NAME", help="generate one reference-feasible case") + action.add_argument("--all-reference", action="store_true", help="generate every reference-feasible case") + parser.add_argument("--output-root", type=Path, default=Path("qsmla_stage0_output")) + return parser.parse_args() + + +def main() -> int: + args = parse_args() + if args.list: + print_cases() + return 0 + + by_name = {case.name: case for case in QSMLA_STAGE0_CASES} + if args.case: + case = by_name.get(args.case) + if case is None: + print(f"unknown QSMLA case: {args.case}", file=sys.stderr) + return 2 + try: + run_reference(case, args.output_root) + except (ValueError, subprocess.CalledProcessError) as error: + print(error, file=sys.stderr) + return 1 + return 0 + + try: + for case in QSMLA_STAGE0_CASES: + if case.reference_feasible: + run_reference(case, args.output_root) + except (ValueError, subprocess.CalledProcessError) as error: + print(error, file=sys.stderr) + return 1 + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/benchmark/one-level-arch/test/kernel/fa/src/qsmla_compare.py b/benchmark/one-level-arch/test/kernel/fa/src/qsmla_compare.py new file mode 100644 index 00000000..de4de163 --- /dev/null +++ b/benchmark/one-level-arch/test/kernel/fa/src/qsmla_compare.py @@ -0,0 +1,108 @@ +#!/usr/bin/env python3 + +import argparse +import math +import struct +import sys +from pathlib import Path + + +HERE = Path(__file__).resolve().parent + + +def parse_args(): + parser = argparse.ArgumentParser( + description="Compare FP16 QSMLA NPU output with FP32 CPU golden output." + ) + parser.add_argument( + "--actual", type=Path, default=HERE / "qsmla_onepass_npu_out.bin", + help="NPU output in little-endian FP16 format", + ) + parser.add_argument( + "--golden", type=Path, default=HERE / "qsmla_golden.bin", + help="CPU golden output in little-endian FP32 format", + ) + parser.add_argument("--b", type=int) + parser.add_argument("--s1", type=int) + parser.add_argument("--n1", type=int) + parser.add_argument("--d", type=int) + parser.add_argument("--atol", type=float, default=1e-3) + parser.add_argument("--rtol", type=float, default=1e-3) + return parser.parse_args() + + +def read_values(path: Path, element_bytes: int, format_code: str): + data = path.read_bytes() + if len(data) % element_bytes != 0: + raise ValueError( + f"{path}: byte size {len(data)} is not divisible by {element_bytes}" + ) + count = len(data) // element_bytes + return struct.unpack(f"<{count}{format_code}", data) + + +def main() -> int: + args = parse_args() + shape_values = (args.b, args.s1, args.n1, args.d) + if any(value is not None for value in shape_values) and not all( + value is not None for value in shape_values + ): + print("--b, --s1, --n1 and --d must be specified together", file=sys.stderr) + return 2 + if args.atol < 0 or args.rtol < 0: + print("--atol and --rtol must be non-negative", file=sys.stderr) + return 2 + + try: + actual = read_values(args.actual, 2, "e") + golden = read_values(args.golden, 4, "f") + except (OSError, ValueError) as error: + print(error, file=sys.stderr) + return 2 + + if len(actual) != len(golden): + print( + f"element count mismatch: actual={len(actual)}, golden={len(golden)}", + file=sys.stderr, + ) + return 2 + + if all(value is not None for value in shape_values): + if any(value <= 0 for value in shape_values): + print("BSND dimensions must be positive", file=sys.stderr) + return 2 + expected = args.b * args.s1 * args.n1 * args.d + if expected != len(actual): + print( + f"shape element count mismatch: BSND[{args.b},{args.s1},{args.n1},{args.d}] " + f"expects {expected}, files contain {len(actual)}", + file=sys.stderr, + ) + return 2 + + errors = [abs(float(a) - float(g)) for a, g in zip(actual, golden)] + passed = [ + math.isfinite(float(a)) + and math.isfinite(float(g)) + and error <= args.atol + args.rtol * abs(float(g)) + for a, g, error in zip(actual, golden, errors) + ] + passed_count = sum(passed) + count = len(actual) + + print(f"actual = {args.actual}") + print(f"golden = {args.golden}") + if all(value is not None for value in shape_values): + print(f"shape = BSND[{args.b},{args.s1},{args.n1},{args.d}]") + print(f"tol = atol={args.atol:g}, rtol={args.rtol:g}") + print(f"passed = {passed_count}/{count} ({100 * passed_count / count:.6f}%)" if count else "passed = 0/0") + print(f"failed = {count - passed_count}") + print(f"max_abs = {max(errors):.9f}" if errors else "max_abs = 0.000000000") + print(f"mean_abs= {sum(errors) / count:.9f}" if count else "mean_abs= 0.000000000") + print(f"nan_npu = {sum(math.isnan(float(value)) for value in actual)}") + print(f"nan_ref = {sum(math.isnan(float(value)) for value in golden)}") + return 0 if passed_count == count else 1 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/benchmark/one-level-arch/test/kernel/fa/src/qsmla_cpu_ref.cpp b/benchmark/one-level-arch/test/kernel/fa/src/qsmla_cpu_ref.cpp new file mode 100644 index 00000000..a465d5ac --- /dev/null +++ b/benchmark/one-level-arch/test/kernel/fa/src/qsmla_cpu_ref.cpp @@ -0,0 +1,233 @@ +// Stage-0 CPU reference for contiguous FP16 QSMLA SWA. +// Inputs are rounded to IEEE FP16, while dot products, softmax and output use FP32. + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#ifndef QB +#define QB 1 +#endif +#ifndef QS1 +#define QS1 64 +#endif +#ifndef QS2 +#define QS2 128 +#endif +#ifndef QN1 +#define QN1 1 +#endif +#ifndef QN2 +#define QN2 1 +#endif +#ifndef QD +#define QD 512 +#endif +#ifndef QK +#define QK 128 +#endif +#ifndef QTM +#define QTM 32 +#endif +#ifndef QTK +#define QTK 32 +#endif +#ifndef QTD +#define QTD 64 +#endif +#ifndef QWIN_LEFT +#define QWIN_LEFT 1 +#endif +#ifndef QWIN_RIGHT +#define QWIN_RIGHT 1 +#endif +#ifndef QSOFTMAX_SCALE +#define QSOFTMAX_SCALE 0.125f +#endif +#ifndef QCASE_NAME +#define QCASE_NAME "baseline_swa" +#endif +#ifndef QOUTPUT_ROOT +#define QOUTPUT_ROOT "." +#endif +#ifndef QLAYOUT_BSND +#define QLAYOUT_BSND 1 +#endif +#ifndef QKV_LAYOUT_BSND +#define QKV_LAYOUT_BSND 1 +#endif + +static_assert(QB > 0, "B must be positive"); +static_assert(QS1 >= 0 && QS2 >= 0, "S1/S2 must be non-negative"); +static_assert(QN1 > 0 && QN2 > 0 && QN1 % QN2 == 0, "N1 must be divisible by N2"); +static_assert(QD > 0 && QK >= 0, "D must be positive and K non-negative"); +static_assert(QTM > 0 && QTK > 0 && QTD > 0, "tile sizes must be positive"); +static_assert(QWIN_LEFT >= -1 && QWIN_RIGHT >= -1, "window bounds must be -1 or non-negative"); + +static uint16_t float_to_half(float value) { + uint32_t bits; + std::memcpy(&bits, &value, sizeof(bits)); + const uint32_t sign = (bits >> 16) & 0x8000u; + uint32_t mantissa = bits & 0x007fffffu; + int exponent = static_cast((bits >> 23) & 0xffu) - 127 + 15; + + if (exponent <= 0) { + if (exponent < -10) return static_cast(sign); + mantissa = (mantissa | 0x00800000u) >> (1 - exponent); + mantissa += 0x00000fffu + ((mantissa >> 13) & 1u); + return static_cast(sign | (mantissa >> 13)); + } + if (exponent >= 31) { + return static_cast(sign | 0x7c00u); + } + mantissa += 0x00000fffu + ((mantissa >> 13) & 1u); + if (mantissa & 0x00800000u) { + mantissa = 0; + ++exponent; + if (exponent >= 31) return static_cast(sign | 0x7c00u); + } + return static_cast(sign | (static_cast(exponent) << 10) | (mantissa >> 13)); +} + +static float half_to_float(uint16_t value) { + const uint32_t sign = static_cast(value & 0x8000u) << 16; + uint32_t exponent = (value >> 10) & 0x1fu; + uint32_t mantissa = value & 0x03ffu; + uint32_t bits; + if (exponent == 0) { + if (mantissa == 0) { + bits = sign; + } else { + int shift = 0; + while ((mantissa & 0x0400u) == 0) { + mantissa <<= 1; + ++shift; + } + mantissa &= 0x03ffu; + bits = sign | (static_cast(127 - 15 - shift) << 23) | (mantissa << 13); + } + } else if (exponent == 31) { + bits = sign | 0x7f800000u | (mantissa << 13); + } else { + bits = sign | ((exponent + 127 - 15) << 23) | (mantissa << 13); + } + float result; + std::memcpy(&result, &bits, sizeof(result)); + return result; +} + +static void init_deterministic_fp16(std::vector& storage, std::vector& values, int seed) { + for (size_t i = 0; i < storage.size(); ++i) { + const float source = static_cast((static_cast(i) * 31 + seed * 17) % 100) / 100.0f - 0.5f; + storage[i] = float_to_half(source); + values[i] = half_to_float(storage[i]); + } +} + +static bool swa_valid(int q_pos, int kv_pos) { + const int threshold = QS2 - QS1 + q_pos + 1; + const int lo = QWIN_LEFT == -1 ? 0 : threshold - QWIN_LEFT - 1; + const int hi = QWIN_RIGHT == -1 ? QS2 : threshold + QWIN_RIGHT; + return kv_pos >= std::max(0, lo) && kv_pos < std::min(QS2, hi); +} + +static size_t q_offset(int b, int head, int token, int dim) { + return (((static_cast(b) * QS1 + token) * QN1 + head) * QD + dim); +} + +static size_t kv_offset(int b, int head, int token, int dim) { + return (((static_cast(b) * QS2 + token) * QN2 + head) * QD + dim); +} + +static bool write_binary(const std::filesystem::path& path, const void* data, size_t bytes) { + FILE* file = std::fopen(path.string().c_str(), "wb"); + if (file == nullptr) return false; + const bool ok = std::fwrite(data, 1, bytes, file) == bytes; + std::fclose(file); + return ok; +} + +int main() { + const size_t q_count = static_cast(QB) * QN1 * QS1 * QD; + const size_t kv_count = static_cast(QB) * QN2 * QS2 * QD; + const size_t out_count = q_count; + std::vector q_storage(q_count), kv_storage(kv_count); + std::vector q(q_count), kv(kv_count), out(out_count, 0.0f); + std::vector scores(QS2); + + init_deterministic_fp16(q_storage, q, 1); + init_deterministic_fp16(kv_storage, kv, 2); + + constexpr int group_size = QN1 / QN2; + for (int b = 0; b < QB; ++b) { + for (int q_head = 0; q_head < QN1; ++q_head) { + const int kv_head = q_head / group_size; + for (int q_pos = 0; q_pos < QS1; ++q_pos) { + float row_max = -std::numeric_limits::infinity(); + for (int kv_pos = 0; kv_pos < QS2; ++kv_pos) { + if (!swa_valid(q_pos, kv_pos)) { + scores[kv_pos] = -std::numeric_limits::infinity(); + continue; + } + float dot = 0.0f; + for (int d = 0; d < QD; ++d) { + dot += q[q_offset(b, q_head, q_pos, d)] * kv[kv_offset(b, kv_head, kv_pos, d)]; + } + scores[kv_pos] = dot * QSOFTMAX_SCALE; + row_max = std::max(row_max, scores[kv_pos]); + } + if (!std::isfinite(row_max)) continue; + + float denominator = 0.0f; + for (int kv_pos = 0; kv_pos < QS2; ++kv_pos) { + if (std::isfinite(scores[kv_pos])) denominator += std::exp(scores[kv_pos] - row_max); + } + for (int d = 0; d < QD; ++d) { + float numerator = 0.0f; + for (int kv_pos = 0; kv_pos < QS2; ++kv_pos) { + if (!std::isfinite(scores[kv_pos])) continue; + const float probability = std::exp(scores[kv_pos] - row_max) / denominator; + numerator += probability * kv[kv_offset(b, kv_head, kv_pos, d)]; + } + out[q_offset(b, q_head, q_pos, d)] = numerator; + } + } + } + } + + const std::filesystem::path output_dir = std::filesystem::path(QOUTPUT_ROOT) / QCASE_NAME; + std::error_code error; + std::filesystem::create_directories(output_dir, error); + if (error) { + std::fprintf(stderr, "Failed to create %s: %s\n", output_dir.string().c_str(), error.message().c_str()); + return 1; + } + if (!write_binary(output_dir / "q.fp16.bin", q_storage.data(), q_storage.size() * sizeof(uint16_t)) || + !write_binary(output_dir / "kv.fp16.bin", kv_storage.data(), kv_storage.size() * sizeof(uint16_t)) || + !write_binary(output_dir / "out.fp32.bin", out.data(), out.size() * sizeof(float))) { + std::fprintf(stderr, "Failed to write reference data under %s\n", output_dir.string().c_str()); + return 1; + } + + FILE* manifest = std::fopen((output_dir / "manifest.txt").string().c_str(), "w"); + if (manifest == nullptr) return 1; + std::fprintf(manifest, + "case=%s\nB=%d S1=%d S2=%d N1=%d N2=%d D=%d K=%d\n" + "Tm=%d Tk=%d Td=%d win_left=%d win_right=%d softmax_scale=%.9g\n" + "q_layout=%s kv_layout=%s\n" + "q_dtype=fp16 kv_dtype=fp16 accumulation=fp32 out_dtype=fp32\n", + QCASE_NAME, QB, QS1, QS2, QN1, QN2, QD, QK, + QTM, QTK, QTD, QWIN_LEFT, QWIN_RIGHT, static_cast(QSOFTMAX_SCALE), + "BSND", "BSND"); + std::fclose(manifest); + + std::printf("QSMLA_STAGE0 case=%s output=%s elements=%zu\n", + QCASE_NAME, output_dir.string().c_str(), out.size()); + return 0; +} diff --git a/benchmark/one-level-arch/test/kernel/fa/src/quant_sparse_flash_mla.cpp b/benchmark/one-level-arch/test/kernel/fa/src/quant_sparse_flash_mla.cpp new file mode 100644 index 00000000..44ad86ca --- /dev/null +++ b/benchmark/one-level-arch/test/kernel/fa/src/quant_sparse_flash_mla.cpp @@ -0,0 +1,212 @@ +#include +#include "benchmark.h" +#include "fileop.h" +// #include "fa/quant_sparse_flash_mla_pto.hpp" +#if defined(QSMLA_USE_TADD) || defined(QSMLA_USE_TADD_4PE) +#include "fa/quant_sparse_flash_mla_tadd_pto.hpp" +#else +#include "fa/quant_sparse_flash_mla_onepass_pto.hpp" +#endif + +#ifndef Tbatch +#define B 1 +#else +#define B Tbatch +#endif + +#ifndef Tn1 +#define N1 1 +#else +#define N1 Tn1 +#endif + +#ifndef Tn2 +#define N2 1 +#else +#define N2 Tn2 +#endif + +#ifndef Ts1 +#define s1 64 +#else +#define s1 Ts1 +#endif + +#ifndef Ts2 +#define s2 128 +#else +#define s2 Ts2 +#endif + +#ifndef Td +#define D 512 +#else +#define D Td +#endif + +#ifndef Tm +#define kTm 32 +#else +#define kTm Tm +#endif + +#ifndef Tk +#define kTk 32 +#else +#define kTk Tk +#endif + +#ifndef Td_block +#define kTd 64 +#else +#define kTd Td_block +#endif + +#ifndef Tsoftmax_scale +#define softmax_scale_val 0.125f +#else +#define softmax_scale_val Tsoftmax_scale +#endif + +#ifndef Twleft +#define win_left 1 +#else +#define win_left Twleft +#endif + +#ifndef Twright +#define win_right 1 +#else +#define win_right Twright +#endif + +#define ALIGN_MASK 0xfffffffffffff000ull +#define ALIGN 4*1024 +#define MAP_MEM_BASE 0x4000802000ULL + +constexpr uint64_t align_up_4k(uint64_t bytes) { + return (bytes + ALIGN - 1) & ALIGN_MASK; +} + +static void init_deterministic(__half* data, int count, int seed) { + for (int i = 0; i < count; ++i) { + float val = ((float)((i * 31 + seed * 17) % 100)) / 100.0f - 0.5f; + data[i] = (__half)val; + } +} + +int main(){ + using qdtype = __half; + using kvdtype = __half; + using odttype = __half; + constexpr int group_size = N1 / N2; + constexpr int g_slice_max = group_size < 64 ? group_size : 64; + using Config = QsmlaConfig< + B, s1, s2, N1, N2, D, 0, kTm, kTk, kTd, g_slice_max>; + + qdtype qp[B*s1*N1*D + 2*ALIGN]; + kvdtype kvp[B*s2*N2*D + 2*ALIGN]; + + qdtype* q = (qdtype*)(((uint64_t)qp & ALIGN_MASK) + ALIGN); + kvdtype* kv = (kvdtype*)(((uint64_t)kvp & ALIGN_MASK) + ALIGN); + + odttype* out = (odttype*)MAP_MEM_BASE; + +#ifdef QSMLA_USE_TADD_4PE + // Cooperative scratch must live in shared GM rather than a PE-private + // function stack. Each PE computes and receives these identical mapped + // addresses, matching the shared-workspace contract used by the 4-PE FA. + constexpr uint64_t out_bytes = + static_cast(B) * s1 * N1 * D * sizeof(odttype); + constexpr uint64_t score_bytes = + static_cast(kTm) * kTk * sizeof(float); + constexpr uint64_t prob_bytes = + static_cast(kTm) * kTk * sizeof(qdtype); + constexpr uint64_t score_addr = MAP_MEM_BASE + align_up_4k(out_bytes); + constexpr uint64_t prob_addr = score_addr + align_up_4k(score_bytes); + constexpr uint64_t pv_addr = prob_addr + align_up_4k(prob_bytes); + float* score_scratch = reinterpret_cast(score_addr); + qdtype* prob_scratch = reinterpret_cast(prob_addr); + float* pv_scratch = reinterpret_cast(pv_addr); +#endif + + init_deterministic(q, B*s1*N1*D, 1); + init_deterministic(kv, B*s2*N2*D, 2); + + BENCHSTART; +#ifdef QSMLA_USE_TADD_4PE + static_assert(N1 > 1, + "tadd_4pe is a BSND G-slice implementation"); + quant_sparse_flash_mla_swa_tadd_4pe_bsnd_pto< + qdtype, kvdtype, odttype, Config>( + out, q, kv, + softmax_scale_val, + win_left, + win_right, + (float*)nullptr, + (float*)nullptr, + (int*)nullptr, + (int*)nullptr, + (int*)nullptr, + (int*)nullptr, + (int*)nullptr, + (int*)nullptr, + (float*)nullptr, + (int*)nullptr, + (float*)nullptr, + score_scratch, + prob_scratch, + pv_scratch); +#else + if constexpr (N1 == 1 && N2 == 1) { +#ifdef QSMLA_USE_TADD + quant_sparse_flash_mla_swa_tadd_config_pto< +#else + quant_sparse_flash_mla_swa_onepass_config_pto< +#endif + qdtype, kvdtype, odttype, Config>( + out, q, kv, + softmax_scale_val, + win_left, + win_right, + (float*)nullptr, // q_descale + (float*)nullptr, // ori_kv_descale + (int*)nullptr, // ori_sparse_indices + (int*)nullptr, // ori_block_table + (int*)nullptr, // cu_seqlens_q + (int*)nullptr, // cu_seqlens_ori_kv + (int*)nullptr, // seqused_q + (int*)nullptr, // seqused_ori_kv + (float*)nullptr, // sinks + (int*)nullptr, // metadata + (float*)nullptr // softmax_lse + ); + } else { +#ifdef QSMLA_USE_TADD + quant_sparse_flash_mla_swa_tadd_bsnd_pto< +#else + quant_sparse_flash_mla_swa_onepass_bsnd_pto< +#endif + qdtype, kvdtype, odttype, Config>( + out, q, kv, + softmax_scale_val, + win_left, + win_right, + (float*)nullptr, // q_descale + (float*)nullptr, // ori_kv_descale + (int*)nullptr, // ori_sparse_indices + (int*)nullptr, // ori_block_table + (int*)nullptr, // cu_seqlens_q + (int*)nullptr, // cu_seqlens_ori_kv + (int*)nullptr, // seqused_q + (int*)nullptr, // seqused_ori_kv + (float*)nullptr, // sinks + (int*)nullptr, // metadata + (float*)nullptr // softmax_lse + ); + } +#endif + BENCHEND; + + return 0; +}