// 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 #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(0, cb_out, c * 2 + j); tile_regs_release(); } } cb_push_back(cb_out, 2 * CH); cb_pop_front(cb_t, NT); } }