superpoint-p150 / code /kernels /sp_nms /nms_pool_dm.cpp
changh95's picture
Optimized build (2026-10-03): trace 3.73 -> 0.50 ms
c699c4c verified
Raw History Blame Contribute Delete
4.15 kB
// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// SuperPoint NMS window max (nms_pool_compute.cpp), data movement. PROC 0: the NROW = ROWS + 2R strip
// rows y0 - R .. y0 + ROWS - 1 + R of P (PW * 64 B each; previous / next core's shard rows for the
// halo, zeros outside the image) into CB_IN pages (2 KB). PROC 1: the ROWS result pages -> M shard rows
// (SW * 64 B each).
// RT args: p_addr, m_addr, prev_x, prev_y, next_x, next_y, has_prev, has_next
// CT args: PROC, cb_in, cb_out, ROWS, R, PW, SW
#include <stdint.h>
#include "api/dataflow/dataflow_api.h"
void kernel_main() {
const uint32_t p_addr = get_arg_val<uint32_t>(0);
const uint32_t m_addr = get_arg_val<uint32_t>(1);
const uint32_t pvx = get_arg_val<uint32_t>(2);
const uint32_t pvy = get_arg_val<uint32_t>(3);
const uint32_t nxx = get_arg_val<uint32_t>(4);
const uint32_t nxy = get_arg_val<uint32_t>(5);
const uint32_t has_prev = get_arg_val<uint32_t>(6);
const uint32_t has_next = get_arg_val<uint32_t>(7);
constexpr uint32_t PROC = get_compile_time_arg_val(0);
constexpr uint32_t cb_in = get_compile_time_arg_val(1);
constexpr uint32_t cb_out = get_compile_time_arg_val(2);
constexpr uint32_t ROWS = get_compile_time_arg_val(3);
constexpr uint32_t R = get_compile_time_arg_val(4);
constexpr uint32_t PW = get_compile_time_arg_val(5);
constexpr uint32_t SW = get_compile_time_arg_val(6);
constexpr uint32_t NROW = ROWS + 2 * R, RB = PW * 64, OB = SW * 64, PAGE = 2048;
if constexpr (PROC == 0) {
cb_reserve_back(cb_in, NROW);
const uint32_t base = get_write_ptr(cb_in);
for (uint32_t j = 0; j < NROW; ++j) {
const int32_t row = (int32_t)j - (int32_t)R; // relative to this core's first row
const uint32_t dst = base + j * PAGE;
if ((row < 0 && !has_prev) || (row >= (int32_t)ROWS && !has_next)) {
#ifdef POOL_ZDMA
// rows outside the image: zeros by DMA from the hardware zero page (the scalar store loop made the
// first / last core ~3 us late)
const uint64_t zsrc = get_noc_addr(my_x[noc_index], my_y[noc_index], MEM_ZEROS_BASE);
for (uint32_t off = 0; off < RB; off += MEM_ZEROS_SIZE) {
noc_async_read(zsrc, dst + off, RB - off < MEM_ZEROS_SIZE ? RB - off : MEM_ZEROS_SIZE);
}
#else
volatile tt_l1_ptr uint32_t* z = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(dst);
for (uint32_t w = 0; w < RB / 4; ++w) {
z[w] = 0;
}
(void)z[RB / 4 - 1];
#endif
} else if (row < 0) {
noc_async_read(get_noc_addr(pvx, pvy, p_addr + (uint32_t)(row + (int32_t)ROWS) * RB), dst, RB);
} else if (row >= (int32_t)ROWS) {
noc_async_read(get_noc_addr(nxx, nxy, p_addr + (uint32_t)(row - (int32_t)ROWS) * RB), dst, RB);
} else {
noc_async_read(get_noc_addr(my_x[noc_index], my_y[noc_index], p_addr + (uint32_t)row * RB), dst, RB);
}
}
noc_async_read_barrier();
cb_push_back(cb_in, NROW);
} else {
#ifndef POOL_UNF
cb_wait_front(cb_out, ROWS);
const uint32_t src = get_read_ptr(cb_out);
for (uint32_t k = 0; k < ROWS; ++k) {
noc_async_write(src + k * PAGE, get_noc_addr(my_x[noc_index], my_y[noc_index], m_addr + k * OB), OB);
}
noc_async_write_barrier();
cb_pop_front(cb_out, ROWS);
#endif
}
#ifdef POOL_UNF
// SP_NMS_POOL_UNF (PMASK): the NMS unfold + keypoint candidates (nms_unfold_fn.inc, prepended) of this RISC's
// rows straight from the ROWS result pages in CB_OUT (both RISCs read them; nothing pops, M is not written).
// Unfold CT args start at 7, its RT-arg block at 8.
#ifndef UNF_ROW_WAIT
cb_wait_front(cb_out, PROC == 0 ? UNF_SPLIT : ROWS); // PROC 0 unfolds rows [0, UNF_SPLIT), PROC 1 the rest
#endif
unfold_rows<7, PROC, PAGE / 4>(reinterpret_cast<const uint32_t*>(get_read_ptr(cb_out)), 8);
#endif
}