Download code/tt_diffusion_planner/tt/kernels/fattn_reader.cpp from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 4.54 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/kernels/fattn_reader.cpp
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/tt/kernels/fattn_reader.cpp
-
curl -L -o fattn_reader.cpp https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/kernels/fattn_reader.cpp
4.54 kB
| // SPDX-License-Identifier: Apache-2.0 | |
| // Fused fp32 matmul attention (tt/fattn_kernel.py), reader (RISCV_0). Once: the two reduce scaler tiles of the stock | |
| // softmax (fp32 1.0 in row 0 of every face). Per unit u of this core's range (head h = u / Mt, query tile row | |
| // r = u % Mt): the Wt mask tiles of row r (+ 1 unused tile when Wt is odd), the Q tile (r, h), the K tiles (j, h) in | |
| // pairs (the last pair of an odd row carries one unused tile), then the V tiles (j, h) one by one. | |
| // Tile (row, head) of a source = base + row * row_stride + head * head_stride (tiles): a [1, 1, S, W] projection | |
| // output read in place (row_stride = W / 32, head_stride = 1, base = the column tile of the first head) or a | |
| // [1, H, S, 32] heads tensor (row_stride = 1, head_stride = S / 32). | |
| // CT args: [0] Wt, [1] Mt, [2] per-core RT-arg count (P3), [3] mask tile bytes, [4..12] (base, row_stride, | |
| // head_stride) of Q, K, V, then the TensorAccessorArgs of q, k, v, mask. | |
| // Common RT args: [q_addr, k_addr, v_addr, m_addr]. Per-core RT args: [u0, n]. | |
| constexpr uint32_t Wt = get_compile_time_arg_val(0); | |
| constexpr uint32_t Mt = get_compile_time_arg_val(1); | |
| constexpr uint32_t TBM = get_compile_time_arg_val(3); | |
| constexpr uint32_t qb = get_compile_time_arg_val(4), qr = get_compile_time_arg_val(5), qh = get_compile_time_arg_val(6); | |
| constexpr uint32_t kb = get_compile_time_arg_val(7), kr = get_compile_time_arg_val(8), kh = get_compile_time_arg_val(9); | |
| constexpr uint32_t vb = get_compile_time_arg_val(10), vr = get_compile_time_arg_val(11), | |
| vh = get_compile_time_arg_val(12); | |
| constexpr auto q_args = TensorAccessorArgs<13>(); | |
| constexpr auto k_args = TensorAccessorArgs<q_args.next_compile_time_args_offset()>(); | |
| constexpr auto v_args = TensorAccessorArgs<k_args.next_compile_time_args_offset()>(); | |
| constexpr auto m_args = TensorAccessorArgs<v_args.next_compile_time_args_offset()>(); | |
| constexpr uint32_t cb_q = 0, cb_m = 1, cb_k = 2, cb_max_scaler = 3, cb_sum_scaler = 4, cb_v = 5; | |
| 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 q_addr = get_common_arg_val<uint32_t>(0); | |
| const uint32_t k_addr = get_common_arg_val<uint32_t>(1); | |
| const uint32_t v_addr = get_common_arg_val<uint32_t>(2); | |
| const uint32_t m_addr = get_common_arg_val<uint32_t>(3); | |
| const uint32_t u0 = get_arg_val<uint32_t>(0); | |
| const uint32_t n = get_arg_val<uint32_t>(1); | |
| if (n == 0) { | |
| return; | |
| } | |
| const auto qa = TensorAccessor(q_args, q_addr, TB); | |
| const auto ka = TensorAccessor(k_args, k_addr, TB); | |
| const auto va = TensorAccessor(v_args, v_addr, TB); | |
| const auto ma = TensorAccessor(m_args, m_addr, TBM); | |
| fill_scaler(cb_max_scaler); | |
| fill_scaler(cb_sum_scaler); | |
| for (uint32_t u = u0; u < u0 + n; ++u) { | |
| const uint32_t h = u / Mt; | |
| const uint32_t r = u % 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(ma.get_noc_addr(r * Wt + j), p, TBM); | |
| p += TBM; | |
| } | |
| } | |
| cb_reserve_back(cb_q, 1); | |
| noc_async_read(qa.get_noc_addr(qb + r * qr + h * qh), get_write_ptr(cb_q), TB); | |
| noc_async_read_barrier(); | |
| cb_push_back(cb_m, 2 * Wp); | |
| cb_push_back(cb_q, 1); | |
| for (uint32_t pp = 0; pp < Wp; ++pp) { | |
| const uint32_t j = 2 * pp; | |
| cb_reserve_back(cb_k, 2); | |
| const uint32_t p = get_write_ptr(cb_k); | |
| noc_async_read(ka.get_noc_addr(kb + j * kr + h * kh), p, TB); | |
| if (j + 1 < Wt) { | |
| noc_async_read(ka.get_noc_addr(kb + (j + 1) * kr + h * kh), p + TB, TB); | |
| } | |
| noc_async_read_barrier(); | |
| cb_push_back(cb_k, 2); | |
| } | |
| for (uint32_t j = 0; j < Wt; ++j) { | |
| cb_reserve_back(cb_v, 1); | |
| noc_async_read(va.get_noc_addr(vb + j * vr + h * vh), get_write_ptr(cb_v), TB); | |
| noc_async_read_barrier(); | |
| cb_push_back(cb_v, 1); | |
| } | |
| } | |
| } | |