// SPDX-License-Identifier: Apache-2.0 // Fused fp32 matmul attention (tt/fattn_kernel.py, ATTN_FUSED), compute. Per unit (head h, query tile row r; Wt key // tiles), the five stock programs of tt/attention.py (ATTN_FAST=1 + ATTN_SMSM) in one pass, the scores and the // probabilities kept in L1: // phase Q: s_j = Q_r K_j^T in DEST (matmul_block, in1 transposed by the unpacker: the stock Q K^T matmul with its // K = 32 in one block), two key tiles at a time; // phase A: x_j = trunc_tf32(s_j * scale) + mask_j on the SFPU (kernels/smsm_compute.cpp, scale mode 1), packed to // cb_x; // phase B: the stock numeric-stable softmax of the row (kernels/smsm_compute.cpp phase B) -> cb_p; // phase C: o = sum_j P_j V_j accumulated in DEST in key order (the stock P V matmul: the whole K in one block), // packed to cb_o; the writer stores it as tile (r, h) of the merged-heads output. // (kcat_out, KCAT_EMIT) the output tile as the operand of the next K-concatenated split linear: o_hi = bf16(o) // and o_lo = o - o_hi (kernels/kcat_compute.cpp's LLK calls on the packed fp32 tile) -> cb_o / cb_lo. // CT args: [0] Wt, [1] per-core RT-arg count (P3), [2] trunc_tf32, [3] ndst, [4] P pad tiles per row (to a multiple // of ndst), [5] scale bits (fp32), [6] kcat_out. Per-core RT args: [nunits, 0]. #include #include "api/compute/common.h" #include "api/compute/compute_kernel_api.h" #include "api/compute/eltwise_binary.h" #include "api/compute/eltwise_binary_sfpu.h" #include "api/compute/eltwise_unary/binop_with_scalar.h" #include "api/compute/eltwise_unary/bitwise.h" #include "api/compute/eltwise_unary/eltwise_unary.h" #include "api/compute/eltwise_unary/typecast.h" #include "api/compute/tile_move_copy.h" #include "api/compute/matmul.h" #include "api/compute/bcast.h" #include "api/compute/softmax.h" #include "api/compute/reduce.h" #include "ttnn/cpp/ttnn/kernel_lib/reduce_helpers_compute.hpp" #include "ttnn/cpp/ttnn/kernel_lib/eltwise/api/chain.hpp" #include "ttnn/cpp/ttnn/kernel_lib/eltwise/api/convenience.hpp" #include "ttnn/cpp/ttnn/kernel_lib/eltwise/unary/math.hpp" #include "ttnn/cpp/ttnn/kernel_lib/eltwise/core/optional.hpp" namespace ckl = compute_kernel_lib; using namespace ckernel; constexpr uint32_t Wt = get_compile_time_arg_val(0); constexpr uint32_t trunc_tf32 = get_compile_time_arg_val(2); constexpr uint32_t ndst = get_compile_time_arg_val(3); constexpr uint32_t p_pad = get_compile_time_arg_val(4); constexpr uint32_t scale_bits = get_compile_time_arg_val(5); constexpr uint32_t kcat_out = get_compile_time_arg_val(6); constexpr uint32_t cb_q = 0, cb_m = 1, cb_k = 2, cb_max_scaler = 3, cb_sum_scaler = 4, cb_v = 5; constexpr uint32_t cb_o = 16, cb_lo = 17, cb_y = 29; constexpr uint32_t cb_x = 24, cb_max = 25, cb_exps = 26, cb_recip = 27, cb_p = 28; constexpr auto RNE = ckernel::DstRoundingMode::NearestEven; constexpr uint32_t Wp = (Wt + 1) / 2; constexpr uint32_t Wr = Wt + p_pad; template void calc_numeric_stable(std::uint32_t W, std::uint32_t nd) { compute_kernel_lib::reduce< PoolType::MAX, ReduceDim::REDUCE_ROW, dfb_in, dfb_max_scaler, dfb_max, compute_kernel_lib::ReduceInputPolicy::WaitUpfrontNoPop, compute_kernel_lib::ReduceDataFormatReconfigMode::INPUT>(compute_kernel_lib::ReduceInputBlockShape::row(W)); ckl::eltwise_chain( ckl::IterationShape::tiles(W).block_size(nd), ckl::BinaryFpu< ckl::BinaryFpuOp::Sub, ckl::input( dfb_in, ckl::WaitPolicy::Upfront, ckl::PopPolicy::AtEnd, ckl::InputTileMapping::Block, ckl::DataFormatReconfig::Disabled), ckl::input(dfb_max, ckl::BroadcastDim::Col, ckl::WaitPolicy::Upfront, ckl::PopPolicy::AtEnd)>{}, ckl::Exp{}, ckl::PackTile{}); cb_wait_front(dfb_out, W); } void kernel_main() { const uint32_t nunits = get_arg_val(0); if (nunits == 0) { return; } compute_kernel_hw_startup(cb_q, cb_k, cb_x); cb_wait_front(cb_max_scaler, 1); cb_wait_front(cb_sum_scaler, 1); for (uint32_t u = 0; u < nunits; ++u) { // ---- phases Q + A: x_j = trunc(Q K_j^T * scale) + mask_j, two key tiles at a time ---- cb_wait_front(cb_q, 1); cb_wait_front(cb_m, 2 * Wp); pack_reconfig_data_format(cb_x); for (uint32_t p = 0; p < Wp; ++p) { const uint32_t j = 2 * p; const bool two = (j + 1) < Wt; cb_wait_front(cb_k, 2); reconfig_data_format(cb_k, cb_q); matmul_block_init(cb_q, cb_k, 1, 1, 1, 1); tile_regs_acquire(); matmul_block(cb_q, cb_k, 0, 0, 0, 1, 1, 1, 1); matmul_block(cb_q, cb_k, 0, 1, 1, 1, 1, 1, 1); binop_with_scalar_tile_init(); mul_unary_tile(0, scale_bits); mul_unary_tile(1, scale_bits); if constexpr (trunc_tf32) { bitwise_and_tile_init(); bitwise_and_tile(0, 0xFFFFE000u); bitwise_and_tile(1, 0xFFFFE000u); } reconfig_data_format_srca(cb_k, cb_m); copy_init(cb_m); copy_tile(cb_m, j, 2); copy_tile(cb_m, j + 1, 3); add_binary_tile_init(); add_binary_tile(0, 2, 0); add_binary_tile(1, 3, 1); const uint32_t np = two ? 2 : 1; cb_reserve_back(cb_x, np); tile_regs_commit(); tile_regs_wait(); pack_tile(0, cb_x); if (two) { pack_tile(1, cb_x); } tile_regs_release(); cb_push_back(cb_x, np); cb_pop_front(cb_k, 2); } cb_pop_front(cb_m, 2 * Wp); cb_pop_front(cb_q, 1); // ---- phase B: the stock numeric-stable softmax of the row -> cb_p ---- reconfig_data_format(cb_x, cb_x); pack_reconfig_data_format(cb_exps); copy_init(cb_x); calc_numeric_stable(Wt, ndst); reconfig_data_format(cb_exps, cb_sum_scaler); compute_kernel_lib::reduce< PoolType::SUM, ReduceDim::REDUCE_ROW, cb_exps, cb_sum_scaler, cb_recip, compute_kernel_lib::ReduceInputPolicy::WaitUpfrontNoPop>( compute_kernel_lib::ReduceInputBlockShape::row(Wt), compute_kernel_lib::ReduceInputMemoryLayout::contiguous(), compute_kernel_lib::NoAccumulation{}, [](std::uint32_t) { if constexpr (DST_ACCUM_MODE) { recip_tile_init(); recip_tile(0); } else { recip_tile_init(); recip_tile(0); } }); ckl::mul< ckl::input(cb_exps, ckl::WaitPolicy::Upfront, ckl::PopPolicy::AtEnd, ckl::InputTileMapping::Block), ckl::input(cb_recip, ckl::BroadcastDim::Col, ckl::WaitPolicy::Upfront, ckl::PopPolicy::AtEnd), ckl::output(cb_p, ckl::ReservePolicy::PerBlockSize, ckl::PushPolicy::PerBlockSize)>( ckl::IterationShape::tiles(Wt).block_size(ndst)); if constexpr (p_pad > 0) { cb_reserve_back(cb_p, p_pad); cb_push_back(cb_p, p_pad); } // ---- phase C: o = sum_j P_j V_j (key order, one DEST accumulation) ---- cb_wait_front(cb_p, Wr); reconfig_data_format(cb_v, cb_p); pack_reconfig_data_format(cb_o); matmul_block_init(cb_p, cb_v, 0, 1, 1, 1); tile_regs_acquire(); for (uint32_t j = 0; j < Wt; ++j) { cb_wait_front(cb_v, 1); matmul_block(cb_p, cb_v, j, 0, 0, 0, 1, 1, 1); cb_pop_front(cb_v, 1); } if constexpr (kcat_out) { cb_reserve_back(cb_y, 1); tile_regs_commit(); tile_regs_wait(); pack_tile(0, cb_y); tile_regs_release(); cb_push_back(cb_y, 1); cb_pop_front(cb_p, Wr); // ---- the split operand of the next linear (kernels/kcat_compute.cpp) ---- cb_wait_front(cb_y, 1); reconfig_data_format_srca(cb_y); tile_regs_acquire(); copy_init(cb_y); copy_tile(cb_y, 0, 0); copy_tile(cb_y, 0, 1); typecast_tile_init<(uint32_t)DataFormat::Float32, (uint32_t)DataFormat::Float16_b>(); typecast_tile<(uint32_t)DataFormat::Float32, (uint32_t)DataFormat::Float16_b>(0); sub_binary_tile_init(); sub_binary_tile(1, 0, 1); cb_reserve_back(cb_o, 1); cb_reserve_back(cb_lo, 1); tile_regs_commit(); tile_regs_wait(); pack_tile(0, cb_o); pack_tile(1, cb_lo); tile_regs_release(); cb_push_back(cb_o, 1); cb_push_back(cb_lo, 1); cb_pop_front(cb_y, 1); } else { cb_reserve_back(cb_o, 1); tile_regs_commit(); tile_regs_wait(); pack_tile(0, cb_o); tile_regs_release(); cb_push_back(cb_o, 1); cb_pop_front(cb_p, Wr); } } cb_pop_front(cb_max_scaler, 1); cb_pop_front(cb_sum_scaler, 1); }