changh95's picture
tt-model push diffusion-planner-p150 (container)
be62f78 verified
Raw History Blame Contribute Delete
8.21 kB
// SPDX-License-Identifier: Apache-2.0
// Split-row fused fp32 LayerNorm (tt/ln_kernel.py, LN_SPLIT), compute: member j of a tile row owns column tile j.
// The same LLK sequence as ln32_compute.cpp (and the stock 9-program decomposition): the folds over the Wt tiles run
// on the root, in tile order, over the gathered tiles (c_10), so every statistic is bit for bit the stock one.
// h = x (+ res (* rgate)) -> c_11 (to the root's gather) [+ c_17 for write_h]
// root: mean = fold(c_10) / W -> c_4 (the writer broadcasts it into every member's c_5 page 0)
// xc = h - mean; sq = xc * xc -> c_6 (kept), sq -> c_11
// root: rstd = rsqrt(fold(c_10) / W + eps) -> c_4 (broadcast into c_5 page 1)
// y = xc * rstd (* gamma) (+ beta) -> c_16
// 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, [8] sfpu_bcast (the root's statistics broadcast over the columns by the
// SFPU reduce, ln32_sfpu.h; else column 0 and the writer fills), [9] kcat_out (KCAT_EMIT: y is emitted
// as the split operand of the next K-concatenated linear, y_hi = bf16(y) -> c_16 and y_lo = y - y_hi ->
// c_19, the LLK calls of kernels/kcat_compute.cpp on the packed fp32 y tile c_18).
// Per-core RT args: [n_rows, is_root].
#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/eltwise_unary/typecast.h"
#include "api/compute/tile_move_copy.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 sfpu_bcast = get_compile_time_arg_val(8) != 0;
constexpr uint32_t kcat_out = get_compile_time_arg_val(9);
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_gat = 10, cb_snd = 11, cb_out = 16, cb_hout = 17, cb_y = 18, cb_lo = 19;
constexpr uint32_t cb_src = has_res ? cb_h : cb_x;
constexpr auto RNE = ckernel::DstRoundingMode::NearestEven;
// fold of the Wt gathered tiles (in order) -> row sum * 1/W in DST 0
ALWI void fold_gathered() {
cb_wait_front(cb_gat, Wt);
copy_init(cb_gat);
add_binary_tile_init();
copy_tile(cb_gat, 0, 0);
for (uint32_t w = 1; w < Wt; ++w) {
copy_tile(cb_gat, w, 1);
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_one(uint32_t dst, uint32_t cb) {
cb_reserve_back(cb, 1);
tile_regs_commit();
tile_regs_wait();
pack_tile(dst, cb);
tile_regs_release();
cb_push_back(cb, 1);
}
void kernel_main() {
const uint32_t n_rows = get_arg_val<uint32_t>(0);
const bool root = get_arg_val<uint32_t>(1) != 0;
if (n_rows == 0) {
return;
}
compute_kernel_hw_startup(cb_x, cb_out);
for (uint32_t r = 0; r < n_rows; ++r) {
// h (this member's tile) -> c_11 for the gather (and c_9 / c_17)
cb_wait_front(cb_x, 1);
tile_regs_acquire();
if constexpr (has_res) {
cb_wait_front(cb_r, 1);
if constexpr (has_rgate) {
cb_wait_front(cb_rg, 1);
}
copy_init(cb_r);
copy_tile(cb_r, 0, 0);
if constexpr (has_rgate) {
copy_tile(cb_rg, 0, 1);
mul_binary_tile_init();
mul_binary_tile(0, 1, 0);
}
copy_tile(cb_x, 0, 1);
add_binary_tile_init();
add_binary_tile<RNE>(1, 0, 1);
cb_reserve_back(cb_h, 1);
cb_reserve_back(cb_snd, 1);
if constexpr (write_h) {
cb_reserve_back(cb_hout, 1);
}
tile_regs_commit();
tile_regs_wait();
pack_tile(1, cb_h);
pack_tile(1, cb_snd);
if constexpr (write_h) {
pack_tile(1, cb_hout);
}
tile_regs_release();
cb_push_back(cb_h, 1);
cb_push_back(cb_snd, 1);
if constexpr (write_h) {
cb_push_back(cb_hout, 1);
}
cb_pop_front(cb_r, 1);
cb_pop_front(cb_x, 1);
cb_wait_front(cb_h, 1);
} else {
copy_init(cb_x);
copy_tile(cb_x, 0, 0);
pack_one(0, cb_snd);
}
// root: mean
if (root) {
tile_regs_acquire();
fold_gathered();
pack_one(0, cb_stat);
cb_pop_front(cb_gat, Wt);
}
// xc = h - mean -> c_6; sq = xc * xc -> c_11
cb_wait_front(cb_bc, 1);
tile_regs_acquire();
copy_init(cb_src);
copy_tile(cb_src, 0, 0);
copy_tile(cb_bc, 0, 1);
sub_binary_tile_init();
sub_binary_tile<RNE>(0, 1, 0);
pack_one(0, cb_xc);
cb_pop_front(cb_bc, 1);
cb_pop_front(cb_src, 1);
cb_wait_front(cb_xc, 1);
tile_regs_acquire();
copy_init(cb_xc);
copy_tile(cb_xc, 0, 0);
copy_tile(cb_xc, 0, 1);
mul_binary_tile_init();
mul_binary_tile(0, 1, 0);
pack_one(0, cb_snd);
// root: rstd
if (root) {
tile_regs_acquire();
fold_gathered();
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_one(0, cb_stat);
cb_pop_front(cb_gat, Wt);
}
// y = xc * rstd (* gamma) (+ beta)
cb_wait_front(cb_bc, 1);
tile_regs_acquire();
copy_init(cb_xc);
copy_tile(cb_xc, 0, 0);
copy_tile(cb_bc, 0, 1);
mul_binary_tile_init();
mul_binary_tile(0, 1, 0);
if constexpr (has_gamma) {
cb_wait_front(cb_g, 1);
copy_tile(cb_g, 0, 1);
mul_binary_tile(0, 1, 0);
}
if constexpr (has_beta) {
cb_wait_front(cb_b, 1);
copy_tile(cb_b, 0, 1);
add_binary_tile_init();
add_binary_tile<RNE>(0, 1, 0);
}
if constexpr (kcat_out) {
pack_one(0, cb_y);
cb_pop_front(cb_bc, 1);
cb_pop_front(cb_xc, 1);
cb_wait_front(cb_y, 1);
tile_regs_acquire();
copy_init(cb_y);
copy_tile(cb_y, 0, 0);
copy_tile(cb_y, 0, 1);
typecast_tile_init<(uint32_t)DataFormat::Float32, (uint32_t)DataFormat::Float16_b>();
typecast_tile<(uint32_t)DataFormat::Float32, (uint32_t)DataFormat::Float16_b>(0);
sub_binary_tile_init();
sub_binary_tile<RNE>(1, 0, 1);
cb_reserve_back(cb_out, 1);
cb_reserve_back(cb_lo, 1);
tile_regs_commit();
tile_regs_wait();
pack_tile(0, cb_out);
pack_tile(1, cb_lo);
tile_regs_release();
cb_push_back(cb_out, 1);
cb_push_back(cb_lo, 1);
cb_pop_front(cb_y, 1);
} else {
pack_one(0, cb_out);
cb_pop_front(cb_bc, 1);
cb_pop_front(cb_xc, 1);
}
}
}