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();
}