File size: 8,076 Bytes
c699c4c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 | // 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);
}
|