Download code/tt_diffusion_planner/tt/kernels/ln32s_reader.cpp from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 5.34 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/kernels/ln32s_reader.cpp
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/tt/kernels/ln32s_reader.cpp
-
curl -L -o ln32s_reader.cpp https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/kernels/ln32s_reader.cpp
5.34 kB
| // SPDX-License-Identifier: Apache-2.0 | |
| // Split-row fused fp32 LayerNorm (tt/ln_kernel.py, LN_SPLIT), reader (RISCV_0). One tile row is spread over Wt | |
| // cores (member j owns column tile j); member 0 (the root) gathers the Wt tiles that the stock reduction folds in | |
| // order and broadcasts the statistics back, so the arithmetic stays the stock decomposition's (bit-identical). | |
| // Per row this kernel: reads x (r, j) (+ res (r, j)); on the root, waits for the Wt gathered tiles (semaphore 0, | |
| // monotonic count) and hands them to the compute (c_10); then waits for the two statistic broadcasts (semaphore 1, | |
| // monotonic count) and hands each to the compute (c_5, 2 pages: mean, rstd). | |
| // CT args: [0] Wt, [1] has_gamma, [2] has_beta, [3] eps bits, [4] per-core RT-arg count (P3), [5] has_res, | |
| // [6] has_rgate, then the TensorAccessorArgs of x, gamma, beta, res, rgate. | |
| // Common RT args: [x_addr, gamma_addr, beta_addr, res_addr, rgate_addr]. | |
| // Per-core RT args: [row0, n_rows, row_stride, j, root_x, root_y, (member x, y) x Wt]. | |
| 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 eps_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 auto x_args = TensorAccessorArgs<7>(); | |
| constexpr auto g_args = TensorAccessorArgs<x_args.next_compile_time_args_offset()>(); | |
| constexpr auto b_args = TensorAccessorArgs<g_args.next_compile_time_args_offset()>(); | |
| constexpr auto r_args = TensorAccessorArgs<b_args.next_compile_time_args_offset()>(); | |
| constexpr auto rg_args = TensorAccessorArgs<r_args.next_compile_time_args_offset()>(); | |
| constexpr uint32_t cb_x = 0, cb_g = 1, cb_b = 2, cb_eps = 3, cb_bc = 5, cb_r = 7, cb_rg = 8, cb_gat = 10; | |
| constexpr uint32_t TB = 4096; | |
| FORCE_INLINE void fill_first_row(uint32_t p) { | |
| constexpr uint32_t kFace = 1024, kRow = 64; | |
| const uint32_t f0 = p, f1 = p + kFace; | |
| for (uint32_t n = kRow; n < kFace; n <<= 1) { | |
| noc_async_read(get_noc_addr(f0), f0 + n, n); | |
| noc_async_read(get_noc_addr(f1), f1 + n, n); | |
| noc_async_read_barrier(); | |
| } | |
| noc_async_read(get_noc_addr(f0), f0 + 2 * kFace, 2 * kFace); | |
| noc_async_read_barrier(); | |
| } | |
| template <typename A> | |
| FORCE_INLINE void read_row_tile(uint32_t cb, const A& acc, uint32_t j) { | |
| cb_reserve_back(cb, 1); | |
| const uint32_t p = get_write_ptr(cb); | |
| noc_async_read(acc.get_noc_addr(j), p, TB); | |
| noc_async_read_barrier(); | |
| fill_first_row(p); | |
| cb_push_back(cb, 1); | |
| } | |
| void kernel_main() { | |
| const uint32_t x_addr = get_common_arg_val<uint32_t>(0); | |
| const uint32_t g_addr = get_common_arg_val<uint32_t>(1); | |
| const uint32_t b_addr = get_common_arg_val<uint32_t>(2); | |
| const uint32_t r_addr = get_common_arg_val<uint32_t>(3); | |
| const uint32_t rg_addr = get_common_arg_val<uint32_t>(4); | |
| const uint32_t row0 = get_arg_val<uint32_t>(0); | |
| const uint32_t n_rows = get_arg_val<uint32_t>(1); | |
| const uint32_t stride = get_arg_val<uint32_t>(2); | |
| const uint32_t j = get_arg_val<uint32_t>(3); | |
| if (n_rows == 0) { | |
| return; | |
| } | |
| const bool root = j == 0; | |
| const auto x = TensorAccessor(x_args, x_addr, TB); | |
| const auto res = TensorAccessor(r_args, r_addr, TB); | |
| volatile tt_l1_ptr uint32_t* sem_g = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(get_semaphore(0)); | |
| volatile tt_l1_ptr uint32_t* sem_b = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(get_semaphore(1)); | |
| uint32_t n_gat = 0, n_bc = 0; | |
| auto read_x = [&](uint32_t r) { | |
| cb_reserve_back(cb_x, 1); | |
| noc_async_read(x.get_noc_addr(r * Wt + j), get_write_ptr(cb_x), TB); | |
| if constexpr (has_res) { | |
| cb_reserve_back(cb_r, 1); | |
| noc_async_read(res.get_noc_addr(r * Wt + j), get_write_ptr(cb_r), TB); | |
| noc_async_read_barrier(); | |
| cb_push_back(cb_r, 1); | |
| } else { | |
| noc_async_read_barrier(); | |
| } | |
| cb_push_back(cb_x, 1); | |
| }; | |
| auto gather = [&]() { | |
| cb_reserve_back(cb_gat, Wt); | |
| n_gat += Wt; | |
| noc_semaphore_wait_min(sem_g, n_gat); | |
| cb_push_back(cb_gat, Wt); | |
| }; | |
| auto bcast_in = [&]() { | |
| cb_reserve_back(cb_bc, 1); | |
| n_bc += 1; | |
| noc_semaphore_wait_min(sem_b, n_bc); | |
| cb_push_back(cb_bc, 1); | |
| }; | |
| read_x(row0); | |
| if constexpr (has_rgate) { | |
| read_row_tile(cb_rg, TensorAccessor(rg_args, rg_addr, TB), j); | |
| } | |
| cb_reserve_back(cb_eps, 1); | |
| { | |
| auto* e = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(get_write_ptr(cb_eps)); | |
| for (uint32_t i = 0; i < 1024; ++i) { | |
| e[i] = eps_bits; | |
| } | |
| } | |
| cb_push_back(cb_eps, 1); | |
| if constexpr (has_gamma) { | |
| read_row_tile(cb_g, TensorAccessor(g_args, g_addr, TB), j); | |
| } | |
| if constexpr (has_beta) { | |
| read_row_tile(cb_b, TensorAccessor(b_args, b_addr, TB), j); | |
| } | |
| for (uint32_t i = 0; i < n_rows; ++i) { | |
| if (i > 0) { | |
| read_x(row0 + i * stride); | |
| } | |
| if (root) { | |
| gather(); | |
| } | |
| bcast_in(); | |
| if (root) { | |
| gather(); | |
| } | |
| bcast_in(); | |
| } | |
| } | |