// 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 #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(0); const uint32_t t0 = get_arg_val(0); const uint32_t n = get_arg_val(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(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(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(); }