changh95's picture
tt-model push diffusion-planner-p150 (container)
be62f78 verified
Raw History Blame Contribute Delete
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].
#include <cstdint>
#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<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);
}
}
}