superpoint-p150 / code /kernels /sp_nms /sample_direct_writer.cpp
changh95's picture
Optimized build (2026-10-03): trace 3.73 -> 0.50 ms
c699c4c verified
Raw History Blame Contribute Delete
9.16 kB
// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// SuperPoint descriptor sampling, pipelined variant with its own keypoint list (SP_SF_DIRECT=1, see
// sample_direct_reader.cpp), writer. Also writes the keypoint header that kp_compact3.cpp wrote into the bucket: the
// unit of tile row tr (channel chunk 0) its 32 entries (zeros beyond the kept count), unit 0 the 16 header words
// (total, overflow, kept, bucket b), also into the speculative bucket when b differs. Otherwise as
// sample_pipe_writer.cpp (SP_SF_PIPE=1): compact tap-weight pages of keypoint groups
// 1..3 of the core's first unit (0..3 of further units), pushed per group (2 pages, layout see sample_pipe_reader.cpp),
// then the fp32 result pages in the compute order (first unit 1, 2, 3, 0) into ONE bucket tensor as
// sample_fused_writer.cpp (bucket b = max(ceil(n / BSTEP) - 1, HDR[3]) in single-D2H mode, rows after SF_HR header rows).
// RT args: cnt_addr, wtab_addr, nunits, unit ids. Common RT args: rec_addr, NMS core coordinates, then from SF_COMB:
// NB bucket addresses, prm_addr (SF_TAIL: NB tail addresses, NB split rows): sent once, not per core. CT args: cb_w, cb_o, cb_scratch, C, KV, BSTEP, CPU, W, accessors (counts, wtab, bucket, prm).
#include <stdint.h>
#include "api/dataflow/dataflow_api.h"
static uint16_t pre[SF_NSLOT + 1]; // kept-prefix of the slots (RISC-local)
void kernel_main() {
const uint32_t cnt_addr = get_arg_val<uint32_t>(0);
const uint32_t rec_addr = get_common_arg_val<uint32_t>(0);
const uint32_t wtab_addr = get_arg_val<uint32_t>(1);
const uint32_t nunits = get_arg_val<uint32_t>(2);
constexpr uint32_t cb_w = get_compile_time_arg_val(0);
constexpr uint32_t cb_o = get_compile_time_arg_val(1);
constexpr uint32_t cb_scratch = get_compile_time_arg_val(2);
constexpr uint32_t C = get_compile_time_arg_val(3);
constexpr uint32_t KV = get_compile_time_arg_val(4);
constexpr uint32_t BSTEP = get_compile_time_arg_val(5);
constexpr uint32_t CPU = get_compile_time_arg_val(6);
constexpr uint32_t W = get_compile_time_arg_val(7);
constexpr uint32_t NB = KV / BSTEP;
constexpr uint32_t NQ = C / CPU;
constexpr uint32_t KPT = 1024 / CPU;
#ifndef SF_KT
#define SF_KT 4
#endif
constexpr uint32_t KT = SF_KT; // keypoint groups per unit (SP_SF_KPU = 16: 2, a unit is half a tile row)
constexpr uint32_t KPU = KT * KPT; // keypoints per unit
static_assert(KPT == 8 && (KT == 4 || KT == 2), "SF_PIPE: 128 channels per unit");
constexpr auto n_args = TensorAccessorArgs<8>();
constexpr auto wt_args = TensorAccessorArgs<n_args.next_compile_time_args_offset()>();
constexpr auto o_args = TensorAccessorArgs<wt_args.next_compile_time_args_offset()>();
constexpr auto q_args = TensorAccessorArgs<o_args.next_compile_time_args_offset()>();
const auto nacc = TensorAccessor(n_args, cnt_addr, SF_CNT_PAGE);
const auto wtacc = TensorAccessor(wt_args, wtab_addr, W * 16);
const auto qacc = TensorAccessor(q_args, get_common_arg_val<uint32_t>(SF_COMB + NB), 64);
const uint32_t cnt_l1 = get_write_ptr(cb_scratch); // slot counts (NSLOT x 16 B)
const uint32_t hb = cnt_l1 + (SF_NSLOT * 16 + SF_CNT_PAGE - 1) / SF_CNT_PAGE * SF_CNT_PAGE; // 16 header words (64 B) + parameter page (64 B)
const uint32_t kps = hb + 128; // 32 keypoint entries (512 B)
const uint32_t wblk = kps + 512; // 32 x 64 B weight blocks
for (uint32_t pg = 0; pg * SF_CNT_PAGE < SF_NSLOT * 16; ++pg) { // L1-interleaved count pages
noc_async_read(nacc.get_noc_addr(pg), cnt_l1 + pg * SF_CNT_PAGE, SF_CNT_PAGE);
}
noc_async_read(qacc.get_noc_addr(0), hb + 64, 64);
noc_async_read_barrier();
const KpListInfo info = kplist_scan<SF_NSLOT, SF_CAP, KV>(cnt_l1, pre);
const uint32_t n = info.kept;
static_assert(KV % BSTEP == 0, "buckets");
// bucket: b = max(ceil(n / BSTEP) - 1, spec), as kp_compact3.cpp
const uint32_t spec = reinterpret_cast<volatile uint32_t*>(hb + 64)[2];
uint32_t b = n == 0 ? 0 : (n + BSTEP - 1) / BSTEP - 1;
if (spec > b) {
b = spec;
}
if (b > NB - 1) {
b = NB - 1;
}
constexpr uint32_t HR = SF_HR;
constexpr uint32_t ROWB = C * 4;
const uint32_t o_addr = get_common_arg_val<uint32_t>(SF_COMB + b);
const auto oacc = TensorAccessor(o_args, o_addr, ROWB);
#ifdef SF_TAIL
// SP_KPC_SPLIT: descriptor rows >= S of bucket b go to its tail tensor (row r -> tail page r - S)
const uint32_t targ = SF_COMB + NB + 1;
const auto tacc = TensorAccessor(o_args, get_common_arg_val<uint32_t>(targ + b), ROWB);
const uint32_t S = get_common_arg_val<uint32_t>(targ + NB + b);
#endif
// header words (unit 0 of core 0 writes them; also to the speculative bucket when b != spec)
auto write_bytes = [&](const auto& acc, uint32_t src, uint32_t off, uint32_t bytes) {
while (bytes) {
const uint32_t p = off / ROWB, o = off % ROWB;
const uint32_t sz = ROWB - o < bytes ? ROWB - o : bytes;
noc_async_write(src, acc.get_noc_addr(p) + o, sz);
src += sz;
off += sz;
bytes -= sz;
}
};
uint32_t cur_tr = 0xFFFFFFFF;
for (uint32_t ui = 0; ui < nunits; ++ui) {
const uint32_t u = get_arg_val<uint32_t>(3 + ui);
const uint32_t tr = (u / NQ) * KPU, q = u % NQ; // tr: first keypoint row of the unit
const bool active = tr < n, first = ui == 0;
const uint32_t* kp = reinterpret_cast<const uint32_t*>(kps);
if (ui == 0 && u == 0) {
volatile uint32_t* h = reinterpret_cast<volatile uint32_t*>(hb);
h[0] = info.total;
h[1] = info.overflow;
h[2] = n;
h[3] = b;
for (uint32_t w = 4; w < 16; ++w) {
h[w] = 0;
}
write_bytes(oacc, hb, 0, 64);
if (b != spec && spec < NB) {
const auto sacc = TensorAccessor(o_args, get_common_arg_val<uint32_t>(SF_COMB + spec), ROWB);
write_bytes(sacc, hb, 0, 64);
}
}
if (active && tr != cur_tr) {
kplist_gather<SF_NSLOT, SF_CAP, 1, KPU>(pre, rec_addr, tr, n, kps);
noc_async_read_barrier();
if (q == 0) {
write_bytes(oacc, kps, 64 + 16 * tr, 16 * KPU); // this unit's header entries
}
for (uint32_t r = 0; r < KPU; ++r) {
const uint32_t yx = kp[4 * r];
const uint32_t y = yx >> 16, x = yx & 0xFFFF;
noc_async_read(wtacc.get_noc_addr(y) + ((x * 16) & ~63u), wblk + r * 64, 64);
}
noc_async_read_barrier();
cur_tr = tr;
}
for (uint32_t g = first ? 1 : 0; g < KT; ++g) {
cb_reserve_back(cb_w, 2);
if (active) {
uint32_t* w0 = reinterpret_cast<uint32_t*>(get_write_ptr(cb_w));
for (uint32_t i = 0; i < KPT; ++i) {
const uint32_t r = g * KPT + i;
const uint32_t x = kp[4 * r] & 0xFFFF;
const uint32_t* wv = reinterpret_cast<const uint32_t*>(wblk + r * 64 + ((x * 16) & 63));
for (uint32_t t = 0; t < 4; ++t) {
const uint32_t v = wv[t];
uint32_t* d = w0 + (t >> 1) * 1024 + ((t & 1) * KPT + i) * 64;
#ifdef SF_WC16
// SF_WC16: only DST row 4 b of the block (16 words); the compute kernel broadcasts it (SFPTRANSP)
for (uint32_t c = 0; c < 16; c += 8) {
d[c] = v; d[c + 1] = v; d[c + 2] = v; d[c + 3] = v;
d[c + 4] = v; d[c + 5] = v; d[c + 6] = v; d[c + 7] = v;
}
#else
for (uint32_t c = 0; c < 64; c += 8) {
d[c] = v; d[c + 1] = v; d[c + 2] = v; d[c + 3] = v;
d[c + 4] = v; d[c + 5] = v; d[c + 6] = v; d[c + 7] = v;
}
#endif
}
}
}
cb_push_back(cb_w, 2);
}
for (uint32_t gi = 0; gi < KT; ++gi) {
const uint32_t g = first ? ((gi + 1) % KT) : gi;
cb_wait_front(cb_o, 1);
if (active) {
const uint32_t src = get_read_ptr(cb_o);
const uint32_t r0 = tr + g * KPT;
for (uint32_t i = 0; i < KPT; ++i) {
#ifdef SF_TAIL
const uint64_t dst = r0 + i < S ? oacc.get_noc_addr(HR + r0 + i) : tacc.get_noc_addr(r0 + i - S);
noc_async_write(src + i * CPU * 4, dst + q * CPU * 4, CPU * 4);
#else
noc_async_write(src + i * CPU * 4, oacc.get_noc_addr(HR + r0 + i) + q * CPU * 4, CPU * 4);
#endif
}
noc_async_write_barrier();
}
cb_pop_front(cb_o, 1);
}
}
noc_async_write_barrier();
}