// SPDX-License-Identifier: Apache-2.0 // Fused fp32 LayerNorm (tt/ln_kernel.py), reader (RISCV_0): the affine rows once, then this core's tile rows of x. // // CT args: [0] Wt, [1] has_gamma, [2] has_beta, [3] eps bits (fp32), [4] per-core RT-arg count (P3: part of the // program hash), [5] has_res, [6] has_rgate, [7] res_t (LN_TR: res is stored transposed per entity, i.e. // [.., E, W, T]: tile (r, w) of res is read from tile (e, w, tr) of that layout, e = r / Tt, tr = r % Tt; // the compute transposes it back), [8] Tt (tile rows per entity), then the TensorAccessorArgs of x, gamma, // beta, res, rgate (x's when absent). // Common RT args: [x_addr, gamma_addr, beta_addr, res_addr, rgate_addr]. Per-core RT args: [row0, n_rows]. // CBs: c_0 x (Wt tiles per row, double-buffered), c_1 gamma rows (Wt tiles, row 0 replicated over the tile), // c_2 beta rows (same), c_3 the eps tile (every element eps), c_7 res (as c_0), c_8 the residual gate rows. #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 uint32_t res_t = get_compile_time_arg_val(7); constexpr uint32_t Tt = get_compile_time_arg_val(8); constexpr auto x_args = TensorAccessorArgs<9>(); 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_r = 7, cb_rg = 8; constexpr uint32_t TB = 4096; // fp32 tile bytes // Replicate row 0 of each of the n fp32 tiles at p over its 32 rows (binary_ng fill_tile_with_first_row: local // NoC doubling), all tiles per doubling step under one barrier. FORCE_INLINE void fill_first_rows(uint32_t p, uint32_t n_tiles) { constexpr uint32_t kFace = 1024, kRow = 64; for (uint32_t n = kRow; n < kFace; n <<= 1) { for (uint32_t t = 0; t < n_tiles; ++t) { const uint32_t f0 = p + t * TB, f1 = f0 + kFace; noc_async_read(get_noc_addr(f0), f0 + n, n); noc_async_read(get_noc_addr(f1), f1 + n, n); } noc_async_read_barrier(); } for (uint32_t t = 0; t < n_tiles; ++t) { const uint32_t f0 = p + t * TB; noc_async_read(get_noc_addr(f0), f0 + 2 * kFace, 2 * kFace); } noc_async_read_barrier(); } template FORCE_INLINE void read_rows(uint32_t cb, const A& acc) { cb_reserve_back(cb, Wt); uint32_t p = get_write_ptr(cb); for (uint32_t w = 0; w < Wt; ++w) { noc_async_read(acc.get_noc_addr(w), p + w * TB, TB); } noc_async_read_barrier(); fill_first_rows(p, Wt); cb_push_back(cb, Wt); } template FORCE_INLINE void read_x_row(uint32_t r, const A& x, const R& res) { cb_reserve_back(cb_x, Wt); uint32_t p = get_write_ptr(cb_x); for (uint32_t w = 0; w < Wt; ++w) { noc_async_read(x.get_noc_addr(r * Wt + w), p + w * TB, TB); } if constexpr (has_res) { cb_reserve_back(cb_r, Wt); uint32_t q = get_write_ptr(cb_r); for (uint32_t w = 0; w < Wt; ++w) { const uint32_t ri = res_t ? (r / Tt) * (Wt * Tt) + w * Tt + (r % Tt) : r * Wt + w; noc_async_read(res.get_noc_addr(ri), q + w * TB, TB); } noc_async_read_barrier(); cb_push_back(cb_r, Wt); } else { noc_async_read_barrier(); } cb_push_back(cb_x, Wt); } 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); if (n_rows == 0) { return; } const auto x = TensorAccessor(x_args, x_addr, TB); const auto res = TensorAccessor(r_args, r_addr, TB); // the first x row goes first (the compute starts on it); the affine rows are needed only by its last phase read_x_row(row0, x, res); // eps tile 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_rgate) { read_rows(cb_rg, TensorAccessor(rg_args, rg_addr, TB)); } if constexpr (has_gamma) { read_rows(cb_g, TensorAccessor(g_args, g_addr, TB)); } if constexpr (has_beta) { read_rows(cb_b, TensorAccessor(b_args, b_addr, TB)); } for (uint32_t r = row0 + 1; r < row0 + n_rows; ++r) { read_x_row(r, x, res); } }