// 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) #include #include "api/dataflow/dataflow_api.h" 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 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]; #pragma GCC unroll 16 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; #pragma GCC unroll 16 for (uint32_t j = 0; j < KH; ++j) { ss += (int32_t)s[j] * k[j]; } *o = (uint8_t)clip8(ss); } } } template static inline void vpass(const int32_t* v, const uint8_t* t0, uint32_t* orow) { int32_t k[KV]; #pragma GCC unroll 24 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(t0 + xx); #pragma GCC unroll 24 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(0); const uint32_t coef_addr = get_arg_val(1); const uint32_t W = get_arg_val(2); const uint32_t kh = get_arg_val(3); const uint32_t kv = get_arg_val(4); const uint32_t vt_off = get_arg_val(5); const uint32_t Y0 = get_arg_val(6); const uint32_t r0 = get_arg_val(7); // core's source-row span [r0, r0 + nrows) const uint32_t nrows = get_arg_val(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(); 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(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(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(get_semaphore(PROC)); volatile tt_l1_ptr uint32_t* s_ot = reinterpret_cast(get_semaphore(1 - PROC)); #ifdef RSZ5 // 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(rows_l1 + RBMAX + TMP_BYTES); if constexpr (PROC == 0) { flags[8] = get_write_ptr(cb + 1); asm volatile("fence" ::: "memory"); } #endif *s_me = 1; wait_ge(s_ot, 1); // the whole HT is in L1 #ifdef RSZ5 if constexpr (PROC == 0) { cb_reserve_back(cb, 1); cb_push_back(cb, 1); } rsz5_work(PROC, reinterpret_cast(ht_l1), vt, kv + 1, reinterpret_cast(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; #endif // horizontal pass of rows [ra, rz) const int32_t* ht = reinterpret_cast(ht_l1); const uint8_t* rows = reinterpret_cast(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(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; } }