// 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 #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(); 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(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(0); const uint32_t m_addr = get_common_arg_val(1); const uint32_t r0 = get_arg_val(0); const uint32_t n = get_arg_val(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(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; } } }