// SPDX-License-Identifier: Apache-2.0 // Fused residual add + LayerNorm on 2 cores per tile row (moge-2 opt kernel, round 7): compute. // Per core (HALF tiles of one tile row, the partner core holds the other half): // z = x + y (bf16, packed to c_16 for the LN and to c_17 for the writer) // partial mean = rowsum(z) / W -> c_18 -> exchange -> c_19 [own, partner] -> mean (fp32 SFPU add) // xmm = z - mean (fp32 c_22); partial var = rowsum(xmm^2) / W -> c_18 -> exchange -> c_20 -> var // rstd = rsqrt(var + eps); out = xmm * rstd (bf16 c_26) // Mirrors ttnn's interleaved LayerNorm kernel (layernorm.cpp: two-pass variance, fp32 intermediates, // reduce scaler 1 + scale by 1/W, eps truncated to bf16, rsqrt non-legacy); the halves' partial sums are // added in fp32 on the SFPU (exact unpack to DEST). // Compile args: HALF, BLK, W, EPS_BITS (fp32 bits, bf16-truncated). #include #define BCAST_LLKOP EltwiseBinaryType::ELWMUL #define BCAST_DIM BroadcastType::COL #include "api/compute/compute_kernel_api.h" #include "api/compute/bcast.h" #include "api/compute/eltwise_binary.h" #include "api/compute/eltwise_binary_sfpu.h" #include "api/compute/eltwise_unary/sfpu_split_includes.h" #include "api/compute/eltwise_unary/rsqrt.h" #include "api/compute/eltwise_unary/binop_with_scalar.h" #include "api/compute/tile_move_copy.h" #include "api/compute/eltwise_unary/eltwise_unary.h" #include "api/dataflow/dataflow_buffer.h" #include "ttnn/operations/normalization/kernel_util/compute/numeric.h" namespace numeric = norm::kernel_util::compute::numeric; namespace policies = norm::kernel_util::compute::policies; #ifdef ADDLN_PROF #include "tools/profiler/kernel_profiler.hpp" #define PZONE(n) DeviceZoneScopedN(n) #else #define PZONE(n) #endif void kernel_main() { constexpr uint32_t HALF = get_compile_time_arg_val(0); constexpr uint32_t BLK = get_compile_time_arg_val(1); constexpr uint32_t W = get_compile_time_arg_val(2); constexpr uint32_t EPS_BITS = get_compile_time_arg_val(3); constexpr uint32_t cb_x = 0, cb_y = 1, cb_scaler = 2, cb_z = 16, cb_zout = 17, cb_send = 18, cb_pair1 = 19, cb_pair2 = 20, cb_mean = 21, cb_xmm = 22, cb_xmm2 = 23, cb_rstd = 25, cb_out = 26, cb_out2 = 27; DataflowBuffer dx(cb_x), dy(cb_y), dscaler(cb_scaler), dz(cb_z), dzout(cb_zout), dsend(cb_send), dpair1(cb_pair1), dpair2(cb_pair2), dmean(cb_mean), dxmm(cb_xmm), dxmm2(cb_xmm2), drstd(cb_rstd), dout(cb_out); compute_kernel_hw_startup(cb_x, cb_y, cb_z); // mean = own + partner (fp32, SFPU) auto pair_sum = [&](DataflowBuffer& dpair, uint32_t cb_pair, uint32_t cb_res, bool add_eps) { dpair.wait_front(2); reconfig_data_format_srca(cb_pair); copy_tile_to_dst_init_short(cb_pair); tile_regs_acquire(); copy_tile(cb_pair, 0, 0); copy_tile(cb_pair, 1, 1); add_binary_tile_init(); add_binary_tile(0, 1, 0); if (add_eps) { binop_with_scalar_tile_init(); add_unary_tile(0, EPS_BITS); rsqrt_tile_init(); rsqrt_tile(0); } tile_regs_commit(); dpair.pop_front(2); DataflowBuffer dres(cb_res); dres.reserve_back(1); pack_reconfig_data_format(cb_res); tile_regs_wait(); pack_tile(0, cb_res); tile_regs_release(); dres.push_back(1); }; { PZONE("ZADD"); // ---- z = x + y reconfig_data_format(cb_x, cb_y); pack_reconfig_data_format(cb_z); add_init(cb_x, cb_y); for (uint32_t b = 0; b < HALF; b += BLK) { dx.wait_front(BLK); dy.wait_front(BLK); tile_regs_acquire(); for (uint32_t i = 0; i < BLK; ++i) { add_tiles(cb_x, cb_y, i, i, i); } tile_regs_commit(); dx.pop_front(BLK); dy.pop_front(BLK); dz.reserve_back(BLK); dzout.reserve_back(BLK); tile_regs_wait(); for (uint32_t i = 0; i < BLK; ++i) { pack_tile(i, cb_z); } for (uint32_t i = 0; i < BLK; ++i) { pack_tile(i, cb_zout); } tile_regs_release(); dz.push_back(BLK); dzout.push_back(BLK); } } { PZONE("ZMEAN"); // ---- partial mean -> exchange numeric::row_wise_mean(dz, dscaler, dsend, W, HALF, BLK); } { PZONE("ZPAIR1"); pair_sum(dpair1, cb_pair1, cb_mean, false); } { PZONE("ZSUB"); // ---- xmm = z - mean reconfig_data_format(cb_z, cb_mean); pack_reconfig_data_format(cb_xmm); dmean.wait_front(1); sub_bcast_cols_init(cb_z, cb_mean); for (uint32_t b = 0; b < HALF; b += BLK) { tile_regs_acquire(); for (uint32_t i = 0; i < BLK; ++i) { sub_tiles_bcast_cols(cb_z, cb_mean, b + i, 0, i); } tile_regs_commit(); dxmm.reserve_back(BLK); tile_regs_wait(); for (uint32_t i = 0; i < BLK; ++i) { pack_tile(i, cb_xmm); } tile_regs_release(); dxmm.push_back(BLK); } dz.pop_front(HALF); dmean.pop_front(1); } { PZONE("ZSQ"); // ---- xmm^2 -> partial var -> exchange reconfig_data_format(cb_xmm, cb_xmm); pack_reconfig_data_format(cb_xmm2); mul_init(cb_xmm, cb_xmm); for (uint32_t b = 0; b < HALF; b += BLK) { dxmm.wait_front(b + BLK); tile_regs_acquire(); for (uint32_t i = 0; i < BLK; ++i) { mul_tiles(cb_xmm, cb_xmm, b + i, b + i, i); } tile_regs_commit(); dxmm2.reserve_back(BLK); tile_regs_wait(); for (uint32_t i = 0; i < BLK; ++i) { pack_tile(i, cb_xmm2); } tile_regs_release(); dxmm2.push_back(BLK); } } { PZONE("ZVAR"); numeric::row_wise_mean(dxmm2, dscaler, dsend, W, HALF, BLK); } { PZONE("ZPAIR2"); // rstd = rsqrt(own + partner + eps) pair_sum(dpair2, cb_pair2, cb_rstd, true); } { PZONE("ZOUT"); // ---- out = xmm * rstd reconfig_data_format(cb_xmm, cb_rstd); pack_reconfig_data_format(cb_out); drstd.wait_front(1); mul_bcast_cols_init(cb_xmm, cb_rstd); for (uint32_t b = 0; b < HALF; b += BLK) { tile_regs_acquire(); for (uint32_t i = 0; i < BLK; ++i) { mul_tiles_bcast_cols(cb_xmm, cb_rstd, b + i, 0, i); } tile_regs_commit(); // even blocks -> c_26 (BRISC writes them), odd blocks -> c_27 (NCRISC) const uint32_t cbo = ((b / BLK) & 1) ? cb_out2 : cb_out; DataflowBuffer dcbo(cbo); dcbo.reserve_back(BLK); tile_regs_wait(); for (uint32_t i = 0; i < BLK; ++i) { pack_tile(i, cbo); } tile_regs_release(); dcbo.push_back(BLK); } } dxmm.pop_front(HALF); drstd.pop_front(1); dscaler.pop_front(1); }