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();
}
|