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();
}