File size: 1,356 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
// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// SuperPoint NMS fold on 5 RISCs per core (SP_NMS_FOLD5), the TRISCs: plain scalar code (no FPU/SFPU). UNPACK waits for
// the S tiles in CB_S and hands their address to MATH and PACK through the TRISC mailboxes (get_tile_address), then
// each folds its rows (nms_fold5_row.inc, prepended). RT args: p_addr, y0. CT args: cb_s, WC, SW, PAD, ROWS, NT.
#include <cstdint>
#include "api/compute/compute_kernel_api.h"
#include "api/compute/common.h"
#include "api/compute/cb_api.h"

void kernel_main() {
    const uint32_t p_addr = get_arg_val<uint32_t>(0);
    const uint32_t y0 = get_arg_val<uint32_t>(1);
    constexpr uint32_t cb_s = get_compile_time_arg_val(0);
    constexpr uint32_t WC = get_compile_time_arg_val(1);
    constexpr uint32_t SW = get_compile_time_arg_val(2);
    constexpr uint32_t PAD = get_compile_time_arg_val(3);
    constexpr uint32_t ROWS = get_compile_time_arg_val(4);
    constexpr uint32_t NT = get_compile_time_arg_val(5);
    cb_wait_front(cb_s, NT);
    const uint32_t tiles = get_tile_address(cb_s, 0);
#if defined(TRISC_UNPACK)
    constexpr uint32_t who = 2;
#elif defined(TRISC_MATH)
    constexpr uint32_t who = 3;
#else
    constexpr uint32_t who = 4;
#endif
    fold5_rows<WC, SW, PAD, ROWS>(who, tiles, p_addr, y0);
}