File size: 8,076 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
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// Block-0 conv_b compute (layout: models/tt/conv_cell.py). Per output tile row TR and output pixel
// p: S[2 (p & 1) + n] (fp32 DST, n = output channel half) = sum over taps (ky, kx) and input channel
// halves h of X(ky, kx, h) @ W[t, h][n], K order = ttnn's im2col order (ky, kx, channel).
// The bias is added exactly like ttnn's conv (fuse_bias path of conv_bmm_tilize.cpp): the fp32
// partial sums are packed to CB_P, then add_tiles_bcast_rows(partials, bias) and ReLU on pack, so the
// conv values are those of ttnn.conv2d. The horizontal half of the following 2x2 max pool is fused:
// for each pixel pair (2j, 2j + 1) the SFPU max of the two fp32 sums is taken in DST right after the
// matmuls, BEFORE the bias phase: every later step (unpack to tf32, + bias, bf16 rounding, ReLU) is
// monotone non-decreasing, so max-then-bias equals bias-then-max bit for bit, and only half the
// partials are packed / bias-added. The max goes to output tile (TR, 2 j + n) of the half-pooled
// [cells, 4 px * 64 ch] map. The bias phase (which needs an unpack/pack data-format reconfig and a
// matmul re-init afterwards) runs once per TRB output tile rows instead of once per row.
// CB_W: W[ky, kx, h][n] at 12 (ky + 1) + 6 h + 2 (1 - kx) + n (kx descending, so the taps of one
// input tile for consecutive output pixels are consecutive in1 tiles); CB_B: bias tiles (row 0) n = 0, 1.
#include <cstdint>
#include "api/compute/compute_kernel_api.h"
#include "api/compute/matmul.h"
#include "api/compute/compute_kernel_hw_startup.h"
#include "api/compute/pack.h"
#include "api/compute/bcast.h"
#include "api/compute/binary_max_min.h"
#include "api/compute/reconfig_data_format.h"
#ifdef PROFZ
#include "tools/profiler/kernel_profiler.hpp"
#define ZONE(n) DeviceZoneScopedN(n)
#else
#define ZONE(n)
#endif

void kernel_main() {
    constexpr uint32_t cb_x = get_compile_time_arg_val(0);
    constexpr uint32_t cb_s0 = get_compile_time_arg_val(1);
    constexpr uint32_t cb_s1 = get_compile_time_arg_val(2);
    constexpr uint32_t cb_w = get_compile_time_arg_val(3);
    constexpr uint32_t cb_out = get_compile_time_arg_val(4);
    constexpr uint32_t W_TILES = get_compile_time_arg_val(5);
    constexpr uint32_t cb_p = get_compile_time_arg_val(6);
    constexpr uint32_t cb_b = get_compile_time_arg_val(7);
    constexpr uint32_t TRB = get_compile_time_arg_val(8);  // output tile rows per bias phase
    constexpr uint32_t PIX = get_compile_time_arg_val(9);  // output pixels per DST block (2: half sync, 4: full sync)
    constexpr uint32_t P = get_compile_time_arg_val(10);      // pixels per cell
    constexpr uint32_t TRS = get_compile_time_arg_val(11);    // output tile rows per core
    constexpr uint32_t GROUP = get_compile_time_arg_val(12);  // shifted-tap tiles per tile row (2 QT + 12)
    constexpr bool HMAX = get_compile_time_arg_val(13) == 1;  // fused horizontal half of the 2x2 pool
    constexpr uint32_t QT = 2 * P;
    constexpr uint32_t OQT = HMAX ? P : QT;  // output tiles per tile row
    constexpr uint32_t E0 = 2 * QT;          // first edge slot of a group
    static_assert(TRS % TRB == 0, "TRB must divide TRS");
    static_assert(P % PIX == 0, "PIX must divide P");
    constexpr uint32_t PT = OQT * TRB;  // fp32 partial tiles per bias phase

    compute_kernel_hw_startup<SrcOrder::Reverse>(cb_x, cb_w, cb_p);
    cb_wait_front(cb_x, TRS * QT);
    cb_wait_front(cb_w, W_TILES);
    cb_wait_front(cb_b, 2);
    cb_reserve_back(cb_out, TRS * OQT);
    for (uint32_t tr0 = 0; tr0 < TRS; tr0 += TRB) {
        // ---- matmul phase: per row and pixel pair, fp32 max of the 2 pixels x 2 channel halves -> CB_P
        {
        ZONE("CB0_MM");
        matmul_block_init(cb_x, cb_w, false, 2, 1, 1);
        if constexpr (HMAX) {
            binary_max_tile_init();
        }
        cb_reserve_back(cb_p, PT);
        for (uint32_t tri = 0; tri < TRB; ++tri) {
            const uint32_t tr = tr0 + tri;
            const uint32_t cb_s = ((BMASK_DEF >> tr) & 1) ? cb_s1 : cb_s0;  // group built by PROC 1 / PROC 0
            cb_wait_front(cb_s, GROUP);
            for (uint32_t p0 = 0; p0 < P; p0 += PIX) {
                tile_regs_acquire();
                for (int32_t ky = -1; ky <= 1; ++ky) {
                    // one in0 tile (source pixel pp, channel half h) feeds every output pixel of this
                    // block it is a tap of: ct = 2 x (1..3) output tiles in one matmul_block call.
                    // Per output pixel the accumulation order stays (ky, kx, h) as in af7b261.
                    for (int32_t pp = (int32_t)p0 - 1; pp <= (int32_t)(p0 + PIX); ++pp) {
                        const int32_t lo = pp - 1 > (int32_t)p0 ? pp - 1 : (int32_t)p0;
                        const int32_t hi = pp + 1 < (int32_t)(p0 + PIX - 1) ? pp + 1 : (int32_t)(p0 + PIX - 1);
                        const uint32_t ct = 2 * (uint32_t)(hi - lo + 1);
                        const uint32_t w0 = 12 * (uint32_t)(ky + 1) + 2 * (uint32_t)(1 - pp + lo);
                        const uint32_t d = 2 * (uint32_t)(lo - (int32_t)p0);
                        uint32_t cb, idx;
                        if (pp >= 0 && pp <= (int32_t)P - 1) {
                            if (ky == 0) {
                                cb = cb_x;
                                idx = tr * QT + 2 * (uint32_t)pp;
                            } else {
                                cb = cb_s;
                                idx = (ky < 0 ? 0 : QT) + 2 * (uint32_t)pp;
                            }
                        } else {
                            cb = cb_s;
                            const uint32_t e = ky == 0 ? (pp < 0 ? 0 : 1) : (ky < 0 ? (pp < 0 ? 2 : 3) : (pp < 0 ? 4 : 5));
                            idx = E0 + 2 * e;
                        }
                        matmul_block(cb, cb_w, idx, w0, d, false, ct, 1, 1);
                        matmul_block(cb, cb_w, idx + 1, w0 + 6, d, false, ct, 1, 1);
                    }
                }
                if constexpr (HMAX) {
                    for (uint32_t q = 0; q < PIX / 2; ++q) {
                        binary_max_tile(4 * q, 4 * q + 2, 4 * q);
                        binary_max_tile(4 * q + 1, 4 * q + 3, 4 * q + 1);
                    }
                }
                tile_regs_commit();
                tile_regs_wait();
                if constexpr (HMAX) {
                    for (uint32_t q = 0; q < PIX / 2; ++q) {
                        pack_tile<true>(4 * q, cb_p, OQT * tri + p0 + 2 * q);
                        pack_tile<true>(4 * q + 1, cb_p, OQT * tri + p0 + 2 * q + 1);
                    }
                } else {
                    for (uint32_t d = 0; d < 2 * PIX; ++d) {
                        pack_tile<true>(d, cb_p, OQT * tri + 2 * p0 + d);
                    }
                }
                tile_regs_release();
            }
            cb_pop_front(cb_s, GROUP);
        }
        cb_push_back(cb_p, PT);
        }
        ZONE("CB0_BIAS");
        // ---- bias phase (ttnn conv order): S + bias (row broadcast), ReLU pack
        cb_wait_front(cb_p, PT);
        PACK((llk_pack_relu_config(ReluConfig::zero())));
        pack_reconfig_data_format(cb_p, cb_out);
        reconfig_data_format(cb_w, cb_p, cb_x, cb_b);
        add_bcast_rows_init(cb_p, cb_b);
        for (uint32_t k = 0; k < PT; k += 4) {
            tile_regs_acquire();
            for (uint32_t d = 0; d < 4; ++d) {
                add_tiles_bcast_rows(cb_p, cb_b, k + d, d & 1, d);
            }
            tile_regs_commit();
            tile_regs_wait();
            for (uint32_t d = 0; d < 4; ++d) {
                pack_tile<true>(d, cb_out, tr0 * OQT + k + d);
            }
            tile_regs_release();
        }
        cb_pop_front(cb_p, PT);
        PACK((llk_pack_relu_config(ReluConfig::none())));
        reconfig_data_format(cb_p, cb_w, cb_b, cb_x);
        pack_reconfig_data_format(cb_out, cb_p);
    }
    cb_push_back(cb_out, TRS * QT / 2);
}