changh95's picture
Optimized build (2026-10-03): trace 3.73 -> 0.50 ms
c699c4c verified
Raw History Blame Contribute Delete
8.58 kB
// 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();
}