Download code/tt_diffusion_planner/tt/kernels/ln32_compute.cpp from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 10.7 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/kernels/ln32_compute.cpp
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/tt/kernels/ln32_compute.cpp
-
curl -L -o ln32_compute.cpp https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/kernels/ln32_compute.cpp
10.7 kB
| // 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]. | |
| 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); | |
| } | |
| } | |
| 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); | |
| } | |
| } | |