// SPDX-License-Identifier: Apache-2.0 // Fused residual add + LayerNorm on 2 cores per tile row (moge-2 opt kernel, round 7): NCRISC. // 1. Reads this core's half row of x (HALF tiles) into c_0, RB tiles per NoC barrier (116 cores with 32 reads // each in flight saturate the NoC: 21 us vs 7 us for the same bytes). // 2. Two partial-statistic exchanges with the partner core of the same tile row: round k (1: mean, 2: // variance): wait for the compute's partial tile in c_18, copy it to slot 0 of the local pair CB (c_19 / // c_20) and to slot 1 of the partner's pair CB (same L1 address on both cores), increment the partner's // semaphore, wait until the own semaphore reaches k, push the pair (2 tiles). The semaphore is reset to 0 // before exit (the partner has finished both increments by then). // 3. Writes the odd BLK-blocks of LN(z) (c_27); the writer (BRISC) writes the even ones. // Runtime args: x_addr, out_addr, row, col0, partner_noc_x, partner_noc_y. // Compile args: HALF, WT, BLK, SEM_ID, TensorAccessorArgs x, out. #include "api/dataflow/dataflow_api.h" #include "api/dataflow/noc.h" #include "api/dataflow/dataflow_buffer.h" #include "api/tensor/noc_traits.h" #ifdef ADDLN_PROF #include "tools/profiler/kernel_profiler.hpp" #define PZONE(n) DeviceZoneScopedN(n) #else #define PZONE(n) #endif #ifndef RB #define RB 2 #endif void kernel_main() { const uint32_t x_addr = get_arg_val(0); const uint32_t o_addr = get_arg_val(1); const uint32_t row = get_arg_val(2); const uint32_t col0 = get_arg_val(3); const uint32_t px = get_arg_val(4); const uint32_t py = get_arg_val(5); constexpr uint32_t HALF = get_compile_time_arg_val(0); constexpr uint32_t WT = get_compile_time_arg_val(1); constexpr uint32_t BLK = get_compile_time_arg_val(2); constexpr uint32_t SEM_ID = get_compile_time_arg_val(3); constexpr auto x_args = TensorAccessorArgs<4>(); constexpr auto o_args = TensorAccessorArgs(); constexpr uint32_t cb_x = 0, cb_scaler = 2, cb_send = 18, cb_pair1 = 19, cb_pair2 = 20, cb_out2 = 27; constexpr uint32_t PAGE = 2048; // bf16 tile constexpr uint32_t FPAGE = 4096; // fp32 tile const auto sx = TensorAccessor(x_args, x_addr); const auto so = TensorAccessor(o_args, o_addr); DataflowBuffer dx(cb_x); const uint32_t base = row * WT + col0; { PZONE("RREAD"); for (uint32_t b = 0; b < HALF; b += RB) { dx.reserve_back(RB); const uint32_t lx = dx.get_write_ptr(); for (uint32_t j = 0; j < RB; ++j) { noc_async_read(sx.get_noc_addr(base + b + j), lx + j * PAGE, PAGE); } noc_async_read_barrier(); dx.push_back(RB); } } const uint32_t sem_addr = get_semaphore(SEM_ID); volatile tt_l1_ptr uint32_t* sem_ptr = reinterpret_cast(sem_addr); DataflowBuffer send(cb_send); for (uint32_t k = 1; k <= 2; ++k) { DataflowBuffer pair(k == 1 ? cb_pair1 : cb_pair2); { PZONE("RWAIT"); send.wait_front(1); } PZONE("REXCH"); pair.reserve_back(2); const uint32_t src = send.get_read_ptr(); const uint32_t dst = pair.get_write_ptr(); noc_async_write(src, get_noc_addr(dst), FPAGE); noc_async_write(src, get_noc_addr(px, py, dst + FPAGE), FPAGE); noc_async_write_barrier(); send.pop_front(1); noc_semaphore_inc(get_noc_addr(px, py, sem_addr), 1); noc_semaphore_wait_min(sem_ptr, k); pair.push_back(2); } DataflowBuffer dout(cb_out2); for (uint32_t b = BLK; b < HALF; b += 2 * BLK) { dout.wait_front(BLK); const uint32_t l = dout.get_read_ptr(); for (uint32_t j = 0; j < BLK; ++j) { noc_async_write(l + j * PAGE, so.get_noc_addr(base + b + j), PAGE); } noc_async_writes_flushed(); dout.pop_front(BLK); } noc_async_write_barrier(); noc_async_atomic_barrier(); *sem_ptr = 0; }