Download code/tt_diffusion_planner/tt/kernels/smask_reader.cpp from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 2.15 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/kernels/smask_reader.cpp
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/tt/kernels/smask_reader.cpp
-
curl -L -o smask_reader.cpp https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/kernels/smask_reader.cpp
2.15 kB
| // SPDX-License-Identifier: Apache-2.0 | |
| // Attention score scale + mask (tt/smask_kernel.py), reader (RISCV_0): the scale tile once, then per mask tile m of | |
| // this core's range the bf16 mask tile and the H fp32 score tiles (h, m) (the mask broadcasts over the heads). | |
| // CT args: [0] H, [1] mask tiles per head plane (Mt * Kt), [2] scale bits (fp32), [3] per-core RT-arg count (P3), | |
| // [4] mask tile bytes (bf16 2048 / fp32 4096), then the TensorAccessorArgs of scores and mask. | |
| // Common RT args: [s_addr, m_addr]. Per-core RT args: [m0, n]. | |
| constexpr uint32_t H = get_compile_time_arg_val(0); | |
| constexpr uint32_t plane = get_compile_time_arg_val(1); | |
| constexpr uint32_t scale_bits = get_compile_time_arg_val(2); | |
| constexpr uint32_t TBM = get_compile_time_arg_val(4); | |
| constexpr auto s_args = TensorAccessorArgs<5>(); | |
| constexpr auto m_args = TensorAccessorArgs<s_args.next_compile_time_args_offset()>(); | |
| constexpr uint32_t cb_s = 0, cb_m = 1, cb_c = 2; | |
| constexpr uint32_t TB = 4096; | |
| void kernel_main() { | |
| const uint32_t s_addr = get_common_arg_val<uint32_t>(0); | |
| const uint32_t m_addr = get_common_arg_val<uint32_t>(1); | |
| const uint32_t m0 = get_arg_val<uint32_t>(0); | |
| const uint32_t n = get_arg_val<uint32_t>(1); | |
| if (n == 0) { | |
| return; | |
| } | |
| const auto s = TensorAccessor(s_args, s_addr, TB); | |
| const auto mk = TensorAccessor(m_args, m_addr, TBM); | |
| cb_reserve_back(cb_c, 1); | |
| { | |
| auto* p = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(get_write_ptr(cb_c)); | |
| for (uint32_t i = 0; i < 1024; ++i) { | |
| p[i] = scale_bits; | |
| } | |
| } | |
| cb_push_back(cb_c, 1); | |
| for (uint32_t m = m0; m < m0 + n; ++m) { | |
| cb_reserve_back(cb_m, 1); | |
| noc_async_read(mk.get_noc_addr(m), get_write_ptr(cb_m), TBM); | |
| cb_reserve_back(cb_s, H); | |
| const uint32_t p = get_write_ptr(cb_s); | |
| for (uint32_t h = 0; h < H; ++h) { | |
| noc_async_read(s.get_noc_addr(h * plane + m), p + h * TB, TB); | |
| } | |
| noc_async_read_barrier(); | |
| cb_push_back(cb_m, 1); | |
| cb_push_back(cb_s, H); | |
| } | |
| } | |