superpoint-p150 / code /kernels /sp_nms /nms_unfold.cpp
changh95's picture
Optimized build (2026-10-03): trace 3.73 -> 0.50 ms
c699c4c verified
Raw History Blame Contribute Delete
2.58 kB
// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// SuperPoint NMS, step 3 (data movement only, runs via ttnn.generic_op).
// M = 9x9 window max in strip layout [H, SW, 32] (this core's L1 shard), P = padded strip scores
// [H, SW + 2*PAD, 32] (this core's L1 shard). Writes the natural dense NMS map row by row:
// out[y, x] = P[y, xw + PAD, l] if it equals M[y, xw, l] else 0, x = l*SW + xw
// (bit-identical to where(s == maxpool9x9(s), s, 0) for s >= 0). out: [H, W] ROW_MAJOR interleaved.
#include <stdint.h>
#include "api/dataflow/dataflow_api.h"
void kernel_main() {
const uint32_t m_addr = get_arg_val<uint32_t>(0);
const uint32_t p_addr = get_arg_val<uint32_t>(1);
const uint32_t o_addr = get_arg_val<uint32_t>(2);
const uint32_t y0 = get_arg_val<uint32_t>(3);
const uint32_t yl0 = get_arg_val<uint32_t>(4);
const uint32_t yl1 = get_arg_val<uint32_t>(5);
constexpr uint32_t cb_scratch = get_compile_time_arg_val(0);
constexpr uint32_t SW = get_compile_time_arg_val(1);
constexpr uint32_t PAD = get_compile_time_arg_val(2);
constexpr uint32_t LANES = 32;
constexpr uint32_t PW = SW + 2 * PAD;
constexpr uint32_t W = SW * LANES;
constexpr auto o_args = TensorAccessorArgs<3>();
const auto oacc = TensorAccessor(o_args, o_addr, W * 2);
const uint32_t scratch = get_write_ptr(cb_scratch);
const uint32_t* mm = reinterpret_cast<const uint32_t*>(m_addr);
const uint32_t* pp = reinterpret_cast<const uint32_t*>(p_addr);
uint32_t buf = 0;
for (uint32_t yl = yl0; yl < yl1; ++yl) {
const uint32_t row_l1 = scratch + buf * W * 2;
uint16_t* orow = reinterpret_cast<uint16_t*>(row_l1);
const uint32_t* mrow = mm + yl * SW * (LANES / 2);
const uint32_t* prow = pp + (yl * PW + PAD) * (LANES / 2);
for (uint32_t xw = 0; xw < SW; ++xw) {
const uint32_t* mw = mrow + xw * (LANES / 2);
const uint32_t* pw = prow + xw * (LANES / 2);
uint16_t* o = orow + xw;
#pragma GCC unroll 16
for (uint32_t l = 0; l < LANES; l += 2) {
const uint32_t d = pw[l >> 1];
const uint32_t x = d ^ mw[l >> 1];
o[l * SW] = (x & 0xFFFF) ? 0 : (uint16_t)d;
o[(l + 1) * SW] = (x >> 16) ? 0 : (uint16_t)(d >> 16);
}
}
noc_async_write(row_l1, oacc.get_noc_addr(y0 + yl), W * 2);
buf ^= 1;
if (buf == 0) {
noc_async_write_barrier();
}
}
noc_async_write_barrier();
}