Download code/kernels/sp_resize/resize_r8.cpp from changh95/superpoint-p150: direct link, hf CLI and curl.
- Browser
- Download file 10.2 kB
-
https://huggingface.co/changh95/superpoint-p150/resolve/main/code/kernels/sp_resize/resize_r8.cpp
- Command line
-
hf download hf://changh95/superpoint-p150/code/kernels/sp_resize/resize_r8.cpp
-
curl -L -o resize_r8.cpp https://huggingface.co/changh95/superpoint-p150/resolve/main/code/kernels/sp_resize/resize_r8.cpp
10.2 kB
| // SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc. | |
| // SPDX-License-Identifier: Apache-2.0 | |
| // | |
| // Bit-exact Pillow bilinear resize of one 8-bit plane (H x W, DRAM interleaved, one page per source | |
| // row) to the OUT_H x OUT_W uint8 network input (L1 height-sharded, 2 * OUT_ROWS rows per core), | |
| // written straight into the core's input shard. Integer arithmetic of Pillow's Resample.c (8bpc): | |
| // out = clamp(((1 << 21) + sum(px * k)) >> 22, 0, 255), horizontal pass then vertical pass, with the | |
| // 22-bit fixed-point coefficients computed on the host (models/tt/resize_r8.py). | |
| // | |
| // Coefficient tensor (uint32, 256-byte pages): | |
| // HT at word 0: per output column xx: [xmin, k_0 .. k_{kh-1}] (stride kh + 1) | |
| // VT at word vt_off: per output row yy: [ymin, k_0 .. k_{kv-1}] (stride kv + 1) | |
| // Taps past Pillow's count carry k = 0: they read padding / spare rows and add nothing, so every | |
| // output uses exactly kh (kv) taps and the loops are unrolled per (odd) tap count. | |
| // | |
| // Both data-movement RISCs of every core run this kernel (PROC 0 / 1): | |
| // 1. each reads half of the HT and the source rows of its half of the core's source-row span | |
| // [r0, r1) (the rows the core's 2 * OUT_ROWS output rows read); sync | |
| // 2. horizontal pass of its rows (coefficients of a column held in registers across the rows) | |
| // into the shared [rows, OUT_W] buffer; sync | |
| // 3. vertical pass of its OUT_ROWS output rows (4 pixels per 32-bit load) into the shard. | |
| // Semaphores 0 / 1 are the per-RISC progress counters (reset by PROC 0 at the end). | |
| // | |
| // RT args: src_addr, coef_addr, W, kh, kv, vt_off_words, Y0 (first output row of the core), r0, nrows | |
| // CT args: cb, PROC, OUT_W, OUT_ROWS, KHMAX, KVMAX, NRMAX, RBMAX, then TensorAccessorArgs(src), (coef) | |
| static inline uint32_t clip8(int32_t ss) { | |
| if (ss >= (1 << 30)) { | |
| return 255; | |
| } | |
| if (ss <= 0) { | |
| return 0; | |
| } | |
| return (uint32_t)ss >> 22; | |
| } | |
| template <uint32_t KH, uint32_t OUT_W> | |
| static inline void hpass(const int32_t* ht, const uint8_t* rows, uint32_t rb, uint32_t nrows, uint8_t* tmp) { | |
| for (uint32_t xx = 0; xx < OUT_W; ++xx, ht += KH + 1) { | |
| int32_t k[KH]; | |
| for (uint32_t j = 0; j < KH; ++j) { | |
| k[j] = ht[1 + j]; | |
| } | |
| const uint8_t* s = rows + ht[0]; | |
| uint8_t* o = tmp + xx; | |
| for (uint32_t r = 0; r < nrows; ++r, s += rb, o += OUT_W) { | |
| int32_t ss = 1 << 21; | |
| for (uint32_t j = 0; j < KH; ++j) { | |
| ss += (int32_t)s[j] * k[j]; | |
| } | |
| *o = (uint8_t)clip8(ss); | |
| } | |
| } | |
| } | |
| template <uint32_t KV, uint32_t OUT_W> | |
| static inline void vpass(const int32_t* v, const uint8_t* t0, uint32_t* orow) { | |
| int32_t k[KV]; | |
| for (uint32_t j = 0; j < KV; ++j) { | |
| k[j] = v[j]; | |
| } | |
| for (uint32_t xx = 0; xx < OUT_W; xx += 4) { | |
| int32_t a0 = 1 << 21, a1 = 1 << 21, a2 = 1 << 21, a3 = 1 << 21; | |
| const uint32_t* t = reinterpret_cast<const uint32_t*>(t0 + xx); | |
| for (uint32_t j = 0; j < KV; ++j) { | |
| const uint32_t w = t[j * (OUT_W / 4)]; | |
| a0 += (int32_t)(w & 0xFF) * k[j]; | |
| a1 += (int32_t)((w >> 8) & 0xFF) * k[j]; | |
| a2 += (int32_t)((w >> 16) & 0xFF) * k[j]; | |
| a3 += (int32_t)(w >> 24) * k[j]; | |
| } | |
| orow[xx / 4] = clip8(a0) | (clip8(a1) << 8) | (clip8(a2) << 16) | (clip8(a3) << 24); | |
| } | |
| } | |
| static inline void wait_ge(volatile tt_l1_ptr uint32_t* s, uint32_t v) { | |
| do { | |
| invalidate_l1_cache(); | |
| } while (*s < v); | |
| } | |
| void kernel_main() { | |
| const uint32_t src_addr = get_arg_val<uint32_t>(0); | |
| const uint32_t coef_addr = get_arg_val<uint32_t>(1); | |
| const uint32_t W = get_arg_val<uint32_t>(2); | |
| const uint32_t kh = get_arg_val<uint32_t>(3); | |
| 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); // core's source-row span [r0, r0 + nrows) | |
| const uint32_t nrows = get_arg_val<uint32_t>(8); | |
| constexpr uint32_t cb = get_compile_time_arg_val(0); | |
| constexpr uint32_t PROC = get_compile_time_arg_val(1); | |
| 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); // source-row buffer bytes | |
| 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; | |
| static_assert(OUT_W % 4 == 0, "vertical pass works on 32-bit words"); | |
| constexpr auto s_args = TensorAccessorArgs<8>(); | |
| constexpr auto c_args = TensorAccessorArgs<s_args.next_compile_time_args_offset()>(); | |
| const uint32_t rb = ((W + 63) & ~63u) + 64; // row stride: aligned row + room for padded taps | |
| const auto sacc = TensorAccessor(s_args, src_addr, W); | |
| const auto cacc = TensorAccessor(c_args, coef_addr, PAGE); | |
| const uint32_t base = get_write_ptr(cb); | |
| const uint32_t ht_l1 = base; | |
| const uint32_t vt_l1 = ht_l1 + HT_BYTES + PROC * VT_BYTES; | |
| const uint32_t rows_l1 = ht_l1 + HT_BYTES + 2 * VT_BYTES; | |
| uint8_t* tmp = reinterpret_cast<uint8_t*>(rows_l1 + RBMAX); | |
| // vertical entries of the core's rows: words [vt_off + Y0 * (kv + 1), + CORE_ROWS * (kv + 1)) | |
| const uint32_t vw0 = vt_off + Y0 * (kv + 1); | |
| const uint32_t vw1 = vw0 + CORE_ROWS * (kv + 1); | |
| const uint32_t vp0 = vw0 / 64, vp1 = (vw1 + 63) / 64; | |
| for (uint32_t p = vp0; p < vp1; ++p) { | |
| noc_async_read(cacc.get_noc_addr(p), vt_l1 + (p - vp0) * PAGE, PAGE); | |
| } | |
| const uint32_t hp = (OUT_W * (kh + 1) + 63) / 64; | |
| const uint32_t hsplit = hp / 2; | |
| for (uint32_t p = PROC == 0 ? 0 : hsplit; p < (PROC == 0 ? hsplit : hp); ++p) { | |
| noc_async_read(cacc.get_noc_addr(p), ht_l1 + p * PAGE, PAGE); | |
| } | |
| noc_async_read_barrier(); | |
| const int32_t* vt = reinterpret_cast<const int32_t*>(vt_l1) + (vw0 - vp0 * 64); | |
| const uint32_t half = (nrows + 1) / 2; | |
| const uint32_t ra = PROC == 0 ? 0 : half, rz = PROC == 0 ? half : nrows; | |
| const uint32_t row_bytes = (W + 63) & ~63u; | |
| for (uint32_t r = ra; r < rz; ++r) { | |
| noc_async_read(sacc.get_noc_addr(r0 + r), rows_l1 + r * rb, row_bytes); | |
| } | |
| noc_async_read_barrier(); | |
| volatile tt_l1_ptr uint32_t* s_me = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(get_semaphore(PROC)); | |
| volatile tt_l1_ptr uint32_t* s_ot = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(get_semaphore(1 - PROC)); | |
| // SP_RSZ5: 5 RISCs share the work by output columns (resize_r8_fn.inc rsz5_work). PROC 0 stores the output shard | |
| // address for the TRISCs BEFORE the DM sync below, then (after it, all rows in L1) pushes the scratch CB that starts them. | |
| constexpr uint32_t TMP_BYTES = (NRMAX + KVMAX + 1) * OUT_W; | |
| volatile uint32_t* flags = reinterpret_cast<volatile uint32_t*>(rows_l1 + RBMAX + TMP_BYTES); | |
| if constexpr (PROC == 0) { | |
| flags[8] = get_write_ptr(cb + 1); | |
| asm volatile("fence" ::: "memory"); | |
| } | |
| *s_me = 1; | |
| wait_ge(s_ot, 1); // the whole HT is in L1 | |
| if constexpr (PROC == 0) { | |
| cb_reserve_back(cb, 1); | |
| cb_push_back(cb, 1); | |
| } | |
| rsz5_work<RSZ_KH, RSZ_KV, OUT_W, CORE_ROWS>(PROC, reinterpret_cast<const int32_t*>(ht_l1), vt, kv + 1, | |
| reinterpret_cast<const uint8_t*>(rows_l1), rb, nrows, tmp, r0, | |
| get_write_ptr(cb + 1)); | |
| *s_me = 3; | |
| if constexpr (PROC == 0) { | |
| wait_ge(s_ot, 3); | |
| *s_me = 0; | |
| *s_ot = 0; | |
| } | |
| return; | |
| // horizontal pass of rows [ra, rz) | |
| const int32_t* ht = reinterpret_cast<const int32_t*>(ht_l1); | |
| const uint8_t* rows = reinterpret_cast<const uint8_t*>(rows_l1) + ra * rb; | |
| uint8_t* tp = tmp + ra * OUT_W; | |
| const uint32_t n = rz - ra; | |
| switch (kh) { | |
| case 3: hpass<3, OUT_W>(ht, rows, rb, n, tp); break; | |
| case 5: hpass<5, OUT_W>(ht, rows, rb, n, tp); break; | |
| case 7: hpass<7, OUT_W>(ht, rows, rb, n, tp); break; | |
| case 9: hpass<9, OUT_W>(ht, rows, rb, n, tp); break; | |
| case 11: hpass<11, OUT_W>(ht, rows, rb, n, tp); break; | |
| case 13: hpass<13, OUT_W>(ht, rows, rb, n, tp); break; | |
| case 15: hpass<15, OUT_W>(ht, rows, rb, n, tp); break; | |
| default: break; | |
| } | |
| *s_me = 2; | |
| wait_ge(s_ot, 2); // all horizontal rows done | |
| // vertical pass of my OUT_ROWS output rows, straight into the shard | |
| uint32_t* out = reinterpret_cast<uint32_t*>(get_write_ptr(cb + 1) + PROC * OUT_ROWS * OUT_W); | |
| for (uint32_t i = 0; i < OUT_ROWS; ++i) { | |
| const int32_t* v = vt + (PROC * OUT_ROWS + i) * (kv + 1); | |
| const uint8_t* t0 = tmp + ((uint32_t)v[0] - r0) * OUT_W; | |
| uint32_t* orow = out + i * (OUT_W / 4); | |
| switch (kv) { | |
| case 3: vpass<3, OUT_W>(v + 1, t0, orow); break; | |
| case 5: vpass<5, OUT_W>(v + 1, t0, orow); break; | |
| case 7: vpass<7, OUT_W>(v + 1, t0, orow); break; | |
| case 9: vpass<9, OUT_W>(v + 1, t0, orow); break; | |
| case 11: vpass<11, OUT_W>(v + 1, t0, orow); break; | |
| case 13: vpass<13, OUT_W>(v + 1, t0, orow); break; | |
| case 15: vpass<15, OUT_W>(v + 1, t0, orow); break; | |
| case 17: vpass<17, OUT_W>(v + 1, t0, orow); break; | |
| case 19: vpass<19, OUT_W>(v + 1, t0, orow); break; | |
| case 21: vpass<21, OUT_W>(v + 1, t0, orow); break; | |
| case 23: vpass<23, OUT_W>(v + 1, t0, orow); break; | |
| default: break; | |
| } | |
| } | |
| *s_me = 3; | |
| if constexpr (PROC == 0) { | |
| wait_ge(s_ot, 3); | |
| *s_me = 0; | |
| *s_ot = 0; | |
| } | |
| } | |