File size: 8,577 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
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// RowConv data movement (models/tt/row_conv.py). Both RISCs run this kernel (PROC 0 / 1):
//  - build this RISC's half of the in0 tiles (channel tiles PROC*CH .. PROC*CH+CH-1) of each K block
//    ky into its CB_A from the host copy schedule (32-byte units, entries for channel tile 0; +64 units
//    per channel tile on both sides); sources are up to 6 input shards (NoC) and a local zero page;
//  - weights: the group's sender core (WPROC) reads K block ky (3*CT*NL tiles) from DRAM into CB_W and
//    multicasts it to the group rectangle, then multicasts semaphore 1 = ky + 1. PROC 0 of every core
//    pushes CB_W chunk ky once semaphore 1 >= ky + 1. Receivers first increment the sender's
//    semaphore 0 (they run the program: its semaphores are initialised);
//  - PROC 0 reads the bias tiles; PROC 1 writes the output rows (16-row halves) to the output shards.
// RT: x, out, w, b, sched, core index, is_sender, sender x, y, mcast x0, y0, x1, y1, n_recv, w tile0,
//     b tile0, 6 x (src x, y), 3 x (dst x, y)
// CT: PROC, cb_a, cb_w, cb_b, cb_o, cb_t, NL, CH, CT, PAGE_BYTES, BLK_H, WPROC, accessors (w, b, sched)
#include <stdint.h>
#include "api/dataflow/dataflow_api.h"

void kernel_main() {
    const uint32_t x_addr = get_arg_val<uint32_t>(0);
    const uint32_t out_addr = get_arg_val<uint32_t>(1);
    const uint32_t w_addr = get_arg_val<uint32_t>(2);
    const uint32_t b_addr = get_arg_val<uint32_t>(3);
    const uint32_t s_addr = get_arg_val<uint32_t>(4);
    const uint32_t core_k = get_arg_val<uint32_t>(5);
    const uint32_t is_snd = get_arg_val<uint32_t>(6);
    const uint32_t snd_x = get_arg_val<uint32_t>(7);
    const uint32_t snd_y = get_arg_val<uint32_t>(8);
    uint32_t mx0 = get_arg_val<uint32_t>(9);
    uint32_t my0 = get_arg_val<uint32_t>(10);
    uint32_t mx1 = get_arg_val<uint32_t>(11);
    uint32_t my1 = get_arg_val<uint32_t>(12);
    const uint32_t nrecv = get_arg_val<uint32_t>(13);
    const uint32_t wt0 = get_arg_val<uint32_t>(14);
    const uint32_t bt0 = get_arg_val<uint32_t>(15);
    constexpr uint32_t PROC = get_compile_time_arg_val(0);
    constexpr uint32_t cb_a = get_compile_time_arg_val(1);
    constexpr uint32_t cb_w = get_compile_time_arg_val(2);
    constexpr uint32_t cb_b = get_compile_time_arg_val(3);
    constexpr uint32_t cb_o = get_compile_time_arg_val(4);
    constexpr uint32_t cb_t = get_compile_time_arg_val(5);
    constexpr uint32_t NL = get_compile_time_arg_val(6);
    constexpr uint32_t CH = get_compile_time_arg_val(7);
    constexpr uint32_t CT = get_compile_time_arg_val(8);
    constexpr uint32_t PAGE_BYTES = get_compile_time_arg_val(9);
    constexpr uint32_t BLK_H = get_compile_time_arg_val(10);
    constexpr uint32_t WPROC = get_compile_time_arg_val(11);
    constexpr uint32_t TILE = 2048, ZSEL = 7;
    constexpr uint32_t CHUNK = 3 * CT * NL;  // weight tiles per K block
    constexpr auto w_args = TensorAccessorArgs<12>();
    constexpr auto b_args = TensorAccessorArgs<w_args.next_compile_time_args_offset()>();
    constexpr auto s_args = TensorAccessorArgs<b_args.next_compile_time_args_offset()>();
    const auto wacc = TensorAccessor(w_args, w_addr, TILE);
    const auto bacc = TensorAccessor(b_args, b_addr, TILE);
    const auto sacc = TensorAccessor(s_args, s_addr, PAGE_BYTES);

    const uint32_t sched_l1 = get_write_ptr(cb_t) + PROC * (PAGE_BYTES + 1024);
    const uint32_t zero_l1 = sched_l1 + PAGE_BYTES;
    noc_async_read(sacc.get_noc_addr(core_k), sched_l1, PAGE_BYTES);
    {
        volatile tt_l1_ptr uint32_t* z = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(zero_l1);
        for (uint32_t i = 0; i < 256; ++i) {
            z[i] = 0;
        }
        (void)z[255];
    }
    volatile tt_l1_ptr uint32_t* rdy = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(get_semaphore(0));
    volatile tt_l1_ptr uint32_t* wsem = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(get_semaphore(1));
    const bool wsender = is_snd != 0 && PROC == WPROC;
    const uint32_t wl = get_write_ptr(cb_w);
    if constexpr (PROC == 0) {
        cb_reserve_back(cb_b, NL);
        const uint32_t bl = get_write_ptr(cb_b);
        for (uint32_t n = 0; n < NL; ++n) {
            noc_async_read(bacc.get_noc_addr(bt0 + n), bl + n * TILE, TILE);
        }
    }
    if (is_snd == 0 && PROC == WPROC) {
        noc_semaphore_inc(get_noc_addr(snd_x, snd_y, (uint32_t)rdy), 1);
    }
    if (wsender && noc_index == 1) {
        const uint32_t tx = mx0, ty = my0;
        mx0 = mx1;
        my0 = my1;
        mx1 = tx;
        my1 = ty;
    }
    noc_async_read_barrier();
    if constexpr (PROC == 0) {
        cb_push_back(cb_b, NL);
    }
    uint64_t base[8];
    for (uint32_t s = 0; s < 6; ++s) {
        base[s] = get_noc_addr(get_arg_val<uint32_t>(16 + 2 * s), get_arg_val<uint32_t>(17 + 2 * s), 0);
    }
    base[6] = base[0];
    base[ZSEL] = get_noc_addr(my_x[noc_index], my_y[noc_index], 0);
    const uint32_t* e = reinterpret_cast<const uint32_t*>(sched_l1);
    uint32_t off = 4;
    for (uint32_t ky = 0; ky < 3; ++ky) {
        if (wsender) {
            for (uint32_t t = 0; t < CHUNK; ++t) {
                noc_async_read(wacc.get_noc_addr(wt0 + ky * CHUNK + t), wl + (ky * CHUNK + t) * TILE, TILE);
            }
        }
        cb_reserve_back(cb_a, BLK_H);
        const uint32_t a = get_write_ptr(cb_a);
        const uint32_t n = e[ky];
        uint32_t cur = 0xFFFFFFFFu;
        for (uint32_t j = 0; j < n; ++j) {
            const uint32_t w0 = e[off + 2 * j], w1 = e[off + 2 * j + 1];
            const uint32_t sel = w0 >> 16, su = w0 & 0xFFFF, du = w1 >> 16, len = w1 & 0xFFFF;
            const uint32_t key = (sel << 16) | len;
            if (key != cur) {
                noc_async_read_one_packet_set_state(base[sel], len << 5);
                cur = key;
            }
            for (uint32_t cl = 0; cl < CH; ++cl) {
                const uint32_t src = sel == ZSEL ? zero_l1 : x_addr + ((su + (PROC * CH + cl) * 64) << 5);
                noc_async_read_one_packet_with_state(src, a + ((du + cl * 64) << 5));
            }
        }
        off += 2 * n;
        noc_async_read_barrier();
        cb_push_back(cb_a, BLK_H);
        if (wsender) {
            if (ky == 0) {
                noc_semaphore_wait(rdy, nrecv);
                noc_semaphore_set(rdy, 0);
            }
            const uint32_t src = wl + ky * CHUNK * TILE;
            noc_async_write_multicast(src, get_noc_multicast_addr(mx0, my0, mx1, my1, src), CHUNK * TILE, nrecv);
            noc_async_writes_flushed();
            *wsem = ky + 1;
            noc_semaphore_set_multicast((uint32_t)wsem, get_noc_multicast_addr(mx0, my0, mx1, my1, (uint32_t)wsem), nrecv);
            noc_async_writes_flushed();
        }
        if constexpr (PROC == 0) {
            noc_semaphore_wait_min(wsem, ky + 1);
            cb_reserve_back(cb_w, CHUNK);
            cb_push_back(cb_w, CHUNK);
        }
    }
    if constexpr (PROC == 1) {
        const uint32_t no = e[3];
        cb_wait_front(cb_o, 3 * NL);
        const uint32_t o = get_read_ptr(cb_o);
        for (uint32_t j = 0; j < no; ++j) {
            const uint32_t w0 = e[off + 2 * j], w1 = e[off + 2 * j + 1];
            const uint32_t sel = w0 >> 16, su = w0 & 0xFFFF, du = w1 >> 16, len = w1 & 0xFFFF;
            const uint32_t dx = get_arg_val<uint32_t>(28 + 2 * sel), dy = get_arg_val<uint32_t>(29 + 2 * sel);
            for (uint32_t nn = 0; nn < NL; ++nn) {
                noc_async_write(o + ((su + nn * 64) << 5), get_noc_addr(dx, dy, out_addr + ((du + nn * 64) << 5)), len << 5);
            }
        }
        noc_async_write_barrier();
        cb_pop_front(cb_o, 3 * NL);
    }
#ifndef RC_NO_DONE
    // end-of-op barrier: every receiver's weight RISC reports that it saw all chunks (semaphore 2 on
    // the sender); the sender does not exit before all of them did, so no weight / semaphore traffic
    // of this launch can be in flight when the program ends.
    {
        volatile tt_l1_ptr uint32_t* done = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(get_semaphore(2));
        if (is_snd == 0 && PROC == 0) {  // PROC 0 waited for semaphore 1 >= 3 (all chunks) above
            noc_semaphore_inc(get_noc_addr(snd_x, snd_y, (uint32_t)done), 1);
        } else if (wsender) {
            noc_semaphore_wait(done, nrecv);
            noc_semaphore_set(done, 0);
        }
    }
#endif
    if (wsender) {
        noc_async_write_barrier();
    }
    noc_async_atomic_barrier();
}