File size: 6,002 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 | // SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// SuperPoint descriptor sampling, pipelined variant (SP_SF_PIPE=1, CPU = 128 channels per unit, KPT = 8 keypoints per
// page), reader. Same gather as sample_fused_reader.cpp, but pushed per keypoint group g (8 keypoints: 4 tap pages of
// 2 KB, page t holds G_t[k, c] for k = 8 g .. 8 g + 7, row-major) so the compute kernel starts on the first group.
// Group order: the core's first unit goes 1, 2, 3, 0 (the writer fills the tap weights of groups 1..3, this RISC those
// of group 0 after its gathers, into CB_W0), every further unit 0, 1, 2, 3 (writer fills all). Compact weight pages:
// block (t, i) = 64 words (4 DST rows) at page t / 2, word ((t % 2) * 8 + i) * 64, every word w(i, t); the compute
// kernel replicates it over the 128 channels. Units at or beyond n push their pages unfilled. Pure copies -> exact.
// RT args: d_addr, hdr_addr, nunits, unit ids, wtab_addr. CT args: cb_g, cb_scratch, C, KV, CPU, accessors (d, hdr,
// wtab). Define SF_CBW0: CB of the group-0 weight pages.
#include <stdint.h>
#include "api/dataflow/dataflow_api.h"
void kernel_main() {
const uint32_t d_addr = get_arg_val<uint32_t>(0);
const uint32_t hdr_addr = get_arg_val<uint32_t>(1);
const uint32_t nunits = get_arg_val<uint32_t>(2);
const uint32_t wtab_addr = get_arg_val<uint32_t>(3 + nunits);
constexpr uint32_t cb_g = get_compile_time_arg_val(0);
constexpr uint32_t cb_scratch = get_compile_time_arg_val(1);
constexpr uint32_t C = get_compile_time_arg_val(2);
constexpr uint32_t KV = get_compile_time_arg_val(3);
constexpr uint32_t CPU = get_compile_time_arg_val(4);
constexpr uint32_t NQ = C / CPU;
constexpr uint32_t KPT = 1024 / CPU;
constexpr uint32_t KT = 32 / KPT;
static_assert(KPT == 8 && KT == 4, "SF_PIPE: 128 channels per unit");
constexpr uint32_t cb_w0 = SF_CBW0;
constexpr auto d_args = TensorAccessorArgs<5>();
constexpr auto hdr_args = TensorAccessorArgs<d_args.next_compile_time_args_offset()>();
constexpr auto wt_args = TensorAccessorArgs<hdr_args.next_compile_time_args_offset()>();
const auto dacc = TensorAccessor(d_args, d_addr, C * 2);
const auto hacc = TensorAccessor(hdr_args, hdr_addr, (16 + 4 * KV) * 4);
const auto wtacc = TensorAccessor(wt_args, wtab_addr, SF_W * 16);
const uint32_t hdr0 = get_write_ptr(cb_scratch); // HDR[0..15] (64 B)
const uint32_t kps = hdr0 + 64; // 32 HDR keypoint entries (512 B)
const uint32_t wblk = kps + 512; // 8 x 64 B tap-weight blocks (group 0)
noc_async_read(hacc.get_noc_addr(0), hdr0, 64);
noc_async_read_barrier();
const uint32_t n = reinterpret_cast<volatile uint32_t*>(hdr0)[2];
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, q = u % NQ;
const bool active = tr * 32 < n, first = ui == 0;
const uint32_t* kp = reinterpret_cast<const uint32_t*>(kps);
if (active) {
if (tr != cur_tr) {
noc_async_read(hacc.get_noc_addr(0) + (16 + 128 * tr) * 4, kps, 512);
noc_async_read_barrier();
cur_tr = tr;
}
if (first) { // tap weights of group 0 (waited for by the first group's barrier)
for (uint32_t r = 0; r < KPT; ++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);
}
}
}
for (uint32_t gi = 0; gi < KT; ++gi) {
const uint32_t g = first ? ((gi + 1) & 3) : gi;
cb_reserve_back(cb_g, 4);
if (active) {
const uint32_t g0 = get_write_ptr(cb_g);
for (uint32_t t = 0; t < 4; ++t) {
const uint32_t xs = (t & 1) ? 0 : 16, ys = (t >> 1) ? 0 : 16;
for (uint32_t i = 0; i < KPT; ++i) {
const uint32_t r = g * KPT + i;
const uint32_t cell = ((kp[4 * r + 2] >> ys) & 0xFFFF) + ((kp[4 * r + 3] >> xs) & 0xFFFF);
noc_async_read(dacc.get_noc_addr(cell) + q * CPU * 2, g0 + t * 2048 + i * CPU * 2, CPU * 2);
}
}
noc_async_read_barrier();
}
cb_push_back(cb_g, 4);
}
if (first) {
cb_reserve_back(cb_w0, 2);
if (active) {
uint32_t* w0 = reinterpret_cast<uint32_t*>(get_write_ptr(cb_w0));
for (uint32_t i = 0; i < KPT; ++i) {
const uint32_t x = kp[4 * i] & 0xFFFF;
const uint32_t* wv = reinterpret_cast<const uint32_t*>(wblk + i * 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_w0, 2);
}
}
}
|