File size: 3,798 Bytes
be62f78 | 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 | // SPDX-License-Identifier: Apache-2.0
// Fused fp32 LayerNorm (tt/ln_kernel.py), writer (RISCV_1): per tile row, broadcasts the two statistic tiles the
// compute packs (column 0 = the row's mean, then its rsqrt(var + eps)) over all 32 columns, then drains the Wt
// output tiles to DRAM.
//
// With write_h (LN_RESID) the Wt tiles of the new stream h = x + res (* rgate) (c_17) go first, to h_addr.
//
// CT args: [0] Wt, [1] per-core RT-arg count (P3), [2] write_h, [3] sfpu_bcast (no fill here), [4] out_t (LN_TR:
// the output tiles, transposed by the compute, go to the per-entity transposed layout [.., E, W, T]: tile
// (r, w) -> (e, w, tr), e = r / Tt, tr = r % Tt), [5] Tt, then the TensorAccessorArgs of out and h (out's
// when absent).
// Common RT args: [out_addr, h_addr]. Per-core RT args: [row0, n_rows].
// CBs: c_4 statistic (from compute, column 0 valid), c_5 broadcast statistic (to compute), c_16 out, c_17 h.
#include <cstdint>
#include "api/dataflow/dataflow_api.h"
constexpr uint32_t Wt = get_compile_time_arg_val(0);
constexpr uint32_t write_h = get_compile_time_arg_val(2);
constexpr uint32_t sfpu_bcast = get_compile_time_arg_val(3); // the compute broadcasts the statistics itself
constexpr uint32_t out_t = get_compile_time_arg_val(4);
constexpr uint32_t Tt = get_compile_time_arg_val(5);
constexpr auto out_args = TensorAccessorArgs<6>();
constexpr auto h_args = TensorAccessorArgs<out_args.next_compile_time_args_offset()>();
constexpr uint32_t cb_stat = 4, cb_bc = 5, cb_out = 16, cb_hout = 17;
constexpr uint32_t TB = 4096;
// dst tile = src tile's column 0 replicated over the 32 columns (binary_ng fill_tile_with_first_column, fp32;
// faces 0/1 take face 0's column 0, faces 2/3 face 2's).
FORCE_INLINE void bcast_col(uint32_t src, uint32_t dst) {
auto* s = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(src);
auto* d = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(dst);
for (uint32_t fo = 0; fo < 1024; fo += 512) {
for (uint32_t ro = 0; ro < 256; ro += 16) {
const uint32_t v = s[fo + ro];
volatile tt_l1_ptr uint32_t* l = d + fo + ro;
volatile tt_l1_ptr uint32_t* r = l + 256;
for (uint32_t c = 0; c < 16; ++c) {
l[c] = v;
r[c] = v;
}
}
}
}
FORCE_INLINE void stat() {
cb_wait_front(cb_stat, 1);
cb_reserve_back(cb_bc, 1);
bcast_col(get_read_ptr(cb_stat), get_write_ptr(cb_bc));
cb_push_back(cb_bc, 1);
cb_pop_front(cb_stat, 1);
}
void kernel_main() {
const uint32_t out_addr = get_common_arg_val<uint32_t>(0);
const uint32_t h_addr = get_common_arg_val<uint32_t>(1);
const auto hacc = TensorAccessor(h_args, h_addr, TB);
const uint32_t row0 = get_arg_val<uint32_t>(0);
const uint32_t n_rows = get_arg_val<uint32_t>(1);
const auto out = TensorAccessor(out_args, out_addr, TB);
for (uint32_t r = row0; r < row0 + n_rows; ++r) {
if constexpr (write_h) {
for (uint32_t w = 0; w < Wt; ++w) {
cb_wait_front(cb_hout, 1);
noc_async_write(get_read_ptr(cb_hout), hacc.get_noc_addr(r * Wt + w), TB);
noc_async_writes_flushed();
cb_pop_front(cb_hout, 1);
}
}
if constexpr (!sfpu_bcast) {
stat(); // mean
stat(); // rsqrt(var + eps)
}
for (uint32_t w = 0; w < Wt; ++w) {
cb_wait_front(cb_out, 1);
const uint32_t oi = out_t ? (r / Tt) * (Wt * Tt) + w * Tt + (r % Tt) : r * Wt + w;
noc_async_write(get_read_ptr(cb_out), out.get_noc_addr(oi), TB);
noc_async_writes_flushed();
cb_pop_front(cb_out, 1);
}
}
noc_async_write_barrier();
}
|