Download code/kernels/sp_conv/cb0_compute.cpp from changh95/superpoint-p150: direct link, hf CLI and curl.
- Browser
- Download file 8.08 kB
-
https://huggingface.co/changh95/superpoint-p150/resolve/main/code/kernels/sp_conv/cb0_compute.cpp
- Command line
-
hf download hf://changh95/superpoint-p150/code/kernels/sp_conv/cb0_compute.cpp
-
curl -L -o cb0_compute.cpp https://huggingface.co/changh95/superpoint-p150/resolve/main/code/kernels/sp_conv/cb0_compute.cpp
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. | |
| 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); | |
| } | |