// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc. // SPDX-License-Identifier: Apache-2.0 // uint8 image shard -> bf16 shard through a 256-entry table (bf16(fp32(u) / 255), built on the host: // bit-identical to the host fp32 /255 + bf16 cast). Both data-movement RISCs, each one half of the // core's pixels, local L1 only (input and output are height-sharded on the same cores, same pixel split). // CT args: cb_in, cb_out, n_px (pixels per core, multiple of 8), proc (0 BRISC / 1 NCRISC) // The table is compiled in: the host prepends `#define SP_U8_LUT <256 comma-separated bf16 bit patterns>` // (no runtime args, nothing to dispatch per replay). #include #include "api/dataflow/dataflow_api.h" static const uint16_t lut[256] = {SP_U8_LUT}; void kernel_main() { constexpr uint32_t cb_in = get_compile_time_arg_val(0); constexpr uint32_t cb_out = get_compile_time_arg_val(1); constexpr uint32_t n_px = get_compile_time_arg_val(2); constexpr uint32_t proc = get_compile_time_arg_val(3); constexpr uint32_t n_words = n_px / 4; // 4 pixels per input word constexpr uint32_t half = (n_words / 2) & ~3u; constexpr uint32_t w0 = proc == 0 ? 0 : half; constexpr uint32_t w1 = proc == 0 ? half : n_words; volatile tt_l1_ptr uint32_t* in = reinterpret_cast(get_read_ptr(cb_in)); volatile tt_l1_ptr uint32_t* out = reinterpret_cast(get_write_ptr(cb_out)); static_assert((w1 - w0) % 4 == 0, "4-word unroll"); for (uint32_t w = w0; w < w1; w += 4) { // issue the four L1 loads back to back, then the table lookups / stores uint32_t a = in[w], b = in[w + 1], c = in[w + 2], d = in[w + 3]; volatile tt_l1_ptr uint32_t* o = out + 2 * w; o[0] = (uint32_t)lut[a & 0xFF] | ((uint32_t)lut[(a >> 8) & 0xFF] << 16); o[1] = (uint32_t)lut[(a >> 16) & 0xFF] | ((uint32_t)lut[a >> 24] << 16); o[2] = (uint32_t)lut[b & 0xFF] | ((uint32_t)lut[(b >> 8) & 0xFF] << 16); o[3] = (uint32_t)lut[(b >> 16) & 0xFF] | ((uint32_t)lut[b >> 24] << 16); o[4] = (uint32_t)lut[c & 0xFF] | ((uint32_t)lut[(c >> 8) & 0xFF] << 16); o[5] = (uint32_t)lut[(c >> 16) & 0xFF] | ((uint32_t)lut[c >> 24] << 16); o[6] = (uint32_t)lut[d & 0xFF] | ((uint32_t)lut[(d >> 8) & 0xFF] << 16); o[7] = (uint32_t)lut[(d >> 16) & 0xFF] | ((uint32_t)lut[d >> 24] << 16); } }