File size: 3,936 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
// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// Cell-tile 2x2 pool, vertical half (models/tt/conv_cell.py PoolCell), data movement. Input: the
// half-pooled TILE map [cells, QTI tiles per tile row] of this core (CB_A bound to it), 80 cells per
// image row. Pooled row Y (Y % 2 == PROC builds its operands) reads the even image row (cells
// 160Y + 32k .., k = 0..2) and the odd row below (+80 cells): CB_B gets, per k, the QTI tiles of the
// odd row = rows 16..31 of tile row 5Y + k + 2 (faces 2, 3 -> 0, 1) and rows 0..15 of tile row
// 5Y + k + 3 (faces 0, 1 -> 2, 3), the latter only when it exists (k = 2 of the last pooled row
// reads past the shard: those output rows are not used).
// PROC 1 also moves each result block of CB_U (compute order: Y, k) into the output (CB_O bound):
//   MODE 0 (rm):   untilized [32 cells, QTI * 32] block, valid rows (32, or 16 for k = 2) at cell
//                  Y*80 + 32k of the ROW_MAJOR output (QTI * 64 B per cell)
//   MODE 1 (cell): QTI pooled tiles, 16-row halves copied to pooled cell Y*80 + 32k (+16) of the
//                  TILE output (QTI tiles per tile row)
// CT args: PROC, cb_a, cb_b, cb_u, cb_o, QTI, TRS, PROW, MODE
#include <stdint.h>
#include "api/dataflow/dataflow_api.h"

void kernel_main() {
    constexpr uint32_t PROC = get_compile_time_arg_val(0);
    constexpr uint32_t cb_a = get_compile_time_arg_val(1);
    constexpr uint32_t cb_b = get_compile_time_arg_val(2);
    constexpr uint32_t cb_u = get_compile_time_arg_val(3);
    constexpr uint32_t cb_o = get_compile_time_arg_val(4);
    constexpr uint32_t QTI = get_compile_time_arg_val(5);
    constexpr uint32_t TRS = get_compile_time_arg_val(6);
    constexpr uint32_t PROW = get_compile_time_arg_val(7);
    constexpr uint32_t MODE = get_compile_time_arg_val(8);
    constexpr uint32_t TILE = 2048;
    if constexpr (PROC == 0) {
        cb_reserve_back(cb_a, TRS * QTI);
        cb_push_back(cb_a, TRS * QTI);
    }
    const uint64_t a_noc = get_noc_addr(my_x[noc_index], my_y[noc_index], get_write_ptr(cb_a));
    for (uint32_t Y = PROC; Y < PROW; Y += 2) {
        cb_reserve_back(cb_b, 3 * QTI);
        const uint32_t b_l1 = get_write_ptr(cb_b);
        for (uint32_t k = 0; k < 3; ++k) {
            const uint32_t t0 = 5 * Y + k + 2;
            for (uint32_t q = 0; q < QTI; ++q) {
                const uint32_t dst = b_l1 + (k * QTI + q) * TILE;
                noc_async_read(a_noc + (t0 * QTI + q) * TILE + 1024, dst, 1024);
                if (t0 + 1 < TRS) {
                    noc_async_read(a_noc + ((t0 + 1) * QTI + q) * TILE, dst + 1024, 1024);
                }
            }
        }
        noc_async_read_barrier();
        cb_push_back(cb_b, 3 * QTI);
    }
    if constexpr (PROC == 1) {
        const uint32_t o_l1 = get_write_ptr(cb_o);
        const uint64_t me = get_noc_addr(my_x[noc_index], my_y[noc_index], 0);
        for (uint32_t Y = 0; Y < PROW; ++Y) {
            for (uint32_t k = 0; k < 3; ++k) {
                cb_wait_front(cb_u, QTI);
                const uint32_t u = get_read_ptr(cb_u);
                const uint32_t rows = k == 2 ? 16 : 32;
                const uint32_t c0 = Y * 80 + 32 * k;
                if constexpr (MODE == 0) {
                    noc_async_read(me | u, o_l1 + c0 * QTI * 64, rows * QTI * 64);
                } else {
                    for (uint32_t hs = 0; hs < rows / 16; ++hs) {
                        const uint32_t dc = c0 + 16 * hs;
                        for (uint32_t q = 0; q < QTI; ++q) {
                            noc_async_read(me | (u + q * TILE + hs * 1024),
                                           o_l1 + ((dc >> 5) * QTI + q) * TILE + ((dc >> 4) & 1) * 1024, 1024);
                        }
                    }
                }
                noc_async_read_barrier();
                cb_pop_front(cb_u, QTI);
            }
        }
    }
}