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);
    }
}