File size: 4,812 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
// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// SuperPoint descriptor sampling, pipelined variant (SP_SF_PIPE=1), compute: per keypoint group (order: the core's first
// unit 1, 2, 3, 0 with group 0's weights from CB_W0, further units 0..3)
//   OUT = ((G0*W0 + G1*W1) + G2*W2) + G3*W3     (fp32 in DEST, SFPU mul/add)
// exactly the operations and order of sample_fused_compute.cpp; W_t is built in DST from the compact weight page by
// SFPU register copies (sf_replicate).
#include <cstdint>
#include "api/compute/compute_kernel_api.h"
#include "api/compute/tile_move_copy.h"
#include "api/compute/eltwise_binary_sfpu.h"
#include "api/compute/eltwise_unary/eltwise_unary.h"

#ifdef TRISC_MATH
// DST slot WC holds compact page t / 2 (block (t % 2) * 8 + i = 4 DST rows = SFPU rows 2 b, 2 b + 1, every element
// w(i, t)); write w(i, t) into the 4 SFPU rows (8 DST rows = 128 elements = keypoint i's channels) of keypoint i in slot WS.
template <uint32_t WC, uint32_t WS>
inline void sf_replicate(uint32_t t) {
#pragma GCC unroll 8
    for (uint32_t i = 0; i < 8; ++i) {
        sfpi::vFloat w = sfpi::dst_reg[WC * 32 + ((t & 1) * 8 + i) * 2];
        sfpi::dst_reg[WS * 32 + 4 * i + 0] = w;
        sfpi::dst_reg[WS * 32 + 4 * i + 1] = w;
        sfpi::dst_reg[WS * 32 + 4 * i + 2] = w;
        sfpi::dst_reg[WS * 32 + 4 * i + 3] = w;
    }
}
// SF_WC16: the compact block holds w(i, t) only in DST row 4 b (16 words); load that 4-row group (even columns) into
// LREG0..3 and transpose (SFPTRANSP: LREG k row j <-> LREG j row k) so LREG0 holds row 0 of all four = w in every lane,
// then store it to the 4 SFPU rows (DST rows 8 i .. 8 i + 7, both column parities) of keypoint i. Register moves only.
template <uint32_t WC, uint32_t WS>
inline void sf_replicate16(uint32_t t) {
    for (uint32_t i = 0; i < 8; ++i) {
        const uint32_t src = WC * 64 + 4 * ((t & 1) * 8 + i);
        TT_SFPLOAD(p_sfpu::LREG0, InstrModLoadStore::DEFAULT, ADDR_MOD_7, src);
        TT_SFPLOAD(p_sfpu::LREG1, InstrModLoadStore::DEFAULT, ADDR_MOD_7, src);
        TT_SFPLOAD(p_sfpu::LREG2, InstrModLoadStore::DEFAULT, ADDR_MOD_7, src);
        TT_SFPLOAD(p_sfpu::LREG3, InstrModLoadStore::DEFAULT, ADDR_MOD_7, src);
        TTI_SFPTRANSP(0, 0, 0, 0);
        const uint32_t dst = WS * 64 + 8 * i;
        TT_SFPSTORE(p_sfpu::LREG0, InstrModLoadStore::DEFAULT, ADDR_MOD_7, dst);
        TT_SFPSTORE(p_sfpu::LREG0, InstrModLoadStore::DEFAULT, ADDR_MOD_7, dst + 2);
        TT_SFPSTORE(p_sfpu::LREG0, InstrModLoadStore::DEFAULT, ADDR_MOD_7, dst + 4);
        TT_SFPSTORE(p_sfpu::LREG0, InstrModLoadStore::DEFAULT, ADDR_MOD_7, dst + 6);
    }
}
#endif

#ifndef SF_KT
#define SF_KT 4
#endif

void kernel_main() {
    const uint32_t nunits = get_arg_val<uint32_t>(0);
    constexpr uint32_t cb_g = get_compile_time_arg_val(0);
    constexpr uint32_t cb_w = get_compile_time_arg_val(1);
    constexpr uint32_t cb_o = get_compile_time_arg_val(2);
    constexpr uint32_t cb_w0 = get_compile_time_arg_val(4);

    unary_op_init_common(cb_g, cb_o);
    for (uint32_t ui = 0; ui < nunits; ++ui) {
        for (uint32_t gi = 0; gi < SF_KT; ++gi) {
            const uint32_t g = ui == 0 ? ((gi + 1) % SF_KT) : gi;
            const uint32_t cbw = (ui == 0 && g == 0) ? cb_w0 : cb_w;
            cb_wait_front(cb_g, 4);
            cb_wait_front(cbw, 2);
            tile_regs_acquire();
            for (uint32_t t = 0; t < 4; ++t) {
                const uint32_t dg = t == 0 ? 0 : 1;
                if ((t & 1) == 0) {
                    copy_tile_to_dst_init_short_with_dt(cb_g, cbw);
                    copy_tile(cbw, t >> 1, 3);
                }
                copy_tile_to_dst_init_short_with_dt(cbw, cb_g);
                copy_tile(cb_g, t, dg);
                mul_binary_tile_init();
                MATH((_llk_math_eltwise_sfpu_start_(0)));
#ifdef SF_WC16
                if (dg == 0) {
                    MATH((sf_replicate16<3, 1>(t)));
                } else {
                    MATH((sf_replicate16<3, 2>(t)));
                }
#else
                if (dg == 0) {
                    MATH((sf_replicate<3, 1>(t)));
                } else {
                    MATH((sf_replicate<3, 2>(t)));
                }
#endif
                MATH((_llk_math_eltwise_sfpu_done_()));
                mul_binary_tile(dg, dg + 1, dg);
                if (t > 0) {
                    add_binary_tile_init();
                    add_binary_tile(0, 1, 0);
                }
            }
            tile_regs_commit();
            tile_regs_wait();
            cb_reserve_back(cb_o, 1);
            pack_tile(0, cb_o);
            cb_push_back(cb_o, 1);
            tile_regs_release();
            cb_pop_front(cb_g, 4);
            cb_pop_front(cbw, 2);
        }
    }
}