// SPDX-License-Identifier: Apache-2.0 // Split-matmul operand build (tt/kcat_kernel.py), compute: per fp32 x tile, x_hi = bf16(x) (typecast_tile // Float32 -> Float16_b, the stock ttnn.typecast LLK) kept as fp32, and x_lo = x - x_hi (sub_binary_tile, exact). // With act (KCAT_ACT) the previous linear's activation is applied first, as the stock unary program does it // (math_approx_mode false): 1 = gelu_tile<0> (ttnn.gelu, accurate), 2 = gelu_tanh_tile (GeluVariant.Tanh). // CT args: [0] per-core RT-arg count (P3), [1] act, [2] once (KCAT_ACT_ONCE: the activation on one DEST copy, then // copy_dest_values to the second; the same values). Per-core RT args: [n, 0]. #include #include "api/compute/common.h" #include "api/compute/compute_kernel_api.h" #include "api/compute/copy_dest_values.h" #include "api/compute/eltwise_binary_sfpu.h" #include "api/compute/eltwise_unary/eltwise_unary.h" #include "api/compute/eltwise_unary/gelu.h" #include "api/compute/eltwise_unary/typecast.h" #include "api/compute/tile_move_copy.h" using namespace ckernel; constexpr uint32_t act = get_compile_time_arg_val(1); constexpr uint32_t once = get_compile_time_arg_val(2); constexpr uint32_t cb_x = 0, cb_hi = 16, cb_lo = 17; void kernel_main() { const uint32_t n = get_arg_val(0); if (n == 0) { return; } compute_kernel_hw_startup(cb_x, cb_hi); for (uint32_t t = 0; t < n; ++t) { cb_wait_front(cb_x, 1); tile_regs_acquire(); copy_init(cb_x); copy_tile(cb_x, 0, 0); if constexpr (once && act != 0) { if constexpr (act == 1) { gelu_tile_init<0u>(); gelu_tile<0u>(0); } else { gelu_tanh_tile_init(); gelu_tanh_tile(0); } copy_dest_values_init(); copy_dest_values(0, 1); } else { copy_tile(cb_x, 0, 1); } if constexpr (once && act != 0) { } else if constexpr (act == 1) { gelu_tile_init<0u>(); gelu_tile<0u>(0); gelu_tile<0u>(1); } else if constexpr (act == 2) { gelu_tanh_tile_init(); gelu_tanh_tile(0); gelu_tanh_tile(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(1, 0, 1); cb_reserve_back(cb_hi, 1); cb_reserve_back(cb_lo, 1); tile_regs_commit(); tile_regs_wait(); pack_tile(0, cb_hi); pack_tile(1, cb_lo); tile_regs_release(); cb_push_back(cb_hi, 1); cb_push_back(cb_lo, 1); cb_pop_front(cb_x, 1); } }