superpoint-p150 / code /kernels /sp_nms /nms_fold.cpp
changh95's picture
Optimized build (2026-10-03): trace 3.73 -> 0.50 ms
c699c4c verified
Raw History Blame Contribute Delete
5.92 kB
// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// SuperPoint NMS, step 1 (data movement only, runs via ttnn.generic_op).
// Scatter the softmaxed cell scores S [B*h*w, 65(96)] (TILE, interleaved; cell row r = cy*WC + cx,
// column c = i*8 + j) into the zero-padded "strip" layout
// P[y, xw, l] = Sdense[y, l*SW - PAD + xw] (0 outside the image), P: [H, PW = SW + 2*PAD, 32]
// stored as this core's L1 height shard (ROWS image rows per core). Lane l holds the vertical strip
// of SW consecutive columns starting at x = l*SW, padded by PAD on both sides, so a [9,1] + [1,9]
// max-pool over (H, W) of P gives the exact 9x9 window max of every pixel.
#include <stdint.h>
#include "api/dataflow/dataflow_api.h"
#ifdef PROFZ
#include "tools/profiler/kernel_profiler.hpp"
#define ZONE(n) DeviceZoneScopedN(n)
#else
#define ZONE(n)
#endif
void kernel_main() {
const uint32_t s_addr = get_arg_val<uint32_t>(0);
const uint32_t p_addr = get_arg_val<uint32_t>(1);
const uint32_t y0 = get_arg_val<uint32_t>(2); // first image row of this core's shard
const uint32_t yl0 = get_arg_val<uint32_t>(3); // local rows [yl0, yl1) handled by this RISC
const uint32_t yl1 = get_arg_val<uint32_t>(4);
constexpr uint32_t cb_scratch = get_compile_time_arg_val(0);
constexpr uint32_t WC = get_compile_time_arg_val(1); // cells per image row (w/8)
constexpr uint32_t TCOLS = get_compile_time_arg_val(2); // tile columns of S
constexpr uint32_t SW = get_compile_time_arg_val(3); // strip width (W / 32)
constexpr uint32_t PAD = get_compile_time_arg_val(4); // nms radius
constexpr uint32_t LANES = 32;
constexpr uint32_t PW = SW + 2 * PAD;
constexpr uint32_t W = WC * 8;
constexpr auto s_args = TensorAccessorArgs<5>();
const auto s = TensorAccessor(s_args, s_addr, 2048);
// Scratch: up to (WC+31)/32+1 S tiles, followed by one zero-padded dense row R[PAD + W + PAD].
const uint32_t scratch = get_write_ptr(cb_scratch);
constexpr uint32_t NT = (WC + 31) / 32 + 1;
const uint16_t* src = reinterpret_cast<const uint16_t*>(scratch);
// odd PAD: shift the row buffer by 2 B so that rbuf + PAD stays 4-byte aligned
#ifdef FOLD_LOCAL
// the dense row in RISC-local data memory (single-cycle loads) instead of L1
static uint32_t rloc[(PAD + W + PAD + 2) / 2 + 1];
uint16_t* rbuf = reinterpret_cast<uint16_t*>(rloc) + (PAD & 1);
#else
uint16_t* rbuf = reinterpret_cast<uint16_t*>(scratch + NT * 2048 + (PAD & 1) * 2);
#endif
for (uint32_t k = 0; k < PAD; ++k) {
rbuf[k] = 0;
rbuf[PAD + W + k] = 0;
}
uint32_t* dst = reinterpret_cast<uint32_t*>(p_addr);
uint32_t loaded = 0xFFFFFFFF;
uint32_t roff = 0;
for (uint32_t yl = yl0; yl < yl1; ++yl) {
const uint32_t y = y0 + yl;
const uint32_t cy = y >> 3;
const uint32_t i = y & 7;
const uint32_t tc = (i * 8) >> 5;
const uint32_t key = cy * TCOLS + tc;
const uint32_t r0 = cy * WC;
const uint32_t tr0 = r0 >> 5;
if (key != loaded) {
ZONE("FOLD_RD");
const uint32_t tr1 = (r0 + WC - 1) >> 5;
for (uint32_t tr = tr0; tr <= tr1; ++tr) {
noc_async_read(s.get_noc_addr(tr * TCOLS + tc), scratch + (tr - tr0) * 2048, 2048);
}
noc_async_read_barrier();
loaded = key;
roff = r0 - tr0 * 32;
}
// 1) gather the dense row: cell cx contributes 8 contiguous values (16 B) of one face row.
{
ZONE("FOLD_GATHER");
const uint32_t cl = (i * 8) & 31;
const uint32_t cbase = ((cl >> 4) * 256) + (cl & 15);
uint32_t* r32 = reinterpret_cast<uint32_t*>(rbuf + PAD); // 4 B aligned (see rbuf)
#ifdef FOLD_SEG
// runs of consecutive cells inside one 16-row face half: the source steps by one face row (8 words)
for (uint32_t cx = 0; cx < WC;) {
const uint32_t rr = roff + cx;
const uint32_t rin = rr & 31;
uint32_t n = 16 - (rin & 15);
if (n > WC - cx) {
n = WC - cx;
}
const uint32_t* c32 = reinterpret_cast<const uint32_t*>(src + (rr >> 5) * 1024 + ((rin >> 4) << 9) + (rin & 15) * 16 + cbase);
uint32_t* o = r32 + cx * 4;
cx += n;
#pragma GCC unroll 2
for (; n >= 2; n -= 2) {
const uint32_t a0 = c32[0], a1 = c32[1], a2 = c32[2], a3 = c32[3];
const uint32_t b0 = c32[8], b1 = c32[9], b2 = c32[10], b3 = c32[11];
o[0] = a0; o[1] = a1; o[2] = a2; o[3] = a3;
o[4] = b0; o[5] = b1; o[6] = b2; o[7] = b3;
c32 += 16;
o += 8;
}
if (n) {
o[0] = c32[0]; o[1] = c32[1]; o[2] = c32[2]; o[3] = c32[3];
}
}
}
#else
for (uint32_t cx = 0; cx < WC; ++cx) {
const uint32_t rr = roff + cx;
const uint32_t rin = rr & 31;
const uint32_t off = (rr >> 5) * 1024 + ((rin >> 4) << 9) + (rin & 15) * 16 + cbase;
const uint32_t* c32 = reinterpret_cast<const uint32_t*>(src + off);
r32[cx * 4 + 0] = c32[0];
r32[cx * 4 + 1] = c32[1];
r32[cx * 4 + 2] = c32[2];
r32[cx * 4 + 3] = c32[3];
}
}
#endif
// 2) strip scatter: P[yl, xw, l] = R[l*SW + xw]; two lanes per 32-bit store.
uint32_t* drow = dst + yl * PW * (LANES / 2);
ZONE("FOLD_SCAT");
for (uint32_t xw = 0; xw < PW; ++xw) {
const uint16_t* b = rbuf + xw;
uint32_t* d = drow + xw * (LANES / 2);
#pragma GCC unroll 16
for (uint32_t l = 0; l < LANES; l += 2) {
d[l >> 1] = (uint32_t)b[l * SW] | ((uint32_t)b[(l + 1) * SW] << 16);
}
}
}
}