File size: 5,202 Bytes
c699c4c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 | // 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 <cstdint>
#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<EXP_APPROX>();
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<EXP_APPROX>(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<uint32_t>(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<block_w, num_subblocks_w, subblock_w, dfb_in0, dfb_max_scaler, dfb_max, dfb_exps>();
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<BroadcastType::COL>(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);
}
}
|