Download code/tt_diffusion_planner/tt/kernels/smsm_writer.cpp from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 1.8 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/kernels/smsm_writer.cpp
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/tt/kernels/smsm_writer.cpp
-
curl -L -o smsm_writer.cpp https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/kernels/smsm_writer.cpp
1.8 kB
| // SPDX-License-Identifier: Apache-2.0 | |
| // Attention score scale + mask + softmax (tt/smsm_kernel.py), writer (RISCV_1): per tile row r of this core's range | |
| // the Wt probability tiles in the compute's blocks of ndst (the last one clamped), then the row's pad tiles (pushed | |
| // by the compute to keep cb_out aligned; not written). | |
| // CT args: [0] Wt, [1] ndst, [2] per-core RT-arg count (P3), [3] out pad tiles per row, then the TensorAccessorArgs of | |
| // out. Common RT args: [out_addr]. Per-core RT args: [r0, n]. | |
| constexpr uint32_t Wt = get_compile_time_arg_val(0); | |
| constexpr uint32_t ndst = get_compile_time_arg_val(1); | |
| constexpr uint32_t out_pad = get_compile_time_arg_val(3); | |
| constexpr auto o_args = TensorAccessorArgs<4>(); | |
| constexpr uint32_t cb_out = 16; | |
| constexpr uint32_t TB = 4096; | |
| void kernel_main() { | |
| const uint32_t o_addr = get_common_arg_val<uint32_t>(0); | |
| const uint32_t r0 = 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); | |
| for (uint32_t r = r0; r < r0 + n; ++r) { | |
| uint32_t t = r * Wt; | |
| for (uint32_t j = 0; j < Wt; j += ndst) { | |
| const uint32_t b = (j + ndst > Wt) ? (Wt - j) : ndst; | |
| cb_wait_front(cb_out, b); | |
| uint32_t p = get_read_ptr(cb_out); | |
| for (uint32_t i = 0; i < b; ++i) { | |
| noc_async_write(p, o.get_noc_addr(t + i), TB); | |
| p += TB; | |
| } | |
| noc_async_writes_flushed(); | |
| cb_pop_front(cb_out, b); | |
| t += b; | |
| } | |
| if constexpr (out_pad > 0) { | |
| cb_wait_front(cb_out, out_pad); | |
| cb_pop_front(cb_out, out_pad); | |
| } | |
| } | |
| noc_async_write_barrier(); | |
| } | |