changh95's picture
tt-model push diffusion-planner-p150 (container)
be62f78 verified
Raw History Blame Contribute Delete
1.29 kB
// SPDX-License-Identifier: Apache-2.0
// Attention score scale + mask (tt/smask_kernel.py), writer (RISCV_1): the H output tiles (h, m) per mask tile m.
// CT args: [0] H, [1] plane (Mt * Kt), [2] per-core RT-arg count (P3), then the TensorAccessorArgs of out.
// Common RT args: [out_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 auto o_args = TensorAccessorArgs<3>();
constexpr uint32_t cb_out = 16;
constexpr uint32_t TB = 4096;
void kernel_main() {
const uint32_t o_addr = get_common_arg_val<uint32_t>(0);
const uint32_t m0 = get_arg_val<uint32_t>(0);
const uint32_t n = get_arg_val<uint32_t>(1);
const auto o = TensorAccessor(o_args, o_addr, TB);
for (uint32_t m = m0; m < m0 + n; ++m) {
for (uint32_t h = 0; h < H; h += 2) {
cb_wait_front(cb_out, 2);
const uint32_t p = get_read_ptr(cb_out);
noc_async_write(p, o.get_noc_addr(h * plane + m), TB);
noc_async_write(p + TB, o.get_noc_addr((h + 1) * plane + m), TB);
noc_async_writes_flushed();
cb_pop_front(cb_out, 2);
}
}
noc_async_write_barrier();
}