File size: 4,145 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
// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// SuperPoint NMS window max (nms_pool_compute.cpp), data movement. PROC 0: the NROW = ROWS + 2R strip
// rows y0 - R .. y0 + ROWS - 1 + R of P (PW * 64 B each; previous / next core's shard rows for the
// halo, zeros outside the image) into CB_IN pages (2 KB). PROC 1: the ROWS result pages -> M shard rows
// (SW * 64 B each).
// RT args: p_addr, m_addr, prev_x, prev_y, next_x, next_y, has_prev, has_next
// CT args: PROC, cb_in, cb_out, ROWS, R, PW, SW
#include <stdint.h>
#include "api/dataflow/dataflow_api.h"

void kernel_main() {
    const uint32_t p_addr = get_arg_val<uint32_t>(0);
    const uint32_t m_addr = get_arg_val<uint32_t>(1);
    const uint32_t pvx = get_arg_val<uint32_t>(2);
    const uint32_t pvy = get_arg_val<uint32_t>(3);
    const uint32_t nxx = get_arg_val<uint32_t>(4);
    const uint32_t nxy = get_arg_val<uint32_t>(5);
    const uint32_t has_prev = get_arg_val<uint32_t>(6);
    const uint32_t has_next = get_arg_val<uint32_t>(7);
    constexpr uint32_t PROC = get_compile_time_arg_val(0);
    constexpr uint32_t cb_in = get_compile_time_arg_val(1);
    constexpr uint32_t cb_out = get_compile_time_arg_val(2);
    constexpr uint32_t ROWS = get_compile_time_arg_val(3);
    constexpr uint32_t R = get_compile_time_arg_val(4);
    constexpr uint32_t PW = get_compile_time_arg_val(5);
    constexpr uint32_t SW = get_compile_time_arg_val(6);
    constexpr uint32_t NROW = ROWS + 2 * R, RB = PW * 64, OB = SW * 64, PAGE = 2048;
    if constexpr (PROC == 0) {
        cb_reserve_back(cb_in, NROW);
        const uint32_t base = get_write_ptr(cb_in);
        for (uint32_t j = 0; j < NROW; ++j) {
            const int32_t row = (int32_t)j - (int32_t)R;  // relative to this core's first row
            const uint32_t dst = base + j * PAGE;
            if ((row < 0 && !has_prev) || (row >= (int32_t)ROWS && !has_next)) {
#ifdef POOL_ZDMA
                // rows outside the image: zeros by DMA from the hardware zero page (the scalar store loop made the
                // first / last core ~3 us late)
                const uint64_t zsrc = get_noc_addr(my_x[noc_index], my_y[noc_index], MEM_ZEROS_BASE);
                for (uint32_t off = 0; off < RB; off += MEM_ZEROS_SIZE) {
                    noc_async_read(zsrc, dst + off, RB - off < MEM_ZEROS_SIZE ? RB - off : MEM_ZEROS_SIZE);
                }
#else
                volatile tt_l1_ptr uint32_t* z = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(dst);
                for (uint32_t w = 0; w < RB / 4; ++w) {
                    z[w] = 0;
                }
                (void)z[RB / 4 - 1];
#endif
            } else if (row < 0) {
                noc_async_read(get_noc_addr(pvx, pvy, p_addr + (uint32_t)(row + (int32_t)ROWS) * RB), dst, RB);
            } else if (row >= (int32_t)ROWS) {
                noc_async_read(get_noc_addr(nxx, nxy, p_addr + (uint32_t)(row - (int32_t)ROWS) * RB), dst, RB);
            } else {
                noc_async_read(get_noc_addr(my_x[noc_index], my_y[noc_index], p_addr + (uint32_t)row * RB), dst, RB);
            }
        }
        noc_async_read_barrier();
        cb_push_back(cb_in, NROW);
    } else {
#ifndef POOL_UNF
        cb_wait_front(cb_out, ROWS);
        const uint32_t src = get_read_ptr(cb_out);
        for (uint32_t k = 0; k < ROWS; ++k) {
            noc_async_write(src + k * PAGE, get_noc_addr(my_x[noc_index], my_y[noc_index], m_addr + k * OB), OB);
        }
        noc_async_write_barrier();
        cb_pop_front(cb_out, ROWS);
#endif
    }
#ifdef POOL_UNF
    // SP_NMS_POOL_UNF (PMASK): the NMS unfold + keypoint candidates (nms_unfold_fn.inc, prepended) of this RISC's
    // rows straight from the ROWS result pages in CB_OUT (both RISCs read them; nothing pops, M is not written).
    // Unfold CT args start at 7, its RT-arg block at 8.
#ifndef UNF_ROW_WAIT
    cb_wait_front(cb_out, PROC == 0 ? UNF_SPLIT : ROWS);  // PROC 0 unfolds rows [0, UNF_SPLIT), PROC 1 the rest
#endif
    unfold_rows<7, PROC, PAGE / 4>(reinterpret_cast<const uint32_t*>(get_read_ptr(cb_out)), 8);
#endif
}