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