superpoint-p150 / code /kernels /sp_conv /pc0_compute.cpp
changh95's picture
Optimized build (2026-10-03): trace 3.73 -> 0.50 ms
c699c4c verified
Raw History Blame Contribute Delete
3.92 kB
// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// Cell-tile 2x2 pool, vertical half: for pooled row Y and tile row k, dst[o] = max(A(5Y + k, o),
// B(k, o)) for the QTI tiles o (SFPU max on bf16: exact; full-sync bf16 DST: A in 0..QTI-1, B in
// QTI..2QTI-1), then either pack-untilized as one [32 rows, QTI tiles] row-major block (MODE 0)
// or packed as QTI tiles (MODE 1) into CB_U.
#include <cstdint>
#include "api/compute/compute_kernel_api.h"
#include "api/compute/tile_move_copy.h"
#include "api/compute/binary_max_min.h"
#include "api/compute/eltwise_unary/eltwise_unary.h"
#include "api/compute/pack_untilize.h"
void kernel_main() {
constexpr uint32_t cb_a = get_compile_time_arg_val(0);
constexpr uint32_t cb_b0 = get_compile_time_arg_val(1);
constexpr uint32_t cb_b1 = get_compile_time_arg_val(2);
constexpr uint32_t cb_u = get_compile_time_arg_val(3);
constexpr uint32_t QTI = get_compile_time_arg_val(4);
constexpr uint32_t TRS = get_compile_time_arg_val(5);
constexpr uint32_t PROW = get_compile_time_arg_val(6);
constexpr uint32_t MODE = get_compile_time_arg_val(7);
static_assert(2 * QTI <= 16, "full-sync bf16 DST holds 16 tiles");
unary_op_init_common(cb_a, cb_u);
if constexpr (MODE == 0) {
pack_untilize_dest_init<QTI>(cb_u);
}
cb_wait_front(cb_a, TRS * QTI);
for (uint32_t Y = 0; Y < PROW; ++Y) {
const uint32_t cb_b = (Y & 1) ? cb_b1 : cb_b0;
cb_wait_front(cb_b, 3 * QTI);
for (uint32_t k = 0; k < 3; ++k) {
#ifdef PC_HALF
// MODE 1, half-sync DST: two blocks of QTI / 2 tile pairs, the pack of one overlaps the next block's
// unpack + max (same SFPU max per tile pair: bit-identical)
static_assert(MODE == 1 && QTI % 2 == 0, "PC_HALF: tile output only");
constexpr uint32_t HQ = QTI / 2;
for (uint32_t h = 0; h < 2; ++h) {
tile_regs_acquire();
copy_tile_to_dst_init_short(cb_a);
for (uint32_t o = 0; o < HQ; ++o) {
copy_tile(cb_a, (5 * Y + k) * QTI + h * HQ + o, o);
}
copy_tile_to_dst_init_short(cb_b);
for (uint32_t o = 0; o < HQ; ++o) {
copy_tile(cb_b, k * QTI + h * HQ + o, HQ + o);
}
binary_max_tile_init();
for (uint32_t o = 0; o < HQ; ++o) {
binary_max_tile(o, HQ + o, o);
}
tile_regs_commit();
cb_reserve_back(cb_u, HQ);
tile_regs_wait();
for (uint32_t o = 0; o < HQ; ++o) {
pack_tile(o, cb_u);
}
tile_regs_release();
cb_push_back(cb_u, HQ);
}
continue;
#endif
tile_regs_acquire();
copy_tile_to_dst_init_short(cb_a);
for (uint32_t o = 0; o < QTI; ++o) {
copy_tile(cb_a, (5 * Y + k) * QTI + o, o);
}
copy_tile_to_dst_init_short(cb_b);
for (uint32_t o = 0; o < QTI; ++o) {
copy_tile(cb_b, k * QTI + o, QTI + o);
}
binary_max_tile_init();
for (uint32_t o = 0; o < QTI; ++o) {
binary_max_tile(o, QTI + o, o);
}
tile_regs_commit();
cb_reserve_back(cb_u, QTI);
tile_regs_wait();
if constexpr (MODE == 0) {
pack_untilize_dest<QTI>(cb_u);
} else {
for (uint32_t o = 0; o < QTI; ++o) {
pack_tile(o, cb_u);
}
}
tile_regs_release();
cb_push_back(cb_u, QTI);
}
cb_pop_front(cb_b, 3 * QTI);
}
if constexpr (MODE == 0) {
pack_untilize_uninit(cb_u);
}
}