File size: 2,842 Bytes
be62f78
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
// SPDX-License-Identifier: Apache-2.0
// Split-matmul operand build (tt/kcat_kernel.py), writer (RISCV_1): x tile (r, j) of [.., M, K] goes to
// X' (r, j) and (r, Kt + j) as x_hi and to (r, 2 Kt + j) as x_lo; the core holding (r, 0) also writes the ones tile
// (r, 3 Kt) (columns 0 and 1 = 1.0: the bias rows of the concatenated weight) and the zero tiles
// (r, 3 Kt + 1 .. Ktp - 1) that pad K' to a multiple of the matmul's K block.
// CT args: [0] Kt, [1] per-core RT-arg count (P3), [2] Ktp (row tiles of X'), then the TensorAccessorArgs of X'.
// Common RT args: [out_addr]. Per-core RT args: [t0, n].
#include <cstdint>

#include "api/dataflow/dataflow_api.h"

constexpr uint32_t Kt = get_compile_time_arg_val(0);
constexpr uint32_t OKt = get_compile_time_arg_val(2);
constexpr auto o_args = TensorAccessorArgs<3>();
constexpr uint32_t cb_hi = 16, cb_lo = 17, cb_one = 18, cb_zero = 19;
constexpr uint32_t TB = 4096;

void kernel_main() {
    const uint32_t o_addr = get_common_arg_val<uint32_t>(0);
    const uint32_t t0 = get_arg_val<uint32_t>(0);
    const uint32_t n = get_arg_val<uint32_t>(1);
    const auto o = TensorAccessor(o_args, o_addr, TB);
    // ones tile: element (i, c) at face (i / 16) * 2 + (c / 16), offset (i % 16) * 16 + c % 16
    cb_reserve_back(cb_one, 1);
    const uint32_t one_ptr = get_write_ptr(cb_one);
    {
        auto* p = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(one_ptr);
        for (uint32_t i = 0; i < 1024; ++i) {
            p[i] = 0;
        }
        for (uint32_t i = 0; i < 32; ++i) {
            const uint32_t base = (i / 16) * 512 + (i % 16) * 16;
            p[base] = 0x3F800000u;
            p[base + 1] = 0x3F800000u;
        }
    }
    uint32_t zero_ptr = 0;
    if constexpr (OKt > 3 * Kt + 1) {
        cb_reserve_back(cb_zero, 1);
        zero_ptr = get_write_ptr(cb_zero);
        auto* z = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(zero_ptr);
        for (uint32_t i = 0; i < 1024; ++i) {
            z[i] = 0;
        }
    }
    for (uint32_t t = t0; t < t0 + n; ++t) {
        const uint32_t r = t / Kt, j = t - r * Kt;
        const uint32_t ob = r * OKt;
        cb_wait_front(cb_hi, 1);
        const uint32_t hp = get_read_ptr(cb_hi);
        noc_async_write(hp, o.get_noc_addr(ob + j), TB);
        noc_async_write(hp, o.get_noc_addr(ob + Kt + j), TB);
        cb_wait_front(cb_lo, 1);
        noc_async_write(get_read_ptr(cb_lo), o.get_noc_addr(ob + 2 * Kt + j), TB);
        if (j == 0) {
            noc_async_write(one_ptr, o.get_noc_addr(ob + 3 * Kt), TB);
            for (uint32_t z = 3 * Kt + 1; z < OKt; ++z) {
                noc_async_write(zero_ptr, o.get_noc_addr(ob + z), TB);
            }
        }
        noc_async_writes_flushed();
        cb_pop_front(cb_hi, 1);
        cb_pop_front(cb_lo, 1);
    }
    noc_async_write_barrier();
}