superpoint-p150 / code /kernels /sp_resize /resize_r8_trisc.cpp
changh95's picture
Optimized build (2026-10-03): trace 3.73 -> 0.50 ms
c699c4c verified
Raw History Blame Contribute Delete
2.73 kB
// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// SP_RSZ5: the three TRISCs take 3 of the 5 shares of the core's resize work (resize_r8_fn.inc, prepended; plain scalar
// code, no FPU/SFPU). They start when the data-movement PROC 0 pushes the scratch CB (coefficients and source rows in
// L1, output shard address stored in word 8 of the flag area). RT args as resize_r8.cpp.
#include <cstdint>
#include "api/compute/compute_kernel_api.h"
#include "api/compute/common.h"
#include "api/compute/cb_api.h"
void kernel_main() {
const uint32_t W = get_arg_val<uint32_t>(2);
const uint32_t kv = get_arg_val<uint32_t>(4);
const uint32_t vt_off = get_arg_val<uint32_t>(5);
const uint32_t Y0 = get_arg_val<uint32_t>(6);
const uint32_t r0 = get_arg_val<uint32_t>(7);
const uint32_t nrows = get_arg_val<uint32_t>(8);
constexpr uint32_t cb = get_compile_time_arg_val(0);
constexpr uint32_t OUT_W = get_compile_time_arg_val(2);
constexpr uint32_t OUT_ROWS = get_compile_time_arg_val(3);
constexpr uint32_t KHMAX = get_compile_time_arg_val(4);
constexpr uint32_t KVMAX = get_compile_time_arg_val(5);
constexpr uint32_t NRMAX = get_compile_time_arg_val(6);
constexpr uint32_t RBMAX = get_compile_time_arg_val(7);
constexpr uint32_t PAGE = 256;
constexpr uint32_t CORE_ROWS = 2 * OUT_ROWS;
constexpr uint32_t HT_BYTES = ((OUT_W * (KHMAX + 1) * 4 + PAGE - 1) / PAGE) * PAGE;
constexpr uint32_t VT_BYTES = ((CORE_ROWS * (KVMAX + 1) * 4 + 2 * PAGE - 1) / PAGE + 1) * PAGE;
constexpr uint32_t TMP_BYTES = (NRMAX + KVMAX + 1) * OUT_W;
#if defined(TRISC_UNPACK)
constexpr uint32_t w = 2;
#elif defined(TRISC_MATH)
constexpr uint32_t w = 3;
#else
constexpr uint32_t w = 4;
#endif
cb_wait_front(cb, 1);
const uint32_t base = get_tile_address(cb, 0);
const uint32_t vt_l1 = base + HT_BYTES; // PROC 0's copy of the vertical entries
const uint32_t rows_l1 = base + HT_BYTES + 2 * VT_BYTES;
uint8_t* tmp = reinterpret_cast<uint8_t*>(rows_l1 + RBMAX);
volatile uint32_t* flags = reinterpret_cast<volatile uint32_t*>(rows_l1 + RBMAX + TMP_BYTES);
asm volatile("fence" ::: "memory");
const uint32_t out_l1 = flags[8]; // output shard address, stored by PROC 0 before its push
const uint32_t rb = ((W + 63) & ~63u) + 64;
const uint32_t vw0 = vt_off + Y0 * (kv + 1);
const int32_t* vt = reinterpret_cast<const int32_t*>(vt_l1) + (vw0 - (vw0 / 64) * 64);
rsz5_work<RSZ_KH, RSZ_KV, OUT_W, CORE_ROWS>(w, reinterpret_cast<const int32_t*>(base), vt, kv + 1,
reinterpret_cast<const uint8_t*>(rows_l1), rb, nrows, tmp, r0, out_l1);
}