File size: 5,652 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
// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// SuperPoint descriptor sampling, pipelined variant (SP_SF_PIPE=1), writer: 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: hdr_addr, wtab_addr, nunits, NB bucket addresses, unit ids. CT args as sample_fused_writer.cpp.
#include <stdint.h>
#include "api/dataflow/dataflow_api.h"

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;
    static_assert(KPT == 8 && KT == 4, "SF_PIPE: 128 channels per unit");
    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
    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, first = ui == 0;
        const uint32_t* kp = reinterpret_cast<const uint32_t*>(kps);
        if (active && tr != cur_tr) {
            noc_async_read(hacc.get_noc_addr(0) + (16 + 128 * tr) * 4, kps, 512);
            noc_async_read_barrier();
            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;
        }
        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) & 3) : gi;
            cb_wait_front(cb_o, 1);
            if (active) {
                const uint32_t src = get_read_ptr(cb_o);
                const uint32_t r0 = tr * 32 + g * 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);
        }
    }
}