// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc. // SPDX-License-Identifier: Apache-2.0 // // Block-0 conv_b compute (layout: models/tt/conv_cell.py). Per output tile row TR and output pixel // p: S[2 (p & 1) + n] (fp32 DST, n = output channel half) = sum over taps (ky, kx) and input channel // halves h of X(ky, kx, h) @ W[t, h][n], K order = ttnn's im2col order (ky, kx, channel). // The bias is added exactly like ttnn's conv (fuse_bias path of conv_bmm_tilize.cpp): the fp32 // partial sums are packed to CB_P, then add_tiles_bcast_rows(partials, bias) and ReLU on pack, so the // conv values are those of ttnn.conv2d. The horizontal half of the following 2x2 max pool is fused: // for each pixel pair (2j, 2j + 1) the SFPU max of the two fp32 sums is taken in DST right after the // matmuls, BEFORE the bias phase: every later step (unpack to tf32, + bias, bf16 rounding, ReLU) is // monotone non-decreasing, so max-then-bias equals bias-then-max bit for bit, and only half the // partials are packed / bias-added. The max goes to output tile (TR, 2 j + n) of the half-pooled // [cells, 4 px * 64 ch] map. The bias phase (which needs an unpack/pack data-format reconfig and a // matmul re-init afterwards) runs once per TRB output tile rows instead of once per row. // CB_W: W[ky, kx, h][n] at 12 (ky + 1) + 6 h + 2 (1 - kx) + n (kx descending, so the taps of one // input tile for consecutive output pixels are consecutive in1 tiles); CB_B: bias tiles (row 0) n = 0, 1. #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/binary_max_min.h" #include "api/compute/reconfig_data_format.h" #ifdef PROFZ #include "tools/profiler/kernel_profiler.hpp" #define ZONE(n) DeviceZoneScopedN(n) #else #define ZONE(n) #endif void kernel_main() { constexpr uint32_t cb_x = get_compile_time_arg_val(0); constexpr uint32_t cb_s0 = get_compile_time_arg_val(1); constexpr uint32_t cb_s1 = get_compile_time_arg_val(2); constexpr uint32_t cb_w = get_compile_time_arg_val(3); constexpr uint32_t cb_out = get_compile_time_arg_val(4); constexpr uint32_t W_TILES = get_compile_time_arg_val(5); constexpr uint32_t cb_p = get_compile_time_arg_val(6); constexpr uint32_t cb_b = get_compile_time_arg_val(7); constexpr uint32_t TRB = get_compile_time_arg_val(8); // output tile rows per bias phase constexpr uint32_t PIX = get_compile_time_arg_val(9); // output pixels per DST block (2: half sync, 4: full sync) constexpr uint32_t P = get_compile_time_arg_val(10); // pixels per cell constexpr uint32_t TRS = get_compile_time_arg_val(11); // output tile rows per core constexpr uint32_t GROUP = get_compile_time_arg_val(12); // shifted-tap tiles per tile row (2 QT + 12) constexpr bool HMAX = get_compile_time_arg_val(13) == 1; // fused horizontal half of the 2x2 pool constexpr uint32_t QT = 2 * P; constexpr uint32_t OQT = HMAX ? P : QT; // output tiles per tile row constexpr uint32_t E0 = 2 * QT; // first edge slot of a group static_assert(TRS % TRB == 0, "TRB must divide TRS"); static_assert(P % PIX == 0, "PIX must divide P"); constexpr uint32_t PT = OQT * TRB; // fp32 partial tiles per bias phase compute_kernel_hw_startup(cb_x, cb_w, cb_p); cb_wait_front(cb_x, TRS * QT); cb_wait_front(cb_w, W_TILES); cb_wait_front(cb_b, 2); cb_reserve_back(cb_out, TRS * OQT); for (uint32_t tr0 = 0; tr0 < TRS; tr0 += TRB) { // ---- matmul phase: per row and pixel pair, fp32 max of the 2 pixels x 2 channel halves -> CB_P { ZONE("CB0_MM"); matmul_block_init(cb_x, cb_w, false, 2, 1, 1); if constexpr (HMAX) { binary_max_tile_init(); } cb_reserve_back(cb_p, PT); for (uint32_t tri = 0; tri < TRB; ++tri) { const uint32_t tr = tr0 + tri; const uint32_t cb_s = ((BMASK_DEF >> tr) & 1) ? cb_s1 : cb_s0; // group built by PROC 1 / PROC 0 cb_wait_front(cb_s, GROUP); for (uint32_t p0 = 0; p0 < P; p0 += PIX) { tile_regs_acquire(); for (int32_t ky = -1; ky <= 1; ++ky) { // one in0 tile (source pixel pp, channel half h) feeds every output pixel of this // block it is a tap of: ct = 2 x (1..3) output tiles in one matmul_block call. // Per output pixel the accumulation order stays (ky, kx, h) as in af7b261. for (int32_t pp = (int32_t)p0 - 1; pp <= (int32_t)(p0 + PIX); ++pp) { const int32_t lo = pp - 1 > (int32_t)p0 ? pp - 1 : (int32_t)p0; const int32_t hi = pp + 1 < (int32_t)(p0 + PIX - 1) ? pp + 1 : (int32_t)(p0 + PIX - 1); const uint32_t ct = 2 * (uint32_t)(hi - lo + 1); const uint32_t w0 = 12 * (uint32_t)(ky + 1) + 2 * (uint32_t)(1 - pp + lo); const uint32_t d = 2 * (uint32_t)(lo - (int32_t)p0); uint32_t cb, idx; if (pp >= 0 && pp <= (int32_t)P - 1) { if (ky == 0) { cb = cb_x; idx = tr * QT + 2 * (uint32_t)pp; } else { cb = cb_s; idx = (ky < 0 ? 0 : QT) + 2 * (uint32_t)pp; } } else { cb = cb_s; const uint32_t e = ky == 0 ? (pp < 0 ? 0 : 1) : (ky < 0 ? (pp < 0 ? 2 : 3) : (pp < 0 ? 4 : 5)); idx = E0 + 2 * e; } matmul_block(cb, cb_w, idx, w0, d, false, ct, 1, 1); matmul_block(cb, cb_w, idx + 1, w0 + 6, d, false, ct, 1, 1); } } if constexpr (HMAX) { for (uint32_t q = 0; q < PIX / 2; ++q) { binary_max_tile(4 * q, 4 * q + 2, 4 * q); binary_max_tile(4 * q + 1, 4 * q + 3, 4 * q + 1); } } tile_regs_commit(); tile_regs_wait(); if constexpr (HMAX) { for (uint32_t q = 0; q < PIX / 2; ++q) { pack_tile(4 * q, cb_p, OQT * tri + p0 + 2 * q); pack_tile(4 * q + 1, cb_p, OQT * tri + p0 + 2 * q + 1); } } else { for (uint32_t d = 0; d < 2 * PIX; ++d) { pack_tile(d, cb_p, OQT * tri + 2 * p0 + d); } } tile_regs_release(); } cb_pop_front(cb_s, GROUP); } cb_push_back(cb_p, PT); } ZONE("CB0_BIAS"); // ---- bias phase (ttnn conv order): S + bias (row broadcast), ReLU pack cb_wait_front(cb_p, PT); PACK((llk_pack_relu_config(ReluConfig::zero()))); pack_reconfig_data_format(cb_p, cb_out); reconfig_data_format(cb_w, cb_p, cb_x, cb_b); add_bcast_rows_init(cb_p, cb_b); for (uint32_t k = 0; k < PT; k += 4) { tile_regs_acquire(); for (uint32_t d = 0; d < 4; ++d) { add_tiles_bcast_rows(cb_p, cb_b, k + d, d & 1, d); } tile_regs_commit(); tile_regs_wait(); for (uint32_t d = 0; d < 4; ++d) { pack_tile(d, cb_out, tr0 * OQT + k + d); } tile_regs_release(); } cb_pop_front(cb_p, PT); PACK((llk_pack_relu_config(ReluConfig::none()))); reconfig_data_format(cb_p, cb_w, cb_b, cb_x); pack_reconfig_data_format(cb_out, cb_p); } cb_push_back(cb_out, TRS * QT / 2); }