File size: 9,157 Bytes
c699c4c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 | // 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();
}
|