moge-2-p150 / code /tt_moge /kernels /addln_compute.cpp
changh95's picture
Optimized build (2026-10-03): model call 19.0 ms, trace 17.7 ms
13b4736 verified
Raw History Blame Contribute Delete
7.22 kB
// 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 <cstdint>
#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<false>();
rsqrt_tile<false>(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<PoolType::SUM, ReduceDim::REDUCE_ROW, true, policies::FullBlockWithoutPopPolicy,
policies::WaitAtEndPolicy::NO_WAIT>(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<PoolType::SUM, ReduceDim::REDUCE_ROW, true, policies::FullBlockWithPopPolicy,
policies::WaitAtEndPolicy::NO_WAIT>(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);
}