File size: 6,007 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 | // SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// SuperPoint descriptor sampling, fused op, writer (RISCV_1).
// (1) Builds the fp32 weight pages of a unit: page (p, t) element (i, c) = WTAB[y_k, 4*x_k + t] for
// keypoint k = 32*tr + p*KPT + i (the tap weight repeated over the CPU channels, row-major pseudo
// tile, same order as the reader's G pages).
// (2) Writes the fp32 result pages (KPT rows x CPU channels, row-major) into ONE bucket tensor
// [BSTEP * (b + 1), C], b = ceil(n / BSTEP) - 1 (n = HDR[2]); tile rows >= n are dropped.
// SF_HR (single-D2H mode): bucket b = max(ceil(n / BSTEP) - 1, HDR[3]) (kp_compact2 put the header
// copy there), rows start at row SF_HR of the bucket (the header rows come first).
// Runtime args: hdr_addr, wtab_addr, nunits, NB bucket addresses, then the nunits unit ids.
#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 hdr_addr = get_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;
constexpr uint32_t KT = 32 / KPT;
constexpr auto h_args = TensorAccessorArgs<8>();
constexpr auto wt_args = TensorAccessorArgs<h_args.next_compile_time_args_offset()>();
constexpr auto o_args = TensorAccessorArgs<wt_args.next_compile_time_args_offset()>();
const auto hacc = TensorAccessor(h_args, hdr_addr, (16 + 4 * KV) * 4);
const auto wtacc = TensorAccessor(wt_args, wtab_addr, W * 16);
const uint32_t hb = get_write_ptr(cb_scratch);
const uint32_t kps = hb + 64;
const uint32_t wblk = kps + 512; // 32 x 64 B weight blocks
{
ZONE("SF_W_HDR");
noc_async_read(hacc.get_noc_addr(0), hb, 64);
noc_async_read_barrier();
}
uint32_t n = reinterpret_cast<volatile uint32_t*>(hb)[2];
if (n > KV) {
n = KV;
}
uint32_t nb = n == 0 ? 1 : (n + BSTEP - 1) / BSTEP;
#ifdef SF_HR
{
const uint32_t spec = reinterpret_cast<volatile uint32_t*>(hb)[3] + 1;
if (spec > nb) {
nb = spec > NB ? NB : spec;
}
}
constexpr uint32_t HR = SF_HR;
#else
constexpr uint32_t HR = 0;
#endif
const uint32_t o_addr = get_arg_val<uint32_t>(3 + nb - 1);
const auto oacc = TensorAccessor(o_args, o_addr, C * 4);
uint32_t cur_tr = 0xFFFFFFFF;
for (uint32_t ui = 0; ui < nunits; ++ui) {
const uint32_t u = get_arg_val<uint32_t>(3 + NB + ui);
const uint32_t tr = u / NQ, q = u % NQ;
const bool active = tr * 32 < n;
cb_reserve_back(cb_w, 4 * KT);
if (active) {
if (tr != cur_tr) {
ZONE("SF_W_KPS");
noc_async_read(hacc.get_noc_addr(0) + (16 + 128 * tr) * 4, kps, 512);
noc_async_read_barrier();
const uint32_t* kp = reinterpret_cast<const uint32_t*>(kps);
for (uint32_t r = 0; r < 32; ++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;
}
const uint32_t* kp = reinterpret_cast<const uint32_t*>(kps);
uint32_t* w0 = reinterpret_cast<uint32_t*>(get_write_ptr(cb_w));
#ifdef SF_SPLIT
const uint32_t r_first = ui == 0 ? SF_SPLIT : 0; // the reader fills keypoints 0..SF_SPLIT-1 of the first unit (CB_W slot 0)
#else
const uint32_t r_first = 0;
#endif
#ifndef SF_NO_WFILL
for (uint32_t r = r_first; r < 32; ++r) {
const uint32_t x = kp[4 * r] & 0xFFFF;
const uint32_t* wv = reinterpret_cast<const uint32_t*>(wblk + r * 64 + ((x * 16) & 63));
const uint32_t p = r / KPT, i = r % KPT;
for (uint32_t t = 0; t < 4; ++t) {
const uint32_t v = wv[t];
uint32_t* d = w0 + (p * 4 + t) * 1024 + i * CPU;
for (uint32_t c = 0; c < CPU; 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
#ifdef SF_SPLIT
ZONE("SF_W_SEM");
if (ui == 0) {
volatile tt_l1_ptr uint32_t* fill_sem = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(get_semaphore(0));
noc_semaphore_wait(fill_sem, 1);
noc_semaphore_set(fill_sem, 0);
}
#endif
}
cb_push_back(cb_w, 4 * KT);
ZONE("SF_W_OUT");
for (uint32_t p = 0; p < KT; ++p) {
cb_wait_front(cb_o, 1);
if (active) {
const uint32_t src = get_read_ptr(cb_o);
const uint32_t r0 = tr * 32 + p * KPT;
for (uint32_t i = 0; i < KPT; ++i) {
noc_async_write(src + i * CPU * 4, oacc.get_noc_addr(HR + r0 + i) + q * CPU * 4, CPU * 4);
}
noc_async_write_barrier();
}
cb_pop_front(cb_o, 1);
}
}
}
|