// SPDX-License-Identifier: Apache-2.0 // Split-row fused fp32 LayerNorm (tt/ln_kernel.py, LN_SPLIT), reader (RISCV_0). One tile row is spread over Wt // cores (member j owns column tile j); member 0 (the root) gathers the Wt tiles that the stock reduction folds in // order and broadcasts the statistics back, so the arithmetic stays the stock decomposition's (bit-identical). // Per row this kernel: reads x (r, j) (+ res (r, j)); on the root, waits for the Wt gathered tiles (semaphore 0, // monotonic count) and hands them to the compute (c_10); then waits for the two statistic broadcasts (semaphore 1, // monotonic count) and hands each to the compute (c_5, 2 pages: mean, rstd). // CT args: [0] Wt, [1] has_gamma, [2] has_beta, [3] eps bits, [4] per-core RT-arg count (P3), [5] has_res, // [6] has_rgate, then the TensorAccessorArgs of x, gamma, beta, res, rgate. // Common RT args: [x_addr, gamma_addr, beta_addr, res_addr, rgate_addr]. // Per-core RT args: [row0, n_rows, row_stride, j, root_x, root_y, (member x, y) x Wt]. #include #include "api/dataflow/dataflow_api.h" constexpr uint32_t Wt = get_compile_time_arg_val(0); constexpr uint32_t has_gamma = get_compile_time_arg_val(1); constexpr uint32_t has_beta = get_compile_time_arg_val(2); constexpr uint32_t eps_bits = get_compile_time_arg_val(3); constexpr uint32_t has_res = get_compile_time_arg_val(5); constexpr uint32_t has_rgate = get_compile_time_arg_val(6); constexpr auto x_args = TensorAccessorArgs<7>(); constexpr auto g_args = TensorAccessorArgs(); constexpr auto b_args = TensorAccessorArgs(); constexpr auto r_args = TensorAccessorArgs(); constexpr auto rg_args = TensorAccessorArgs(); constexpr uint32_t cb_x = 0, cb_g = 1, cb_b = 2, cb_eps = 3, cb_bc = 5, cb_r = 7, cb_rg = 8, cb_gat = 10; constexpr uint32_t TB = 4096; FORCE_INLINE void fill_first_row(uint32_t p) { constexpr uint32_t kFace = 1024, kRow = 64; const uint32_t f0 = p, f1 = p + kFace; for (uint32_t n = kRow; n < kFace; n <<= 1) { noc_async_read(get_noc_addr(f0), f0 + n, n); noc_async_read(get_noc_addr(f1), f1 + n, n); noc_async_read_barrier(); } noc_async_read(get_noc_addr(f0), f0 + 2 * kFace, 2 * kFace); noc_async_read_barrier(); } template FORCE_INLINE void read_row_tile(uint32_t cb, const A& acc, uint32_t j) { cb_reserve_back(cb, 1); const uint32_t p = get_write_ptr(cb); noc_async_read(acc.get_noc_addr(j), p, TB); noc_async_read_barrier(); fill_first_row(p); cb_push_back(cb, 1); } void kernel_main() { const uint32_t x_addr = get_common_arg_val(0); const uint32_t g_addr = get_common_arg_val(1); const uint32_t b_addr = get_common_arg_val(2); const uint32_t r_addr = get_common_arg_val(3); const uint32_t rg_addr = get_common_arg_val(4); const uint32_t row0 = get_arg_val(0); const uint32_t n_rows = get_arg_val(1); const uint32_t stride = get_arg_val(2); const uint32_t j = get_arg_val(3); if (n_rows == 0) { return; } const bool root = j == 0; const auto x = TensorAccessor(x_args, x_addr, TB); const auto res = TensorAccessor(r_args, r_addr, TB); volatile tt_l1_ptr uint32_t* sem_g = reinterpret_cast(get_semaphore(0)); volatile tt_l1_ptr uint32_t* sem_b = reinterpret_cast(get_semaphore(1)); uint32_t n_gat = 0, n_bc = 0; auto read_x = [&](uint32_t r) { cb_reserve_back(cb_x, 1); noc_async_read(x.get_noc_addr(r * Wt + j), get_write_ptr(cb_x), TB); if constexpr (has_res) { cb_reserve_back(cb_r, 1); noc_async_read(res.get_noc_addr(r * Wt + j), get_write_ptr(cb_r), TB); noc_async_read_barrier(); cb_push_back(cb_r, 1); } else { noc_async_read_barrier(); } cb_push_back(cb_x, 1); }; auto gather = [&]() { cb_reserve_back(cb_gat, Wt); n_gat += Wt; noc_semaphore_wait_min(sem_g, n_gat); cb_push_back(cb_gat, Wt); }; auto bcast_in = [&]() { cb_reserve_back(cb_bc, 1); n_bc += 1; noc_semaphore_wait_min(sem_b, n_bc); cb_push_back(cb_bc, 1); }; read_x(row0); if constexpr (has_rgate) { read_row_tile(cb_rg, TensorAccessor(rg_args, rg_addr, TB), j); } cb_reserve_back(cb_eps, 1); { auto* e = reinterpret_cast(get_write_ptr(cb_eps)); for (uint32_t i = 0; i < 1024; ++i) { e[i] = eps_bits; } } cb_push_back(cb_eps, 1); if constexpr (has_gamma) { read_row_tile(cb_g, TensorAccessor(g_args, g_addr, TB), j); } if constexpr (has_beta) { read_row_tile(cb_b, TensorAccessor(b_args, b_addr, TB), j); } for (uint32_t i = 0; i < n_rows; ++i) { if (i > 0) { read_x(row0 + i * stride); } if (root) { gather(); } bcast_in(); if (root) { gather(); } bcast_in(); } }