File size: 6,411 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
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// SuperPoint NMS window max, ONE op (replaces halo + max_pool2d [2r+1, 1] + halo + max_pool2d [1, 2r+1]).
// Each strip row (PW positions x 32 lanes bf16 = PW * 64 B) is one "pseudo tile" in DST (bf16, full
// sync): position i = elements 32 i .. 32 i + 31 = DST rows 2 i, 2 i + 1 of the tile, so an SFPU load at
// row offset 2 i reads position i of every lane. Input tiles 0 .. NROW-1 = image rows y0 - R .. y0 + ROWS
// - 1 + R (zero outside the image; scores >= 0, so zeros never change a window max); per output row k:
//   V(k, i) = max_{dy} T(k + dy, i)       (i < PW, into tile NROW + k)
//   M(k, i) = max_{dj} V(k, i + dj)       (i < SW, in place: reads only positions >= i)
// SFPSWAP min/max is exact, so M equals the two ttnn max pools bit for bit.
// PMASK (SP_NMS_PMASK=1): the output is N(k, i) = P(k, i) if bits(P) == bits(M) else 0 (the NMS map in strip
// layout, the compare of nms_unfold_kp.cpp done on the SFPU), so the unfold reads one tensor instead of two.
#include <cstdint>
#include "api/compute/compute_kernel_api.h"
#include "api/compute/tile_move_copy.h"
#include "api/compute/binary_max_min.h"
#include "api/compute/eltwise_unary/eltwise_unary.h"

#ifdef TRISC_MATH
template <uint32_t NROW, uint32_t ROWS, uint32_t R, uint32_t PW, uint32_t SW>
inline void nms_window_max() {
    constexpr uint32_t T = 64;  // DST rows per tile
    for (uint32_t k = 0; k < ROWS; ++k) {
        const uint32_t out = (NROW + k) * T;
        for (uint32_t i = 0; i < PW; ++i) {
            TT_SFPLOAD(p_sfpu::LREG1, InstrModLoadStore::DEFAULT, ADDR_MOD_7, k * T + 2 * i);
            for (uint32_t dy = 1; dy <= 2 * R; ++dy) {
                TT_SFPLOAD(p_sfpu::LREG0, InstrModLoadStore::DEFAULT, ADDR_MOD_7, (k + dy) * T + 2 * i);
                TTI_SFPSWAP(0, p_sfpu::LREG1, p_sfpu::LREG0, sfpi::SFPSWAP_MOD1_VEC_MIN_MAX);
            }
            TT_SFPSTORE(p_sfpu::LREG1, InstrModLoadStore::DEFAULT, ADDR_MOD_7, out + 2 * i);
        }
    }
}
// PMASK: C(k, i) = P(R + k, PAD + i), the centre scores of output row k aligned with M (into tile k; tile R + k is
// read at step k and only tiles < k were written before, R >= 1)
template <uint32_t ROWS, uint32_t R, uint32_t SW>
inline void nms_center() {
    constexpr uint32_t T = 64;
    for (uint32_t k = 0; k < ROWS; ++k) {
        for (uint32_t i = 0; i < SW; ++i) {
            TT_SFPLOAD(p_sfpu::LREG0, InstrModLoadStore::DEFAULT, ADDR_MOD_7, (R + k) * T + 2 * (R + i));
            TT_SFPSTORE(p_sfpu::LREG0, InstrModLoadStore::DEFAULT, ADDR_MOD_7, k * T + 2 * i);
        }
    }
}
// PMASK: tile 0 = M (window max), tile SC = C: tile 0 = (bits(C) == bits(M)) ? M : 0 (= the NMS output N)
template <uint32_t SC, uint32_t SW>
inline void nms_mask() {
    for (uint32_t i = 0; i < SW; ++i) {
        sfpi::vInt m = sfpi::as<sfpi::vInt>(sfpi::vFloat(sfpi::dst_reg[i]));
        sfpi::vInt c = sfpi::as<sfpi::vInt>(sfpi::vFloat(sfpi::dst_reg[SC * 32 + i]));
        v_if(c != m) { sfpi::dst_reg[i] = 0.0f; }
        v_endif;
    }
}
#endif

void kernel_main() {
    constexpr uint32_t cb_in = get_compile_time_arg_val(0);
    constexpr uint32_t cb_out = get_compile_time_arg_val(1);
    constexpr uint32_t ROWS = get_compile_time_arg_val(2);
    constexpr uint32_t R = get_compile_time_arg_val(3);
    constexpr uint32_t PW = get_compile_time_arg_val(4);
    constexpr uint32_t SW = get_compile_time_arg_val(5);
    constexpr uint32_t cb_vw = get_compile_time_arg_val(6);  // V rows, 2 KB pages (pack)
    constexpr uint32_t cb_vr = get_compile_time_arg_val(7);  // same memory, 64 B pages (unpack at position offsets)
    constexpr uint32_t NROW = ROWS + 2 * R;
#ifdef PMASK
    constexpr uint32_t cb_c = get_compile_time_arg_val(8);  // centre rows C (PMASK)
    static_assert(2 * R + 2 <= 16 && R >= 1, "PMASK DST slots");
#endif
    static_assert(NROW + ROWS <= 16, "DST holds 16 bf16 tiles");
    static_assert(PW <= 32, "one strip row per tile");

    unary_op_init_common(cb_in, cb_out);
    // ---- vertical: V(k) = max of strip rows k .. k + 2R (SFPU, in DST)
    cb_wait_front(cb_in, NROW);
    tile_regs_acquire();
    copy_tile_to_dst_init_short(cb_in);
    for (uint32_t j = 0; j < NROW; ++j) {
        copy_tile(cb_in, j, j);
    }
    binary_max_tile_init();  // SFPU config; ADDR_MOD_7 = no auto increment
    MATH((_llk_math_eltwise_sfpu_start_(0)));
    MATH((nms_window_max<NROW, ROWS, R, PW, SW>()));
#ifdef PMASK
    MATH((nms_center<ROWS, R, SW>()));
#endif
    MATH((_llk_math_eltwise_sfpu_done_()));
    tile_regs_commit();
    cb_pop_front(cb_in, NROW);
    cb_reserve_back(cb_vw, ROWS);
#ifdef PMASK
    cb_reserve_back(cb_c, ROWS);
#endif
    tile_regs_wait();
    for (uint32_t k = 0; k < ROWS; ++k) {
        pack_tile(NROW + k, cb_vw, k);
    }
#ifdef PMASK
    for (uint32_t k = 0; k < ROWS; ++k) {
        pack_tile(k, cb_c, k);
    }
#endif
    tile_regs_release();
    cb_push_back(cb_vw, ROWS);
#ifdef PMASK
    cb_push_back(cb_c, ROWS);
    cb_wait_front(cb_c, ROWS);
#endif
    // ---- horizontal: M(k, i) = max_dj V(k, i + dj): V re-read at +64 B (one position) offsets
    cb_wait_front(cb_vw, ROWS);
#ifndef POOL_UNF
    cb_reserve_back(cb_out, ROWS);
#endif
    for (uint32_t k = 0; k < ROWS; ++k) {
        tile_regs_acquire();
        copy_tile_to_dst_init_short(cb_vr);
        for (uint32_t dj = 0; dj <= 2 * R; ++dj) {
            copy_tile(cb_vr, k * 32 + dj, dj);
        }
#ifdef PMASK
        copy_tile_to_dst_init_short(cb_c);
        copy_tile(cb_c, k, 2 * R + 1);
#endif
        binary_max_tile_init();
        for (uint32_t dj = 1; dj <= 2 * R; ++dj) {
            binary_max_tile(0, dj, 0);
        }
#ifdef PMASK
        MATH((_llk_math_eltwise_sfpu_start_(0)));
        MATH((nms_mask<2 * R + 1, SW>()));
        MATH((_llk_math_eltwise_sfpu_done_()));
#endif
        tile_regs_commit();
        tile_regs_wait();
#ifdef POOL_UNF
        // the unfold (on the data-movement RISCs) starts on its first rows while the rest are computed
        cb_reserve_back(cb_out, 1);
        pack_tile(0, cb_out);
        tile_regs_release();
        cb_push_back(cb_out, 1);
#else
        pack_tile(0, cb_out, k);
        tile_regs_release();
#endif
    }
#ifndef POOL_UNF
    cb_push_back(cb_out, ROWS);
#endif
    cb_pop_front(cb_vw, ROWS);
#ifdef PMASK
    cb_pop_front(cb_c, ROWS);
#endif
}