Download code/tt_diffusion_planner/tt/kernels/ln32s_compute.cpp from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 8.21 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/kernels/ln32s_compute.cpp
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/tt/kernels/ln32s_compute.cpp
-
curl -L -o ln32s_compute.cpp https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/kernels/ln32s_compute.cpp
8.21 kB
| // SPDX-License-Identifier: Apache-2.0 | |
| // Split-row fused fp32 LayerNorm (tt/ln_kernel.py, LN_SPLIT), compute: member j of a tile row owns column tile j. | |
| // The same LLK sequence as ln32_compute.cpp (and the stock 9-program decomposition): the folds over the Wt tiles run | |
| // on the root, in tile order, over the gathered tiles (c_10), so every statistic is bit for bit the stock one. | |
| // h = x (+ res (* rgate)) -> c_11 (to the root's gather) [+ c_17 for write_h] | |
| // root: mean = fold(c_10) / W -> c_4 (the writer broadcasts it into every member's c_5 page 0) | |
| // xc = h - mean; sq = xc * xc -> c_6 (kept), sq -> c_11 | |
| // root: rstd = rsqrt(fold(c_10) / W + eps) -> c_4 (broadcast into c_5 page 1) | |
| // y = xc * rstd (* gamma) (+ beta) -> c_16 | |
| // 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, [8] sfpu_bcast (the root's statistics broadcast over the columns by the | |
| // SFPU reduce, ln32_sfpu.h; else column 0 and the writer fills), [9] kcat_out (KCAT_EMIT: y is emitted | |
| // as the split operand of the next K-concatenated linear, y_hi = bf16(y) -> c_16 and y_lo = y - y_hi -> | |
| // c_19, the LLK calls of kernels/kcat_compute.cpp on the packed fp32 y tile c_18). | |
| // Per-core RT args: [n_rows, is_root]. | |
| 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 sfpu_bcast = get_compile_time_arg_val(8) != 0; | |
| constexpr uint32_t kcat_out = get_compile_time_arg_val(9); | |
| 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_gat = 10, cb_snd = 11, cb_out = 16, cb_hout = 17, cb_y = 18, cb_lo = 19; | |
| constexpr uint32_t cb_src = has_res ? cb_h : cb_x; | |
| constexpr auto RNE = ckernel::DstRoundingMode::NearestEven; | |
| // fold of the Wt gathered tiles (in order) -> row sum * 1/W in DST 0 | |
| ALWI void fold_gathered() { | |
| cb_wait_front(cb_gat, Wt); | |
| copy_init(cb_gat); | |
| add_binary_tile_init(); | |
| copy_tile(cb_gat, 0, 0); | |
| for (uint32_t w = 1; w < Wt; ++w) { | |
| copy_tile(cb_gat, w, 1); | |
| 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_one(uint32_t dst, uint32_t cb) { | |
| cb_reserve_back(cb, 1); | |
| tile_regs_commit(); | |
| tile_regs_wait(); | |
| pack_tile(dst, cb); | |
| tile_regs_release(); | |
| cb_push_back(cb, 1); | |
| } | |
| void kernel_main() { | |
| const uint32_t n_rows = get_arg_val<uint32_t>(0); | |
| const bool root = get_arg_val<uint32_t>(1) != 0; | |
| if (n_rows == 0) { | |
| return; | |
| } | |
| compute_kernel_hw_startup(cb_x, cb_out); | |
| for (uint32_t r = 0; r < n_rows; ++r) { | |
| // h (this member's tile) -> c_11 for the gather (and c_9 / c_17) | |
| cb_wait_front(cb_x, 1); | |
| tile_regs_acquire(); | |
| if constexpr (has_res) { | |
| cb_wait_front(cb_r, 1); | |
| if constexpr (has_rgate) { | |
| cb_wait_front(cb_rg, 1); | |
| } | |
| copy_init(cb_r); | |
| copy_tile(cb_r, 0, 0); | |
| if constexpr (has_rgate) { | |
| copy_tile(cb_rg, 0, 1); | |
| mul_binary_tile_init(); | |
| mul_binary_tile(0, 1, 0); | |
| } | |
| copy_tile(cb_x, 0, 1); | |
| add_binary_tile_init(); | |
| add_binary_tile<RNE>(1, 0, 1); | |
| cb_reserve_back(cb_h, 1); | |
| cb_reserve_back(cb_snd, 1); | |
| if constexpr (write_h) { | |
| cb_reserve_back(cb_hout, 1); | |
| } | |
| tile_regs_commit(); | |
| tile_regs_wait(); | |
| pack_tile(1, cb_h); | |
| pack_tile(1, cb_snd); | |
| if constexpr (write_h) { | |
| pack_tile(1, cb_hout); | |
| } | |
| tile_regs_release(); | |
| cb_push_back(cb_h, 1); | |
| cb_push_back(cb_snd, 1); | |
| if constexpr (write_h) { | |
| cb_push_back(cb_hout, 1); | |
| } | |
| cb_pop_front(cb_r, 1); | |
| cb_pop_front(cb_x, 1); | |
| cb_wait_front(cb_h, 1); | |
| } else { | |
| copy_init(cb_x); | |
| copy_tile(cb_x, 0, 0); | |
| pack_one(0, cb_snd); | |
| } | |
| // root: mean | |
| if (root) { | |
| tile_regs_acquire(); | |
| fold_gathered(); | |
| pack_one(0, cb_stat); | |
| cb_pop_front(cb_gat, Wt); | |
| } | |
| // xc = h - mean -> c_6; sq = xc * xc -> c_11 | |
| cb_wait_front(cb_bc, 1); | |
| tile_regs_acquire(); | |
| copy_init(cb_src); | |
| copy_tile(cb_src, 0, 0); | |
| copy_tile(cb_bc, 0, 1); | |
| sub_binary_tile_init(); | |
| sub_binary_tile<RNE>(0, 1, 0); | |
| pack_one(0, cb_xc); | |
| cb_pop_front(cb_bc, 1); | |
| cb_pop_front(cb_src, 1); | |
| cb_wait_front(cb_xc, 1); | |
| tile_regs_acquire(); | |
| copy_init(cb_xc); | |
| copy_tile(cb_xc, 0, 0); | |
| copy_tile(cb_xc, 0, 1); | |
| mul_binary_tile_init(); | |
| mul_binary_tile(0, 1, 0); | |
| pack_one(0, cb_snd); | |
| // root: rstd | |
| if (root) { | |
| tile_regs_acquire(); | |
| fold_gathered(); | |
| 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_one(0, cb_stat); | |
| cb_pop_front(cb_gat, Wt); | |
| } | |
| // y = xc * rstd (* gamma) (+ beta) | |
| cb_wait_front(cb_bc, 1); | |
| tile_regs_acquire(); | |
| copy_init(cb_xc); | |
| copy_tile(cb_xc, 0, 0); | |
| copy_tile(cb_bc, 0, 1); | |
| mul_binary_tile_init(); | |
| mul_binary_tile(0, 1, 0); | |
| if constexpr (has_gamma) { | |
| cb_wait_front(cb_g, 1); | |
| copy_tile(cb_g, 0, 1); | |
| mul_binary_tile(0, 1, 0); | |
| } | |
| if constexpr (has_beta) { | |
| cb_wait_front(cb_b, 1); | |
| copy_tile(cb_b, 0, 1); | |
| add_binary_tile_init(); | |
| add_binary_tile<RNE>(0, 1, 0); | |
| } | |
| if constexpr (kcat_out) { | |
| pack_one(0, cb_y); | |
| cb_pop_front(cb_bc, 1); | |
| cb_pop_front(cb_xc, 1); | |
| cb_wait_front(cb_y, 1); | |
| tile_regs_acquire(); | |
| copy_init(cb_y); | |
| copy_tile(cb_y, 0, 0); | |
| copy_tile(cb_y, 0, 1); | |
| typecast_tile_init<(uint32_t)DataFormat::Float32, (uint32_t)DataFormat::Float16_b>(); | |
| typecast_tile<(uint32_t)DataFormat::Float32, (uint32_t)DataFormat::Float16_b>(0); | |
| sub_binary_tile_init(); | |
| sub_binary_tile<RNE>(1, 0, 1); | |
| cb_reserve_back(cb_out, 1); | |
| cb_reserve_back(cb_lo, 1); | |
| tile_regs_commit(); | |
| tile_regs_wait(); | |
| pack_tile(0, cb_out); | |
| pack_tile(1, cb_lo); | |
| tile_regs_release(); | |
| cb_push_back(cb_out, 1); | |
| cb_push_back(cb_lo, 1); | |
| cb_pop_front(cb_y, 1); | |
| } else { | |
| pack_one(0, cb_out); | |
| cb_pop_front(cb_bc, 1); | |
| cb_pop_front(cb_xc, 1); | |
| } | |
| } | |
| } | |