// SPDX-License-Identifier: Apache-2.0 // Split-row fused fp32 LayerNorm (tt/ln_kernel.py, LN_SPLIT), compute: member j of a tile row owns column tile j. // The same LLK sequence as ln32_compute.cpp (and the stock 9-program decomposition): the folds over the Wt tiles run // on the root, in tile order, over the gathered tiles (c_10), so every statistic is bit for bit the stock one. // h = x (+ res (* rgate)) -> c_11 (to the root's gather) [+ c_17 for write_h] // root: mean = fold(c_10) / W -> c_4 (the writer broadcasts it into every member's c_5 page 0) // xc = h - mean; sq = xc * xc -> c_6 (kept), sq -> c_11 // root: rstd = rsqrt(fold(c_10) / W + eps) -> c_4 (broadcast into c_5 page 1) // y = xc * rstd (* gamma) (+ beta) -> c_16 // CT args: [0] Wt, [1] has_gamma, [2] has_beta, [3] 1/W bits, [4] per-core RT-arg count (P3), [5] has_res, // [6] has_rgate, [7] write_h, [8] sfpu_bcast (the root's statistics broadcast over the columns by the // SFPU reduce, ln32_sfpu.h; else column 0 and the writer fills), [9] kcat_out (KCAT_EMIT: y is emitted // as the split operand of the next K-concatenated linear, y_hi = bf16(y) -> c_16 and y_lo = y - y_hi -> // c_19, the LLK calls of kernels/kcat_compute.cpp on the packed fp32 y tile c_18). // Per-core RT args: [n_rows, is_root]. #include #include "api/compute/common.h" #include "api/compute/compute_kernel_api.h" #include "api/compute/eltwise_binary_sfpu.h" #include "api/compute/eltwise_unary/binop_with_scalar.h" #include "api/compute/eltwise_unary/eltwise_unary.h" #include "api/compute/eltwise_unary/rsqrt.h" #include "api/compute/eltwise_unary/typecast.h" #include "api/compute/tile_move_copy.h" #include "ln32_sfpu.h" using namespace ckernel; 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 inv_w_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 write_h = get_compile_time_arg_val(7); constexpr bool sfpu_bcast = get_compile_time_arg_val(8) != 0; constexpr uint32_t kcat_out = get_compile_time_arg_val(9); constexpr uint32_t cb_x = 0, cb_g = 1, cb_b = 2, cb_eps = 3, cb_stat = 4, cb_bc = 5, cb_xc = 6, cb_r = 7, cb_rg = 8, cb_h = 9, cb_gat = 10, cb_snd = 11, cb_out = 16, cb_hout = 17, cb_y = 18, cb_lo = 19; constexpr uint32_t cb_src = has_res ? cb_h : cb_x; constexpr auto RNE = ckernel::DstRoundingMode::NearestEven; // fold of the Wt gathered tiles (in order) -> row sum * 1/W in DST 0 ALWI void fold_gathered() { cb_wait_front(cb_gat, Wt); copy_init(cb_gat); add_binary_tile_init(); copy_tile(cb_gat, 0, 0); for (uint32_t w = 1; w < Wt; ++w) { copy_tile(cb_gat, w, 1); add_binary_tile(0, 1, 0); } sfpu_reduce_init(); if constexpr (sfpu_bcast) { ln_row_sum_bcast_tile(0); } else { sfpu_reduce(0, 1, 1); } binop_with_scalar_tile_init(); mul_unary_tile(0, inv_w_bits); } ALWI void pack_one(uint32_t dst, uint32_t cb) { cb_reserve_back(cb, 1); tile_regs_commit(); tile_regs_wait(); pack_tile(dst, cb); tile_regs_release(); cb_push_back(cb, 1); } void kernel_main() { const uint32_t n_rows = get_arg_val(0); const bool root = get_arg_val(1) != 0; if (n_rows == 0) { return; } compute_kernel_hw_startup(cb_x, cb_out); for (uint32_t r = 0; r < n_rows; ++r) { // h (this member's tile) -> c_11 for the gather (and c_9 / c_17) cb_wait_front(cb_x, 1); tile_regs_acquire(); if constexpr (has_res) { cb_wait_front(cb_r, 1); if constexpr (has_rgate) { cb_wait_front(cb_rg, 1); } copy_init(cb_r); copy_tile(cb_r, 0, 0); if constexpr (has_rgate) { copy_tile(cb_rg, 0, 1); mul_binary_tile_init(); mul_binary_tile(0, 1, 0); } copy_tile(cb_x, 0, 1); add_binary_tile_init(); add_binary_tile(1, 0, 1); cb_reserve_back(cb_h, 1); cb_reserve_back(cb_snd, 1); if constexpr (write_h) { cb_reserve_back(cb_hout, 1); } tile_regs_commit(); tile_regs_wait(); pack_tile(1, cb_h); pack_tile(1, cb_snd); if constexpr (write_h) { pack_tile(1, cb_hout); } tile_regs_release(); cb_push_back(cb_h, 1); cb_push_back(cb_snd, 1); if constexpr (write_h) { cb_push_back(cb_hout, 1); } cb_pop_front(cb_r, 1); cb_pop_front(cb_x, 1); cb_wait_front(cb_h, 1); } else { copy_init(cb_x); copy_tile(cb_x, 0, 0); pack_one(0, cb_snd); } // root: mean if (root) { tile_regs_acquire(); fold_gathered(); pack_one(0, cb_stat); cb_pop_front(cb_gat, Wt); } // xc = h - mean -> c_6; sq = xc * xc -> c_11 cb_wait_front(cb_bc, 1); tile_regs_acquire(); copy_init(cb_src); copy_tile(cb_src, 0, 0); copy_tile(cb_bc, 0, 1); sub_binary_tile_init(); sub_binary_tile(0, 1, 0); pack_one(0, cb_xc); cb_pop_front(cb_bc, 1); cb_pop_front(cb_src, 1); cb_wait_front(cb_xc, 1); tile_regs_acquire(); copy_init(cb_xc); copy_tile(cb_xc, 0, 0); copy_tile(cb_xc, 0, 1); mul_binary_tile_init(); mul_binary_tile(0, 1, 0); pack_one(0, cb_snd); // root: rstd if (root) { tile_regs_acquire(); fold_gathered(); cb_wait_front(cb_eps, 1); copy_init(cb_eps); copy_tile(cb_eps, 0, 1); add_binary_tile_init(); add_binary_tile(0, 1, 0); rsqrt_tile_init(); rsqrt_tile(0); pack_one(0, cb_stat); cb_pop_front(cb_gat, Wt); } // y = xc * rstd (* gamma) (+ beta) cb_wait_front(cb_bc, 1); tile_regs_acquire(); copy_init(cb_xc); copy_tile(cb_xc, 0, 0); copy_tile(cb_bc, 0, 1); mul_binary_tile_init(); mul_binary_tile(0, 1, 0); if constexpr (has_gamma) { cb_wait_front(cb_g, 1); copy_tile(cb_g, 0, 1); mul_binary_tile(0, 1, 0); } if constexpr (has_beta) { cb_wait_front(cb_b, 1); copy_tile(cb_b, 0, 1); add_binary_tile_init(); add_binary_tile(0, 1, 0); } if constexpr (kcat_out) { pack_one(0, cb_y); cb_pop_front(cb_bc, 1); cb_pop_front(cb_xc, 1); cb_wait_front(cb_y, 1); tile_regs_acquire(); copy_init(cb_y); copy_tile(cb_y, 0, 0); copy_tile(cb_y, 0, 1); typecast_tile_init<(uint32_t)DataFormat::Float32, (uint32_t)DataFormat::Float16_b>(); typecast_tile<(uint32_t)DataFormat::Float32, (uint32_t)DataFormat::Float16_b>(0); sub_binary_tile_init(); sub_binary_tile(1, 0, 1); cb_reserve_back(cb_out, 1); cb_reserve_back(cb_lo, 1); tile_regs_commit(); tile_regs_wait(); pack_tile(0, cb_out); pack_tile(1, cb_lo); tile_regs_release(); cb_push_back(cb_out, 1); cb_push_back(cb_lo, 1); cb_pop_front(cb_y, 1); } else { pack_one(0, cb_out); cb_pop_front(cb_bc, 1); cb_pop_front(cb_xc, 1); } } }