superpoint-p150 / code /kernels /sp_conv /cb0_reader.cpp
changh95's picture
Optimized build (2026-10-03): trace 3.73 -> 0.50 ms
c699c4c verified
Raw History Blame Contribute Delete
12.1 kB
// 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
}