File size: 3,610 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
// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// RowConv compute (models/tt/row_conv.py): ttnn conv_bmm_tilize.cpp's height-sharded order with
// packer_l1_acc and fused bias. Per K block ky (3*CT tiles in K order kx, c): fp32 DST accumulates the
// tile products of output tile row m and NL output channel tiles; the block result is packed to fp32
// partials (ky = 0 overwrite, else packer L1 accumulate). Then partials + bias (row broadcast) with
// ReLU on pack -> bf16 output tiles (m, n) at m * NL + n.
#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/reconfig_data_format.h"

void kernel_main() {
    constexpr uint32_t cb_a0 = get_compile_time_arg_val(0);
    constexpr uint32_t cb_a1 = get_compile_time_arg_val(1);
    constexpr uint32_t cb_w = get_compile_time_arg_val(2);
    constexpr uint32_t cb_b = get_compile_time_arg_val(3);
    constexpr uint32_t cb_p = get_compile_time_arg_val(4);
    constexpr uint32_t cb_o = get_compile_time_arg_val(5);
    constexpr uint32_t NL = get_compile_time_arg_val(6);
    constexpr uint32_t CH = get_compile_time_arg_val(7);
    constexpr uint32_t CT = get_compile_time_arg_val(8);
    constexpr uint32_t BLK_H = get_compile_time_arg_val(9);
    constexpr bool RELU = get_compile_time_arg_val(10) == 1;
    constexpr uint32_t KB = 3 * CT;  // K tiles per block
    static_assert(NL <= 4, "fp32 half-sync DST holds 4 tiles");

    compute_kernel_hw_startup<SrcOrder::Reverse>(cb_a0, cb_w, cb_p);
    matmul_block_init(cb_a0, cb_w, false, NL, 1, 1);
    cb_reserve_back(cb_p, 3 * NL);
    for (uint32_t ky = 0; ky < 3; ++ky) {
        cb_wait_front(cb_a0, (ky + 1) * BLK_H);
        cb_wait_front(cb_a1, (ky + 1) * BLK_H);
        cb_wait_front(cb_w, (ky + 1) * KB * NL);
        pack_reconfig_l1_acc(ky > 0 ? 1 : 0);
        for (uint32_t m = 0; m < 3; ++m) {
            tile_regs_acquire();
            for (uint32_t kx = 0; kx < 3; ++kx) {
                for (uint32_t c = 0; c < CT; ++c) {
                    const uint32_t cb = c < CH ? cb_a0 : cb_a1;
                    const uint32_t idx = ky * BLK_H + (m * 3 + kx) * CH + (c < CH ? c : c - CH);
                    matmul_block(cb, cb_w, idx, (ky * KB + kx * CT + c) * NL, 0, false, NL, 1, 1);
                }
            }
            tile_regs_commit();
            tile_regs_wait();
            for (uint32_t n = 0; n < NL; ++n) {
                pack_tile<true>(n, cb_p, m * NL + n);
            }
            tile_regs_release();
        }
    }
    pack_reconfig_l1_acc(0);
    cb_push_back(cb_p, 3 * NL);

    cb_wait_front(cb_p, 3 * NL);
    cb_wait_front(cb_b, NL);
    cb_reserve_back(cb_o, 3 * NL);
    if constexpr (RELU) {
        PACK((llk_pack_relu_config(ReluConfig::zero())));
    }
    pack_reconfig_data_format(cb_p, cb_o);
    reconfig_data_format(cb_w, cb_p, cb_a0, cb_b);
    add_bcast_rows_init(cb_p, cb_b);
    for (uint32_t m = 0; m < 3; ++m) {
        tile_regs_acquire();
        for (uint32_t n = 0; n < NL; ++n) {
            add_tiles_bcast_rows(cb_p, cb_b, m * NL + n, n, n);
        }
        tile_regs_commit();
        tile_regs_wait();
        for (uint32_t n = 0; n < NL; ++n) {
            pack_tile<true>(n, cb_o, m * NL + n);
        }
        tile_regs_release();
    }
    if constexpr (RELU) {
        PACK((llk_pack_relu_config(ReluConfig::none())));
    }
    cb_pop_front(cb_p, 3 * NL);
    cb_push_back(cb_o, 3 * NL);
}