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