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
// Fused fp32 LayerNorm (tt/ln_kernel.py), reader (RISCV_0): the affine rows once, then this core's tile rows of x.
//
// CT args: [0] Wt, [1] has_gamma, [2] has_beta, [3] eps bits (fp32), [4] per-core RT-arg count (P3: part of the
// program hash), [5] has_res, [6] has_rgate, [7] res_t (LN_TR: res is stored transposed per entity, i.e.
// [.., E, W, T]: tile (r, w) of res is read from tile (e, w, tr) of that layout, e = r / Tt, tr = r % Tt;
// the compute transposes it back), [8] Tt (tile rows per entity), then the TensorAccessorArgs of x, gamma,
// beta, res, rgate (x's when absent).
// Common RT args: [x_addr, gamma_addr, beta_addr, res_addr, rgate_addr]. Per-core RT args: [row0, n_rows].
// CBs: c_0 x (Wt tiles per row, double-buffered), c_1 gamma rows (Wt tiles, row 0 replicated over the tile),
// c_2 beta rows (same), c_3 the eps tile (every element eps), c_7 res (as c_0), c_8 the residual gate rows.
#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 uint32_t res_t = get_compile_time_arg_val(7);
constexpr uint32_t Tt = get_compile_time_arg_val(8);
constexpr auto x_args = TensorAccessorArgs<9>();
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_r = 7, cb_rg = 8;
constexpr uint32_t TB = 4096; // fp32 tile bytes
// Replicate row 0 of each of the n fp32 tiles at p over its 32 rows (binary_ng fill_tile_with_first_row: local
// NoC doubling), all tiles per doubling step under one barrier.
FORCE_INLINE void fill_first_rows(uint32_t p, uint32_t n_tiles) {
constexpr uint32_t kFace = 1024, kRow = 64;
for (uint32_t n = kRow; n < kFace; n <<= 1) {
for (uint32_t t = 0; t < n_tiles; ++t) {
const uint32_t f0 = p + t * TB, f1 = f0 + kFace;
noc_async_read(get_noc_addr(f0), f0 + n, n);
noc_async_read(get_noc_addr(f1), f1 + n, n);
}
noc_async_read_barrier();
}
for (uint32_t t = 0; t < n_tiles; ++t) {
const uint32_t f0 = p + t * TB;
noc_async_read(get_noc_addr(f0), f0 + 2 * kFace, 2 * kFace);
}
noc_async_read_barrier();
}
template <typename A>
FORCE_INLINE void read_rows(uint32_t cb, const A& acc) {
cb_reserve_back(cb, Wt);
uint32_t p = get_write_ptr(cb);
for (uint32_t w = 0; w < Wt; ++w) {
noc_async_read(acc.get_noc_addr(w), p + w * TB, TB);
}
noc_async_read_barrier();
fill_first_rows(p, Wt);
cb_push_back(cb, Wt);
}
template <typename A, typename R>
FORCE_INLINE void read_x_row(uint32_t r, const A& x, const R& res) {
cb_reserve_back(cb_x, Wt);
uint32_t p = get_write_ptr(cb_x);
for (uint32_t w = 0; w < Wt; ++w) {
noc_async_read(x.get_noc_addr(r * Wt + w), p + w * TB, TB);
}
if constexpr (has_res) {
cb_reserve_back(cb_r, Wt);
uint32_t q = get_write_ptr(cb_r);
for (uint32_t w = 0; w < Wt; ++w) {
const uint32_t ri = res_t ? (r / Tt) * (Wt * Tt) + w * Tt + (r % Tt) : r * Wt + w;
noc_async_read(res.get_noc_addr(ri), q + w * TB, TB);
}
noc_async_read_barrier();
cb_push_back(cb_r, Wt);
} else {
noc_async_read_barrier();
}
cb_push_back(cb_x, Wt);
}
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);
if (n_rows == 0) {
return;
}
const auto x = TensorAccessor(x_args, x_addr, TB);
const auto res = TensorAccessor(r_args, r_addr, TB);
// the first x row goes first (the compute starts on it); the affine rows are needed only by its last phase
read_x_row(row0, x, res);
// eps tile
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_rgate) {
read_rows(cb_rg, TensorAccessor(rg_args, rg_addr, TB));
}
if constexpr (has_gamma) {
read_rows(cb_g, TensorAccessor(g_args, g_addr, TB));
}
if constexpr (has_beta) {
read_rows(cb_b, TensorAccessor(b_args, b_addr, TB));
}
for (uint32_t r = row0 + 1; r < row0 + n_rows; ++r) {
read_x_row(r, x, res);
}
}