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