File size: 10,722 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
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
// SPDX-License-Identifier: Apache-2.0
// Fused fp32 LayerNorm (tt/ln_kernel.py), compute: the exact op sequence of tt/layers.py layer_norm_fp32 (stock
// ttnn programs: mean -> sub -> mul -> mean -> add eps -> rsqrt -> mul -> mul gamma -> add beta), with the same SFPU
// LLK calls in the same order, so the result is meant to be bit-identical to the 9-program decomposition:
//   - ttnn.mean (accurate fp32 path, reduce_op.cpp use_sfpu_fp32_mean): copy tile 0, add_binary_tile fold of tiles
//     1..Wt-1 in order, sfpu_reduce<SUM, Float32, REDUCE_ROW>, mul_unary_tile(1/W) (the AVG post-mul);
//   - binary_ng fp32 ops: sub / add as sub_binary_tile / add_binary_tile<NearestEven>, mul as mul_binary_tile,
//     the broadcast operand built by the dataflow (column / row fill), lhs in DST 0;
//   - ttnn.rsqrt(fast_and_approximate_mode=False): rsqrt_tile<RsqrtMode::Default> (math_approx_mode false).
// Everything stays fp32 (UnpackToDestFp32 copies, fp32 DST, fp32 CBs): every intermediate the stock graph packs to
// an fp32 DRAM tensor round-trips losslessly.
//
// With has_res (LN_RESID, the residual add of the stream fused in front): h = x + res (* rgate), the stock
// `ttnn.add(x, ttnn.multiply(res, rgate))` / `ttnn.add(x, res)` as mul_binary_tile / add_binary_tile, packed to
// c_9 (the LayerNorm input) and to c_17 (written out as the new stream when write_h).
//
// 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, [9] sfpu_bcast (the statistics broadcast over the columns by the SFPU
//          reduce itself, ln32_sfpu.h, and packed straight to c_5; else column 0 packed to c_4 and broadcast by the
//          writer), [8] lean (one copy_init / SFPU-binary init per phase instead of one per
//          tile: every CB copied from is an fp32 UnpackToDestFp32 CB, and add / sub / mul share one init).
//          [10] res_t (LN_TR: the residual tiles arrive transposed and are transposed back with transpose_tile, the
//          stock ttnn.transpose LLK, exact with the fp32 unpack-to-dest; no rgate), [11] out_t (the y tiles are
//          transposed the same way, via c_18, before they are packed for the writer).
// Per-core RT args: [n_rows].
#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/tile_move_copy.h"
#include "api/compute/transpose.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 lean = get_compile_time_arg_val(8) != 0;
constexpr bool sfpu_bcast = get_compile_time_arg_val(9) != 0;
constexpr uint32_t res_t = get_compile_time_arg_val(10);
constexpr uint32_t out_t = get_compile_time_arg_val(11);

// per-tile (re-)inits, skipped in lean mode (done once at the start of the phase instead)
ALWI void ci(uint32_t cb) {
    if constexpr (!lean) {
        copy_init(cb);
    }
}
#define BIN_INIT(op)                  \
    do {                              \
        if constexpr (!lean) {        \
            op##_binary_tile_init();  \
        }                             \
    } while (0)

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_out = 16, cb_hout = 17, cb_yt = 18;
constexpr uint32_t cb_src = has_res ? cb_h : cb_x;
constexpr auto RNE = ckernel::DstRoundingMode::NearestEven;

// row mean of the Wt tiles of `cb` (already waited) into DST 0; `square` folds xc * xc instead of x
template <bool square>
ALWI void row_sum_to_dst0(uint32_t cb) {
    copy_init(cb);
    if constexpr (lean) {
        add_binary_tile_init();
    }
    if constexpr (square) {
        copy_tile(cb, 0, 0);
        copy_tile(cb, 0, 1);
        BIN_INIT(mul);
        mul_binary_tile(0, 1, 0);
    } else {
        copy_tile(cb, 0, 0);
    }
    for (uint32_t w = 1; w < Wt; ++w) {
        ci(cb);
        copy_tile(cb, w, 1);
        if constexpr (square) {
            copy_tile(cb, w, 2);
            BIN_INIT(mul);
            mul_binary_tile(1, 2, 1);
        }
        BIN_INIT(add);
        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_stat() {
    constexpr uint32_t cb = sfpu_bcast ? cb_bc : cb_stat;
    cb_reserve_back(cb, 1);
    tile_regs_commit();
    tile_regs_wait();
    pack_tile(0, cb);
    tile_regs_release();
    cb_push_back(cb, 1);
}

void kernel_main() {
    const uint32_t n_rows = get_arg_val<uint32_t>(0);
    if (n_rows == 0) {
        return;
    }
    compute_kernel_hw_startup(cb_x, cb_out);
    // the constant tiles are waited for where they are first used (the reader sends x row 0 first)
    for (uint32_t r = 0; r < n_rows; ++r) {
        cb_wait_front(cb_x, Wt);
        if constexpr (has_res) {
            // h = x + res (* rgate)
            cb_wait_front(cb_r, Wt);
            if constexpr (has_rgate) {
                cb_wait_front(cb_rg, Wt);
            }
            if constexpr (lean) {
                copy_init(cb_r);
                add_binary_tile_init();
            }
            for (uint32_t w = 0; w < Wt; ++w) {
                tile_regs_acquire();
                if constexpr (res_t) {
                    transpose_init(cb_r);
                    transpose_tile(cb_r, w, 0);
                    copy_init(cb_x);
                    add_binary_tile_init();
                } else {
                    ci(cb_r);
                    copy_tile(cb_r, w, 0);
                }
                if constexpr (has_rgate) {
                    ci(cb_rg);
                    copy_tile(cb_rg, w, 1);
                    BIN_INIT(mul);
                    mul_binary_tile(0, 1, 0);
                }
                ci(cb_x);
                copy_tile(cb_x, w, 1);
                BIN_INIT(add);
                add_binary_tile<RNE>(1, 0, 1);
                cb_reserve_back(cb_h, 1);
                if constexpr (write_h) {
                    cb_reserve_back(cb_hout, 1);
                }
                tile_regs_commit();
                tile_regs_wait();
                pack_tile(1, cb_h);
                if constexpr (write_h) {
                    pack_tile(1, cb_hout);
                }
                tile_regs_release();
                cb_push_back(cb_h, 1);
                if constexpr (write_h) {
                    cb_push_back(cb_hout, 1);
                }
            }
            cb_pop_front(cb_r, Wt);
            cb_pop_front(cb_x, Wt);
            cb_wait_front(cb_h, Wt);
        }
        // mean
        tile_regs_acquire();
        row_sum_to_dst0<false>(cb_src);
        pack_stat();

        // xc = x - mean
        cb_wait_front(cb_bc, 1);
        if constexpr (lean) {
            copy_init(cb_src);
            sub_binary_tile_init();
        }
        for (uint32_t w = 0; w < Wt; ++w) {
            tile_regs_acquire();
            ci(cb_src);
            copy_tile(cb_src, w, 0);
            ci(cb_bc);
            copy_tile(cb_bc, 0, 1);
            BIN_INIT(sub);
            sub_binary_tile<RNE>(0, 1, 0);
            cb_reserve_back(cb_xc, 1);
            tile_regs_commit();
            tile_regs_wait();
            pack_tile(0, cb_xc);
            tile_regs_release();
            cb_push_back(cb_xc, 1);
        }
        cb_pop_front(cb_bc, 1);
        cb_pop_front(cb_src, Wt);

        // rstd = rsqrt(mean(xc * xc) + eps)
        cb_wait_front(cb_xc, Wt);
        tile_regs_acquire();
        row_sum_to_dst0<true>(cb_xc);
        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_stat();

        // y = xc * rstd (* gamma) (+ beta)
        cb_wait_front(cb_bc, 1);
        if constexpr (has_gamma) {
            cb_wait_front(cb_g, Wt);
        }
        if constexpr (has_beta) {
            cb_wait_front(cb_b, Wt);
        }
        if constexpr (lean) {
            copy_init(cb_xc);
            mul_binary_tile_init();
        }
        for (uint32_t w = 0; w < Wt; ++w) {
            tile_regs_acquire();
            ci(cb_xc);
            copy_tile(cb_xc, w, 0);
            ci(cb_bc);
            copy_tile(cb_bc, 0, 1);
            BIN_INIT(mul);
            mul_binary_tile(0, 1, 0);
            if constexpr (has_gamma) {
                ci(cb_g);
                copy_tile(cb_g, w, 1);
                BIN_INIT(mul);
                mul_binary_tile(0, 1, 0);
            }
            if constexpr (has_beta) {
                ci(cb_b);
                copy_tile(cb_b, w, 1);
                BIN_INIT(add);
                add_binary_tile<RNE>(0, 1, 0);
            }
            if constexpr (out_t) {
                cb_reserve_back(cb_yt, 1);
                tile_regs_commit();
                tile_regs_wait();
                pack_tile(0, cb_yt);
                tile_regs_release();
                cb_push_back(cb_yt, 1);
                cb_wait_front(cb_yt, 1);
                tile_regs_acquire();
                transpose_init(cb_yt);
                transpose_tile(cb_yt, 0, 0);
                cb_reserve_back(cb_out, 1);
                tile_regs_commit();
                tile_regs_wait();
                pack_tile(0, cb_out);
                tile_regs_release();
                cb_push_back(cb_out, 1);
                cb_pop_front(cb_yt, 1);
                copy_init(cb_xc);
                mul_binary_tile_init();
            } else {
                cb_reserve_back(cb_out, 1);
                tile_regs_commit();
                tile_regs_wait();
                pack_tile(0, cb_out);
                tile_regs_release();
                cb_push_back(cb_out, 1);
            }
        }
        cb_pop_front(cb_bc, 1);
        cb_pop_front(cb_xc, Wt);
    }
}