superpoint-p150 / code /kernels /sp_conv /cb0_compute.cpp
changh95's picture
Optimized build (2026-10-03): trace 3.73 -> 0.50 ms
c699c4c verified
Raw History Blame Contribute Delete
8.08 kB
// 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 <cstdint>
#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<SrcOrder::Reverse>(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<true>(4 * q, cb_p, OQT * tri + p0 + 2 * q);
pack_tile<true>(4 * q + 1, cb_p, OQT * tri + p0 + 2 * q + 1);
}
} else {
for (uint32_t d = 0; d < 2 * PIX; ++d) {
pack_tile<true>(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<true>(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);
}