File size: 12,086 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
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// Block-0 conv_b, data movement (see models/tt/conv_cell.py for the layout). Each RISC builds the
// GROUP-tile shifted-tap group of its output tile rows (BMASK: PROC 1's rows) into its CB
// by executing a host-built copy schedule: per TR a list of (source, src unit, dst unit, units)
// NoC reads in 32-byte units; source 0 = this core's activation shard, 1 / 2 = the previous / next
// core's shard (same L1 address), 3 = a local zero page. PROC 0 also publishes the local shard
// (CB_X bound to it) and provides the weights in CB_W / CB_B: with MCAST, only the sender core
// (logical core 0, PROC 1) reads them from DRAM and multicasts them to the other cores
// (120 cores reading the same 76 KB from DRAM took up to ~35 us on the far cores); receivers first
// signal the sender that
// they are running (semaphore 0 on the sender), then wait for semaphore 1 (set by the sender's
// multicast after the data, same NoC / VC). Multicasting from several sender cores at once, or tile
// by tile, measured much slower (concurrent multicast path reservations).
//
// RT args: x_addr, core_class (0 first / 1 mid / 2 last), prev_x, prev_y, next_x, next_y, w_addr, sched_addr,
//          is_sender, sender x, sender y, mcast x_start, y_start, x_end, y_end, n_receivers
// CT args: PROC, cb_s, cb_x, cb_w, cb_t, W_TILES, MAXE, TRS, SCHED_BYTES, cb_b, TensorAccessorArgs(w), (sched)
#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 x_addr = get_arg_val<uint32_t>(0);
    const uint32_t cls = get_arg_val<uint32_t>(1);
    const uint32_t pnx = get_arg_val<uint32_t>(2);
    const uint32_t pny = get_arg_val<uint32_t>(3);
    const uint32_t nnx = get_arg_val<uint32_t>(4);
    const uint32_t nny = get_arg_val<uint32_t>(5);
    const uint32_t w_addr = get_arg_val<uint32_t>(6);
    const uint32_t s_addr = get_arg_val<uint32_t>(7);
    const uint32_t is_sender = get_arg_val<uint32_t>(8);
    const uint32_t snd_x = get_arg_val<uint32_t>(9);
    const uint32_t snd_y = get_arg_val<uint32_t>(10);
    uint32_t mx0 = get_arg_val<uint32_t>(11);
    uint32_t my0 = get_arg_val<uint32_t>(12);
    uint32_t mx1 = get_arg_val<uint32_t>(13);
    uint32_t my1 = get_arg_val<uint32_t>(14);
    const uint32_t nrecv = get_arg_val<uint32_t>(15);
    constexpr uint32_t PROC = get_compile_time_arg_val(0);
    constexpr uint32_t cb_s = get_compile_time_arg_val(1);
    constexpr uint32_t cb_x = get_compile_time_arg_val(2);
    constexpr uint32_t cb_w = get_compile_time_arg_val(3);
    constexpr uint32_t cb_t = get_compile_time_arg_val(4);
    constexpr uint32_t W_TILES = get_compile_time_arg_val(5);
    constexpr uint32_t MAXE = get_compile_time_arg_val(6);
    constexpr uint32_t TRS = get_compile_time_arg_val(7);
    constexpr uint32_t SCHED_BYTES = get_compile_time_arg_val(8);
    constexpr uint32_t cb_b = get_compile_time_arg_val(9);
    constexpr uint32_t GROUP = GROUP_DEF, QT = QT_DEF, TILE = 2048, PAGE = MAXE * 8;
#ifndef SCHED_RD
#define SCHED_RD PAGE  // bytes of each schedule page actually read (SP_SCHED_REP: only the used entries)
#endif
    // output tile rows whose group PROC 1 builds (bit tr set); the rest are PROC 0's. PROC 0 (NOC 1) builds
    // ~2x slower than PROC 1 here and also provides the weights, so it gets fewer rows.
    constexpr uint32_t BMASK = BMASK_DEF;
    constexpr auto w_args = TensorAccessorArgs<10>();
    constexpr auto s_args = TensorAccessorArgs<w_args.next_compile_time_args_offset()>();
    const auto wacc = TensorAccessor(w_args, w_addr, TILE);
    const auto sacc = TensorAccessor(s_args, s_addr, PAGE);

    const uint32_t scratch = get_write_ptr(cb_t) + PROC * (SCHED_BYTES + 1024);
    const uint32_t sched_l1 = scratch;
    const uint32_t zero_l1 = scratch + SCHED_BYTES;
    {
        volatile tt_l1_ptr uint32_t* z = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(zero_l1);
        for (uint32_t k = 0; k < 256; ++k) {
            z[k] = 0;
        }
        (void)z[255];
    }
    // copy schedule of this RISC's TRs (the weight sender's PROC 0 issues it after the multicast: these
    // DRAM pages are read by every core and are slow to arrive)
    uint32_t nk = 0;
#if defined(MCAST) && !defined(W_PRE)
    const bool sched_late = is_sender != 0;  // both RISCs of the sender: weights first
#else
    const bool sched_late = false;
#endif
    if (!sched_late) {
        for (uint32_t tr = 0; tr < TRS; ++tr) {
            if (((BMASK >> tr) & 1) == PROC) {
                noc_async_read(sacc.get_noc_addr(cls * TRS + tr), sched_l1 + nk * PAGE, SCHED_RD);
                ++nk;
            }
        }
    }
    if constexpr (PROC == 0) {
        cb_reserve_back(cb_x, TRS * QT);
        cb_push_back(cb_x, TRS * QT);
    }
    // weights (W_TILES tiles -> CB_W) and bias (2 tiles -> CB_B). With MCAST only the sender core (logical
    // core 0) reads them from DRAM (PROC 1) and multicasts them to all other cores (PROC 0).
    constexpr uint32_t NT = W_TILES + 2;
    const uint32_t wl = get_write_ptr(cb_w);
    const uint32_t bl = get_write_ptr(cb_b);
#if defined(W_PRE)
    // SP_CW_PF: the weights were multicast into the tensors behind CB_W / CB_B by the previous cell op
    // (PF_NEXT below); nothing to read
    const bool reads_w = false;
#elif defined(MCAST)
    volatile tt_l1_ptr uint32_t* rdy_sem = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(get_semaphore(0));
    volatile tt_l1_ptr uint32_t* rcv_sem = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(get_semaphore(1));
    volatile tt_l1_ptr uint32_t* loc_sem = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(get_semaphore(2));
    const bool reads_w = is_sender != 0 && PROC == 1;  // PROC 1 (NOC 0) reads ~2x faster here
#else
    const bool reads_w = PROC == 0;
#endif
    if (reads_w) {
        ZONE("CB0_WDRAM");
        for (uint32_t t = 0; t < NT; ++t) {
            noc_async_read(wacc.get_noc_addr(t), t < W_TILES ? wl + t * TILE : bl + (t - W_TILES) * TILE, TILE);
        }
        noc_async_read_barrier();
#if defined(MCAST) && !defined(W_PRE)
        // local handoff to PROC 0, which multicasts (on its NoC: the other cores' PROC 1 are already
        // building groups on the other NoC, which made a multicast from PROC 1 ~5x slower)
        noc_semaphore_inc(get_noc_addr(my_x[noc_index], my_y[noc_index], (uint32_t)loc_sem), 1);
        noc_async_atomic_barrier();
#endif
    }
    if constexpr (PROC == 0) {
        cb_reserve_back(cb_w, W_TILES);
        cb_reserve_back(cb_b, 2);
#if defined(MCAST) && !defined(W_PRE)
        if (is_sender) {
            ZONE("CB0_WMC");
            noc_semaphore_wait(loc_sem, 1);
            noc_semaphore_set(loc_sem, 0);
            if (nrecv > 0) {
                if (noc_index == 1) {
                    const uint32_t tx = mx0, ty = my0;
                    mx0 = mx1;
                    my0 = my1;
                    mx1 = tx;
                    my1 = ty;
                }
                noc_semaphore_wait(rdy_sem, nrecv);
                noc_semaphore_set(rdy_sem, 0);
                noc_async_write_multicast(wl, get_noc_multicast_addr(mx0, my0, mx1, my1, wl), W_TILES * TILE, nrecv);
                noc_async_write_multicast(bl, get_noc_multicast_addr(mx0, my0, mx1, my1, bl), 2 * TILE, nrecv);
                noc_async_writes_flushed();
                *rcv_sem = 1;
                noc_semaphore_set_multicast((uint32_t)rcv_sem, get_noc_multicast_addr(mx0, my0, mx1, my1, (uint32_t)rcv_sem), nrecv);
                noc_async_write_barrier();
            }
        } else {
            ZONE("CB0_WRCV");
            noc_semaphore_inc(get_noc_addr(snd_x, snd_y, (uint32_t)rdy_sem), 1);
            noc_semaphore_wait(rcv_sem, 1);
            noc_semaphore_set(rcv_sem, 0);
            noc_async_atomic_barrier();
        }
#endif
        cb_push_back(cb_w, W_TILES);
        cb_push_back(cb_b, 2);
    }
    if (sched_late) {
        for (uint32_t tr = 0; tr < TRS; ++tr) {
            if (((BMASK >> tr) & 1) == PROC) {
                noc_async_read(sacc.get_noc_addr(cls * TRS + tr), sched_l1 + nk * PAGE, SCHED_RD);
                ++nk;
            }
        }
    }
    noc_async_read_barrier();
    uint64_t base[4];
    base[0] = get_noc_addr(my_x[noc_index], my_y[noc_index], x_addr);
    base[1] = get_noc_addr(pnx, pny, x_addr);
    base[2] = get_noc_addr(nnx, nny, x_addr);
    base[3] = get_noc_addr(my_x[noc_index], my_y[noc_index], zero_l1);
    nk = 0;
    for (uint32_t tr = 0; tr < TRS; ++tr) {
        if (((BMASK >> tr) & 1) != PROC) {
            continue;
        }
        cb_reserve_back(cb_s, GROUP);
        const uint32_t g = get_write_ptr(cb_s);
#ifndef NO_BUILD
        const uint32_t* e = reinterpret_cast<const uint32_t*>(sched_l1 + nk * PAGE);
        const uint32_t cnt = e[0];
        uint32_t cur = 0xFFFFFFFFu;
        ZONE("CB0_BUILD");
        for (uint32_t j = 1; j <= cnt; ++j) {
            const uint32_t w0 = e[2 * j], w1 = e[2 * j + 1];
#ifdef ABL_ONLY_BIG
            if ((w1 & 0xFFFF) < 32) continue;  // ablation (timing only): skip the small copies
#endif
#ifdef ABL_ONLY_SMALL
            if ((w1 & 0xFFFF) >= 32) continue;  // ablation (timing only): skip the 1 KB copies
#endif
            const uint32_t sel = w0 >> 16;
            const uint32_t key = (sel << 16) | (w1 & 0xFFFF);
            if (key != cur) {
                // entries are sorted by (source, length): new source core / size only on a change
                noc_async_read_one_packet_set_state(base[sel], (w1 & 0xFFFF) << 5);
                cur = key;
            }
            const uint32_t src = (sel == 3) ? zero_l1 : x_addr + ((w0 & 0xFFFF) << 5);
            noc_async_read_one_packet_with_state(src, g + ((w1 >> 16) << 5));
        }
        noc_async_read_barrier();
#endif
        cb_push_back(cb_s, GROUP);
        ++nk;
    }
#ifdef PF_NEXT
    // SP_CW_PF: after its own work, the sender's PROC 0 reads the weights (+ bias) of later cell ops from DRAM into its
    // own shard of the in-graph L1 tensors behind those ops' CB_W / CB_B and multicasts them to the same address on
    // every other core, one multicast at a time (the target tensors are not used by any core during this op, so no
    // handshake; the later ops start only after this program, i.e. after the write barrier, has completed).
    // RT args 16: number of targets m, then m x (DRAM weight tensor, target W tensor, target bias tensor).
    if constexpr (PROC == 0) {
        if (is_sender) {
            ZONE("CB0_PF");
            uint32_t ax0 = get_arg_val<uint32_t>(11), ay0 = get_arg_val<uint32_t>(12);
            uint32_t ax1 = get_arg_val<uint32_t>(13), ay1 = get_arg_val<uint32_t>(14);
            if (noc_index == 1) {
                const uint32_t tx = ax0, ty = ay0;
                ax0 = ax1;
                ay0 = ay1;
                ax1 = tx;
                ay1 = ty;
            }
            const uint32_t m = get_arg_val<uint32_t>(16);
            for (uint32_t i = 0; i < m; ++i) {
                const uint32_t nw = get_arg_val<uint32_t>(17 + 3 * i);
                const uint32_t pw = get_arg_val<uint32_t>(18 + 3 * i);
                const uint32_t pb = get_arg_val<uint32_t>(19 + 3 * i);
                const auto nacc = TensorAccessor(w_args, nw, TILE);
                for (uint32_t t = 0; t < NT; ++t) {
                    noc_async_read(nacc.get_noc_addr(t), t < W_TILES ? pw + t * TILE : pb + (t - W_TILES) * TILE, TILE);
                }
                noc_async_read_barrier();
                if (nrecv > 0) {
                    noc_async_write_multicast(pw, get_noc_multicast_addr(ax0, ay0, ax1, ay1, pw), W_TILES * TILE, nrecv);
                    noc_async_write_multicast(pb, get_noc_multicast_addr(ax0, ay0, ax1, ay1, pb), 2 * TILE, nrecv);
                    noc_async_write_barrier();
                }
            }
        }
    }
#endif
}