// SPDX-License-Identifier: Apache-2.0 // Fused fp32 LayerNorm (tt/ln_kernel.py), compute: the exact op sequence of tt/layers.py layer_norm_fp32 (stock // ttnn programs: mean -> sub -> mul -> mean -> add eps -> rsqrt -> mul -> mul gamma -> add beta), with the same SFPU // LLK calls in the same order, so the result is meant to be bit-identical to the 9-program decomposition: // - ttnn.mean (accurate fp32 path, reduce_op.cpp use_sfpu_fp32_mean): copy tile 0, add_binary_tile fold of tiles // 1..Wt-1 in order, sfpu_reduce, mul_unary_tile(1/W) (the AVG post-mul); // - binary_ng fp32 ops: sub / add as sub_binary_tile / add_binary_tile, mul as mul_binary_tile, // the broadcast operand built by the dataflow (column / row fill), lhs in DST 0; // - ttnn.rsqrt(fast_and_approximate_mode=False): rsqrt_tile (math_approx_mode false). // Everything stays fp32 (UnpackToDestFp32 copies, fp32 DST, fp32 CBs): every intermediate the stock graph packs to // an fp32 DRAM tensor round-trips losslessly. // // With has_res (LN_RESID, the residual add of the stream fused in front): h = x + res (* rgate), the stock // `ttnn.add(x, ttnn.multiply(res, rgate))` / `ttnn.add(x, res)` as mul_binary_tile / add_binary_tile, packed to // c_9 (the LayerNorm input) and to c_17 (written out as the new stream when write_h). // // 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, [9] sfpu_bcast (the statistics broadcast over the columns by the SFPU // reduce itself, ln32_sfpu.h, and packed straight to c_5; else column 0 packed to c_4 and broadcast by the // writer), [8] lean (one copy_init / SFPU-binary init per phase instead of one per // tile: every CB copied from is an fp32 UnpackToDestFp32 CB, and add / sub / mul share one init). // [10] res_t (LN_TR: the residual tiles arrive transposed and are transposed back with transpose_tile, the // stock ttnn.transpose LLK, exact with the fp32 unpack-to-dest; no rgate), [11] out_t (the y tiles are // transposed the same way, via c_18, before they are packed for the writer). // Per-core RT args: [n_rows]. #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/tile_move_copy.h" #include "api/compute/transpose.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 lean = get_compile_time_arg_val(8) != 0; constexpr bool sfpu_bcast = get_compile_time_arg_val(9) != 0; constexpr uint32_t res_t = get_compile_time_arg_val(10); constexpr uint32_t out_t = get_compile_time_arg_val(11); // per-tile (re-)inits, skipped in lean mode (done once at the start of the phase instead) ALWI void ci(uint32_t cb) { if constexpr (!lean) { copy_init(cb); } } #define BIN_INIT(op) \ do { \ if constexpr (!lean) { \ op##_binary_tile_init(); \ } \ } while (0) 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_out = 16, cb_hout = 17, cb_yt = 18; constexpr uint32_t cb_src = has_res ? cb_h : cb_x; constexpr auto RNE = ckernel::DstRoundingMode::NearestEven; // row mean of the Wt tiles of `cb` (already waited) into DST 0; `square` folds xc * xc instead of x template ALWI void row_sum_to_dst0(uint32_t cb) { copy_init(cb); if constexpr (lean) { add_binary_tile_init(); } if constexpr (square) { copy_tile(cb, 0, 0); copy_tile(cb, 0, 1); BIN_INIT(mul); mul_binary_tile(0, 1, 0); } else { copy_tile(cb, 0, 0); } for (uint32_t w = 1; w < Wt; ++w) { ci(cb); copy_tile(cb, w, 1); if constexpr (square) { copy_tile(cb, w, 2); BIN_INIT(mul); mul_binary_tile(1, 2, 1); } BIN_INIT(add); 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_stat() { constexpr uint32_t cb = sfpu_bcast ? cb_bc : cb_stat; cb_reserve_back(cb, 1); tile_regs_commit(); tile_regs_wait(); pack_tile(0, cb); tile_regs_release(); cb_push_back(cb, 1); } void kernel_main() { const uint32_t n_rows = get_arg_val(0); if (n_rows == 0) { return; } compute_kernel_hw_startup(cb_x, cb_out); // the constant tiles are waited for where they are first used (the reader sends x row 0 first) for (uint32_t r = 0; r < n_rows; ++r) { cb_wait_front(cb_x, Wt); if constexpr (has_res) { // h = x + res (* rgate) cb_wait_front(cb_r, Wt); if constexpr (has_rgate) { cb_wait_front(cb_rg, Wt); } if constexpr (lean) { copy_init(cb_r); add_binary_tile_init(); } for (uint32_t w = 0; w < Wt; ++w) { tile_regs_acquire(); if constexpr (res_t) { transpose_init(cb_r); transpose_tile(cb_r, w, 0); copy_init(cb_x); add_binary_tile_init(); } else { ci(cb_r); copy_tile(cb_r, w, 0); } if constexpr (has_rgate) { ci(cb_rg); copy_tile(cb_rg, w, 1); BIN_INIT(mul); mul_binary_tile(0, 1, 0); } ci(cb_x); copy_tile(cb_x, w, 1); BIN_INIT(add); add_binary_tile(1, 0, 1); cb_reserve_back(cb_h, 1); if constexpr (write_h) { cb_reserve_back(cb_hout, 1); } tile_regs_commit(); tile_regs_wait(); pack_tile(1, cb_h); if constexpr (write_h) { pack_tile(1, cb_hout); } tile_regs_release(); cb_push_back(cb_h, 1); if constexpr (write_h) { cb_push_back(cb_hout, 1); } } cb_pop_front(cb_r, Wt); cb_pop_front(cb_x, Wt); cb_wait_front(cb_h, Wt); } // mean tile_regs_acquire(); row_sum_to_dst0(cb_src); pack_stat(); // xc = x - mean cb_wait_front(cb_bc, 1); if constexpr (lean) { copy_init(cb_src); sub_binary_tile_init(); } for (uint32_t w = 0; w < Wt; ++w) { tile_regs_acquire(); ci(cb_src); copy_tile(cb_src, w, 0); ci(cb_bc); copy_tile(cb_bc, 0, 1); BIN_INIT(sub); sub_binary_tile(0, 1, 0); cb_reserve_back(cb_xc, 1); tile_regs_commit(); tile_regs_wait(); pack_tile(0, cb_xc); tile_regs_release(); cb_push_back(cb_xc, 1); } cb_pop_front(cb_bc, 1); cb_pop_front(cb_src, Wt); // rstd = rsqrt(mean(xc * xc) + eps) cb_wait_front(cb_xc, Wt); tile_regs_acquire(); row_sum_to_dst0(cb_xc); 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_stat(); // y = xc * rstd (* gamma) (+ beta) cb_wait_front(cb_bc, 1); if constexpr (has_gamma) { cb_wait_front(cb_g, Wt); } if constexpr (has_beta) { cb_wait_front(cb_b, Wt); } if constexpr (lean) { copy_init(cb_xc); mul_binary_tile_init(); } for (uint32_t w = 0; w < Wt; ++w) { tile_regs_acquire(); ci(cb_xc); copy_tile(cb_xc, w, 0); ci(cb_bc); copy_tile(cb_bc, 0, 1); BIN_INIT(mul); mul_binary_tile(0, 1, 0); if constexpr (has_gamma) { ci(cb_g); copy_tile(cb_g, w, 1); BIN_INIT(mul); mul_binary_tile(0, 1, 0); } if constexpr (has_beta) { ci(cb_b); copy_tile(cb_b, w, 1); BIN_INIT(add); add_binary_tile(0, 1, 0); } if constexpr (out_t) { cb_reserve_back(cb_yt, 1); tile_regs_commit(); tile_regs_wait(); pack_tile(0, cb_yt); tile_regs_release(); cb_push_back(cb_yt, 1); cb_wait_front(cb_yt, 1); tile_regs_acquire(); transpose_init(cb_yt); transpose_tile(cb_yt, 0, 0); cb_reserve_back(cb_out, 1); tile_regs_commit(); tile_regs_wait(); pack_tile(0, cb_out); tile_regs_release(); cb_push_back(cb_out, 1); cb_pop_front(cb_yt, 1); copy_init(cb_xc); mul_binary_tile_init(); } else { cb_reserve_back(cb_out, 1); tile_regs_commit(); tile_regs_wait(); pack_tile(0, cb_out); tile_regs_release(); cb_push_back(cb_out, 1); } } cb_pop_front(cb_bc, 1); cb_pop_front(cb_xc, Wt); } }