superpoint-p150 / code /kernels /sp_nms /untilize_chunks.cpp
changh95's picture
Optimized build (2026-10-03): trace 3.73 -> 0.50 ms
c699c4c verified
Raw History Blame Contribute Delete
2.82 kB
// 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 <stdint.h>
#include "api/dataflow/dataflow_api.h"
void kernel_main() {
const uint32_t s_addr = get_arg_val<uint32_t>(0);
const uint32_t hdr_addr = get_arg_val<uint32_t>(1);
const uint32_t nunits = get_arg_val<uint32_t>(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<s_args.next_compile_time_args_offset()>();
constexpr auto o_args = TensorAccessorArgs<h_args.next_compile_time_args_offset()>();
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<const uint32_t*>(hb)[2];
if (n == 0) {
return;
}
if (n > KMAX) {
n = KMAX;
}
const uint32_t o_addr = get_arg_val<uint32_t>(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<uint32_t>(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();
}
}