// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc. // SPDX-License-Identifier: Apache-2.0 // // RowConv compute (models/tt/row_conv.py): ttnn conv_bmm_tilize.cpp's height-sharded order with // packer_l1_acc and fused bias. Per K block ky (3*CT tiles in K order kx, c): fp32 DST accumulates the // tile products of output tile row m and NL output channel tiles; the block result is packed to fp32 // partials (ky = 0 overwrite, else packer L1 accumulate). Then partials + bias (row broadcast) with // ReLU on pack -> bf16 output tiles (m, n) at m * NL + n. #include #include "api/compute/compute_kernel_api.h" #include "api/compute/matmul.h" #include "api/compute/compute_kernel_hw_startup.h" #include "api/compute/pack.h" #include "api/compute/bcast.h" #include "api/compute/reconfig_data_format.h" void kernel_main() { constexpr uint32_t cb_a0 = get_compile_time_arg_val(0); constexpr uint32_t cb_a1 = get_compile_time_arg_val(1); constexpr uint32_t cb_w = get_compile_time_arg_val(2); constexpr uint32_t cb_b = get_compile_time_arg_val(3); constexpr uint32_t cb_p = get_compile_time_arg_val(4); constexpr uint32_t cb_o = get_compile_time_arg_val(5); constexpr uint32_t NL = get_compile_time_arg_val(6); constexpr uint32_t CH = get_compile_time_arg_val(7); constexpr uint32_t CT = get_compile_time_arg_val(8); constexpr uint32_t BLK_H = get_compile_time_arg_val(9); constexpr bool RELU = get_compile_time_arg_val(10) == 1; constexpr uint32_t KB = 3 * CT; // K tiles per block static_assert(NL <= 4, "fp32 half-sync DST holds 4 tiles"); compute_kernel_hw_startup(cb_a0, cb_w, cb_p); matmul_block_init(cb_a0, cb_w, false, NL, 1, 1); cb_reserve_back(cb_p, 3 * NL); for (uint32_t ky = 0; ky < 3; ++ky) { cb_wait_front(cb_a0, (ky + 1) * BLK_H); cb_wait_front(cb_a1, (ky + 1) * BLK_H); cb_wait_front(cb_w, (ky + 1) * KB * NL); pack_reconfig_l1_acc(ky > 0 ? 1 : 0); for (uint32_t m = 0; m < 3; ++m) { tile_regs_acquire(); for (uint32_t kx = 0; kx < 3; ++kx) { for (uint32_t c = 0; c < CT; ++c) { const uint32_t cb = c < CH ? cb_a0 : cb_a1; const uint32_t idx = ky * BLK_H + (m * 3 + kx) * CH + (c < CH ? c : c - CH); matmul_block(cb, cb_w, idx, (ky * KB + kx * CT + c) * NL, 0, false, NL, 1, 1); } } tile_regs_commit(); tile_regs_wait(); for (uint32_t n = 0; n < NL; ++n) { pack_tile(n, cb_p, m * NL + n); } tile_regs_release(); } } pack_reconfig_l1_acc(0); cb_push_back(cb_p, 3 * NL); cb_wait_front(cb_p, 3 * NL); cb_wait_front(cb_b, NL); cb_reserve_back(cb_o, 3 * NL); if constexpr (RELU) { PACK((llk_pack_relu_config(ReluConfig::zero()))); } pack_reconfig_data_format(cb_p, cb_o); reconfig_data_format(cb_w, cb_p, cb_a0, cb_b); add_bcast_rows_init(cb_p, cb_b); for (uint32_t m = 0; m < 3; ++m) { tile_regs_acquire(); for (uint32_t n = 0; n < NL; ++n) { add_tiles_bcast_rows(cb_p, cb_b, m * NL + n, n, n); } tile_regs_commit(); tile_regs_wait(); for (uint32_t n = 0; n < NL; ++n) { pack_tile(n, cb_o, m * NL + n); } tile_regs_release(); } if constexpr (RELU) { PACK((llk_pack_relu_config(ReluConfig::none()))); } cb_pop_front(cb_p, 3 * NL); cb_push_back(cb_o, 3 * NL); }