Download code/kernels/sp_conv/rc_dm.cpp from changh95/superpoint-p150: direct link, hf CLI and curl.
- Browser
- Download file 8.58 kB
-
https://huggingface.co/changh95/superpoint-p150/resolve/main/code/kernels/sp_conv/rc_dm.cpp
- Command line
-
hf download hf://changh95/superpoint-p150/code/kernels/sp_conv/rc_dm.cpp
-
curl -L -o rc_dm.cpp https://huggingface.co/changh95/superpoint-p150/resolve/main/code/kernels/sp_conv/rc_dm.cpp
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) | |
| 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); | |
| } | |
| // 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); | |
| } | |
| } | |
| if (wsender) { | |
| noc_async_write_barrier(); | |
| } | |
| noc_async_atomic_barrier(); | |
| } | |