// SPDX-License-Identifier: Apache-2.0 // Attention score scale + mask + softmax (tt/smsm_kernel.py, ATTN_SMSM), compute. Per tile row (one head, 32 query // rows, Wt key tiles): // phase A: x = trunc_tf32(s * scale) + mask, the smask kernel's SFPU sequence (kernels/smask_compute.cpp: the two // stock binary_ng programs bit for bit), packed to cb_x (L1) instead of DRAM; // phase B: the stock ttnn.softmax(numeric_stable=True) of that row, i.e. the no-mask path of // ttnn/.../softmax/device/kernels/attention/compute/softmax.cpp with the same kernel_lib calls: row max // (FPU reduce), exp(x - max) (FPU bcast sub + SFPU exp), row sum (FPU reduce) + precise fp32 // reciprocal, x * 1/sum (FPU bcast mul). // The stock softmax unpacks its fp32 input to SrcA (TF32); x holds TF32 values already (phase A truncates them and // the mask is 0 / -inf), so the stock program saw exactly these operands. // CT args: [0] Wt, [1] per-core RT-arg count (P3), [2] trunc_tf32, [3] ndst (block), [4] out pad tiles per row, // [5] scale mode: 0 = multiply by a tile filled with the scale (mul_binary_tile, as the stock binary_ng // program), 1 = mul_unary_tile with the scale bits (the same SFPU fp32 multiply, no scale tile copy). // Per-core RT args: [nrows, 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/tile_move_copy.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 out_pad = get_compile_time_arg_val(4); constexpr uint32_t scale_mode = get_compile_time_arg_val(5); constexpr uint32_t scale_bits = get_compile_time_arg_val(6); constexpr uint32_t cb_s = 0, cb_m = 1, cb_c = 2, cb_max_scaler = 3, cb_sum_scaler = 4; constexpr uint32_t cb_out = 16; constexpr uint32_t cb_x = 24, cb_max = 25, cb_exps = 26, cb_recip = 27; constexpr auto RNE = ckernel::DstRoundingMode::NearestEven; constexpr uint32_t Wp = (Wt + 1) / 2; // tile pairs per row (the last one half-used when Wt is odd) // the stock calc_numeric_stable (softmax.cpp), CB ids instead of DFB handles 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 nrows = get_arg_val(0); if (nrows == 0) { return; } compute_kernel_hw_startup(cb_s, cb_c, cb_x); cb_wait_front(cb_c, 1); cb_wait_front(cb_max_scaler, 1); cb_wait_front(cb_sum_scaler, 1); for (uint32_t r = 0; r < nrows; ++r) { // ---- phase A: x = trunc(s * scale) + mask (pairs of tiles) ---- reconfig_data_format(cb_s, cb_c); pack_reconfig_data_format(cb_x); cb_wait_front(cb_m, 2 * Wp); for (uint32_t p = 0; p < Wp; ++p) { const uint32_t j = 2 * p; const bool two = (j + 1) < Wt; cb_wait_front(cb_s, 2); tile_regs_acquire(); copy_init(cb_s); copy_tile(cb_s, 0, 0); copy_tile(cb_s, 1, 1); if constexpr (scale_mode == 0) { copy_init(cb_c); copy_tile(cb_c, 0, 2); mul_binary_tile_init(); mul_binary_tile(0, 2, 0); mul_binary_tile(1, 2, 1); } else { 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(scale_mode == 0 ? cb_c : cb_s, cb_m); copy_init(cb_m); copy_tile(cb_m, j, 2); copy_tile(cb_m, j + 1, 3); reconfig_data_format_srca(cb_m, cb_s); 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_s, 2); } cb_pop_front(cb_m, 2 * Wp); // ---- phase B: the stock numeric-stable softmax of the row ---- 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_out, ckl::ReservePolicy::PerBlockSize, ckl::PushPolicy::PerBlockSize)>( ckl::IterationShape::tiles(Wt).block_size(ndst)); if constexpr (out_pad > 0) { cb_reserve_back(cb_out, out_pad); cb_push_back(cb_out, out_pad); } } cb_pop_front(cb_c, 1); cb_pop_front(cb_max_scaler, 1); cb_pop_front(cb_sum_scaler, 1); }