changh95's picture
tt-model push diffusion-planner-p150 (container)
be62f78 verified
Raw History Blame Contribute Delete
2.84 kB
// 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();
}