// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc. // SPDX-License-Identifier: Apache-2.0 // // Descriptor L2 normalisation + untilize in one op (models/tt/desc_norm.py). Per tile row of the // height-sharded [rows, 256] TILE descriptor map: S = sum_c x^2 (x*x on the FPU, packed fp32, row-reduced // with a ones scaler), r = rsqrt(S) (SFPU), y = x * r (column broadcast), pack-untilized into a // [32 rows, 256] row-major block for the writer. #include #include "api/compute/compute_kernel_api.h" #include "api/compute/compute_kernel_hw_startup.h" #include "api/compute/eltwise_binary.h" #include "api/compute/bcast.h" #include "api/compute/reduce.h" #include "api/compute/eltwise_unary/rsqrt.h" #include "api/compute/pack_untilize.h" #include "api/compute/reconfig_data_format.h" void kernel_main() { constexpr uint32_t cb_in = get_compile_time_arg_val(0); constexpr uint32_t cb_one = get_compile_time_arg_val(1); constexpr uint32_t cb_sq = get_compile_time_arg_val(2); constexpr uint32_t cb_rs = get_compile_time_arg_val(3); constexpr uint32_t cb_out = get_compile_time_arg_val(4); constexpr uint32_t TR = get_compile_time_arg_val(5); // tile rows per core constexpr uint32_t WT = 8, HB = 4; // tiles per row (256 ch), tiles per DST block compute_kernel_hw_startup(cb_in, cb_in, cb_sq); cb_wait_front(cb_in, TR * WT); cb_wait_front(cb_one, 1); for (uint32_t r = 0; r < TR; ++r) { // ---- x^2 -> CB_SQ (fp32) reconfig_data_format(cb_in, cb_in); pack_reconfig_data_format(cb_sq); mul_init(cb_in, cb_in); cb_reserve_back(cb_sq, WT); for (uint32_t b = 0; b < WT / HB; ++b) { tile_regs_acquire(); for (uint32_t k = 0; k < HB; ++k) { const uint32_t t = r * WT + b * HB + k; mul_tiles(cb_in, cb_in, t, t, k); } tile_regs_commit(); tile_regs_wait(); for (uint32_t k = 0; k < HB; ++k) { pack_tile(k, cb_sq); } tile_regs_release(); } cb_push_back(cb_sq, WT); // ---- row sum, rsqrt -> CB_RS (fp32, column 0) cb_wait_front(cb_sq, WT); reconfig_data_format(cb_one, cb_sq); pack_reconfig_data_format(cb_rs); reduce_init(cb_sq, cb_one, cb_rs); cb_reserve_back(cb_rs, 1); tile_regs_acquire(); for (uint32_t t = 0; t < WT; ++t) { reduce_tile(cb_sq, cb_one, t, 0, 0); } rsqrt_tile_init(); rsqrt_tile(0); tile_regs_commit(); tile_regs_wait(); pack_tile(0, cb_rs); tile_regs_release(); reduce_uninit(cb_sq); cb_push_back(cb_rs, 1); cb_pop_front(cb_sq, WT); // ---- x * r (column broadcast), pack-untilize -> CB_OUT cb_wait_front(cb_rs, 1); reconfig_data_format(cb_in, cb_rs); pack_reconfig_data_format(cb_out); mul_bcast_cols_init(cb_in, cb_rs); pack_untilize_dest_init(cb_out); cb_reserve_back(cb_out, WT); for (uint32_t b = 0; b < WT / HB; ++b) { tile_regs_acquire(); for (uint32_t k = 0; k < HB; ++k) { mul_tiles_bcast_cols(cb_in, cb_rs, r * WT + b * HB + k, 0, k); } tile_regs_commit(); tile_regs_wait(); pack_untilize_dest(cb_out, 1, b); tile_regs_release(); } pack_untilize_uninit(cb_out); cb_push_back(cb_out, WT); cb_pop_front(cb_rs, 1); } }