Download code/tt_diffusion_planner/tt/kernels/kcat_writer.cpp from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 2.84 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/kernels/kcat_writer.cpp
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/tt/kernels/kcat_writer.cpp
-
curl -L -o kcat_writer.cpp https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/kernels/kcat_writer.cpp
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]. | |
| 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(); | |
| } | |