// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc. // SPDX-License-Identifier: Apache-2.0 // // Descriptor head in one op (models/tt/desc_head.py): 1x1 conv (256 -> 256) + bias, then the L2 norm and // untilize of DescNormRM (dn_compute.cpp), per tile row of the height-sharded [rows, 256] TILE input. // 1) z_f32 = x @ W, fp32 DST, K in ttnn order (k = 0..7), at HiFi2 (the fidelity of the ttnn 1x1 conv; // the kernel itself is compiled at HiFi4 for the norm steps) -> CB_P (fp32) // 2) z = bf16(z_f32 + bias) (row broadcast; ttnn's fused-bias epilogue) -> CB_Z (bf16, the ttnn conv output) // 3) S = sum_c z^2 (fp32), r = rsqrt(S), y = z * r, pack-untilized -> CB_OUT (exactly dn_compute.cpp) // CB_W tile t = nb * 32 + k * 4 + j holds W tile (k, 4 nb + j); tiles 64 + n hold the bias (row 0) of // output tile n. #include #include "api/compute/compute_kernel_api.h" #include "api/compute/compute_kernel_hw_startup.h" #include "api/compute/matmul.h" #include "api/compute/pack.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" #ifndef DH_SKIP #define DH_SKIP 0 #endif #ifndef DH_KIN #define DH_KIN 8 // input tiles per tile row (16: merged head conv output, descriptor channels first) #endif #ifndef DH_NW #define DH_NW 72 #endif #ifndef MM_FID #define MM_FID ckernel::MathFidelity::HiFi2 #endif // matmul_block_init / matmul_block with an explicit math fidelity (no dynamic throttle) ALWI void mmf_init(uint32_t in0, uint32_t in1, uint32_t ct) { state_configure(in1, in0, __builtin_LINE()); UNPACK((llk_unpack_AB_matmul_init(in0, in1, 0, ct, 1, 1))); MATH((llk_math_matmul_init(in0, in1, 0, ct, 1))); } ALWI void mmf_block(uint32_t in0, uint32_t in1, uint32_t i0, uint32_t i1, uint32_t d, uint32_t ct) { state_configure(in1, in0, __builtin_LINE()); UNPACK((llk_unpack_AB_matmul(in0, in1, i0, i1, ct, 1, 1))); MATH((llk_math_matmul(d, ct, 1))); } 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 cb_w = get_compile_time_arg_val(6); constexpr uint32_t cb_p = get_compile_time_arg_val(7); constexpr uint32_t cb_z = get_compile_time_arg_val(8); constexpr uint32_t WT = 8, HB = 4, KT = 8, NW = DH_NW, BIAS0 = 64, KIN = DH_KIN; compute_kernel_hw_startup(cb_in, cb_w, cb_p); cb_wait_front(cb_in, TR * KIN); constexpr uint32_t PS = 0; cb_wait_front(cb_w, NW); cb_wait_front(cb_one, 1); #ifdef DH_SCORE_CB // ---- merged head: score 1x1 first (s = bf16(x[:, 256:] @ Ws + bs), the ttnn 1x1 conv steps: fp32 DST over K // in order at MM_FID, bias row-broadcast on the fp32 partials); one tile row at a time packed in place into the // logits shard (CB bound to it) and pushed, so the writer can signal the softmax cores early { constexpr uint32_t cb_s = DH_SCORE_CB, cb_sp = DH_SP_CB, ST = 3, SW0 = 72, SB0 = 96; for (uint32_t r = 0; r < TR; ++r) { reconfig_data_format(cb_in, cb_w); pack_reconfig_data_format(cb_sp); mmf_init(cb_in, cb_w, ST); cb_reserve_back(cb_sp, ST); tile_regs_acquire(); for (uint32_t k = 0; k < KT; ++k) { mmf_block(cb_in, cb_w, r * KIN + KT + k, SW0 + k * ST, 0, ST); } tile_regs_commit(); tile_regs_wait(); for (uint32_t j = 0; j < ST; ++j) { pack_tile(j, cb_sp, j); } tile_regs_release(); cb_push_back(cb_sp, ST); cb_wait_front(cb_sp, ST); reconfig_data_format(cb_sp, cb_w); pack_reconfig_data_format(cb_s); add_bcast_rows_init(cb_sp, cb_w); cb_reserve_back(cb_s, ST); tile_regs_acquire(); for (uint32_t j = 0; j < ST; ++j) { add_tiles_bcast_rows(cb_sp, cb_w, j, SB0 + j, j); } tile_regs_commit(); tile_regs_wait(); for (uint32_t j = 0; j < ST; ++j) { pack_tile(j, cb_s); } tile_regs_release(); cb_push_back(cb_s, ST); cb_pop_front(cb_sp, ST); } } #endif for (uint32_t r = 0; r < TR; ++r) { // ---- 1) matmul -> CB_P (fp32) reconfig_data_format(cb_in, cb_w); pack_reconfig_data_format(cb_p); mmf_init(cb_in, cb_w, HB); cb_reserve_back(cb_p, WT + PS); for (uint32_t nb = 0; nb < WT / HB; ++nb) { tile_regs_acquire(); for (uint32_t k = 0; k < KT; ++k) { if constexpr (!(DH_SKIP & 1)) mmf_block(cb_in, cb_w, r * KIN + k, nb * (KT * HB) + k * HB, 0, HB); } tile_regs_commit(); tile_regs_wait(); for (uint32_t j = 0; j < HB; ++j) { pack_tile(j, cb_p, nb * HB + j); } tile_regs_release(); } cb_push_back(cb_p, WT + PS); // ---- 2) + bias -> CB_Z (bf16) cb_wait_front(cb_p, WT + PS); reconfig_data_format(cb_p, cb_w); pack_reconfig_data_format(cb_z); add_bcast_rows_init(cb_p, cb_w); cb_reserve_back(cb_z, WT); for (uint32_t b = 0; b < WT / HB; ++b) { tile_regs_acquire(); for (uint32_t k = 0; k < HB; ++k) { if constexpr (!(DH_SKIP & 2)) add_tiles_bcast_rows(cb_p, cb_w, b * HB + k, BIAS0 + b * HB + k, k); } tile_regs_commit(); tile_regs_wait(); for (uint32_t k = 0; k < HB; ++k) { pack_tile(k, cb_z); } tile_regs_release(); } cb_push_back(cb_z, WT); cb_pop_front(cb_p, WT + PS); cb_wait_front(cb_z, WT); // ---- 3) x^2 -> CB_SQ (fp32) reconfig_data_format(cb_z, cb_z); pack_reconfig_data_format(cb_sq); mul_init(cb_z, cb_z); 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) { if constexpr (!(DH_SKIP & 4)) mul_tiles(cb_z, cb_z, b * HB + k, b * HB + k, 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) { if constexpr (!(DH_SKIP & 8)) reduce_tile(cb_sq, cb_one, t, 0, 0); } rsqrt_tile_init(); if constexpr (!(DH_SKIP & 16)) 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_z, cb_rs); pack_reconfig_data_format(cb_out); mul_bcast_cols_init(cb_z, 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) { if constexpr (!(DH_SKIP & 32)) mul_tiles_bcast_cols(cb_z, cb_rs, 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); cb_pop_front(cb_z, WT); } }