// 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]. #include #include "api/dataflow/dataflow_api.h" 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(); constexpr auto v_args = TensorAccessorArgs(); constexpr auto m_args = TensorAccessorArgs(); 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(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(0); const uint32_t k_addr = get_common_arg_val(1); const uint32_t v_addr = get_common_arg_val(2); const uint32_t m_addr = get_common_arg_val(3); const uint32_t u0 = get_arg_val(0); const uint32_t n = get_arg_val(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); } } }