File size: 7,220 Bytes
13b4736 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 | // 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);
}
|