// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc. // SPDX-License-Identifier: Apache-2.0 // // Score softmax on the otherwise idle cores of the merged head op (models/tt/desc_head.py, SP_HEAD_SM=1). // The per-row math is ttnn's sharded attention softmax kernel (softmax_sharded.cpp, NUMERIC_STABLE, no mask) // verbatim, with the compute config ttnn.softmax uses (HiFi4, fp32 DST off, approx mode on -> bf16 // intermediates): max reduce, x - max, exp, sum reduce + recip, x * (1 / sum). The score logits arrive with // -inf in the padding columns (65..95), so this equals ttnn.softmax's padded-column-masked path bit for bit // (verified: scratch/r8_sm_test.py, test_superpoint_merged_head_softmax). #include #include "api/compute/eltwise_binary.h" #include "api/compute/tile_move_copy.h" #include "api/compute/bcast.h" #include "api/compute/softmax.h" #include "api/compute/reduce.h" #include "api/dataflow/dataflow_buffer.h" #include "ttnn/cpp/ttnn/kernel_lib/reduce_helpers_compute.hpp" template < std::uint32_t block_w, std::uint32_t num_subblocks_w, std::uint32_t subblock_w, std::uint32_t dfb_in_id, std::uint32_t dfb_max_scaler_id, std::uint32_t dfb_max_id, std::uint32_t dfb_out_id> ALWI void calc_numeric_stable() { DataflowBuffer dfb_in_obj(dfb_in_id); DataflowBuffer dfb_max_obj(dfb_max_id); DataflowBuffer dfb_out_obj(dfb_out_id); compute_kernel_lib::reduce< PoolType::MAX, ReduceDim::REDUCE_ROW, dfb_in_id, dfb_max_scaler_id, dfb_max_id, compute_kernel_lib::ReduceInputPolicy::NoWaitNoPop>(compute_kernel_lib::ReduceInputBlockShape::row(block_w)); exp_tile_init(); reconfig_data_format_srcb(dfb_max_id); dfb_max_obj.wait_front(1); sub_bcast_cols_init(dfb_in_id, dfb_max_id); std::uint32_t index_subblock_w_offset = 0; for (std::uint32_t j = 0; j < num_subblocks_w; j++) { tile_regs_acquire(); dfb_out_obj.reserve_back(subblock_w); for (std::uint32_t w = 0; w < subblock_w; w++) { std::uint32_t index = w + index_subblock_w_offset; sub_tiles_bcast_cols(dfb_in_id, dfb_max_id, index, 0, w); } dfb_out_obj.reserve_back(subblock_w); for (std::uint32_t w = 0; w < subblock_w; w++) { exp_tile(w); } tile_regs_commit(); tile_regs_wait(); for (std::uint32_t w = 0; w < subblock_w; w++) { pack_tile(w, dfb_out_id); } tile_regs_release(); dfb_out_obj.push_back(subblock_w); index_subblock_w_offset += subblock_w; } dfb_in_obj.pop_front(block_w); dfb_max_obj.pop_front(1); dfb_out_obj.wait_front(block_w); } void kernel_main() { const std::uint32_t nrows = get_arg_val(0); constexpr std::uint32_t dfb_in0 = get_compile_time_arg_val(0); constexpr std::uint32_t dfb_max_scaler = get_compile_time_arg_val(1); constexpr std::uint32_t dfb_sum_scaler = get_compile_time_arg_val(2); constexpr std::uint32_t dfb_exps = get_compile_time_arg_val(3); constexpr std::uint32_t dfb_recipsumexps = get_compile_time_arg_val(4); constexpr std::uint32_t dfb_out0 = get_compile_time_arg_val(5); constexpr std::uint32_t dfb_max = get_compile_time_arg_val(6); constexpr std::uint32_t block_w = 3, subblock_w = 3, num_subblocks_w = 1; compute_kernel_hw_startup(dfb_in0, dfb_max_scaler, dfb_exps); DataflowBuffer dfb_in0_obj(dfb_in0); DataflowBuffer dfb_exps_obj(dfb_exps); DataflowBuffer dfb_recipsumexps_obj(dfb_recipsumexps); DataflowBuffer dfb_out0_obj(dfb_out0); for (std::uint32_t i = 0; i < nrows; i++) { dfb_in0_obj.wait_front(block_w); calc_numeric_stable(); dfb_exps_obj.wait_front(block_w); compute_kernel_lib::reduce< PoolType::SUM, ReduceDim::REDUCE_ROW, dfb_exps, dfb_sum_scaler, dfb_recipsumexps, compute_kernel_lib::ReduceInputPolicy::NoWaitNoPop>( compute_kernel_lib::ReduceInputBlockShape::row(block_w), compute_kernel_lib::ReduceInputMemoryLayout::contiguous(), compute_kernel_lib::NoAccumulation{}, [](std::uint32_t) { recip_tile_init(); recip_tile(0); }); reconfig_data_format(dfb_exps, dfb_recipsumexps); pack_reconfig_data_format(dfb_out0); dfb_recipsumexps_obj.wait_front(1); mul_bcast_cols_init(dfb_exps, dfb_recipsumexps); tile_regs_acquire(); dfb_out0_obj.reserve_back(subblock_w); for (std::uint32_t w = 0; w < subblock_w; w++) { mul_tiles_bcast(dfb_exps, dfb_recipsumexps, w, 0, w); } tile_regs_commit(); tile_regs_wait(); for (std::uint32_t w = 0; w < subblock_w; w++) { pack_tile(w, dfb_out0); } tile_regs_release(); dfb_out0_obj.push_back(subblock_w); dfb_recipsumexps_obj.pop_front(1); dfb_exps_obj.pop_front(block_w); } }