changh95's picture
Optimized build (2026-10-03): trace 3.73 -> 0.50 ms
c699c4c verified
Raw History Blame Contribute Delete
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)
#include <stdint.h>
#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 <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];
#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 <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];
#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<const uint32_t*>(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<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));
#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<volatile uint32_t*>(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<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;
#endif
// 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;
}
}