File size: 5,922 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
// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// SuperPoint NMS, step 1 (data movement only, runs via ttnn.generic_op).
// Scatter the softmaxed cell scores S [B*h*w, 65(96)] (TILE, interleaved; cell row r = cy*WC + cx,
// column c = i*8 + j) into the zero-padded "strip" layout
//   P[y, xw, l] = Sdense[y, l*SW - PAD + xw]   (0 outside the image),  P: [H, PW = SW + 2*PAD, 32]
// stored as this core's L1 height shard (ROWS image rows per core). Lane l holds the vertical strip
// of SW consecutive columns starting at x = l*SW, padded by PAD on both sides, so a [9,1] + [1,9]
// max-pool over (H, W) of P gives the exact 9x9 window max of every pixel.
#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 s_addr = get_arg_val<uint32_t>(0);
    const uint32_t p_addr = get_arg_val<uint32_t>(1);
    const uint32_t y0 = get_arg_val<uint32_t>(2);   // first image row of this core's shard
    const uint32_t yl0 = get_arg_val<uint32_t>(3);  // local rows [yl0, yl1) handled by this RISC
    const uint32_t yl1 = get_arg_val<uint32_t>(4);

    constexpr uint32_t cb_scratch = get_compile_time_arg_val(0);
    constexpr uint32_t WC = get_compile_time_arg_val(1);     // cells per image row (w/8)
    constexpr uint32_t TCOLS = get_compile_time_arg_val(2);  // tile columns of S
    constexpr uint32_t SW = get_compile_time_arg_val(3);     // strip width (W / 32)
    constexpr uint32_t PAD = get_compile_time_arg_val(4);    // nms radius
    constexpr uint32_t LANES = 32;
    constexpr uint32_t PW = SW + 2 * PAD;
    constexpr uint32_t W = WC * 8;
    constexpr auto s_args = TensorAccessorArgs<5>();
    const auto s = TensorAccessor(s_args, s_addr, 2048);

    // Scratch: up to (WC+31)/32+1 S tiles, followed by one zero-padded dense row R[PAD + W + PAD].
    const uint32_t scratch = get_write_ptr(cb_scratch);
    constexpr uint32_t NT = (WC + 31) / 32 + 1;
    const uint16_t* src = reinterpret_cast<const uint16_t*>(scratch);
    // odd PAD: shift the row buffer by 2 B so that rbuf + PAD stays 4-byte aligned
#ifdef FOLD_LOCAL
    // the dense row in RISC-local data memory (single-cycle loads) instead of L1
    static uint32_t rloc[(PAD + W + PAD + 2) / 2 + 1];
    uint16_t* rbuf = reinterpret_cast<uint16_t*>(rloc) + (PAD & 1);
#else
    uint16_t* rbuf = reinterpret_cast<uint16_t*>(scratch + NT * 2048 + (PAD & 1) * 2);
#endif
    for (uint32_t k = 0; k < PAD; ++k) {
        rbuf[k] = 0;
        rbuf[PAD + W + k] = 0;
    }
    uint32_t* dst = reinterpret_cast<uint32_t*>(p_addr);

    uint32_t loaded = 0xFFFFFFFF;
    uint32_t roff = 0;
    for (uint32_t yl = yl0; yl < yl1; ++yl) {
        const uint32_t y = y0 + yl;
        const uint32_t cy = y >> 3;
        const uint32_t i = y & 7;
        const uint32_t tc = (i * 8) >> 5;
        const uint32_t key = cy * TCOLS + tc;
        const uint32_t r0 = cy * WC;
        const uint32_t tr0 = r0 >> 5;
        if (key != loaded) {
            ZONE("FOLD_RD");
            const uint32_t tr1 = (r0 + WC - 1) >> 5;
            for (uint32_t tr = tr0; tr <= tr1; ++tr) {
                noc_async_read(s.get_noc_addr(tr * TCOLS + tc), scratch + (tr - tr0) * 2048, 2048);
            }
            noc_async_read_barrier();
            loaded = key;
            roff = r0 - tr0 * 32;
        }
        // 1) gather the dense row: cell cx contributes 8 contiguous values (16 B) of one face row.
        {
        ZONE("FOLD_GATHER");
        const uint32_t cl = (i * 8) & 31;
        const uint32_t cbase = ((cl >> 4) * 256) + (cl & 15);
        uint32_t* r32 = reinterpret_cast<uint32_t*>(rbuf + PAD);  // 4 B aligned (see rbuf)
#ifdef FOLD_SEG
        // runs of consecutive cells inside one 16-row face half: the source steps by one face row (8 words)
        for (uint32_t cx = 0; cx < WC;) {
            const uint32_t rr = roff + cx;
            const uint32_t rin = rr & 31;
            uint32_t n = 16 - (rin & 15);
            if (n > WC - cx) {
                n = WC - cx;
            }
            const uint32_t* c32 = reinterpret_cast<const uint32_t*>(src + (rr >> 5) * 1024 + ((rin >> 4) << 9) + (rin & 15) * 16 + cbase);
            uint32_t* o = r32 + cx * 4;
            cx += n;
#pragma GCC unroll 2
            for (; n >= 2; n -= 2) {
                const uint32_t a0 = c32[0], a1 = c32[1], a2 = c32[2], a3 = c32[3];
                const uint32_t b0 = c32[8], b1 = c32[9], b2 = c32[10], b3 = c32[11];
                o[0] = a0; o[1] = a1; o[2] = a2; o[3] = a3;
                o[4] = b0; o[5] = b1; o[6] = b2; o[7] = b3;
                c32 += 16;
                o += 8;
            }
            if (n) {
                o[0] = c32[0]; o[1] = c32[1]; o[2] = c32[2]; o[3] = c32[3];
            }
        }
        }
#else
        for (uint32_t cx = 0; cx < WC; ++cx) {
            const uint32_t rr = roff + cx;
            const uint32_t rin = rr & 31;
            const uint32_t off = (rr >> 5) * 1024 + ((rin >> 4) << 9) + (rin & 15) * 16 + cbase;
            const uint32_t* c32 = reinterpret_cast<const uint32_t*>(src + off);
            r32[cx * 4 + 0] = c32[0];
            r32[cx * 4 + 1] = c32[1];
            r32[cx * 4 + 2] = c32[2];
            r32[cx * 4 + 3] = c32[3];
        }
        }
#endif
        // 2) strip scatter: P[yl, xw, l] = R[l*SW + xw]; two lanes per 32-bit store.
        uint32_t* drow = dst + yl * PW * (LANES / 2);
        ZONE("FOLD_SCAT");
        for (uint32_t xw = 0; xw < PW; ++xw) {
            const uint16_t* b = rbuf + xw;
            uint32_t* d = drow + xw * (LANES / 2);
#pragma GCC unroll 16
            for (uint32_t l = 0; l < LANES; l += 2) {
                d[l >> 1] = (uint32_t)b[l * SW] | ((uint32_t)b[(l + 1) * SW] << 16);
            }
        }
    }
}