superpoint-p150 / code /kernels /sp_pool /pool2x2_compute.cpp
changh95's picture
Optimized build (2026-10-03): trace 3.73 -> 0.50 ms
c699c4c verified
Raw History Blame Contribute Delete
2.96 kB
// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// 2x2 / stride-2 max pool of an L1 height-sharded ROW_MAJOR NHWC activation whose shard holds whole
// image rows (ROWS_IN = 2 * ROWS_OUT rows of W pixels, C = 64 channels per pixel), written as an L1
// height-sharded TILE tensor [.., C] (W/2 output pixels per output row). No halo, no data movement:
// * CB_IN is bound to the input shard; one page = 64 pixels of one image row (8 KB) which, read
// as 32 rows of 2*C = 128 channels ("pixel pairs"), tilizes into 4 tiles:
// t0 = even pixel ch 0..31, t1 = even pixel ch 32..63, t2 = odd ch 0..31, t3 = odd ch 32..63,
// tile row r = output pixel r of the 32-pixel output block.
// * out tile j (ch 32j..32j+31) = max(even_y0, odd_y0, even_y1, odd_y1) (SFPU binary max on bf16
// values: exact), packed straight into the output shard (CB_OUT bound to it).
#include <cstdint>
#include "api/compute/compute_kernel_api.h"
#include "api/compute/tile_move_copy.h"
#include "api/compute/tilize.h"
#include "api/compute/binary_max_min.h"
#include "api/compute/eltwise_unary/eltwise_unary.h"
void kernel_main() {
constexpr uint32_t cb_in = get_compile_time_arg_val(0);
constexpr uint32_t cb_t = get_compile_time_arg_val(1);
constexpr uint32_t cb_out = get_compile_time_arg_val(2);
constexpr uint32_t ROWS_OUT = get_compile_time_arg_val(3);
constexpr uint32_t CH = get_compile_time_arg_val(4); // 64-pixel chunks per image row
constexpr uint32_t NT = 8 * CH; // tilized tiles per output row
unary_op_init_common(cb_in, cb_out);
cb_wait_front(cb_in, 2 * ROWS_OUT * CH);
for (uint32_t oy = 0; oy < ROWS_OUT; ++oy) {
cb_reserve_back(cb_t, NT);
tilize_init(cb_in, 4, cb_t);
for (uint32_t c = 0; c < CH; ++c) {
tilize_block(cb_in, 4, cb_t, ((2 * oy) * CH + c) * 4, c * 8);
tilize_block(cb_in, 4, cb_t, ((2 * oy + 1) * CH + c) * 4, c * 8 + 4);
}
tilize_uninit(cb_in, cb_t);
cb_push_back(cb_t, NT);
cb_wait_front(cb_t, NT);
copy_tile_to_dst_init_short(cb_t);
binary_max_tile_init();
cb_reserve_back(cb_out, 2 * CH);
for (uint32_t c = 0; c < CH; ++c) {
for (uint32_t j = 0; j < 2; ++j) {
tile_regs_acquire();
copy_tile(cb_t, c * 8 + j, 0);
copy_tile(cb_t, c * 8 + 2 + j, 1);
copy_tile(cb_t, c * 8 + 4 + j, 2);
copy_tile(cb_t, c * 8 + 6 + j, 3);
binary_max_tile(0, 1, 0);
binary_max_tile(2, 3, 2);
binary_max_tile(0, 2, 0);
tile_regs_commit();
tile_regs_wait();
pack_tile<true>(0, cb_out, c * 2 + j);
tile_regs_release();
}
}
cb_push_back(cb_out, 2 * CH);
cb_pop_front(cb_t, NT);
}
}