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);
    }
}