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
}
|