File size: 6,824 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
// 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: no compaction op; the
// list is gathered from the NMS records, sample_kplist.hpp.inc, prepended by nms_kernels.py), reader. Otherwise as
// sample_pipe_reader.cpp (SP_SF_PIPE=1, CPU = 128 channels per unit, KPT = 8 keypoints per page): 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, cnt_addr, nunits, unit ids, wtab_addr. Common RT args: rec_addr, NCORE (noc_x << 16 | noc_y) of the
// NMS cores. CT args: cb_g, cb_scratch, C, KV, CPU, accessors (d, counts, wtab). Defines SF_CBW0, SF_NSLOT, SF_CAP.
#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 d_addr = get_arg_val<uint32_t>(0);
    const uint32_t cnt_addr = get_arg_val<uint32_t>(1);
    const uint32_t rec_addr = get_common_arg_val<uint32_t>(0);
    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;
#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 uint32_t cb_w0 = SF_CBW0;
    constexpr auto d_args = TensorAccessorArgs<5>();
    constexpr auto n_args = TensorAccessorArgs<d_args.next_compile_time_args_offset()>();
    constexpr auto wt_args = TensorAccessorArgs<n_args.next_compile_time_args_offset()>();
    const auto dacc = TensorAccessor(d_args, d_addr, C * 2);
    const auto nacc = TensorAccessor(n_args, cnt_addr, SF_CNT_PAGE);
    const auto wtacc = TensorAccessor(wt_args, wtab_addr, SF_W * 16);

    const uint32_t cnt_l1 = get_write_ptr(cb_scratch);  // slot counts (NSLOT x 16 B)
    const uint32_t kps = cnt_l1 + (SF_NSLOT * 16 + SF_CNT_PAGE - 1) / SF_CNT_PAGE * SF_CNT_PAGE;        // 32 keypoint entries (512 B)
    const uint32_t wblk = kps + 512;                    // 8 x 64 B tap-weight blocks (group 0)

    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_barrier();
    const uint32_t n = kplist_scan<SF_NSLOT, SF_CAP, KV>(cnt_l1, pre).kept;
    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 (active) {
            if (tr != cur_tr) {
                kplist_gather<SF_NSLOT, SF_CAP, 1, KPU>(pre, rec_addr, tr, n, kps);
                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) % KT) : 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);
        }
    }
}