changh95's picture
tt-model push diffusion-planner-p150 (container)
be62f78 verified
Raw History Blame Contribute Delete
10.7 kB
// 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<SUM, Float32, REDUCE_ROW>, mul_unary_tile(1/W) (the AVG post-mul);
// - binary_ng fp32 ops: sub / add as sub_binary_tile / add_binary_tile<NearestEven>, 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<RsqrtMode::Default> (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 <cstdint>
#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 <bool square>
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<PoolType::SUM, DataFormat::Float32>();
if constexpr (sfpu_bcast) {
ln_row_sum_bcast_tile(0);
} else {
sfpu_reduce<PoolType::SUM, DataFormat::Float32, ReduceDim::REDUCE_ROW>(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<uint32_t>(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<RNE>(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<false>(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<RNE>(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<true>(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<RNE>(0, 1, 0);
rsqrt_tile_init();
rsqrt_tile<RsqrtMode::Default>(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<RNE>(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);
}
}