From f92e8d04cc5c4005944dd2256a3c12a67b1184a4 Mon Sep 17 00:00:00 2001 From: chenglongyu Date: Thu, 6 Aug 2026 09:36:35 +0800 Subject: [PATCH 1/9] Add Quant Sparse Flash mla (QSMLA) operator --- .../fa/quant_sparse_flash_mla_onepass_pto.hpp | 255 +++++++++++++ .../kernels/fa/quant_sparse_flash_mla_pto.hpp | 350 ++++++++++++++++++ .../fa/quant_sparse_flash_mla_tadd_pto.hpp | 294 +++++++++++++++ .../one-level-arch/test/kernel/fa/Makefile | 23 ++ .../test/kernel/fa/src/qsmla_compare.py | 122 ++++++ .../test/kernel/fa/src/qsmla_cpu_ref.cpp | 270 ++++++++++++++ .../kernel/fa/src/quant_sparse_flash_mla.cpp | 120 ++++++ .../test/kernel/fa/src/test_memwrite.cpp | 18 + .../test/kernel/fa/src/test_tload_store.cpp | 60 +++ 9 files changed, 1512 insertions(+) create mode 100644 benchmark/one-level-arch/kernels/fa/quant_sparse_flash_mla_onepass_pto.hpp create mode 100644 benchmark/one-level-arch/kernels/fa/quant_sparse_flash_mla_pto.hpp create mode 100644 benchmark/one-level-arch/kernels/fa/quant_sparse_flash_mla_tadd_pto.hpp create mode 100644 benchmark/one-level-arch/test/kernel/fa/src/qsmla_compare.py create mode 100644 benchmark/one-level-arch/test/kernel/fa/src/qsmla_cpu_ref.cpp create mode 100644 benchmark/one-level-arch/test/kernel/fa/src/quant_sparse_flash_mla.cpp create mode 100644 benchmark/one-level-arch/test/kernel/fa/src/test_memwrite.cpp create mode 100644 benchmark/one-level-arch/test/kernel/fa/src/test_tload_store.cpp 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..c4edbe2b --- /dev/null +++ b/benchmark/one-level-arch/kernels/fa/quant_sparse_flash_mla_onepass_pto.hpp @@ -0,0 +1,255 @@ +#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 方式】 +// TADD: mask=float[s1*s2], 0.0(有效)/-1e30(无效), score += mask +// +// 【切换方法】 +// test 文件中: +// #include "fa/quant_sparse_flash_mla_onepass_pto.hpp" +// quant_sparse_flash_mla_swa_onepass_pto<...>(...) +// ============================================================================= + +#include +#include "template_asm.h" + +using namespace pto; + +static inline void build_swa_mask_onepass( + float* mask, 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); + mask[q * s2 + kv] = valid ? 0.0f : -1e30f; + } + } +} + +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) +{ + constexpr int Db = D / kTd; + + float mask_buf[s1 * s2]; + build_swa_mask_onepass(mask_buf, s1, s2, ori_win_left, ori_win_right); + + using gmQ = global_tensor>; + using gmKV = global_tensor>; + using gmO = global_tensor>; + using gmMask = global_tensor>; + + using tileQ = TileLeft; + using tileKV = TileRight; + using tileW_out = TileAcc; + 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 itKV = global_iterator; + using itV = global_iterator; + using itO = global_iterator; + using itMask = global_iterator; + + itQ gIterQ(q_ptr); + itKV gIterKV(ori_kv_ptr); + itV gIterV(ori_kv_ptr); + itO gIterO(out_ptr); + itMask gIterMask(mask_buf); + + 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) { + + // ============================================================ + // 一遍式 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 < Kb; ++j) { + + // --- Step 1: QK^T 沿全 D 累加 --- + tileW_out tW_out; + bool first_d = true; + #pragma clang loop unroll(full) + for (int dd = 0; dd < Db; ++dd) { + tileQ tQ; + auto gQ = gIterQ(i, dd); + TLOAD(tQ, gQ); + + tileKV tK; + auto gK = gIterKV(j, dd); + TLOAD(tK, gK); + + if (first_d) { + TMATMUL(tW_out, tQ, tK); + first_d = false; + } else { + TMATMUL_ACC(tW_out, tQ, tK); + } + } + + // --- Step 2: scale + mask --- + tileW tW; + ACCCVT(tW, tW_out); + TMULS(tW, tW, scale); + + tileMask tMask; + auto gMask = gIterMask(i, 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_cast tExpW; + TCVT(tExpW, tW); + tileW_left tW_left; + TCVT(tW_left, tExpW); + + #pragma clang loop unroll(full) + for (int dd = 0; dd < Db; ++dd) { + tileV tV; + auto gV = gIterV(j, dd); + TLOAD(tV, gV); + + tileO_out tPV_out; + TMATMUL(tPV_out, tW_left, tV); + tileO tPV; + ACCCVT(tPV, tPV_out); + + 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); + } + } +} + +#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..e9cc83f0 --- /dev/null +++ b/benchmark/one-level-arch/kernels/fa/quant_sparse_flash_mla_pto.hpp @@ -0,0 +1,350 @@ +#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 位打包格式 (每 32 列 1 个 uint32, bit=1 表示无效/被 mask) +// TSEL(dst, mask, neg_inf): mask bit=1 → dst=-1e30 (无效), bit=0 → 保持原 score +// 窗口范围 (对第 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 + TMATMUL_ACC) +// PV: 每个 D 分块独立计算并存储 +// +// 【两遍式】 +// Pass 1: online softmax 归约 (m, l), 含 mask +// Pass 2: 归一化 P, 计算 P@V, 含 mask +// +// 【mask tile 格式】 +// TSEL 的 mask 必须是 UINT32 位打包格式: +// maskWordsPerRow = ceil(kTk / 32) +// mask tile shape: [kTm, maskWordsPerRow], dtype=uint32_t, RowMajor +// bit (1< +#include "template_asm.h" + +using namespace pto; + +// CPU 侧 mask 预计算 (在 kernel 内部调用, 放在 stack 上) +// 生成 UINT32 位打包格式的 mask: +// 对 [s1, s2] 的每个 (q, kv), 如果 kv 在窗口外则对应 bit=1 +// 每 32 个 kv 列打包成 1 个 uint32 +// maskBuf 行大小 = ceil(s2 / 32) 个 uint32 +static inline void build_swa_mask_bitpacked( + uint32_t* maskBuf, int s1, int s2, int win_left, int win_right) +{ + const int wordsPerRow = (s2 + 31) / 32; + 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; + uint32_t* row = maskBuf + q * wordsPerRow; + for (int w = 0; w < wordsPerRow; ++w) row[w] = 0; + for (int kv = 0; kv < s2; ++kv) { + bool valid = (kv >= lo) && (kv <= hi); + if (!valid) { + row[kv / 32] |= (uint32_t{1} << (kv % 32)); + } + } + } +} + +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; + + // === mask 预计算 (UINT32 位打包) === + // 全局 mask: [s1, ceil(s2/32)] uint32 + constexpr int maskWordsPerRowGlobal = (s2 + 31) / 32; + uint32_t maskBuf[s1 * maskWordsPerRowGlobal]; //[64 * 4] + build_swa_mask_bitpacked(maskBuf, s1, s2, ori_win_left, ori_win_right); + + // 每个 [kTm, kTk] block 对应的 mask tile: + // maskWordsPerRowBlock = ceil(kTk / 32) + // mask tile: [kTm, maskWordsPerRowBlock], uint32, RowMajor + constexpr int maskWordsPerRowBlock = (kTk + 31) / 32; // 1 + + using gmQ = global_tensor>; + using gmKV = global_tensor>; + using gmO = global_tensor>; + + using tileQ = TileLeft; + using tileKV = TileRight; + using tileW_out = TileAcc; + + // score tile: RowMajor float + using tileW = Tile; + + // mask tile: RowMajor uint32, shape [kTm, maskWordsPerRowBlock] + // TSEL 要求 dst/mask/src 同 tile_shape, 但模拟器内部 mask 按 uint32 位打包读 + // 这里用 float 类型满足编译器约束, 实际数据是 uint32 位掩码 + // tile Cols 设为 kTk (与 score tile 同形状), 模拟器只读前 maskWordsPerRowBlock 个 uint32 + 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 itKV = global_iterator; + using itV = global_iterator; + using itO = global_iterator; + + itQ gIterQ(q_ptr); + itKV gIterKV(ori_kv_ptr); + itV gIterV(ori_kv_ptr); + itO gIterO(out_ptr); + + 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_out tW_out; + bool first_d = true; + #pragma clang loop unroll(full) + for (int dd = 0; dd < Db; ++dd) { + tileQ tQ; + auto gQ = gIterQ(i, dd); + TLOAD(tQ, gQ); + + tileKV tK; + auto gK = gIterKV(j, dd); + TLOAD(tK, gK); + + if (first_d) { + TMATMUL(tW_out, tQ, tK); + first_d = false; + } else { + TMATMUL_ACC(tW_out, tQ, tK); + } + } + + tileW tW; + ACCCVT(tW, tW_out); + TMULS(tW, tW, scale); + + // 应用 token 级 mask: TSEL(score, mask, neg_inf) + // mask bit=1 → score=-1e30 (无效), bit=0 → 保持原 score + // mask 数据从全局 maskBuf 中提取当前 [kTm, kTk] block 对应的位打包数据 + // 构建局部 mask tile: 从全局 [s1, maskWordsPerRowGlobal] 中提取 + // [kTm 行, kTk 列] 对应的 bit, 重新打包为 [kTm, maskWordsPerRowBlock] + { + // 构建 block mask: [kTm, maskWordsPerRowBlock] uint32 + // 从全局 mask 中提取列 [j*kTk, (j+1)*kTk) 的 bits + uint32_t blockMask[kTm * maskWordsPerRowBlock]; // [32 * 1] + for (int r = 0; r < kTm; ++r) { + int q_idx = i * kTm + r; + const uint32_t* globalRow = maskBuf + q_idx * maskWordsPerRowGlobal; + uint32_t* blockRow = blockMask + r * maskWordsPerRowBlock; + for (int w = 0; w < maskWordsPerRowBlock; ++w) blockRow[w] = 0; + for (int c = 0; c < kTk; ++c) { + int global_col = j * kTk + c; + if (global_col >= s2) { + blockRow[c / 32] |= (uint32_t{1} << (c % 32)); + continue; + } + bool masked = (globalRow[global_col / 32] >> (global_col % 32)) & 1; + if (masked) { + blockRow[c / 32] |= (uint32_t{1} << (c % 32)); + } + } + } + // TLOAD mask tile from blockMask (stack buffer) + using gmBlockMask = global_tensor>; + gmBlockMask gBlockMask(blockMask); + tileMask tMask; + TLOAD(tMask, gBlockMask); + TSEL(tW, tMask, tNegInf); + } + + // 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_out tW_out; + bool first_d = true; + #pragma clang loop unroll(full) + for (int dd2 = 0; dd2 < Db; ++dd2) { + tileQ tQ; + auto gQ = gIterQ(i, dd2); + TLOAD(tQ, gQ); + + tileKV tK; + auto gK = gIterKV(j, dd2); + TLOAD(tK, gK); + + if (first_d) { + TMATMUL(tW_out, tQ, tK); + first_d = false; + } else { + TMATMUL_ACC(tW_out, tQ, tK); + } + } + + tileW tW; + ACCCVT(tW, tW_out); + TMULS(tW, tW, scale); + + // 应用 token 级 mask: TSEL(score, mask, neg_inf) + { + uint32_t blockMask[kTm * maskWordsPerRowBlock]; + for (int r = 0; r < kTm; ++r) { + int q_idx = i * kTm + r; + const uint32_t* globalRow = maskBuf + q_idx * maskWordsPerRowGlobal; + uint32_t* blockRow = blockMask + r * maskWordsPerRowBlock; + for (int w = 0; w < maskWordsPerRowBlock; ++w) blockRow[w] = 0; + for (int c = 0; c < kTk; ++c) { + int global_col = j * kTk + c; + if (global_col >= s2) { + blockRow[c / 32] |= (uint32_t{1} << (c % 32)); + continue; + } + bool masked = (globalRow[global_col / 32] >> (global_col % 32)) & 1; + if (masked) { + blockRow[c / 32] |= (uint32_t{1} << (c % 32)); + } + } + } + using gmBlockMask = global_tensor>; + gmBlockMask gBlockMask(blockMask); + tileMask tMask; + TLOAD(tMask, gBlockMask); + TSEL(tW, tMask, tNegInf); + } + + // 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; + TCVT(tW_left, tExpW); + + // PV = p * V (当前 D 分块) + tileV tV; + auto gV = gIterV(j, dd); + TLOAD(tV, gV); + + tileO_out tPV_out; + TMATMUL(tPV_out, tW_left, tV); + tileO tPV; + ACCCVT(tPV, tPV_out); + + TADD(tO, tO, tPV); + } + + // 写回 O 分块 [kTm, kTd] + tileO_cast tO_cast; + TCVT(tO_cast, tO); + auto gO = gIterO(i, dd); + TSTORE(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..161f8e29 --- /dev/null +++ b/benchmark/one-level-arch/kernels/fa/quant_sparse_flash_mla_tadd_pto.hpp @@ -0,0 +1,294 @@ +#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 + TMATMUL_ACC) +// PV: 每个 D 分块独立计算并存储 +// +// 【两遍式】 +// Pass 1: online softmax 归约 (m, l), 含 mask +// Pass 2: 归一化 P, 计算 P@V, 含 mask +// ============================================================================= + +#include +#include "template_asm.h" + +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( + float* mask, 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); + mask[q * s2 + kv] = valid ? 0.0f : -1e30f; + } + } +} + +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) +{ + constexpr int Db = D / kTd; + + // kernel 内部计算 mask, 放在 stack 上 + // s1*s2*sizeof(float) = 64*128*4 = 32KB (可容纳) + float mask_buf[s1 * s2]; + build_swa_mask(mask_buf, s1, s2, ori_win_left, ori_win_right); + + using gmQ = global_tensor>; + using gmKV = global_tensor>; + using gmO = global_tensor>; + using gmMask = global_tensor>; + + using tileQ = TileLeft; + using tileKV = TileRight; + using tileW_out = TileAcc; + + // score tile 与 mask tile 都用 RowMajor, 保证 TADD 类型一致 + 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 itKV = global_iterator; + using itV = global_iterator; + using itO = global_iterator; + using itMask = global_iterator; + + itQ gIterQ(q_ptr); + itKV gIterKV(ori_kv_ptr); + itV gIterV(ori_kv_ptr); + itO gIterO(out_ptr); + itMask gIterMask(mask_buf); + + 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); + + for (int j = 0; j < Kb; ++j) { + + // QK^T 沿 D 维累加 + tileW_out tW_out; + bool first_d = true; + #pragma clang loop unroll(full) + for (int dd = 0; dd < Db; ++dd) { + tileQ tQ; + auto gQ = gIterQ(i, dd); + TLOAD(tQ, gQ); + + tileKV tK; + auto gK = gIterKV(j, dd); + TLOAD(tK, gK); + + if (first_d) { + TMATMUL(tW_out, tQ, tK); + first_d = false; + } else { + TMATMUL_ACC(tW_out, tQ, tK); + } + } + + tileW tW; + ACCCVT(tW, tW_out); + TMULS(tW, tW, scale); + + // 应用 token 级 mask: score += mask (0 保持原值, -1e30 屏蔽) + tileMask tMask; + auto gMask = gIterMask(i, 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 < Kb; ++j) { + + // 计算完整 QK^T (沿 D 累加, 与 Pass 1 一致) + tileW_out tW_out; + bool first_d = true; + #pragma clang loop unroll(full) + for (int dd2 = 0; dd2 < Db; ++dd2) { + tileQ tQ; + auto gQ = gIterQ(i, dd2); + TLOAD(tQ, gQ); + + tileKV tK; + auto gK = gIterKV(j, dd2); + TLOAD(tK, gK); + + if (first_d) { + TMATMUL(tW_out, tQ, tK); + first_d = false; + } else { + TMATMUL_ACC(tW_out, tQ, tK); + } + } + + tileW tW; + ACCCVT(tW, tW_out); + TMULS(tW, tW, scale); + + // 应用 token 级 mask: score += mask (0 保持原值, -1e30 屏蔽) + tileMask tMask; + auto gMask = gIterMask(i, j); + TLOAD(tMask, gMask); + TADD(tW, tW, tMask); + + // 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; + TCVT(tW_left, tExpW); + + // PV = p * V (当前 D 分块) + tileV tV; + auto gV = gIterV(j, dd); + TLOAD(tV, gV); + + tileO_out tPV_out; + TMATMUL(tPV_out, tW_left, tV); + tileO tPV; + ACCCVT(tPV, tPV_out); + + TADD(tO, tO, tPV); + } + + // 写回 O 分块 [kTm, kTd] + tileO_cast tO_cast; + TCVT(tO_cast, tO); + auto gO = gIterO(i, dd); + TSTORE(gO, tO_cast); + } + } +} + +#endif diff --git a/benchmark/one-level-arch/test/kernel/fa/Makefile b/benchmark/one-level-arch/test/kernel/fa/Makefile index 86e099d6..7666e11f 100644 --- a/benchmark/one-level-arch/test/kernel/fa/Makefile +++ b/benchmark/one-level-arch/test/kernel/fa/Makefile @@ -84,6 +84,29 @@ 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 + s1 ?= 64 + s2 ?= 128 + D ?= 512 + Tm ?= 32 + Tk ?= 32 + Td_block ?= 64 + softmax_scale ?= 0.125 + wleft ?= 1 + wright ?= 1 + DEFINES += -DTs1=$(s1) + DEFINES += -DTs2=$(s2) + 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) + TARGET = $(ELF_HEAD)/$(TESTCASE)_s1$(s1)_s2$(s2)_D$(D)_Tm$(Tm)_Tk$(Tk)_Td$(Td_block).elf +endif + include ../../common/Makefile.common clean: 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..43464b81 --- /dev/null +++ b/benchmark/one-level-arch/test/kernel/fa/src/qsmla_compare.py @@ -0,0 +1,122 @@ +#!/usr/bin/env python3 +""" +Compare NPU output (FP16, dumped from gfrun --dump-memory) with CPU golden output (FP32). +Usage: python3 qsmla_compare.py [atol] [rtol] + +NPU dump format: raw FP16 (__half, IEEE 754 half-precision, 2 bytes per element) +Golden format: raw FP32 (float, IEEE 754 single-precision, 4 bytes per element) +""" + +import sys +import struct +import math +import numpy as np + +def read_fp16(filename, count=None): + """Read binary file as float16 array, convert to float32.""" + with open(filename, 'rb') as f: + data = f.read() + arr = np.frombuffer(data, dtype=np.float16) + if count is not None: + arr = arr[:count] + return arr.astype(np.float32) + +def read_fp32(filename, count=None): + """Read binary file as float32 array.""" + with open(filename, 'rb') as f: + data = f.read() + arr = np.frombuffer(data, dtype=np.float32) + if count is not None: + arr = arr[:count] + return arr + +def compare(npu_out, golden_out, atol=0.01, rtol=0.05): + """Compare two float arrays with atol/rtol tolerance.""" + min_len = min(len(npu_out), len(golden_out)) + if len(npu_out) != len(golden_out): + print(f"WARNING: length mismatch npu={len(npu_out)} golden={len(golden_out)}, comparing first {min_len}") + + max_abs_err = 0.0 + max_rel_err = 0.0 + fail_count = 0 + fail_examples = [] + + for i in range(min_len): + n_val = float(npu_out[i]) + g_val = float(golden_out[i]) + + if math.isnan(n_val) or math.isnan(g_val): + fail_count += 1 + if len(fail_examples) < 10: + fail_examples.append((i, n_val, g_val, "NaN")) + continue + if math.isinf(n_val) or math.isinf(g_val): + fail_count += 1 + if len(fail_examples) < 10: + fail_examples.append((i, n_val, g_val, "Inf")) + continue + + abs_err = abs(n_val - g_val) + rel_err = abs_err / max(abs(g_val), 1e-8) + + if abs_err > max_abs_err: + max_abs_err = abs_err + if rel_err > max_rel_err: + max_rel_err = rel_err + + if abs_err > atol and rel_err > rtol: + fail_count += 1 + if len(fail_examples) < 10: + fail_examples.append((i, n_val, g_val, abs_err, rel_err)) + + print(f"=== QSMLA Precision Verification ===") + print(f"Compared elements: {min_len}") + print(f"Max absolute error: {max_abs_err:.8f}") + print(f"Max relative error: {max_rel_err:.8f}") + print(f"Tolerance: atol={atol}, rtol={rtol}") + print(f"Failed elements: {fail_count}/{min_len}") + + if fail_examples: + print(f"\nFirst {len(fail_examples)} failures:") + for ex in fail_examples: + if len(ex) == 4: + print(f" idx={ex[0]}: npu={ex[1]:.8f} golden={ex[2]:.8f} [{ex[3]}]") + else: + print(f" idx={ex[0]}: npu={ex[1]:.8f} golden={ex[2]:.8f} abs_err={ex[3]:.8f} rel_err={ex[4]:.8f}") + + # Print first 8 values for sanity + print(f"\nFirst 8 values comparison:") + for i in range(min(8, min_len)): + n = float(npu_out[i]) + g = float(golden_out[i]) + print(f" [{i}] npu={n:.8f} golden={g:.8f} diff={abs(n-g):.8f}") + + if fail_count == 0: + print("\n=== RESULT: PASS ===") + return 0 + else: + pass_rate = (min_len - fail_count) / min_len * 100 + print(f"\n=== RESULT: FAIL (pass rate: {pass_rate:.2f}%) ===") + return 1 + +def main(): + if len(sys.argv) < 3: + print("Usage: python3 qsmla_compare.py [atol] [rtol]") + sys.exit(1) + + npu_file = sys.argv[1] + golden_file = sys.argv[2] + atol = float(sys.argv[3]) if len(sys.argv) > 3 else 0.01 + rtol = float(sys.argv[4]) if len(sys.argv) > 4 else 0.05 + + npu_out = read_fp16(npu_file) + golden_out = read_fp32(golden_file) + + print(f"NPU output: {len(npu_out)} FP16 values from {npu_file}") + print(f"Golden output: {len(golden_out)} FP32 values from {golden_file}") + + ret = compare(npu_out, golden_out, atol, rtol) + sys.exit(ret) + +if __name__ == '__main__': + 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..ee6f6a6c --- /dev/null +++ b/benchmark/one-level-arch/test/kernel/fa/src/qsmla_cpu_ref.cpp @@ -0,0 +1,270 @@ +// CPU reference implementation for quant_sparse_flash_mla SWA mode +// Computes: O = softmax(Q @ K^T * softmax_scale + mask) @ V +// Where K=V=ori_kv (MLA shared KV), with token-level sliding window mask + +#include +#include +#include +#include +#include +#include + +#ifndef S1 +#define S1 64 +#endif +#ifndef S2 +#define S2 128 +#endif +#ifndef D +#define D 512 +#endif +#ifndef KTM +#define KTM 32 +#endif +#ifndef KTK +#define KTK 32 +#endif +#ifndef WIN_LEFT +#define WIN_LEFT 1 +#endif +#ifndef WIN_RIGHT +#define WIN_RIGHT 1 +#endif +#ifndef SOFTMAX_SCALE +#define SOFTMAX_SCALE 0.125f +#endif + +typedef float f32_t; + +static void init_deterministic_f32(f32_t* 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] = val; + } +} + +// Build token-level SWA mask (SEL semantics) +// mask[q * s2 + kv] = true if kv 在窗口外 (被 mask, score 置 -1e30) +// = false if kv 在窗口内 (有效, 保持原 score) +// valid: diagonal - win_left <= kv <= diagonal + win_right +// where diagonal = (s2 - s1) + q +// Apply: score = mask ? -1e30 : score (TSEL semantics) +static void build_swa_mask(bool* mask, 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); + mask[q * s2 + kv] = !valid; + } + } +} + +// CPU reference: SWA MLA attention with token-level mask (TSEL semantics) +// Q: [s1, D], KV: [s2, D] (K=V), O: [s1, D] +void cpu_quant_sparse_flash_mla_swa( + f32_t* out, const f32_t* q, const f32_t* kv, + const bool* mask, + int s1, int s2, int d, int kTm, int kTk, + float softmax_scale, int win_left, int win_right) +{ + const int Qb = (s1 + kTm - 1) / kTm; + const int Kb = (s2 + kTk - 1) / kTk; + + f32_t* score = (f32_t*)malloc(kTm * kTk * sizeof(f32_t)); + f32_t* p = (f32_t*)malloc(kTm * kTk * sizeof(f32_t)); + f32_t* pv = (f32_t*)malloc(kTm * d * sizeof(f32_t)); + + for (int qi = 0; qi < Qb; ++qi) { + + // Pass 1: online softmax with mask + f32_t row_max[kTm]; + f32_t row_sum[kTm]; + for (int r = 0; r < kTm; ++r) { + row_max[r] = -1e30f; + row_sum[r] = 0.0f; + } + + for (int j = 0; j < Kb; ++j) { + for (int r = 0; r < kTm; ++r) { + int q_row = qi * kTm + r; + if (q_row >= s1) continue; + for (int c = 0; c < kTk; ++c) { + int kv_row = j * kTk + c; + if (kv_row >= s2) { score[r * kTk + c] = -1e30f; continue; } + f32_t dot = 0.0f; + for (int dd = 0; dd < d; ++dd) { + dot += q[q_row * d + dd] * kv[kv_row * d + dd]; + } + // Apply mask: TSEL semantics (mask=true → -1e30, mask=false → keep score) + f32_t raw_score = dot * softmax_scale; + score[r * kTk + c] = mask[q_row * s2 + kv_row] ? -1e30f : raw_score; + } + } + + for (int r = 0; r < kTm; ++r) { + int q_row = qi * kTm + r; + if (q_row >= s1) continue; + + f32_t local_max = -1e30f; + for (int c = 0; c < kTk; ++c) { + int kv_row = j * kTk + c; + if (kv_row >= s2) continue; + if (score[r * kTk + c] > local_max) + local_max = score[r * kTk + c]; + } + + f32_t new_max = (row_max[r] > local_max) ? row_max[r] : local_max; + f32_t scale_old = expf(row_max[r] - new_max); + row_sum[r] *= scale_old; + + for (int c = 0; c < kTk; ++c) { + int kv_row = j * kTk + c; + if (kv_row >= s2) continue; + row_sum[r] += expf(score[r * kTk + c] - new_max); + } + + row_max[r] = new_max; + } + } + + // Pass 2: compute P @ V with mask + for (int dd = 0; dd < d; ++dd) { + for (int r = 0; r < kTm; ++r) { + pv[r * d + dd] = 0.0f; + } + } + + for (int j = 0; j < Kb; ++j) { + for (int r = 0; r < kTm; ++r) { + int q_row = qi * kTm + r; + if (q_row >= s1) continue; + for (int c = 0; c < kTk; ++c) { + int kv_row = j * kTk + c; + if (kv_row >= s2) { p[r * kTk + c] = 0.0f; continue; } + f32_t dot = 0.0f; + for (int dd = 0; dd < d; ++dd) { + dot += q[q_row * d + dd] * kv[kv_row * d + dd]; + } + // Apply mask: TSEL semantics (mask=true → -1e30, mask=false → keep score) + f32_t raw_score = dot * softmax_scale; + f32_t s = mask[q_row * s2 + kv_row] ? -1e30f : raw_score; + p[r * kTk + c] = expf(s - row_max[r]) / row_sum[r]; + } + } + + for (int r = 0; r < kTm; ++r) { + int q_row = qi * kTm + r; + if (q_row >= s1) continue; + for (int dd = 0; dd < d; ++dd) { + for (int c = 0; c < kTk; ++c) { + int kv_row = j * kTk + c; + if (kv_row >= s2) continue; + pv[r * d + dd] += p[r * kTk + c] * kv[kv_row * d + dd]; + } + } + } + } + + for (int r = 0; r < kTm; ++r) { + int q_row = qi * kTm + r; + if (q_row >= s1) continue; + for (int dd = 0; dd < d; ++dd) { + out[q_row * d + dd] = pv[r * d + dd]; + } + } + } + + free(score); + free(p); + free(pv); +} + +int main() { + int q_count = S1 * D; + int kv_count = S2 * D; + int out_count = S1 * D; + int mask_count = S1 * S2; + + f32_t* q = (f32_t*)malloc(q_count * sizeof(f32_t)); + f32_t* kv = (f32_t*)malloc(kv_count * sizeof(f32_t)); + f32_t* out = (f32_t*)malloc(out_count * sizeof(f32_t)); + bool* mask = (bool*)malloc(mask_count * sizeof(bool)); + + init_deterministic_f32(q, q_count, 1); + init_deterministic_f32(kv, kv_count, 2); + + // for (int i = 0; i < 100; i++){ + // std::cout << q[i] << ", "; + // } + + // std::cout << std::endl; + + // for (int i = 0; i < 100; i++){ + // std::cout << kv[i] << ", "; + // } + + // return 0; + + memset(out, 0, out_count * sizeof(f32_t)); + + // f32_t mask1[5][10]; + + // build_swa_mask(&mask1[0][0], 5, 10, 2, 2); + + // for(int i = 0; i < 5; i++){ + // for(int j = 0; j < 10; j++){ + // std::cout << mask1[i][j] << ", \t"; + // } + // std::cout << std::endl; + // } + + // return 0; + + + build_swa_mask(mask, S1, S2, WIN_LEFT, WIN_RIGHT); + + printf("QSMLA_CPU: s1=%d s2=%d D=%d win_left=%d win_right=%d\n", + S1, S2, D, WIN_LEFT, WIN_RIGHT); + printf("QSMLA_CPU: computing reference with token-level mask...\n"); + fflush(stdout); + + cpu_quant_sparse_flash_mla_swa(out, q, kv, mask, + S1, S2, D, KTM, KTK, + SOFTMAX_SCALE, WIN_LEFT, WIN_RIGHT); + + const char* golden_path = "qsmla_golden.bin"; + FILE* f = fopen(golden_path, "wb"); + if (!f) { + fprintf(stderr, "Failed to open %s\n", golden_path); + return 1; + } + fwrite(out, sizeof(f32_t), out_count, f); + fclose(f); + + printf("QSMLA_CPU: golden output written to %s (%d floats, %d bytes)\n", + golden_path, out_count, (int)(out_count * sizeof(f32_t))); + + printf("QSMLA_CPU: first 8 output values:\n"); + for (int i = 0; i < 8 && i < out_count; ++i) { + printf(" out[%d] = %.6f\n", i, out[i]); + } + + FILE* fq = fopen("qsmla_input_q.bin", "wb"); + fwrite(q, sizeof(f32_t), q_count, fq); + fclose(fq); + + FILE* fkv = fopen("qsmla_input_kv.bin", "wb"); + fwrite(kv, sizeof(f32_t), kv_count, fkv); + fclose(fkv); + + printf("QSMLA_CPU: input data written to qsmla_input_q.bin and qsmla_input_kv.bin\n"); + + free(q); + free(kv); + free(out); + free(mask); + 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..47edbb49 --- /dev/null +++ b/benchmark/one-level-arch/test/kernel/fa/src/quant_sparse_flash_mla.cpp @@ -0,0 +1,120 @@ +#include +#include "benchmark.h" +#include "fileop.h" +// #include "fa/quant_sparse_flash_mla_pto.hpp" +// #include "fa/quant_sparse_flash_mla_tadd_pto.hpp" +#include "fa/quant_sparse_flash_mla_onepass_pto.hpp" + +#define B 1 +#define H 1 + +#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 + +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; + + qdtype qp[B*H*s1*D + 2*ALIGN]; + kvdtype kvp[B*H*s2*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; + + init_deterministic(q, B*H*s1*D, 1); + init_deterministic(kv, B*H*s2*D, 2); + + BENCHSTART; + for(int i=0;i( + out + i*H*s1*D + j*s1*D, + q + i*H*s1*D + j*s1*D, + kv + i*H*s2*D + j*s2*D, + 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 + ); + } + } + BENCHEND; + + return 0; +} diff --git a/benchmark/one-level-arch/test/kernel/fa/src/test_memwrite.cpp b/benchmark/one-level-arch/test/kernel/fa/src/test_memwrite.cpp new file mode 100644 index 00000000..cf97df18 --- /dev/null +++ b/benchmark/one-level-arch/test/kernel/fa/src/test_memwrite.cpp @@ -0,0 +1,18 @@ +#include +#include "benchmark.h" + +#define MAP_MEM_BASE 0x4000802000ULL +#define OUT_COUNT 32768 + +int main(){ + __half* out = (__half*)MAP_MEM_BASE; + + for (int i = 0; i < OUT_COUNT; ++i) { + out[i] = (__half)((float)i * 0.001f); + } + + BENCHSTART; + BENCHEND; + + return 0; +} diff --git a/benchmark/one-level-arch/test/kernel/fa/src/test_tload_store.cpp b/benchmark/one-level-arch/test/kernel/fa/src/test_tload_store.cpp new file mode 100644 index 00000000..e9673164 --- /dev/null +++ b/benchmark/one-level-arch/test/kernel/fa/src/test_tload_store.cpp @@ -0,0 +1,60 @@ +#include +#include "benchmark.h" + +#define S1 64 +#define D 512 +#define KTM 32 +#define KTD 64 +#define ALIGN_MASK 0xfffffffffffff000ull +#define ALIGN 4*1024 +#define MAP_MEM_BASE 0x4000802000ULL + +using namespace pto; + +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 dtype = __half; + + dtype qp[S1*D + 2*ALIGN]; + dtype* q = (dtype*)(((uint64_t)qp & ALIGN_MASK) + ALIGN); + + init_deterministic(q, S1*D, 1); + + dtype* out = (dtype*)MAP_MEM_BASE; + + using gmQ = global_tensor>; + using gmO = global_tensor>; + using tileQ = Tile; + using tileO = Tile; + + using itQ = global_iterator; + using itO = global_iterator; + + itQ gIterQ(q); + itO gIterO(out); + + const int Qb = S1 / KTM; + const int Db = D / KTD; + + BENCHSTART; + for (int i = 0; i < Qb; ++i) { + for (int dd = 0; dd < Db; ++dd) { + tileQ tQ; + auto gQ = gIterQ(i, dd); + TLOAD(tQ, gQ); + + tileO tO; + auto gO = gIterO(i, dd); + TSTORE(gO, tQ); + } + } + BENCHEND; + + return 0; +} From 96c44768716126ed35a48027551300249b26c452 Mon Sep 17 00:00:00 2001 From: chenglongyu Date: Mon, 10 Aug 2026 15:25:35 +0800 Subject: [PATCH 2/9] add qsmla 0810 adaptation --- .../fa/quant_sparse_flash_mla_onepass_pto.hpp | 53 ++--- .../kernels/fa/quant_sparse_flash_mla_pto.hpp | 183 ++++++------------ .../fa/quant_sparse_flash_mla_tadd_pto.hpp | 72 ++++--- .../test/kernel/fa/src/qsmla_compare.py | 127 ++---------- .../kernel/fa/src/quant_sparse_flash_mla.cpp | 6 +- .../test/kernel/fa/src/test_memwrite.cpp | 18 -- .../test/kernel/fa/src/test_tload_store.cpp | 60 ------ 7 files changed, 138 insertions(+), 381 deletions(-) delete mode 100644 benchmark/one-level-arch/test/kernel/fa/src/test_memwrite.cpp delete mode 100644 benchmark/one-level-arch/test/kernel/fa/src/test_tload_store.cpp 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 index c4edbe2b..4d865dca 100644 --- 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 @@ -78,11 +78,16 @@ void quant_sparse_flash_mla_swa_onepass_pto( using gmQ = global_tensor>; using gmKV = global_tensor>; + // Same storage as gmKV, viewed as K^T. RowMajor and + // ColMajor have the same address formula, so this view performs + // the logical transpose without moving data. TCOPYIN below then + // performs the physical DN -> ZN conversion required by Cube SrcR. + using gmKT = global_tensor>; using gmO = global_tensor>; using gmMask = global_tensor>; - using tileQ = TileLeft; - using tileKV = TileRight; + using tileQ = TileLeft; + using tileKRight = TileRight; using tileW_out = TileAcc; using tileW = Tile; using tileMask = Tile; @@ -98,13 +103,13 @@ void quant_sparse_flash_mla_swa_onepass_pto( using tileSum = Tile; using itQ = global_iterator; - using itKV = global_iterator; + using itK = global_iterator; using itV = global_iterator; using itO = global_iterator; using itMask = global_iterator; itQ gIterQ(q_ptr); - itKV gIterKV(ori_kv_ptr); + itK gIterK(ori_kv_ptr); itV gIterV(ori_kv_ptr); itO gIterO(out_ptr); itMask gIterMask(mask_buf); @@ -148,29 +153,29 @@ void quant_sparse_flash_mla_swa_onepass_pto( for (int j = 0; j < Kb; ++j) { // --- Step 1: QK^T 沿全 D 累加 --- - tileW_out tW_out; - bool first_d = true; + // 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); - - tileKV tK; - auto gK = gIterKV(j, dd); - TLOAD(tK, gK); - - if (first_d) { - TMATMUL(tW_out, tQ, tK); - first_d = false; - } else { - TMATMUL_ACC(tW_out, tQ, tK); - } + TCOPYIN(tQ, gQ); // RowMajor Q: ND -> NZ (Cube SrcL) + + tileKRight tK; + auto gK = gIterK(dd, j); + TCOPYIN(tK, gK); // ColumnMajor K^T: DN -> ZN (Cube SrcR) + + tileW_out tW_out; + TMATMUL(tW_out, tQ, tK); + tileW tW_partial; + TCVT_Impl(tW_partial, tW_out); // ACC NZ -> Vec ND + TADD(tW, tW, tW_partial); } // --- Step 2: scale + mask --- - tileW tW; - ACCCVT(tW, tW_out); TMULS(tW, tW, scale); tileMask tMask; @@ -215,18 +220,18 @@ void quant_sparse_flash_mla_swa_onepass_pto( tileW_cast tExpW; TCVT(tExpW, tW); tileW_left tW_left; - TCVT(tW_left, tExpW); + TMOV_ND2NZ(tW_left, tExpW); // Vec ND -> Cube SrcL NZ #pragma clang loop unroll(full) for (int dd = 0; dd < Db; ++dd) { tileV tV; auto gV = gIterV(j, dd); - TLOAD(tV, gV); + TCOPYIN(tV, gV); // RowMajor V: ND -> ZN (Cube SrcR) tileO_out tPV_out; TMATMUL(tPV_out, tW_left, tV); tileO tPV; - ACCCVT(tPV, tPV_out); + TCVT_Impl(tPV, tPV_out); // ACC NZ -> Vec ND TADD(tO[dd], tO[dd], tPV); } @@ -247,7 +252,7 @@ void quant_sparse_flash_mla_swa_onepass_pto( tileO_cast tO_cast; TCVT(tO_cast, tO[dd]); auto gO = gIterO(i, dd); - TSTORE(gO, tO_cast); + TCOPYOUT(gO, tO_cast); } } } 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 index e9cc83f0..6a8c6354 100644 --- 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 @@ -10,8 +10,8 @@ // Q: [s1, D], KV: [s2, D] (MLA shared K=V=ori_kv), O: [s1, D] // // 【SWA 滑动窗口 — kernel 内部 token 级 mask, 使用 TSEL】 -// mask 为 UINT32 位打包格式 (每 32 列 1 个 uint32, bit=1 表示无效/被 mask) -// TSEL(dst, mask, neg_inf): mask bit=1 → dst=-1e30 (无效), bit=0 → 保持原 score +// 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] (闭区间) @@ -26,19 +26,15 @@ // // 【D=512 分块】 // D 超出单 tile 上限, 沿 D 维切分为 Db 块 (kTd) -// QK^T: 沿 D 累加 (TMATMUL + TMATMUL_ACC) +// QK^T: 每个 D 块独立 TMATMUL, 转为 Vec 后用 TADD 累加 // PV: 每个 D 分块独立计算并存储 // // 【两遍式】 // Pass 1: online softmax 归约 (m, l), 含 mask // Pass 2: 归一化 P, 计算 P@V, 含 mask // -// 【mask tile 格式】 -// TSEL 的 mask 必须是 UINT32 位打包格式: -// maskWordsPerRow = ceil(kTk / 32) -// mask tile shape: [kTm, maskWordsPerRow], dtype=uint32_t, RowMajor -// bit (1< @@ -47,26 +43,18 @@ using namespace pto; // CPU 侧 mask 预计算 (在 kernel 内部调用, 放在 stack 上) -// 生成 UINT32 位打包格式的 mask: -// 对 [s1, s2] 的每个 (q, kv), 如果 kv 在窗口外则对应 bit=1 -// 每 32 个 kv 列打包成 1 个 uint32 -// maskBuf 行大小 = ceil(s2 / 32) 个 uint32 -static inline void build_swa_mask_bitpacked( +// 生成 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 wordsPerRow = (s2 + 31) / 32; 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; - uint32_t* row = maskBuf + q * wordsPerRow; - for (int w = 0; w < wordsPerRow; ++w) row[w] = 0; for (int kv = 0; kv < s2; ++kv) { bool valid = (kv >= lo) && (kv <= hi); - if (!valid) { - row[kv / 32] |= (uint32_t{1} << (kv % 32)); - } + maskBuf[q * s2 + kv] = valid ? 0u : 1u; } } } @@ -95,33 +83,25 @@ void quant_sparse_flash_mla_swa_pto( { constexpr int Db = D / kTd; - // === mask 预计算 (UINT32 位打包) === - // 全局 mask: [s1, ceil(s2/32)] uint32 - constexpr int maskWordsPerRowGlobal = (s2 + 31) / 32; - uint32_t maskBuf[s1 * maskWordsPerRowGlobal]; //[64 * 4] - build_swa_mask_bitpacked(maskBuf, s1, s2, ori_win_left, ori_win_right); - - // 每个 [kTm, kTk] block 对应的 mask tile: - // maskWordsPerRowBlock = ceil(kTk / 32) - // mask tile: [kTm, maskWordsPerRowBlock], uint32, RowMajor - constexpr int maskWordsPerRowBlock = (kTk + 31) / 32; // 1 + 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 tileKV = TileRight; + using tileK = TileRight; using tileW_out = TileAcc; // score tile: RowMajor float using tileW = Tile; - // mask tile: RowMajor uint32, shape [kTm, maskWordsPerRowBlock] - // TSEL 要求 dst/mask/src 同 tile_shape, 但模拟器内部 mask 按 uint32 位打包读 - // 这里用 float 类型满足编译器约束, 实际数据是 uint32 位掩码 - // tile Cols 设为 kTk (与 score tile 同形状), 模拟器只读前 maskWordsPerRowBlock 个 uint32 - using tileMask = Tile; + using tileMask = Tile; using tileW_cast = Tile; using tileW_left = TileLeft; @@ -135,14 +115,16 @@ void quant_sparse_flash_mla_swa_pto( using tileSum = Tile; using itQ = global_iterator; - using itKV = global_iterator; + using itK = global_iterator; using itV = global_iterator; using itO = global_iterator; + using itMask = global_iterator; itQ gIterQ(q_ptr); - itKV gIterKV(ori_kv_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; @@ -164,62 +146,35 @@ void quant_sparse_flash_mla_swa_pto( for (int j = 0; j < Kb; ++j) { // QK^T 沿 D 维累加 - tileW_out tW_out; - bool first_d = true; + 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); + TCOPYIN(tQ, gQ); - tileKV tK; - auto gK = gIterKV(j, dd); - TLOAD(tK, gK); + tileK tK; + auto gK = gIterK(dd, j); + TCOPYIN(tK, gK); - if (first_d) { - TMATMUL(tW_out, tQ, tK); - first_d = false; - } else { - TMATMUL_ACC(tW_out, tQ, tK); - } + tileW_out tW_out; + TMATMUL(tW_out, tQ, tK); + tileW tW_partial; + TCVT_Impl(tW_partial, tW_out); + TADD(tW, tW, tW_partial); } - tileW tW; - ACCCVT(tW, tW_out); TMULS(tW, tW, scale); - // 应用 token 级 mask: TSEL(score, mask, neg_inf) - // mask bit=1 → score=-1e30 (无效), bit=0 → 保持原 score - // mask 数据从全局 maskBuf 中提取当前 [kTm, kTk] block 对应的位打包数据 - // 构建局部 mask tile: 从全局 [s1, maskWordsPerRowGlobal] 中提取 - // [kTm 行, kTk 列] 对应的 bit, 重新打包为 [kTm, maskWordsPerRowBlock] + // mask=1 选择 neg_inf,mask=0 保留 score。 { - // 构建 block mask: [kTm, maskWordsPerRowBlock] uint32 - // 从全局 mask 中提取列 [j*kTk, (j+1)*kTk) 的 bits - uint32_t blockMask[kTm * maskWordsPerRowBlock]; // [32 * 1] - for (int r = 0; r < kTm; ++r) { - int q_idx = i * kTm + r; - const uint32_t* globalRow = maskBuf + q_idx * maskWordsPerRowGlobal; - uint32_t* blockRow = blockMask + r * maskWordsPerRowBlock; - for (int w = 0; w < maskWordsPerRowBlock; ++w) blockRow[w] = 0; - for (int c = 0; c < kTk; ++c) { - int global_col = j * kTk + c; - if (global_col >= s2) { - blockRow[c / 32] |= (uint32_t{1} << (c % 32)); - continue; - } - bool masked = (globalRow[global_col / 32] >> (global_col % 32)) & 1; - if (masked) { - blockRow[c / 32] |= (uint32_t{1} << (c % 32)); - } - } - } - // TLOAD mask tile from blockMask (stack buffer) - using gmBlockMask = global_tensor>; - gmBlockMask gBlockMask(blockMask); tileMask tMask; - TLOAD(tMask, gBlockMask); - TSEL(tW, tMask, tNegInf); + 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)) @@ -263,55 +218,35 @@ void quant_sparse_flash_mla_swa_pto( for (int j = 0; j < Kb; ++j) { // 计算完整 QK^T (沿 D 累加, 与 Pass 1 一致) - tileW_out tW_out; - bool first_d = true; + 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); - - tileKV tK; - auto gK = gIterKV(j, dd2); - TLOAD(tK, gK); - - if (first_d) { - TMATMUL(tW_out, tQ, tK); - first_d = false; - } else { - TMATMUL_ACC(tW_out, tQ, tK); - } + 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); } - tileW tW; - ACCCVT(tW, tW_out); TMULS(tW, tW, scale); // 应用 token 级 mask: TSEL(score, mask, neg_inf) { - uint32_t blockMask[kTm * maskWordsPerRowBlock]; - for (int r = 0; r < kTm; ++r) { - int q_idx = i * kTm + r; - const uint32_t* globalRow = maskBuf + q_idx * maskWordsPerRowGlobal; - uint32_t* blockRow = blockMask + r * maskWordsPerRowBlock; - for (int w = 0; w < maskWordsPerRowBlock; ++w) blockRow[w] = 0; - for (int c = 0; c < kTk; ++c) { - int global_col = j * kTk + c; - if (global_col >= s2) { - blockRow[c / 32] |= (uint32_t{1} << (c % 32)); - continue; - } - bool masked = (globalRow[global_col / 32] >> (global_col % 32)) & 1; - if (masked) { - blockRow[c / 32] |= (uint32_t{1} << (c % 32)); - } - } - } - using gmBlockMask = global_tensor>; - gmBlockMask gBlockMask(blockMask); tileMask tMask; - TLOAD(tMask, gBlockMask); - TSEL(tW, tMask, tNegInf); + auto gMask = gIterMask(i, j); + TLOAD(tMask, gMask); + tileW tMasked; + TSELECT_Impl(tMasked, tMask, tNegInf, tW); + tW = tMasked; } // p = exp(score - m) / l @@ -323,17 +258,17 @@ void quant_sparse_flash_mla_swa_pto( tileW_cast tExpW; TCVT(tExpW, tW); tileW_left tW_left; - TCVT(tW_left, tExpW); + TMOV_ND2NZ(tW_left, tExpW); // PV = p * V (当前 D 分块) tileV tV; auto gV = gIterV(j, dd); - TLOAD(tV, gV); + TCOPYIN(tV, gV); tileO_out tPV_out; TMATMUL(tPV_out, tW_left, tV); tileO tPV; - ACCCVT(tPV, tPV_out); + TCVT_Impl(tPV, tPV_out); TADD(tO, tO, tPV); } @@ -342,7 +277,7 @@ void quant_sparse_flash_mla_swa_pto( tileO_cast tO_cast; TCVT(tO_cast, tO); auto gO = gIterO(i, dd); - TSTORE(gO, tO_cast); + TCOPYOUT(gO, tO_cast); } } } 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 index 161f8e29..e1d29214 100644 --- 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 @@ -43,7 +43,7 @@ // // 【D=512 分块】 // D 超出单 tile 上限, 沿 D 维切分为 Db 块 (kTd) -// QK^T: 沿 D 累加 (TMATMUL + TMATMUL_ACC) +// QK^T: 每个 D 块独立 TMATMUL, 转为 Vec 后用 TADD 累加 // PV: 每个 D 分块独立计算并存储 // // 【两遍式】 @@ -107,11 +107,13 @@ void quant_sparse_flash_mla_swa_tadd_pto( 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 tileKV = TileRight; + using tileK = TileRight; using tileW_out = TileAcc; // score tile 与 mask tile 都用 RowMajor, 保证 TADD 类型一致 @@ -129,13 +131,13 @@ void quant_sparse_flash_mla_swa_tadd_pto( using tileSum = Tile; using itQ = global_iterator; - using itKV = global_iterator; + using itK = global_iterator; using itV = global_iterator; using itO = global_iterator; using itMask = global_iterator; itQ gIterQ(q_ptr); - itKV gIterKV(ori_kv_ptr); + itK gIterK(ori_kv_ptr); itV gIterV(ori_kv_ptr); itO gIterO(out_ptr); itMask gIterMask(mask_buf); @@ -157,28 +159,25 @@ void quant_sparse_flash_mla_swa_tadd_pto( for (int j = 0; j < Kb; ++j) { // QK^T 沿 D 维累加 - tileW_out tW_out; - bool first_d = true; + 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); + TCOPYIN(tQ, gQ); - tileKV tK; - auto gK = gIterKV(j, dd); - TLOAD(tK, gK); + tileK tK; + auto gK = gIterK(dd, j); + TCOPYIN(tK, gK); - if (first_d) { - TMATMUL(tW_out, tQ, tK); - first_d = false; - } else { - TMATMUL_ACC(tW_out, tQ, tK); - } + tileW_out tW_out; + TMATMUL(tW_out, tQ, tK); + tileW tW_partial; + TCVT_Impl(tW_partial, tW_out); + TADD(tW, tW, tW_partial); } - tileW tW; - ACCCVT(tW, tW_out); TMULS(tW, tW, scale); // 应用 token 级 mask: score += mask (0 保持原值, -1e30 屏蔽) @@ -228,28 +227,25 @@ void quant_sparse_flash_mla_swa_tadd_pto( for (int j = 0; j < Kb; ++j) { // 计算完整 QK^T (沿 D 累加, 与 Pass 1 一致) - tileW_out tW_out; - bool first_d = true; + 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); - - tileKV tK; - auto gK = gIterKV(j, dd2); - TLOAD(tK, gK); - - if (first_d) { - TMATMUL(tW_out, tQ, tK); - first_d = false; - } else { - TMATMUL_ACC(tW_out, tQ, tK); - } + 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); } - tileW tW; - ACCCVT(tW, tW_out); TMULS(tW, tW, scale); // 应用 token 级 mask: score += mask (0 保持原值, -1e30 屏蔽) @@ -267,17 +263,17 @@ void quant_sparse_flash_mla_swa_tadd_pto( tileW_cast tExpW; TCVT(tExpW, tW); tileW_left tW_left; - TCVT(tW_left, tExpW); + TMOV_ND2NZ(tW_left, tExpW); // PV = p * V (当前 D 分块) tileV tV; auto gV = gIterV(j, dd); - TLOAD(tV, gV); + TCOPYIN(tV, gV); tileO_out tPV_out; TMATMUL(tPV_out, tW_left, tV); tileO tPV; - ACCCVT(tPV, tPV_out); + TCVT_Impl(tPV, tPV_out); TADD(tO, tO, tPV); } @@ -286,7 +282,7 @@ void quant_sparse_flash_mla_swa_tadd_pto( tileO_cast tO_cast; TCVT(tO_cast, tO); auto gO = gIterO(i, dd); - TSTORE(gO, tO_cast); + TCOPYOUT(gO, tO_cast); } } } 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 index 43464b81..d1cba2ac 100644 --- a/benchmark/one-level-arch/test/kernel/fa/src/qsmla_compare.py +++ b/benchmark/one-level-arch/test/kernel/fa/src/qsmla_compare.py @@ -1,122 +1,21 @@ #!/usr/bin/env python3 -""" -Compare NPU output (FP16, dumped from gfrun --dump-memory) with CPU golden output (FP32). -Usage: python3 qsmla_compare.py [atol] [rtol] -NPU dump format: raw FP16 (__half, IEEE 754 half-precision, 2 bytes per element) -Golden format: raw FP32 (float, IEEE 754 single-precision, 4 bytes per element) -""" - -import sys -import struct import math -import numpy as np - -def read_fp16(filename, count=None): - """Read binary file as float16 array, convert to float32.""" - with open(filename, 'rb') as f: - data = f.read() - arr = np.frombuffer(data, dtype=np.float16) - if count is not None: - arr = arr[:count] - return arr.astype(np.float32) - -def read_fp32(filename, count=None): - """Read binary file as float32 array.""" - with open(filename, 'rb') as f: - data = f.read() - arr = np.frombuffer(data, dtype=np.float32) - if count is not None: - arr = arr[:count] - return arr - -def compare(npu_out, golden_out, atol=0.01, rtol=0.05): - """Compare two float arrays with atol/rtol tolerance.""" - min_len = min(len(npu_out), len(golden_out)) - if len(npu_out) != len(golden_out): - print(f"WARNING: length mismatch npu={len(npu_out)} golden={len(golden_out)}, comparing first {min_len}") - - max_abs_err = 0.0 - max_rel_err = 0.0 - fail_count = 0 - fail_examples = [] - - for i in range(min_len): - n_val = float(npu_out[i]) - g_val = float(golden_out[i]) - - if math.isnan(n_val) or math.isnan(g_val): - fail_count += 1 - if len(fail_examples) < 10: - fail_examples.append((i, n_val, g_val, "NaN")) - continue - if math.isinf(n_val) or math.isinf(g_val): - fail_count += 1 - if len(fail_examples) < 10: - fail_examples.append((i, n_val, g_val, "Inf")) - continue - - abs_err = abs(n_val - g_val) - rel_err = abs_err / max(abs(g_val), 1e-8) - - if abs_err > max_abs_err: - max_abs_err = abs_err - if rel_err > max_rel_err: - max_rel_err = rel_err - - if abs_err > atol and rel_err > rtol: - fail_count += 1 - if len(fail_examples) < 10: - fail_examples.append((i, n_val, g_val, abs_err, rel_err)) - - print(f"=== QSMLA Precision Verification ===") - print(f"Compared elements: {min_len}") - print(f"Max absolute error: {max_abs_err:.8f}") - print(f"Max relative error: {max_rel_err:.8f}") - print(f"Tolerance: atol={atol}, rtol={rtol}") - print(f"Failed elements: {fail_count}/{min_len}") - - if fail_examples: - print(f"\nFirst {len(fail_examples)} failures:") - for ex in fail_examples: - if len(ex) == 4: - print(f" idx={ex[0]}: npu={ex[1]:.8f} golden={ex[2]:.8f} [{ex[3]}]") - else: - print(f" idx={ex[0]}: npu={ex[1]:.8f} golden={ex[2]:.8f} abs_err={ex[3]:.8f} rel_err={ex[4]:.8f}") - - # Print first 8 values for sanity - print(f"\nFirst 8 values comparison:") - for i in range(min(8, min_len)): - n = float(npu_out[i]) - g = float(golden_out[i]) - print(f" [{i}] npu={n:.8f} golden={g:.8f} diff={abs(n-g):.8f}") - - if fail_count == 0: - print("\n=== RESULT: PASS ===") - return 0 - else: - pass_rate = (min_len - fail_count) / min_len * 100 - print(f"\n=== RESULT: FAIL (pass rate: {pass_rate:.2f}%) ===") - return 1 - -def main(): - if len(sys.argv) < 3: - print("Usage: python3 qsmla_compare.py [atol] [rtol]") - sys.exit(1) +import struct - npu_file = sys.argv[1] - golden_file = sys.argv[2] - atol = float(sys.argv[3]) if len(sys.argv) > 3 else 0.01 - rtol = float(sys.argv[4]) if len(sys.argv) > 4 else 0.05 +path = "SuperNPUBench/benchmark/one-level-arch/test/kernel/fa/src/" - npu_out = read_fp16(npu_file) - golden_out = read_fp32(golden_file) +with open(path + "qsmla_onepass_npu_out.bin", "rb") as f: + actual = struct.unpack("<32768e", f.read()) - print(f"NPU output: {len(npu_out)} FP16 values from {npu_file}") - print(f"Golden output: {len(golden_out)} FP32 values from {golden_file}") +with open(path + "qsmla_golden.bin", "rb") as f: + golden = struct.unpack("<32768f", f.read()) - ret = compare(npu_out, golden_out, atol, rtol) - sys.exit(ret) +errors = [abs(float(a) - float(g)) for a, g in zip(actual, golden)] +passed = [e <= 1e-3 + 1e-3 * abs(g) for e, g in zip(errors, golden)] -if __name__ == '__main__': - main() +print(f"passed = {sum(passed)}/32768 ({100*sum(passed)/32768:.6f}%)") +print(f"failed = {32768-sum(passed)}") +print(f"max_abs = {max(errors):.9f}") +print(f"mean_abs= {sum(errors)/len(errors):.9f}") +print(f"nan_npu = {sum(math.isnan(float(x)) for x in actual)}") 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 index 47edbb49..881e710b 100644 --- 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 @@ -2,8 +2,8 @@ #include "benchmark.h" #include "fileop.h" // #include "fa/quant_sparse_flash_mla_pto.hpp" -// #include "fa/quant_sparse_flash_mla_tadd_pto.hpp" -#include "fa/quant_sparse_flash_mla_onepass_pto.hpp" +#include "fa/quant_sparse_flash_mla_tadd_pto.hpp" +// #include "fa/quant_sparse_flash_mla_onepass_pto.hpp" #define B 1 #define H 1 @@ -92,7 +92,7 @@ int main(){ BENCHSTART; for(int i=0;i( out + i*H*s1*D + j*s1*D, q + i*H*s1*D + j*s1*D, diff --git a/benchmark/one-level-arch/test/kernel/fa/src/test_memwrite.cpp b/benchmark/one-level-arch/test/kernel/fa/src/test_memwrite.cpp deleted file mode 100644 index cf97df18..00000000 --- a/benchmark/one-level-arch/test/kernel/fa/src/test_memwrite.cpp +++ /dev/null @@ -1,18 +0,0 @@ -#include -#include "benchmark.h" - -#define MAP_MEM_BASE 0x4000802000ULL -#define OUT_COUNT 32768 - -int main(){ - __half* out = (__half*)MAP_MEM_BASE; - - for (int i = 0; i < OUT_COUNT; ++i) { - out[i] = (__half)((float)i * 0.001f); - } - - BENCHSTART; - BENCHEND; - - return 0; -} diff --git a/benchmark/one-level-arch/test/kernel/fa/src/test_tload_store.cpp b/benchmark/one-level-arch/test/kernel/fa/src/test_tload_store.cpp deleted file mode 100644 index e9673164..00000000 --- a/benchmark/one-level-arch/test/kernel/fa/src/test_tload_store.cpp +++ /dev/null @@ -1,60 +0,0 @@ -#include -#include "benchmark.h" - -#define S1 64 -#define D 512 -#define KTM 32 -#define KTD 64 -#define ALIGN_MASK 0xfffffffffffff000ull -#define ALIGN 4*1024 -#define MAP_MEM_BASE 0x4000802000ULL - -using namespace pto; - -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 dtype = __half; - - dtype qp[S1*D + 2*ALIGN]; - dtype* q = (dtype*)(((uint64_t)qp & ALIGN_MASK) + ALIGN); - - init_deterministic(q, S1*D, 1); - - dtype* out = (dtype*)MAP_MEM_BASE; - - using gmQ = global_tensor>; - using gmO = global_tensor>; - using tileQ = Tile; - using tileO = Tile; - - using itQ = global_iterator; - using itO = global_iterator; - - itQ gIterQ(q); - itO gIterO(out); - - const int Qb = S1 / KTM; - const int Db = D / KTD; - - BENCHSTART; - for (int i = 0; i < Qb; ++i) { - for (int dd = 0; dd < Db; ++dd) { - tileQ tQ; - auto gQ = gIterQ(i, dd); - TLOAD(tQ, gQ); - - tileO tO; - auto gO = gIterO(i, dd); - TSTORE(gO, tQ); - } - } - BENCHEND; - - return 0; -} From e286c2e57aa6a746728dbdfb98609ee9e6f4cca3 Mon Sep 17 00:00:00 2001 From: chenglongyu Date: Wed, 12 Aug 2026 16:14:13 +0800 Subject: [PATCH 3/9] add qsmla 0812 golden reference --- .../test/kernel/fa/qsmla_stage0_cases.py | 156 +++++++ .../test/kernel/fa/run_qsmla_stage0.py | 101 +++++ .../test/kernel/fa/src/qsmla_cpu_ref.cpp | 423 ++++++++---------- 3 files changed, 446 insertions(+), 234 deletions(-) create mode 100644 benchmark/one-level-arch/test/kernel/fa/qsmla_stage0_cases.py create mode 100644 benchmark/one-level-arch/test/kernel/fa/run_qsmla_stage0.py 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..f1090b86 --- /dev/null +++ b/benchmark/one-level-arch/test/kernel/fa/qsmla_stage0_cases.py @@ -0,0 +1,156 @@ +#!/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 = "CONTIGUOUS_2D" + kv_layout: str = "CONTIGUOUS_2D" + 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="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"}), + ), +) + + +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..cc99ea96 --- /dev/null +++ b/benchmark/one-level-arch/test/kernel/fa/run_qsmla_stage0.py @@ -0,0 +1,101 @@ +#!/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()}\"', + } + 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_cpu_ref.cpp b/benchmark/one-level-arch/test/kernel/fa/src/qsmla_cpu_ref.cpp index ee6f6a6c..7a018d89 100644 --- 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 @@ -1,270 +1,225 @@ -// CPU reference implementation for quant_sparse_flash_mla SWA mode -// Computes: O = softmax(Q @ K^T * softmax_scale + mask) @ V -// Where K=V=ori_kv (MLA shared KV), with token-level sliding window mask - -#include -#include -#include -#include -#include -#include - -#ifndef S1 -#define S1 64 +// 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 S2 -#define S2 128 +#ifndef QS1 +#define QS1 64 #endif -#ifndef D -#define D 512 +#ifndef QS2 +#define QS2 128 #endif -#ifndef KTM -#define KTM 32 +#ifndef QN1 +#define QN1 1 #endif -#ifndef KTK -#define KTK 32 +#ifndef QN2 +#define QN2 1 #endif -#ifndef WIN_LEFT -#define WIN_LEFT 1 +#ifndef QD +#define QD 512 #endif -#ifndef WIN_RIGHT -#define WIN_RIGHT 1 +#ifndef QK +#define QK 128 #endif -#ifndef SOFTMAX_SCALE -#define SOFTMAX_SCALE 0.125f +#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 -typedef float f32_t; - -static void init_deterministic_f32(f32_t* 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] = val; +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)); } -// Build token-level SWA mask (SEL semantics) -// mask[q * s2 + kv] = true if kv 在窗口外 (被 mask, score 置 -1e30) -// = false if kv 在窗口内 (有效, 保持原 score) -// valid: diagonal - win_left <= kv <= diagonal + win_right -// where diagonal = (s2 - s1) + q -// Apply: score = mask ? -1e30 : score (TSEL semantics) -static void build_swa_mask(bool* mask, 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); - mask[q * s2 + kv] = !valid; +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; } -// CPU reference: SWA MLA attention with token-level mask (TSEL semantics) -// Q: [s1, D], KV: [s2, D] (K=V), O: [s1, D] -void cpu_quant_sparse_flash_mla_swa( - f32_t* out, const f32_t* q, const f32_t* kv, - const bool* mask, - int s1, int s2, int d, int kTm, int kTk, - float softmax_scale, int win_left, int win_right) -{ - const int Qb = (s1 + kTm - 1) / kTm; - const int Kb = (s2 + kTk - 1) / kTk; - - f32_t* score = (f32_t*)malloc(kTm * kTk * sizeof(f32_t)); - f32_t* p = (f32_t*)malloc(kTm * kTk * sizeof(f32_t)); - f32_t* pv = (f32_t*)malloc(kTm * d * sizeof(f32_t)); - - for (int qi = 0; qi < Qb; ++qi) { - - // Pass 1: online softmax with mask - f32_t row_max[kTm]; - f32_t row_sum[kTm]; - for (int r = 0; r < kTm; ++r) { - row_max[r] = -1e30f; - row_sum[r] = 0.0f; - } - - for (int j = 0; j < Kb; ++j) { - for (int r = 0; r < kTm; ++r) { - int q_row = qi * kTm + r; - if (q_row >= s1) continue; - for (int c = 0; c < kTk; ++c) { - int kv_row = j * kTk + c; - if (kv_row >= s2) { score[r * kTk + c] = -1e30f; continue; } - f32_t dot = 0.0f; - for (int dd = 0; dd < d; ++dd) { - dot += q[q_row * d + dd] * kv[kv_row * d + dd]; - } - // Apply mask: TSEL semantics (mask=true → -1e30, mask=false → keep score) - f32_t raw_score = dot * softmax_scale; - score[r * kTk + c] = mask[q_row * s2 + kv_row] ? -1e30f : raw_score; - } - } - - for (int r = 0; r < kTm; ++r) { - int q_row = qi * kTm + r; - if (q_row >= s1) continue; - - f32_t local_max = -1e30f; - for (int c = 0; c < kTk; ++c) { - int kv_row = j * kTk + c; - if (kv_row >= s2) continue; - if (score[r * kTk + c] > local_max) - local_max = score[r * kTk + c]; - } +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]); + } +} - f32_t new_max = (row_max[r] > local_max) ? row_max[r] : local_max; - f32_t scale_old = expf(row_max[r] - new_max); - row_sum[r] *= scale_old; +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); +} - for (int c = 0; c < kTk; ++c) { - int kv_row = j * kTk + c; - if (kv_row >= s2) continue; - row_sum[r] += expf(score[r * kTk + c] - new_max); - } +static size_t q_offset(int b, int head, int token, int dim) { + return (((static_cast(b) * QN1 + head) * QS1 + token) * QD + dim); +} - row_max[r] = new_max; - } - } +static size_t kv_offset(int b, int head, int token, int dim) { + return (((static_cast(b) * QN2 + head) * QS2 + token) * QD + dim); +} - // Pass 2: compute P @ V with mask - for (int dd = 0; dd < d; ++dd) { - for (int r = 0; r < kTm; ++r) { - pv[r * d + dd] = 0.0f; - } - } +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; +} - for (int j = 0; j < Kb; ++j) { - for (int r = 0; r < kTm; ++r) { - int q_row = qi * kTm + r; - if (q_row >= s1) continue; - for (int c = 0; c < kTk; ++c) { - int kv_row = j * kTk + c; - if (kv_row >= s2) { p[r * kTk + c] = 0.0f; continue; } - f32_t dot = 0.0f; - for (int dd = 0; dd < d; ++dd) { - dot += q[q_row * d + dd] * kv[kv_row * d + dd]; +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)]; } - // Apply mask: TSEL semantics (mask=true → -1e30, mask=false → keep score) - f32_t raw_score = dot * softmax_scale; - f32_t s = mask[q_row * s2 + kv_row] ? -1e30f : raw_score; - p[r * kTk + c] = expf(s - row_max[r]) / row_sum[r]; + scores[kv_pos] = dot * QSOFTMAX_SCALE; + row_max = std::max(row_max, scores[kv_pos]); } - } + if (!std::isfinite(row_max)) continue; - for (int r = 0; r < kTm; ++r) { - int q_row = qi * kTm + r; - if (q_row >= s1) continue; - for (int dd = 0; dd < d; ++dd) { - for (int c = 0; c < kTk; ++c) { - int kv_row = j * kTk + c; - if (kv_row >= s2) continue; - pv[r * d + dd] += p[r * kTk + c] * kv[kv_row * d + dd]; + 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; } } } - - for (int r = 0; r < kTm; ++r) { - int q_row = qi * kTm + r; - if (q_row >= s1) continue; - for (int dd = 0; dd < d; ++dd) { - out[q_row * d + dd] = pv[r * d + dd]; - } - } } - free(score); - free(p); - free(pv); -} - -int main() { - int q_count = S1 * D; - int kv_count = S2 * D; - int out_count = S1 * D; - int mask_count = S1 * S2; - - f32_t* q = (f32_t*)malloc(q_count * sizeof(f32_t)); - f32_t* kv = (f32_t*)malloc(kv_count * sizeof(f32_t)); - f32_t* out = (f32_t*)malloc(out_count * sizeof(f32_t)); - bool* mask = (bool*)malloc(mask_count * sizeof(bool)); - - init_deterministic_f32(q, q_count, 1); - init_deterministic_f32(kv, kv_count, 2); - - // for (int i = 0; i < 100; i++){ - // std::cout << q[i] << ", "; - // } - - // std::cout << std::endl; - - // for (int i = 0; i < 100; i++){ - // std::cout << kv[i] << ", "; - // } - - // return 0; - - memset(out, 0, out_count * sizeof(f32_t)); - - // f32_t mask1[5][10]; - - // build_swa_mask(&mask1[0][0], 5, 10, 2, 2); - - // for(int i = 0; i < 5; i++){ - // for(int j = 0; j < 10; j++){ - // std::cout << mask1[i][j] << ", \t"; - // } - // std::cout << std::endl; - // } - - // return 0; - - - build_swa_mask(mask, S1, S2, WIN_LEFT, WIN_RIGHT); - - printf("QSMLA_CPU: s1=%d s2=%d D=%d win_left=%d win_right=%d\n", - S1, S2, D, WIN_LEFT, WIN_RIGHT); - printf("QSMLA_CPU: computing reference with token-level mask...\n"); - fflush(stdout); - - cpu_quant_sparse_flash_mla_swa(out, q, kv, mask, - S1, S2, D, KTM, KTK, - SOFTMAX_SCALE, WIN_LEFT, WIN_RIGHT); - - const char* golden_path = "qsmla_golden.bin"; - FILE* f = fopen(golden_path, "wb"); - if (!f) { - fprintf(stderr, "Failed to open %s\n", golden_path); + 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; } - fwrite(out, sizeof(f32_t), out_count, f); - fclose(f); - - printf("QSMLA_CPU: golden output written to %s (%d floats, %d bytes)\n", - golden_path, out_count, (int)(out_count * sizeof(f32_t))); - - printf("QSMLA_CPU: first 8 output values:\n"); - for (int i = 0; i < 8 && i < out_count; ++i) { - printf(" out[%d] = %.6f\n", i, out[i]); + 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* fq = fopen("qsmla_input_q.bin", "wb"); - fwrite(q, sizeof(f32_t), q_count, fq); - fclose(fq); - - FILE* fkv = fopen("qsmla_input_kv.bin", "wb"); - fwrite(kv, sizeof(f32_t), kv_count, fkv); - fclose(fkv); - - printf("QSMLA_CPU: input data written to qsmla_input_q.bin and qsmla_input_kv.bin\n"); - - free(q); - free(kv); - free(out); - free(mask); + 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_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)); + 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; } From aad72ac0baa7ba3307f33b614a52ea6e20e24002 Mon Sep 17 00:00:00 2001 From: chenglongyu Date: Fri, 14 Aug 2026 18:12:14 +0800 Subject: [PATCH 4/9] add qsmla 0814 golden reference --- .../test/kernel/fa/qsmla_stage0_cases.py | 20 +++++++++++++++++-- .../test/kernel/fa/run_qsmla_stage0.py | 2 ++ .../test/kernel/fa/src/qsmla_cpu_ref.cpp | 16 ++++++++++++++- 3 files changed, 35 insertions(+), 3 deletions(-) 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 index f1090b86..6b886fc1 100644 --- a/benchmark/one-level-arch/test/kernel/fa/qsmla_stage0_cases.py +++ b/benchmark/one-level-arch/test/kernel/fa/qsmla_stage0_cases.py @@ -24,8 +24,8 @@ class QsmlaCase: softmax_scale: float = 0.125 source: str = "supernpubench:stage0-smoke" mode: str = "SWA" - q_layout: str = "CONTIGUOUS_2D" - kv_layout: str = "CONTIGUOUS_2D" + q_layout: str = "BNSD" + kv_layout: str = "BNSD" logical_dtype: str = "fp16" source_storage_dtype: str = "fp16" stage0_compute_dtype: str = "fp16" @@ -134,6 +134,22 @@ class QsmlaCase: 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.", + ), ) 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 index cc99ea96..cc6b7805 100644 --- a/benchmark/one-level-arch/test/kernel/fa/run_qsmla_stage0.py +++ b/benchmark/one-level-arch/test/kernel/fa/run_qsmla_stage0.py @@ -35,6 +35,8 @@ def compile_defines(case: QsmlaCase, output_root: Path): "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()] 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 index 7a018d89..f23a37be 100644 --- 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 @@ -56,6 +56,12 @@ #ifndef QOUTPUT_ROOT #define QOUTPUT_ROOT "." #endif +#ifndef QLAYOUT_BSND +#define QLAYOUT_BSND 0 +#endif +#ifndef QKV_LAYOUT_BSND +#define QKV_LAYOUT_BSND 0 +#endif static_assert(QB > 0, "B must be positive"); static_assert(QS1 >= 0 && QS2 >= 0, "S1/S2 must be non-negative"); @@ -132,10 +138,16 @@ static bool swa_valid(int q_pos, int kv_pos) { } static size_t q_offset(int b, int head, int token, int dim) { + if constexpr (QLAYOUT_BSND != 0) { + return (((static_cast(b) * QS1 + token) * QN1 + head) * QD + dim); + } return (((static_cast(b) * QN1 + head) * QS1 + token) * QD + dim); } static size_t kv_offset(int b, int head, int token, int dim) { + if constexpr (QKV_LAYOUT_BSND != 0) { + return (((static_cast(b) * QS2 + token) * QN2 + head) * QD + dim); + } return (((static_cast(b) * QN2 + head) * QS2 + token) * QD + dim); } @@ -214,9 +226,11 @@ int main() { 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)); + QTM, QTK, QTD, QWIN_LEFT, QWIN_RIGHT, static_cast(QSOFTMAX_SCALE), + QLAYOUT_BSND ? "BSND" : "BNSD", QKV_LAYOUT_BSND ? "BSND" : "BNSD"); std::fclose(manifest); std::printf("QSMLA_STAGE0 case=%s output=%s elements=%zu\n", From 537e51eefcc058c55264c2780c3ea873b7d87035 Mon Sep 17 00:00:00 2001 From: chenglongyu Date: Mon, 17 Aug 2026 17:29:28 +0800 Subject: [PATCH 5/9] add qsmla 0814 golden --- .../test/kernel/fa/qsmla_stage0_cases.py | 64 ++++++++++++++++++- .../test/kernel/fa/src/qsmla_cpu_ref.cpp | 16 ++--- 2 files changed, 67 insertions(+), 13 deletions(-) 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 index 6b886fc1..25508d4d 100644 --- a/benchmark/one-level-arch/test/kernel/fa/qsmla_stage0_cases.py +++ b/benchmark/one-level-arch/test/kernel/fa/qsmla_stage0_cases.py @@ -24,8 +24,8 @@ class QsmlaCase: softmax_scale: float = 0.125 source: str = "supernpubench:stage0-smoke" mode: str = "SWA" - q_layout: str = "BNSD" - kv_layout: str = "BNSD" + q_layout: str = "BSND" + kv_layout: str = "BSND" logical_dtype: str = "fp16" source_storage_dtype: str = "fp16" stage0_compute_dtype: str = "fp16" @@ -122,6 +122,66 @@ class QsmlaCase: 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, 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 index f23a37be..a465d5ac 100644 --- 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 @@ -57,10 +57,10 @@ #define QOUTPUT_ROOT "." #endif #ifndef QLAYOUT_BSND -#define QLAYOUT_BSND 0 +#define QLAYOUT_BSND 1 #endif #ifndef QKV_LAYOUT_BSND -#define QKV_LAYOUT_BSND 0 +#define QKV_LAYOUT_BSND 1 #endif static_assert(QB > 0, "B must be positive"); @@ -138,17 +138,11 @@ static bool swa_valid(int q_pos, int kv_pos) { } static size_t q_offset(int b, int head, int token, int dim) { - if constexpr (QLAYOUT_BSND != 0) { - return (((static_cast(b) * QS1 + token) * QN1 + head) * QD + dim); - } - return (((static_cast(b) * QN1 + head) * QS1 + token) * QD + dim); + return (((static_cast(b) * QS1 + token) * QN1 + head) * QD + dim); } static size_t kv_offset(int b, int head, int token, int dim) { - if constexpr (QKV_LAYOUT_BSND != 0) { - return (((static_cast(b) * QS2 + token) * QN2 + head) * QD + dim); - } - return (((static_cast(b) * QN2 + head) * QS2 + token) * QD + 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) { @@ -230,7 +224,7 @@ int main() { "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), - QLAYOUT_BSND ? "BSND" : "BNSD", QKV_LAYOUT_BSND ? "BSND" : "BNSD"); + "BSND", "BSND"); std::fclose(manifest); std::printf("QSMLA_STAGE0 case=%s output=%s elements=%zu\n", From 1daec488bf9f49cd3cdc2a1e624bea4af16bd701 Mon Sep 17 00:00:00 2001 From: chenglongyu Date: Mon, 17 Aug 2026 20:36:51 +0800 Subject: [PATCH 6/9] add qsmla 0817 onepass --- .../kernels/fa/qsmla_config_pto.hpp | 77 +++++++++++++++++++ .../fa/quant_sparse_flash_mla_onepass_pto.hpp | 48 +++++++++++- .../kernel/fa/src/quant_sparse_flash_mla.cpp | 9 ++- 3 files changed, 126 insertions(+), 8 deletions(-) create mode 100644 benchmark/one-level-arch/kernels/fa/qsmla_config_pto.hpp 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..3efa51b9 --- /dev/null +++ b/benchmark/one-level-arch/kernels/fa/qsmla_config_pto.hpp @@ -0,0 +1,77 @@ +#ifndef QSMLA_CONFIG_PTO_HPP +#define QSMLA_CONFIG_PTO_HPP + +#include + +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 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 index 4d865dca..95aa7b60 100644 --- 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 @@ -31,6 +31,7 @@ #include #include "template_asm.h" +#include "qsmla_config_pto.hpp" using namespace pto; @@ -49,10 +50,8 @@ static inline void build_swa_mask_onepass( } } -template -void quant_sparse_flash_mla_swa_onepass_pto( +template +void quant_sparse_flash_mla_swa_onepass_config_pto( odttype* out_ptr, qdtype* q_ptr, kvdtype* ori_kv_ptr, @@ -71,6 +70,14 @@ void quant_sparse_flash_mla_swa_onepass_pto( int* metadata, float* softmax_lse) { + 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"); constexpr int Db = D / kTd; float mask_buf[s1 * s2]; @@ -257,4 +264,37 @@ void quant_sparse_flash_mla_swa_onepass_pto( } } +// 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); +} + #endif 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 index 881e710b..e8c8a350 100644 --- 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 @@ -2,8 +2,8 @@ #include "benchmark.h" #include "fileop.h" // #include "fa/quant_sparse_flash_mla_pto.hpp" -#include "fa/quant_sparse_flash_mla_tadd_pto.hpp" -// #include "fa/quant_sparse_flash_mla_onepass_pto.hpp" +// #include "fa/quant_sparse_flash_mla_tadd_pto.hpp" +#include "fa/quant_sparse_flash_mla_onepass_pto.hpp" #define B 1 #define H 1 @@ -77,6 +77,7 @@ int main(){ using qdtype = __half; using kvdtype = __half; using odttype = __half; + using Config = QsmlaConfig; qdtype qp[B*H*s1*D + 2*ALIGN]; kvdtype kvp[B*H*s2*D + 2*ALIGN]; @@ -92,8 +93,8 @@ int main(){ BENCHSTART; for(int i=0;i( + quant_sparse_flash_mla_swa_onepass_config_pto< + qdtype, kvdtype, odttype, Config>( out + i*H*s1*D + j*s1*D, q + i*H*s1*D + j*s1*D, kv + i*H*s2*D + j*s2*D, From 6235e4408b75c39bf60a2bed0cc583e5bf437866 Mon Sep 17 00:00:00 2001 From: chenglongyu Date: Thu, 20 Aug 2026 10:16:52 +0800 Subject: [PATCH 7/9] add qsmla 0820 onepass&tadd BSND --- .../kernels/fa/qsmla_config_pto.hpp | 12 + .../fa/quant_sparse_flash_mla_onepass_pto.hpp | 169 ++++++++++--- .../fa/quant_sparse_flash_mla_tadd_pto.hpp | 231 ++++++++++++++---- .../one-level-arch/test/kernel/fa/Makefile | 12 +- .../test/kernel/fa/src/qsmla_compare.py | 111 ++++++++- .../kernel/fa/src/quant_sparse_flash_mla.cpp | 61 +++-- 6 files changed, 498 insertions(+), 98 deletions(-) diff --git a/benchmark/one-level-arch/kernels/fa/qsmla_config_pto.hpp b/benchmark/one-level-arch/kernels/fa/qsmla_config_pto.hpp index 3efa51b9..a01be9dd 100644 --- a/benchmark/one-level-arch/kernels/fa/qsmla_config_pto.hpp +++ b/benchmark/one-level-arch/kernels/fa/qsmla_config_pto.hpp @@ -62,6 +62,18 @@ struct QsmlaConfig { 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; 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 index 95aa7b60..b0ffe00f 100644 --- 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 @@ -36,11 +36,13 @@ using namespace pto; static inline void build_swa_mask_onepass( - float* mask, int s1, int s2, int win_left, int win_right) + float* mask, int s1, int s2, int win_left, int win_right, + int q_position = -1, int q_sequence_length = -1) { - const int causal_offset = s2 - s1; for (int q = 0; q < s1; ++q) { - int diagonal = causal_offset + q; + const int logical_q = q_position >= 0 ? q_position : q; + const int logical_s1 = q_position >= 0 ? q_sequence_length : s1; + int diagonal = s2 - logical_s1 + logical_q; int lo = diagonal - win_left; int hi = diagonal + win_right; for (int kv = 0; kv < s2; ++kv) { @@ -68,7 +70,9 @@ void quant_sparse_flash_mla_swa_onepass_config_pto( int* seqused_ori_kv, float* sinks, int* metadata, - float* softmax_lse) + float* softmax_lse, + int q_position = -1, + int q_sequence_length = -1) { constexpr int s1 = Config::S1; constexpr int s2 = Config::S2; @@ -81,27 +85,22 @@ void quant_sparse_flash_mla_swa_onepass_config_pto( constexpr int Db = D / kTd; float mask_buf[s1 * s2]; - build_swa_mask_onepass(mask_buf, s1, s2, ori_win_left, ori_win_right); + 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>; - // Same storage as gmKV, viewed as K^T. RowMajor and - // ColMajor have the same address formula, so this view performs - // the logical transpose without moving data. TCOPYIN below then - // performs the physical DN -> ZN conversion required by Cube SrcR. - using gmKT = global_tensor>; using gmO = global_tensor>; using gmMask = global_tensor>; using tileQ = TileLeft; + using tileKSrc = Tile; using tileKRight = TileRight; - using tileW_out = TileAcc; 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; @@ -110,13 +109,13 @@ void quant_sparse_flash_mla_swa_onepass_config_pto( using tileSum = Tile; using itQ = global_iterator; - using itK = global_iterator; + using itKSrc = global_iterator; using itV = global_iterator; using itO = global_iterator; using itMask = global_iterator; itQ gIterQ(q_ptr); - itK gIterK(ori_kv_ptr); + itKSrc gIterKSrc(ori_kv_ptr); itV gIterV(ori_kv_ptr); itO gIterO(out_ptr); itMask gIterMask(mask_buf); @@ -169,16 +168,17 @@ void quant_sparse_flash_mla_swa_onepass_config_pto( for (int dd = 0; dd < Db; ++dd) { tileQ tQ; auto gQ = gIterQ(i, dd); - TCOPYIN(tQ, gQ); // RowMajor Q: ND -> NZ (Cube SrcL) + TLOAD(tQ, gQ); // RowMajor Q -> Left tile + + tileKSrc tKSrc; + auto gK = gIterKSrc(j, dd); + TLOAD(tKSrc, gK); // Original K block [Tk, Td] tileKRight tK; - auto gK = gIterK(dd, j); - TCOPYIN(tK, gK); // ColumnMajor K^T: DN -> ZN (Cube SrcR) + TTRANS(tK, tKSrc); // [Tk, Td] -> [Td, Tk] for MM1 SrcR - tileW_out tW_out; - TMATMUL(tW_out, tQ, tK); tileW tW_partial; - TCVT_Impl(tW_partial, tW_out); // ACC NZ -> Vec ND + TMATMUL(tW_partial, tQ, tK); TADD(tW, tW, tW_partial); } @@ -224,21 +224,20 @@ void quant_sparse_flash_mla_swa_onepass_config_pto( } // --- Step 5: P@V for each D block --- - tileW_cast tExpW; - TCVT(tExpW, tW); tileW_left tW_left; - TMOV_ND2NZ(tW_left, tExpW); // Vec ND -> Cube SrcL NZ + // 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); - TCOPYIN(tV, gV); // RowMajor V: ND -> ZN (Cube SrcR) + TLOAD(tV, gV); // RowMajor V -> Right tile - tileO_out tPV_out; - TMATMUL(tPV_out, tW_left, tV); tileO tPV; - TCVT_Impl(tPV, tPV_out); // ACC NZ -> Vec ND + TMATMUL(tPV, tW_left, tV); TADD(tO[dd], tO[dd], tPV); } @@ -259,7 +258,7 @@ void quant_sparse_flash_mla_swa_onepass_config_pto( tileO_cast tO_cast; TCVT(tO_cast, tO[dd]); auto gO = gIterO(i, dd); - TCOPYOUT(gO, tO_cast); + TSTORE(gO, tO_cast); } } } @@ -294,7 +293,117 @@ void quant_sparse_flash_mla_swa_onepass_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); + 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>( + 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>( + 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_tadd_pto.hpp b/benchmark/one-level-arch/kernels/fa/quant_sparse_flash_mla_tadd_pto.hpp index e1d29214..685cf22f 100644 --- 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 @@ -53,6 +53,7 @@ #include #include "template_asm.h" +#include "qsmla_config_pto.hpp" using namespace pto; @@ -61,12 +62,14 @@ using namespace pto; // 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( - float* mask, int s1, int s2, int win_left, int win_right) +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) { - const int causal_offset = s2 - s1; for (int q = 0; q < s1; ++q) { - int diagonal = causal_offset + q; + const int logical_q = q_position >= 0 ? q_position : q; + const int logical_s1 = q_position >= 0 ? q_sequence_length : s1; + int diagonal = s2 - logical_s1 + logical_q; int lo = diagonal - win_left; int hi = diagonal + win_right; for (int kv = 0; kv < s2; ++kv) { @@ -76,10 +79,8 @@ static inline void build_swa_mask( } } -template -void quant_sparse_flash_mla_swa_tadd_pto( +template +void quant_sparse_flash_mla_swa_tadd_config_pto( odttype* out_ptr, qdtype* q_ptr, kvdtype* ori_kv_ptr, @@ -96,33 +97,41 @@ void quant_sparse_flash_mla_swa_tadd_pto( int* seqused_ori_kv, float* sinks, int* metadata, - float* softmax_lse) + 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"); constexpr int Db = D / kTd; // kernel 内部计算 mask, 放在 stack 上 // s1*s2*sizeof(float) = 64*128*4 = 32KB (可容纳) float mask_buf[s1 * s2]; - build_swa_mask(mask_buf, s1, s2, ori_win_left, ori_win_right); + 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>; - // 与 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; + using tileKSrc = Tile; + using tileKRight = TileRight; // score tile 与 mask tile 都用 RowMajor, 保证 TADD 类型一致 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; @@ -131,13 +140,13 @@ void quant_sparse_flash_mla_swa_tadd_pto( using tileSum = Tile; using itQ = global_iterator; - using itK = global_iterator; + using itKSrc = global_iterator; using itV = global_iterator; using itO = global_iterator; using itMask = global_iterator; itQ gIterQ(q_ptr); - itK gIterK(ori_kv_ptr); + itKSrc gIterKSrc(ori_kv_ptr); itV gIterV(ori_kv_ptr); itO gIterO(out_ptr); itMask gIterMask(mask_buf); @@ -165,16 +174,17 @@ void quant_sparse_flash_mla_swa_tadd_pto( for (int dd = 0; dd < Db; ++dd) { tileQ tQ; auto gQ = gIterQ(i, dd); - TCOPYIN(tQ, gQ); + TLOAD(tQ, gQ); + + tileKSrc tKSrc; + auto gK = gIterKSrc(j, dd); + TLOAD(tKSrc, gK); - tileK tK; - auto gK = gIterK(dd, j); - TCOPYIN(tK, gK); + tileKRight tK; + TTRANS(tK, tKSrc); - tileW_out tW_out; - TMATMUL(tW_out, tQ, tK); tileW tW_partial; - TCVT_Impl(tW_partial, tW_out); + TMATMUL(tW_partial, tQ, tK); TADD(tW, tW, tW_partial); } @@ -233,16 +243,17 @@ void quant_sparse_flash_mla_swa_tadd_pto( for (int dd2 = 0; dd2 < Db; ++dd2) { tileQ tQ; auto gQ = gIterQ(i, dd2); - TCOPYIN(tQ, gQ); + TLOAD(tQ, gQ); - tileK tK; - auto gK = gIterK(dd2, j); - TCOPYIN(tK, gK); + tileKSrc tKSrc; + auto gK = gIterKSrc(j, dd2); + TLOAD(tKSrc, gK); + + tileKRight tK; + TTRANS(tK, tKSrc); - tileW_out tW_out; - TMATMUL(tW_out, tQ, tK); tileW tW_partial; - TCVT_Impl(tW_partial, tW_out); + TMATMUL(tW_partial, tQ, tK); TADD(tW, tW, tW_partial); } @@ -259,21 +270,19 @@ void quant_sparse_flash_mla_swa_tadd_pto( 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); + // 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 分块) tileV tV; auto gV = gIterV(j, dd); - TCOPYIN(tV, gV); + TLOAD(tV, gV); - tileO_out tPV_out; - TMATMUL(tPV_out, tW_left, tV); tileO tPV; - TCVT_Impl(tPV, tPV_out); + TMATMUL(tPV, tW_left, tV); TADD(tO, tO, tPV); } @@ -282,7 +291,147 @@ void quant_sparse_flash_mla_swa_tadd_pto( tileO_cast tO_cast; TCVT(tO_cast, tO); auto gO = gIterO(i, dd); - TCOPYOUT(gO, tO_cast); + TSTORE(gO, tO_cast); + } + } +} + +// 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>( + 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>( + 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); + } } } } diff --git a/benchmark/one-level-arch/test/kernel/fa/Makefile b/benchmark/one-level-arch/test/kernel/fa/Makefile index 7666e11f..253765d4 100644 --- a/benchmark/one-level-arch/test/kernel/fa/Makefile +++ b/benchmark/one-level-arch/test/kernel/fa/Makefile @@ -86,8 +86,11 @@ 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 @@ -95,8 +98,12 @@ ifeq ($(TESTCASE), quant_sparse_flash_mla) 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) @@ -104,7 +111,10 @@ ifeq ($(TESTCASE), quant_sparse_flash_mla) DEFINES += -DTsoftmax_scale=$(softmax_scale) DEFINES += -DTwleft=$(wleft) DEFINES += -DTwright=$(wright) - TARGET = $(ELF_HEAD)/$(TESTCASE)_s1$(s1)_s2$(s2)_D$(D)_Tm$(Tm)_Tk$(Tk)_Td$(Td_block).elf + ifeq ($(IMPL), tadd) + DEFINES += -DQSMLA_USE_TADD + endif + 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 include ../../common/Makefile.common 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 index d1cba2ac..de4de163 100644 --- a/benchmark/one-level-arch/test/kernel/fa/src/qsmla_compare.py +++ b/benchmark/one-level-arch/test/kernel/fa/src/qsmla_compare.py @@ -1,21 +1,108 @@ #!/usr/bin/env python3 +import argparse import math import struct +import sys +from pathlib import Path -path = "SuperNPUBench/benchmark/one-level-arch/test/kernel/fa/src/" -with open(path + "qsmla_onepass_npu_out.bin", "rb") as f: - actual = struct.unpack("<32768e", f.read()) +HERE = Path(__file__).resolve().parent -with open(path + "qsmla_golden.bin", "rb") as f: - golden = struct.unpack("<32768f", f.read()) -errors = [abs(float(a) - float(g)) for a, g in zip(actual, golden)] -passed = [e <= 1e-3 + 1e-3 * abs(g) for e, g in zip(errors, golden)] +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() -print(f"passed = {sum(passed)}/32768 ({100*sum(passed)/32768:.6f}%)") -print(f"failed = {32768-sum(passed)}") -print(f"max_abs = {max(errors):.9f}") -print(f"mean_abs= {sum(errors)/len(errors):.9f}") -print(f"nan_npu = {sum(math.isnan(float(x)) for x in actual)}") + +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/quant_sparse_flash_mla.cpp b/benchmark/one-level-arch/test/kernel/fa/src/quant_sparse_flash_mla.cpp index e8c8a350..2a90c0b5 100644 --- 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 @@ -5,8 +5,23 @@ // #include "fa/quant_sparse_flash_mla_tadd_pto.hpp" #include "fa/quant_sparse_flash_mla_onepass_pto.hpp" +#ifndef Tbatch #define B 1 -#define H 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 @@ -77,27 +92,46 @@ int main(){ using qdtype = __half; using kvdtype = __half; using odttype = __half; - using Config = QsmlaConfig; + 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*H*s1*D + 2*ALIGN]; - kvdtype kvp[B*H*s2*D + 2*ALIGN]; + 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; - init_deterministic(q, B*H*s1*D, 1); - init_deterministic(kv, B*H*s2*D, 2); + init_deterministic(q, B*s1*N1*D, 1); + init_deterministic(kv, B*s2*N2*D, 2); BENCHSTART; - for(int i=0;i( - out + i*H*s1*D + j*s1*D, - q + i*H*s1*D + j*s1*D, - kv + i*H*s2*D + j*s2*D, + if constexpr (N1 == 1 && N2 == 1) { + quant_sparse_flash_mla_swa_onepass_config_pto< + 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 { + quant_sparse_flash_mla_swa_onepass_bsnd_pto< + qdtype, kvdtype, odttype, Config>( + out, q, kv, softmax_scale_val, win_left, win_right, @@ -113,7 +147,6 @@ int main(){ (int*)nullptr, // metadata (float*)nullptr // softmax_lse ); - } } BENCHEND; From ced5e571a2ccb8eadfa0951974a867fd4287e87c Mon Sep 17 00:00:00 2001 From: chenglongyu Date: Thu, 27 Aug 2026 11:26:44 +0800 Subject: [PATCH 8/9] add qsmla 0827 swa kv interval --- .../kernels/fa/qsmla_config_pto.hpp | 61 +++++++++ .../fa/quant_sparse_flash_mla_onepass_pto.hpp | 98 ++++++++++---- .../fa/quant_sparse_flash_mla_tadd_pto.hpp | 127 +++++++++++++----- .../kernel/fa/src/quant_sparse_flash_mla.cpp | 13 +- 4 files changed, 245 insertions(+), 54 deletions(-) diff --git a/benchmark/one-level-arch/kernels/fa/qsmla_config_pto.hpp b/benchmark/one-level-arch/kernels/fa/qsmla_config_pto.hpp index a01be9dd..8be7a6fa 100644 --- a/benchmark/one-level-arch/kernels/fa/qsmla_config_pto.hpp +++ b/benchmark/one-level-arch/kernels/fa/qsmla_config_pto.hpp @@ -3,6 +3,67 @@ #include +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; 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 index b0ffe00f..a9e37fd6 100644 --- 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 @@ -21,7 +21,9 @@ // - tile 寄存器压力更大 (O[Db] 数组跨 j 循环存活) // // 【mask 方式】 -// TADD: mask=float[s1*s2], 0.0(有效)/-1e30(无效), score += mask +// BSND 共享 Q token 路径只预生成 first/last/zero 三个 [TileM,TileK] +// mask,内部整块直接复用 zero mask;旧 2D 路径保留 [s1,s2] mask。 +// mask 写入仍放在 Tile/CUBE 流程前,避免当前编码问题。 // // 【切换方法】 // test 文件中: @@ -42,17 +44,16 @@ static inline void build_swa_mask_onepass( 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; - int diagonal = s2 - logical_s1 + logical_q; - int lo = diagonal - win_left; - int hi = diagonal + win_right; + const QsmlaSwaRange range = qsmla_swa_range( + s2, logical_s1, logical_q, win_left, win_right); for (int kv = 0; kv < s2; ++kv) { - bool valid = (kv >= lo) && (kv <= hi); - mask[q * s2 + kv] = valid ? 0.0f : -1e30f; + mask[q * s2 + kv] = qsmla_swa_mask_value(kv, range); } } } -template +template void quant_sparse_flash_mla_swa_onepass_config_pto( odttype* out_ptr, qdtype* q_ptr, @@ -82,17 +83,23 @@ void quant_sparse_flash_mla_swa_onepass_config_pto( 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; - float mask_buf[s1 * s2]; - build_swa_mask_onepass( - mask_buf, s1, s2, ori_win_left, ori_win_right, - q_position, q_sequence_length); + 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 gmMask = global_tensor>; using tileQ = TileLeft; using tileKSrc = Tile; @@ -112,20 +119,46 @@ void quant_sparse_flash_mla_swa_onepass_config_pto( using itKSrc = global_iterator; using itV = global_iterator; using itO = global_iterator; - using itMask = global_iterator; itQ gIterQ(q_ptr); - itKSrc gIterKSrc(ori_kv_ptr); - itV gIterV(ori_kv_ptr); itO gIterO(out_ptr); - itMask gIterMask(mask_buf); 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) { + 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 @@ -156,7 +189,7 @@ void quant_sparse_flash_mla_swa_onepass_config_pto( TEXPANDS(tO[dd], 0.0f); } - for (int j = 0; j < Kb; ++j) { + for (int j = 0; j < kv_block_count; ++j) { // --- Step 1: QK^T 沿全 D 累加 --- // SuperScalarModel main does not preserve the input ACC correctly @@ -186,8 +219,27 @@ void quant_sparse_flash_mla_swa_onepass_config_pto( TMULS(tW, tW, scale); tileMask tMask; - auto gMask = gIterMask(i, j); - TLOAD(tMask, gMask); + 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 --- @@ -333,7 +385,7 @@ void quant_sparse_flash_mla_swa_onepass_bsnd_pto( Config::TileM, Config::TileK, Config::TileD, Config::TileM>; quant_sparse_flash_mla_swa_onepass_config_pto< - qdtype, kvdtype, odttype, WorkConfig>( + 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), @@ -362,7 +414,7 @@ void quant_sparse_flash_mla_swa_onepass_bsnd_pto( 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>( + 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, 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 index 685cf22f..a094da31 100644 --- 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 @@ -69,17 +69,16 @@ static inline void build_swa_mask_tadd( 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; - int diagonal = s2 - logical_s1 + logical_q; - int lo = diagonal - win_left; - int hi = diagonal + win_right; + const QsmlaSwaRange range = qsmla_swa_range( + s2, logical_s1, logical_q, win_left, win_right); for (int kv = 0; kv < s2; ++kv) { - bool valid = (kv >= lo) && (kv <= hi); - mask[q * s2 + kv] = valid ? 0.0f : -1e30f; + mask[q * s2 + kv] = qsmla_swa_mask_value(kv, range); } } } -template +template void quant_sparse_flash_mla_swa_tadd_config_pto( odttype* out_ptr, qdtype* q_ptr, @@ -109,19 +108,25 @@ void quant_sparse_flash_mla_swa_tadd_config_pto( 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; - // kernel 内部计算 mask, 放在 stack 上 - // s1*s2*sizeof(float) = 64*128*4 = 32KB (可容纳) - float mask_buf[s1 * s2]; - build_swa_mask_tadd( - mask_buf, s1, s2, ori_win_left, ori_win_right, - q_position, q_sequence_length); + 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 gmMask = global_tensor>; using tileQ = TileLeft; using tileKSrc = Tile; @@ -143,29 +148,55 @@ void quant_sparse_flash_mla_swa_tadd_config_pto( using itKSrc = global_iterator; using itV = global_iterator; using itO = global_iterator; - using itMask = global_iterator; itQ gIterQ(q_ptr); - itKSrc gIterKSrc(ori_kv_ptr); - itV gIterV(ori_kv_ptr); itO gIterO(out_ptr); - itMask gIterMask(mask_buf); 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) { + 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) - // 遍历全部 KV 块, 用 mask 屏蔽窗口外 token + // 只遍历与 SWA 有效区间相交的 KV 块 // ============================================================ tileMax tMax; TEXPANDS(tMax, -1e30f); tileSum tSum; TEXPANDS(tSum, 0.0f); - for (int j = 0; j < Kb; ++j) { + for (int j = 0; j < kv_block_count; ++j) { // QK^T 沿 D 维累加 tileW tW; @@ -190,10 +221,28 @@ void quant_sparse_flash_mla_swa_tadd_config_pto( TMULS(tW, tW, scale); - // 应用 token 级 mask: score += mask (0 保持原值, -1e30 屏蔽) tileMask tMask; - auto gMask = gIterMask(i, j); - TLOAD(tMask, gMask); + 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)) @@ -234,7 +283,7 @@ void quant_sparse_flash_mla_swa_tadd_config_pto( tileO tO; TEXPANDS(tO, 0.0f); - for (int j = 0; j < Kb; ++j) { + for (int j = 0; j < kv_block_count; ++j) { // 计算完整 QK^T (沿 D 累加, 与 Pass 1 一致) tileW tW; @@ -259,10 +308,28 @@ void quant_sparse_flash_mla_swa_tadd_config_pto( TMULS(tW, tW, scale); - // 应用 token 级 mask: score += mask (0 保持原值, -1e30 屏蔽) tileMask tMask; - auto gMask = gIterMask(i, j); - TLOAD(tMask, gMask); + 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 @@ -361,7 +428,7 @@ void quant_sparse_flash_mla_swa_tadd_bsnd_pto( Config::TileM, Config::TileK, Config::TileD, Config::TileM>; quant_sparse_flash_mla_swa_tadd_config_pto< - qdtype, kvdtype, odttype, WorkConfig>( + 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), @@ -392,7 +459,7 @@ void quant_sparse_flash_mla_swa_tadd_bsnd_pto( 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>( + 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, 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 index 2a90c0b5..a5c3c269 100644 --- 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 @@ -2,8 +2,11 @@ #include "benchmark.h" #include "fileop.h" // #include "fa/quant_sparse_flash_mla_pto.hpp" -// #include "fa/quant_sparse_flash_mla_tadd_pto.hpp" +#ifdef QSMLA_USE_TADD +#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 @@ -110,7 +113,11 @@ int main(){ BENCHSTART; 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, @@ -129,7 +136,11 @@ int main(){ (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, From 99c91ddb6c481c0502090858c14b2a67a109ccee Mon Sep 17 00:00:00 2001 From: chenglongyu Date: Sat, 29 Aug 2026 14:17:28 +0800 Subject: [PATCH 9/9] add qsmla 0829 4-PE --- .../kernels/fa/qsmla_config_pto.hpp | 20 ++ .../fa/quant_sparse_flash_mla_tadd_pto.hpp | 288 +++++++++++++++++- .../one-level-arch/test/kernel/fa/Makefile | 7 +- .../kernel/fa/src/quant_sparse_flash_mla.cpp | 49 ++- 4 files changed, 354 insertions(+), 10 deletions(-) diff --git a/benchmark/one-level-arch/kernels/fa/qsmla_config_pto.hpp b/benchmark/one-level-arch/kernels/fa/qsmla_config_pto.hpp index 8be7a6fa..35b65b10 100644 --- a/benchmark/one-level-arch/kernels/fa/qsmla_config_pto.hpp +++ b/benchmark/one-level-arch/kernels/fa/qsmla_config_pto.hpp @@ -3,6 +3,26 @@ #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; 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 index a094da31..21c61820 100644 --- 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 @@ -206,14 +206,11 @@ void quant_sparse_flash_mla_swa_tadd_config_pto( 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); @@ -293,14 +290,11 @@ void quant_sparse_flash_mla_swa_tadd_config_pto( 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); @@ -344,10 +338,9 @@ void quant_sparse_flash_mla_swa_tadd_config_pto( TCVT(tW_left, tW); // PV = p * V (当前 D 分块) - tileV tV; auto gV = gIterV(j, dd); + tileV tV; TLOAD(tV, gV); - tileO tPV; TMATMUL(tPV, tW_left, tV); @@ -363,6 +356,285 @@ void quant_sparse_flash_mla_swa_tadd_config_pto( } } +// 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 (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< @@ -159,6 +205,7 @@ int main(){ (float*)nullptr // softmax_lse ); } +#endif BENCHEND; return 0;