File size: 2,434 Bytes
c699c4c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 | // 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 <stdint.h>
#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<volatile tt_l1_ptr uint32_t*>(get_read_ptr(cb_in));
volatile tt_l1_ptr uint32_t* out = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(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);
}
}
|