#ifndef CPU_ATTN_VEC_HPP #define CPU_ATTN_VEC_HPP #include "cpu_attn_fp8.hpp" #include "cpu_attn_impl.hpp" namespace cpu_attention { namespace { // Load 32 kv_cache_t elements starting at ptr and return them as two FP32Vec16s // covering the lower 16 and upper 16 positions. // For FP8: both halves come from a single BF16Vec32 dequant of 32 bytes. // For BF16/FP16/FP32: two separate vector loads at ptr and ptr+16. template FORCE_INLINE std::pair load_b_pair_vec( const kv_cache_t* ptr) { if constexpr (std::is_same_v) { // BF16 container, but values are in the FP16 exponent range (bias 15 not // 127). vec_op::BF16Vec32 bf16_b_reg(reinterpret_cast(ptr), vec_op::fp8_e4m3_tag{}); return {vec_op::FP32Vec16(bf16_b_reg, 0), vec_op::FP32Vec16(bf16_b_reg, 1)}; } else if constexpr (std::is_same_v) { vec_op::BF16Vec32 bf16_b_reg(reinterpret_cast(ptr), vec_op::fp8_e5m2_tag{}); return {vec_op::FP32Vec16(bf16_b_reg, 0), vec_op::FP32Vec16(bf16_b_reg, 1)}; } else { using load_vec_t = typename VecTypeTrait::vec_t; return std::make_pair(vec_op::FP32Vec16(load_vec_t(ptr)), vec_op::FP32Vec16(load_vec_t(ptr + 16))); } } // 8-2-16 pattern, 8 regs for A, 2 regs for B, 16 regs for C, [8, K] @ [k, 32] template class TileGemm82 { public: template FORCE_INLINE static void gemm(const int32_t m_size, float* __restrict__ a_tile, kv_cache_t* __restrict__ b_tile, float* __restrict__ c_tile, const int64_t lda, const int64_t ldb, const int64_t ldc, const int32_t block_size, const int32_t dynamic_k_size, const bool accum_c) { switch (m_size) { case 1: gemm_micro<1>(a_tile, b_tile, c_tile, lda, ldb, ldc, block_size, dynamic_k_size, accum_c); break; case 2: gemm_micro<2>(a_tile, b_tile, c_tile, lda, ldb, ldc, block_size, dynamic_k_size, accum_c); break; case 3: case 4: gemm_micro<4>(a_tile, b_tile, c_tile, lda, ldb, ldc, block_size, dynamic_k_size, accum_c); break; case 5: case 6: gemm_micro<6>(a_tile, b_tile, c_tile, lda, ldb, ldc, block_size, dynamic_k_size, accum_c); break; case 7: case 8: gemm_micro<8>(a_tile, b_tile, c_tile, lda, ldb, ldc, block_size, dynamic_k_size, accum_c); break; } } template static void gemm_micro(float* __restrict__ a_tile, kv_cache_t* __restrict__ b_tile, float* __restrict__ c_tile, const int64_t lda, const int64_t ldb, const int64_t ldc, const int32_t block_size, const int32_t dynamic_k_size, const bool accum_c) { static_assert(0 < M && M <= 8); float* __restrict__ curr_c_0 = c_tile; float* __restrict__ curr_c_1 = c_tile + 16; vec_op::FP32Vec16 c_regs[M * 2]; if (accum_c) { float* __restrict__ curr_m_c_0 = curr_c_0; float* __restrict__ curr_m_c_1 = curr_c_1; vec_op::unroll_loop([&](int32_t i) { c_regs[i * 2] = vec_op::FP32Vec16(curr_m_c_0); c_regs[i * 2 + 1] = vec_op::FP32Vec16(curr_m_c_1); // update curr_m_c_0 += ldc; curr_m_c_1 += ldc; }); } float* __restrict__ curr_a = a_tile; kv_cache_t* __restrict__ curr_b = b_tile; for (int32_t k = 0; k < dynamic_k_size; ++k) { auto [fp32_b_0_reg, fp32_b_1_reg] = load_b_pair_vec(curr_b); float* __restrict__ curr_m_a = curr_a; vec_op::unroll_loop([&](int32_t i) { vec_op::FP32Vec16 a_reg(*curr_m_a); c_regs[i * 2] = c_regs[i * 2] + a_reg * fp32_b_0_reg; c_regs[i * 2 + 1] = c_regs[i * 2 + 1] + a_reg * fp32_b_1_reg; // update curr_m_a += lda; }); // update curr_a += 1; curr_b += ldb; } vec_op::unroll_loop([&](int32_t i) { c_regs[i * 2].save(curr_c_0); c_regs[i * 2 + 1].save(curr_c_1); // update curr_c_0 += ldc; curr_c_1 += ldc; }); } }; } // namespace // This is a general but naive implementation based on vector instructions template class AttentionImpl { static constexpr bool fp8_kv = std::is_same_v || std::is_same_v; public: using query_t = scalar_t; using q_buffer_t = float; using kv_cache_t = kv_cache_scalar_t; using logits_buffer_t = float; using partial_output_buffer_t = float; using prob_buffer_t = float; constexpr static int64_t BlockSizeAlignment = 32; // KV token num unit of QK and PV phases constexpr static int64_t HeadDimAlignment = 32; // headdim num unit of PV phase constexpr static int64_t MaxQHeadNumPerIteration = 8; constexpr static int64_t HeadDim = head_dim; constexpr static ISA ISAType = ISA::VEC; constexpr static bool scale_on_logits = fp8_kv; float k_scale = 1.0f; float v_scale = 1.0f; public: void init_from_input(const AttentionInput* input) { if constexpr (fp8_kv) { k_scale = input->k_scale_fp8; v_scale = input->v_scale_fp8; } } float get_output_v_scale() const noexcept { if constexpr (fp8_kv) { // VEC dequant unpacks FP8 into a pseudo-FP16 layout (exponent bias 15). // E4M3 (bias=7) needs correction 2^(15-7) = 2^8; E5M2 bias matches FP16 // so no correction. if constexpr (std::is_same_v) { return v_scale; } else { return v_scale * 0x1p8f; } } return 1.0f; } template