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