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