File size: 3,916 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 | // 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);
}
}
|