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