Download code/tt_diffusion_planner/tt/kernels/smsm_reader.cpp from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 3.4 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/kernels/smsm_reader.cpp
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/tt/kernels/smsm_reader.cpp
-
curl -L -o smsm_reader.cpp https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/kernels/smsm_reader.cpp
3.4 kB
| // SPDX-License-Identifier: Apache-2.0 | |
| // Attention score scale + mask + softmax (tt/smsm_kernel.py), reader (RISCV_0). Once: the scale tile (fp32 filled | |
| // with the scale) and the two reduce scaler tiles of the stock softmax reader (fp32 1.0 in row 0 of every face, the | |
| // rest 0: dataflow_kernel_lib::calculate_and_prepare_reduce_scaler for MAX / SUM, REDUCE_ROW). Per tile row r of | |
| // this core's range (head h = r / Mt, query tile row mr = r % Mt): the Wt mask tiles of row mr (+ 1 unused tile when | |
| // Wt is odd, so the compute can take the mask in pairs), then the Wt fp32 score tiles in pairs (the last pair of an | |
| // odd row carries one unused tile). | |
| // CT args: [0] Wt, [1] Mt, [2] scale bits (fp32), [3] per-core RT-arg count (P3), [4] mask tile bytes, then the | |
| // TensorAccessorArgs of scores and mask. Common RT args: [s_addr, m_addr]. Per-core RT args: [r0, n]. | |
| constexpr uint32_t Wt = get_compile_time_arg_val(0); | |
| constexpr uint32_t Mt = 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, cb_max_scaler = 3, cb_sum_scaler = 4; | |
| constexpr uint32_t TB = 4096; | |
| constexpr uint32_t Wp = (Wt + 1) / 2; | |
| inline void fill_scaler(uint32_t cb) { | |
| cb_reserve_back(cb, 1); | |
| auto* p = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(get_write_ptr(cb)); | |
| for (uint32_t i = 0; i < 1024; ++i) { | |
| p[i] = 0; | |
| } | |
| for (uint32_t f = 0; f < 4; ++f) { | |
| for (uint32_t c = 0; c < 16; ++c) { | |
| p[f * 256 + c] = 0x3F800000u; | |
| } | |
| } | |
| cb_push_back(cb, 1); | |
| } | |
| 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 r0 = 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); | |
| fill_scaler(cb_max_scaler); | |
| fill_scaler(cb_sum_scaler); | |
| for (uint32_t r = r0; r < r0 + n; ++r) { | |
| const uint32_t mr = r % Mt; | |
| cb_reserve_back(cb_m, 2 * Wp); | |
| { | |
| uint32_t p = get_write_ptr(cb_m); | |
| for (uint32_t j = 0; j < Wt; ++j) { | |
| noc_async_read(mk.get_noc_addr(mr * Wt + j), p, TBM); | |
| p += TBM; | |
| } | |
| } | |
| noc_async_read_barrier(); | |
| cb_push_back(cb_m, 2 * Wp); | |
| uint32_t t = r * Wt; | |
| for (uint32_t pp = 0; pp < Wp; ++pp) { | |
| cb_reserve_back(cb_s, 2); | |
| const uint32_t p = get_write_ptr(cb_s); | |
| noc_async_read(s.get_noc_addr(t), p, TB); | |
| if (2 * pp + 1 < Wt) { | |
| noc_async_read(s.get_noc_addr(t + 1), p + TB, TB); | |
| } | |
| noc_async_read_barrier(); | |
| cb_push_back(cb_s, 2); | |
| t += 2; | |
| } | |
| } | |
| } | |