changh95's picture
tt-model push diffusion-planner-p150 (container)
be62f78 verified
Raw History Blame Contribute Delete
5.34 kB
// SPDX-License-Identifier: Apache-2.0
// Split-row fused fp32 LayerNorm (tt/ln_kernel.py, LN_SPLIT), reader (RISCV_0). One tile row is spread over Wt
// cores (member j owns column tile j); member 0 (the root) gathers the Wt tiles that the stock reduction folds in
// order and broadcasts the statistics back, so the arithmetic stays the stock decomposition's (bit-identical).
// Per row this kernel: reads x (r, j) (+ res (r, j)); on the root, waits for the Wt gathered tiles (semaphore 0,
// monotonic count) and hands them to the compute (c_10); then waits for the two statistic broadcasts (semaphore 1,
// monotonic count) and hands each to the compute (c_5, 2 pages: mean, rstd).
// CT args: [0] Wt, [1] has_gamma, [2] has_beta, [3] eps bits, [4] per-core RT-arg count (P3), [5] has_res,
// [6] has_rgate, then the TensorAccessorArgs of x, gamma, beta, res, rgate.
// Common RT args: [x_addr, gamma_addr, beta_addr, res_addr, rgate_addr].
// Per-core RT args: [row0, n_rows, row_stride, j, root_x, root_y, (member x, y) x Wt].
#include <cstdint>
#include "api/dataflow/dataflow_api.h"
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 eps_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 auto x_args = TensorAccessorArgs<7>();
constexpr auto g_args = TensorAccessorArgs<x_args.next_compile_time_args_offset()>();
constexpr auto b_args = TensorAccessorArgs<g_args.next_compile_time_args_offset()>();
constexpr auto r_args = TensorAccessorArgs<b_args.next_compile_time_args_offset()>();
constexpr auto rg_args = TensorAccessorArgs<r_args.next_compile_time_args_offset()>();
constexpr uint32_t cb_x = 0, cb_g = 1, cb_b = 2, cb_eps = 3, cb_bc = 5, cb_r = 7, cb_rg = 8, cb_gat = 10;
constexpr uint32_t TB = 4096;
FORCE_INLINE void fill_first_row(uint32_t p) {
constexpr uint32_t kFace = 1024, kRow = 64;
const uint32_t f0 = p, f1 = p + kFace;
for (uint32_t n = kRow; n < kFace; n <<= 1) {
noc_async_read(get_noc_addr(f0), f0 + n, n);
noc_async_read(get_noc_addr(f1), f1 + n, n);
noc_async_read_barrier();
}
noc_async_read(get_noc_addr(f0), f0 + 2 * kFace, 2 * kFace);
noc_async_read_barrier();
}
template <typename A>
FORCE_INLINE void read_row_tile(uint32_t cb, const A& acc, uint32_t j) {
cb_reserve_back(cb, 1);
const uint32_t p = get_write_ptr(cb);
noc_async_read(acc.get_noc_addr(j), p, TB);
noc_async_read_barrier();
fill_first_row(p);
cb_push_back(cb, 1);
}
void kernel_main() {
const uint32_t x_addr = get_common_arg_val<uint32_t>(0);
const uint32_t g_addr = get_common_arg_val<uint32_t>(1);
const uint32_t b_addr = get_common_arg_val<uint32_t>(2);
const uint32_t r_addr = get_common_arg_val<uint32_t>(3);
const uint32_t rg_addr = get_common_arg_val<uint32_t>(4);
const uint32_t row0 = get_arg_val<uint32_t>(0);
const uint32_t n_rows = get_arg_val<uint32_t>(1);
const uint32_t stride = get_arg_val<uint32_t>(2);
const uint32_t j = get_arg_val<uint32_t>(3);
if (n_rows == 0) {
return;
}
const bool root = j == 0;
const auto x = TensorAccessor(x_args, x_addr, TB);
const auto res = TensorAccessor(r_args, r_addr, TB);
volatile tt_l1_ptr uint32_t* sem_g = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(get_semaphore(0));
volatile tt_l1_ptr uint32_t* sem_b = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(get_semaphore(1));
uint32_t n_gat = 0, n_bc = 0;
auto read_x = [&](uint32_t r) {
cb_reserve_back(cb_x, 1);
noc_async_read(x.get_noc_addr(r * Wt + j), get_write_ptr(cb_x), TB);
if constexpr (has_res) {
cb_reserve_back(cb_r, 1);
noc_async_read(res.get_noc_addr(r * Wt + j), get_write_ptr(cb_r), TB);
noc_async_read_barrier();
cb_push_back(cb_r, 1);
} else {
noc_async_read_barrier();
}
cb_push_back(cb_x, 1);
};
auto gather = [&]() {
cb_reserve_back(cb_gat, Wt);
n_gat += Wt;
noc_semaphore_wait_min(sem_g, n_gat);
cb_push_back(cb_gat, Wt);
};
auto bcast_in = [&]() {
cb_reserve_back(cb_bc, 1);
n_bc += 1;
noc_semaphore_wait_min(sem_b, n_bc);
cb_push_back(cb_bc, 1);
};
read_x(row0);
if constexpr (has_rgate) {
read_row_tile(cb_rg, TensorAccessor(rg_args, rg_addr, TB), j);
}
cb_reserve_back(cb_eps, 1);
{
auto* e = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(get_write_ptr(cb_eps));
for (uint32_t i = 0; i < 1024; ++i) {
e[i] = eps_bits;
}
}
cb_push_back(cb_eps, 1);
if constexpr (has_gamma) {
read_row_tile(cb_g, TensorAccessor(g_args, g_addr, TB), j);
}
if constexpr (has_beta) {
read_row_tile(cb_b, TensorAccessor(b_args, b_addr, TB), j);
}
for (uint32_t i = 0; i < n_rows; ++i) {
if (i > 0) {
read_x(row0 + i * stride);
}
if (root) {
gather();
}
bcast_in();
if (root) {
gather();
}
bcast_in();
}
}