// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc. // SPDX-License-Identifier: Apache-2.0 // // fp32 TILE [KMAX, C] -> ROW_MAJOR rows of ONE bucket tensor [BSTEP * (b + 1), C] (data movement // only), b = ceil(n / BSTEP) - 1 with n = HDR[2]: the host then reads back a single right-sized // tensor (one D2H) instead of several fixed-size chunks. // Unit u = (tile row tr, tile column tc): the tile's 32 rows x 32 columns are written straight from // the tile buffer as 2 x 64 B segments per row (face rows) into bucket row 32*tr + r. // Tile rows at or beyond n are skipped (never read back). Runtime args: s_addr, hdr_addr, nunits, // NB bucket addresses, then the nunits unit ids. #include #include "api/dataflow/dataflow_api.h" void kernel_main() { const uint32_t s_addr = get_arg_val(0); const uint32_t hdr_addr = get_arg_val(1); const uint32_t nunits = get_arg_val(2); constexpr uint32_t cb_scratch = get_compile_time_arg_val(0); constexpr uint32_t C = get_compile_time_arg_val(1); constexpr uint32_t KMAX = get_compile_time_arg_val(2); constexpr uint32_t BSTEP = get_compile_time_arg_val(3); constexpr uint32_t NB = KMAX / BSTEP; constexpr uint32_t TC = C / 32; constexpr auto s_args = TensorAccessorArgs<4>(); constexpr auto h_args = TensorAccessorArgs(); constexpr auto o_args = TensorAccessorArgs(); const auto sacc = TensorAccessor(s_args, s_addr, 4096); const auto hacc = TensorAccessor(h_args, hdr_addr, (16 + 4 * KMAX) * 4); const uint32_t base = get_write_ptr(cb_scratch); const uint32_t hb = base; const uint32_t tbuf = base + 64; noc_async_read(hacc.get_noc_addr(0), hb, 64); noc_async_read_barrier(); uint32_t n = reinterpret_cast(hb)[2]; if (n == 0) { return; } if (n > KMAX) { n = KMAX; } const uint32_t o_addr = get_arg_val(3 + (n + BSTEP - 1) / BSTEP - 1); const auto oacc = TensorAccessor(o_args, o_addr, C * 4); for (uint32_t ui = 0; ui < nunits; ++ui) { const uint32_t u = get_arg_val(3 + NB + ui); const uint32_t tr = u / TC, tc = u % TC; if (tr * 32 >= n) { continue; } noc_async_read(sacc.get_noc_addr(tr * TC + tc), tbuf, 4096); noc_async_read_barrier(); const uint32_t r0 = tr * 32; for (uint32_t r = 0; r < 32; ++r) { const uint32_t src = tbuf + (((r >> 4) * 2) * 256 + (r & 15) * 16) * 4; const uint64_t dst = oacc.get_noc_addr(r0 + r) + tc * 128; noc_async_write(src, dst, 64); noc_async_write(src + 1024, dst + 64, 64); } noc_async_write_barrier(); } }