superpoint-p150 / code /kernels /sp_nms /nms_pool_compute.cpp
changh95's picture
Optimized build (2026-10-03): trace 3.73 -> 0.50 ms
c699c4c verified
Raw History Blame Contribute Delete
6.41 kB
// 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
}