File size: 2,168 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
// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// SuperPoint descriptor sampling, fused op, compute: per unit and keypoint group p
//   OUT = ((G0*W0 + G1*W1) + G2*W2) + G3*W3     (fp32 in DEST, SFPU mul/add)
// the same SFPU functions and operation order as the former ttnn.multiply / ttnn.add chain
// (binary_ng SFPU path: bf16 G unpacked exactly, fp32 W unpacked with UnpackToDestFp32).
// Pages are row-major "pseudo tiles" (see the reader); elementwise math is order-agnostic.
#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"

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 KT = get_compile_time_arg_val(3);

    unary_op_init_common(cb_g, cb_o);
    for (uint32_t ui = 0; ui < nunits; ++ui) {
        cb_wait_front(cb_w, 4 * KT);
        cb_wait_front(cb_g, 4 * KT);
        for (uint32_t p = 0; p < KT; ++p) {
            tile_regs_acquire();
            for (uint32_t t = 0; t < 4; ++t) {
                const uint32_t dg = t == 0 ? 0 : 1;
                copy_tile_to_dst_init_short_with_dt(cb_w, cb_g);
                copy_tile(cb_g, p * 4 + t, dg);
                copy_tile_to_dst_init_short_with_dt(cb_g, cb_w);
                copy_tile(cb_w, p * 4 + t, dg + 1);
#ifndef SF_NO_SFPU
                mul_binary_tile_init();
                mul_binary_tile(dg, dg + 1, dg);
                if (t > 0) {
                    add_binary_tile_init();
                    add_binary_tile(0, 1, 0);
                }
#endif
            }
            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 * KT);
        cb_pop_front(cb_w, 4 * KT);
    }
}